后端模块拆分

This commit is contained in:
2026-05-29 18:10:08 +08:00
parent 5bb9bc84ea
commit 823a387118
93 changed files with 198 additions and 833 deletions
View File
+549
View File
@@ -0,0 +1,549 @@
from fastapi import APIRouter, Depends, HTTPException, status
from fastapi.security import OAuth2PasswordRequestForm
from sqlalchemy.ext.asyncio import AsyncSession
from pydantic import BaseModel
from typing import Optional, List
from datetime import timedelta
from sqlalchemy import select
from sqlalchemy.orm import selectinload
from shared.database.database import get_db_session
from shared.services.auth_service import (
authenticate_user,
create_access_token,
get_current_active_user,
get_password_hash
)
from shared.models.database import User, Role, Permission, UserRole, RolePermission
from shared.config.settings import settings
from shared.utils.logger import get_logger
logger = get_logger(__name__)
router = APIRouter(prefix="/api/auth", tags=["认证"])
class UserResponse(BaseModel):
id: int
username: str
email: str
full_name: Optional[str]
is_active: bool
roles: List[str]
class Config:
from_attributes = True
class Token(BaseModel):
access_token: str
token_type: str
user: UserResponse
class LoginRequest(BaseModel):
username: str
password: str
class RoleCreate(BaseModel):
code: str
name: str
description: Optional[str] = None
class RoleResponse(BaseModel):
id: int
code: str
name: str
description: Optional[str]
is_system: bool
permissions: List[str]
class Config:
from_attributes = True
class PermissionCreate(BaseModel):
code: str
name: str
module: Optional[str] = None
description: Optional[str] = None
class PermissionResponse(BaseModel):
id: int
code: str
name: str
module: Optional[str]
description: Optional[str]
class Config:
from_attributes = True
class UserCreate(BaseModel):
username: str
email: str
password: str
full_name: Optional[str] = None
role_ids: List[int] = []
class UserUpdate(BaseModel):
email: Optional[str] = None
full_name: Optional[str] = None
is_active: Optional[bool] = None
role_ids: Optional[List[int]] = None
def check_admin(user: User) -> bool:
if not user.is_superuser:
raise HTTPException(status_code=403, detail="需要管理员权限")
return True
@router.post("/login", response_model=Token)
async def login(
form_data: OAuth2PasswordRequestForm = Depends(),
db_session: AsyncSession = Depends(get_db_session)
):
user = await authenticate_user(db_session, form_data.username, form_data.password)
if not user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户名或密码错误",
headers={"WWW-Authenticate": "Bearer"},
)
access_token_expires = timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES)
access_token = create_access_token(
data={"sub": user.username}, expires_delta=access_token_expires
)
return Token(
access_token=access_token,
token_type="bearer",
user=UserResponse(
id=user.id,
username=user.username,
email=user.email,
full_name=user.full_name,
is_active=user.is_active,
roles=[r.code for r in user.roles]
)
)
@router.post("/login/json", response_model=Token)
async def login_json(
login_data: LoginRequest,
db_session: AsyncSession = Depends(get_db_session)
):
user = await authenticate_user(db_session, login_data.username, login_data.password)
if not user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户名或密码错误",
)
access_token_expires = timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES)
access_token = create_access_token(
data={"sub": user.username}, expires_delta=access_token_expires
)
return Token(
access_token=access_token,
token_type="bearer",
user=UserResponse(
id=user.id,
username=user.username,
email=user.email,
full_name=user.full_name,
is_active=user.is_active,
roles=[r.code for r in user.roles]
)
)
@router.get("/me", response_model=UserResponse)
async def get_current_user_info(
current_user: User = Depends(get_current_active_user)
):
return UserResponse(
id=current_user.id,
username=current_user.username,
email=current_user.email,
full_name=current_user.full_name,
is_active=current_user.is_active,
roles=[r.code for r in current_user.roles]
)
@router.post("/logout")
async def logout():
return {"message": "已登出"}
@router.get("/users", response_model=List[UserResponse])
async def list_users(
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(
select(User).options(selectinload(User.user_roles).selectinload(UserRole.role))
)
users = result.scalars().all()
return [
UserResponse(
id=u.id,
username=u.username,
email=u.email,
full_name=u.full_name,
is_active=u.is_active,
roles=[r.code for r in u.roles]
) for u in users
]
@router.post("/users", response_model=UserResponse, status_code=201)
async def create_user(
user_data: UserCreate,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
existing = await db_session.execute(
select(User).where(User.username == user_data.username)
)
if existing.scalar_one_or_none():
raise HTTPException(status_code=400, detail="用户名已存在")
existing_email = await db_session.execute(
select(User).where(User.email == user_data.email)
)
if existing_email.scalar_one_or_none():
raise HTTPException(status_code=400, detail="邮箱已存在")
user = User(
username=user_data.username,
email=user_data.email,
hashed_password=get_password_hash(user_data.password),
full_name=user_data.full_name,
is_active=True
)
db_session.add(user)
await db_session.flush()
for role_id in user_data.role_ids:
user_role = UserRole(user_id=user.id, role_id=role_id)
db_session.add(user_role)
await db_session.commit()
await db_session.refresh(user)
logger.info(f"管理员 {current_user.username} 创建了用户 {user.username}")
return UserResponse(
id=user.id,
username=user.username,
email=user.email,
full_name=user.full_name,
is_active=user.is_active,
roles=[r.code for r in user.roles]
)
@router.put("/users/{user_id}", response_model=UserResponse)
async def update_user(
user_id: int,
user_data: UserUpdate,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if not user:
raise HTTPException(status_code=404, detail="用户不存在")
if user_data.email is not None:
user.email = user_data.email
if user_data.full_name is not None:
user.full_name = user_data.full_name
if user_data.is_active is not None:
user.is_active = user_data.is_active
if user_data.role_ids is not None:
await db_session.execute(
select(UserRole).where(UserRole.user_id == user_id)
)
for ur in (await db_session.execute(select(UserRole).where(UserRole.user_id == user_id))).scalars().all():
await db_session.delete(ur)
for role_id in user_data.role_ids:
user_role = UserRole(user_id=user.id, role_id=role_id)
db_session.add(user_role)
await db_session.commit()
await db_session.refresh(user)
logger.info(f"管理员 {current_user.username} 更新了用户 {user.username}")
return UserResponse(
id=user.id,
username=user.username,
email=user.email,
full_name=user.full_name,
is_active=user.is_active,
roles=[r.code for r in user.roles]
)
@router.delete("/users/{user_id}")
async def delete_user(
user_id: int,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if not user:
raise HTTPException(status_code=404, detail="用户不存在")
if user.id == current_user.id:
raise HTTPException(status_code=400, detail="不能删除自己的账户")
username = user.username
await db_session.delete(user)
await db_session.commit()
logger.info(f"管理员 {current_user.username} 删除了用户 {username}")
return {"message": "用户已删除"}
@router.put("/users/{user_id}/reset-password")
async def reset_user_password(
user_id: int,
new_password: str,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if not user:
raise HTTPException(status_code=404, detail="用户不存在")
user.hashed_password = get_password_hash(new_password)
await db_session.commit()
logger.info(f"管理员 {current_user.username} 重置了用户 {user.username} 的密码")
return {"message": "密码已重置"}
@router.get("/roles", response_model=List[RoleResponse])
async def list_roles(
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(Role))
roles = result.scalars().all()
return [
RoleResponse(
id=r.id,
code=r.code,
name=r.name,
description=r.description,
is_system=r.is_system,
permissions=[p.code for p in r.permissions]
) for r in roles
]
@router.post("/roles", response_model=RoleResponse, status_code=201)
async def create_role(
role_data: RoleCreate,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
existing = await db_session.execute(
select(Role).where(Role.code == role_data.code)
)
if existing.scalar_one_or_none():
raise HTTPException(status_code=400, detail="角色编码已存在")
role = Role(
code=role_data.code,
name=role_data.name,
description=role_data.description
)
db_session.add(role)
await db_session.commit()
await db_session.refresh(role)
logger.info(f"管理员 {current_user.username} 创建了角色 {role.code}")
return RoleResponse(
id=role.id,
code=role.code,
name=role.name,
description=role.description,
is_system=role.is_system,
permissions=[]
)
@router.put("/roles/{role_id}", response_model=RoleResponse)
async def update_role(
role_id: int,
role_data: RoleCreate,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(Role).where(Role.id == role_id))
role = result.scalar_one_or_none()
if not role:
raise HTTPException(status_code=404, detail="角色不存在")
if role.is_system:
raise HTTPException(status_code=400, detail="系统角色不能修改")
role.name = role_data.name
role.description = role_data.description
await db_session.commit()
await db_session.refresh(role)
return RoleResponse(
id=role.id,
code=role.code,
name=role.name,
description=role.description,
is_system=role.is_system,
permissions=[p.code for p in role.permissions]
)
@router.delete("/roles/{role_id}")
async def delete_role(
role_id: int,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(Role).where(Role.id == role_id))
role = result.scalar_one_or_none()
if not role:
raise HTTPException(status_code=404, detail="角色不存在")
if role.is_system:
raise HTTPException(status_code=400, detail="系统角色不能删除")
await db_session.delete(role)
await db_session.commit()
logger.info(f"管理员 {current_user.username} 删除了角色 {role.code}")
return {"message": "角色已删除"}
@router.put("/roles/{role_id}/permissions")
async def set_role_permissions(
role_id: int,
permission_ids: List[int],
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(Role).where(Role.id == role_id))
role = result.scalar_one_or_none()
if not role:
raise HTTPException(status_code=404, detail="角色不存在")
for rp in (await db_session.execute(select(RolePermission).where(RolePermission.role_id == role_id))).scalars().all():
await db_session.delete(rp)
for perm_id in permission_ids:
rp = RolePermission(role_id=role_id, permission_id=perm_id)
db_session.add(rp)
await db_session.commit()
logger.info(f"管理员 {current_user.username} 更新了角色 {role.code} 的权限")
return {"message": "权限已更新"}
@router.get("/permissions", response_model=List[PermissionResponse])
async def list_permissions(
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(Permission))
permissions = result.scalars().all()
return [PermissionResponse.from_orm(p) for p in permissions]
@router.post("/permissions", response_model=PermissionResponse, status_code=201)
async def create_permission(
perm_data: PermissionCreate,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
existing = await db_session.execute(
select(Permission).where(Permission.code == perm_data.code)
)
if existing.scalar_one_or_none():
raise HTTPException(status_code=400, detail="权限编码已存在")
permission = Permission(
code=perm_data.code,
name=perm_data.name,
module=perm_data.module,
description=perm_data.description
)
db_session.add(permission)
await db_session.commit()
await db_session.refresh(permission)
logger.info(f"管理员 {current_user.username} 创建了权限 {permission.code}")
return PermissionResponse.from_orm(permission)
@router.delete("/permissions/{permission_id}")
async def delete_permission(
permission_id: int,
db_session: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user)
):
check_admin(current_user)
result = await db_session.execute(select(Permission).where(Permission.id == permission_id))
permission = result.scalar_one_or_none()
if not permission:
raise HTTPException(status_code=404, detail="权限不存在")
await db_session.delete(permission)
await db_session.commit()
logger.info(f"管理员 {current_user.username} 删除了权限 {permission.code}")
return {"message": "权限已删除"}
+168
View File
@@ -0,0 +1,168 @@
from datetime import datetime, timedelta
from typing import Optional
from jose import JWTError, ExpiredSignatureError, jwt
import bcrypt
from fastapi import Depends, HTTPException, status, Request
from fastapi.security import OAuth2PasswordBearer
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from sqlalchemy.orm import selectinload
from shared.config.settings import settings
from shared.database.database import get_db_session
from shared.models.database import User, UserRole
from shared.utils.logger import get_logger
logger = get_logger(__name__)
pwd_context = bcrypt
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login", auto_error=False)
def verify_password(plain_password: str, hashed_password: str) -> bool:
return pwd_context.checkpw(plain_password.encode('utf-8'), hashed_password.encode('utf-8'))
def get_password_hash(password: str) -> str:
if len(password.encode('utf-8')) > 72:
password = password[:72]
return pwd_context.hashpw(password.encode('utf-8'), pwd_context.gensalt()).decode('utf-8')
def create_access_token(data: dict, expires_delta: Optional[timedelta] = None) -> str:
to_encode = data.copy()
if expires_delta:
expire = datetime.utcnow() + expires_delta
else:
expire = datetime.utcnow() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES)
to_encode.update({"exp": expire})
encoded_jwt = jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
return encoded_jwt
async def get_current_user(
token: Optional[str] = Depends(oauth2_scheme),
db_session: AsyncSession = Depends(get_db_session)
) -> Optional[User]:
if not token:
return None
try:
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
username: str = payload.get("sub")
if username is None:
logger.warning(f"[AUTH] Token 中缺少 sub 字段")
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Token 格式无效:缺少用户标识",
headers={"WWW-Authenticate": "Bearer"},
)
except ExpiredSignatureError:
logger.warning(f"[AUTH] Token 已过期")
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="登录已过期,请重新登录",
headers={"WWW-Authenticate": "Bearer"},
)
except JWTError as e:
logger.warning(f"[AUTH] Token 验证失败: {type(e).__name__}")
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Token 无效,请重新登录",
headers={"WWW-Authenticate": "Bearer"},
)
result = await db_session.execute(
select(User).options(selectinload(User.user_roles).selectinload(UserRole.role)).where(User.username == username)
)
user = result.scalar_one_or_none()
if user is None:
logger.warning(f"[AUTH] Token 有效但用户不存在: {username}")
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户账户不存在,请重新登录",
headers={"WWW-Authenticate": "Bearer"},
)
if not user.is_active:
logger.warning(f"[AUTH] 用户已被禁用: {username}")
raise HTTPException(status_code=400, detail="用户已被禁用")
return user
async def get_current_active_user(
current_user: Optional[User] = Depends(get_current_user)
) -> User:
if not current_user:
logger.warning("[AUTH] 未提供认证信息,拒绝访问")
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="请先登录",
headers={"WWW-Authenticate": "Bearer"},
)
return current_user
async def get_current_admin_user(
current_user: User = Depends(get_current_active_user)
) -> User:
if not current_user.is_superuser:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="需要管理员权限"
)
return current_user
async def authenticate_user(db_session: AsyncSession, username: str, password: str) -> Optional[User]:
result = await db_session.execute(
select(User).options(selectinload(User.user_roles).selectinload(UserRole.role)).where(User.username == username)
)
user = result.scalar_one_or_none()
if not user:
return None
if not verify_password(password, user.hashed_password):
return None
user.last_login = datetime.utcnow()
await db_session.commit()
return user
async def create_user(
db_session: AsyncSession,
username: str,
email: str,
password: str,
full_name: Optional[str] = None
) -> User:
hashed_password = get_password_hash(password)
user = User(
username=username,
email=email,
hashed_password=hashed_password,
full_name=full_name,
is_active=True
)
db_session.add(user)
await db_session.commit()
await db_session.refresh(user)
return user
async def get_user_by_username(db_session: AsyncSession, username: str) -> Optional[User]:
result = await db_session.execute(
select(User).where(User.username == username)
)
return result.scalar_one_or_none()
async def get_user_by_email(db_session: AsyncSession, email: str) -> Optional[User]:
result = await db_session.execute(
select(User).where(User.email == email)
)
return result.scalar_one_or_none()
+224
View File
@@ -0,0 +1,224 @@
# services/redis_task_manager.py
"""Redis 任务管理器 - 替代内存字典,支持 TTL 自动清理"""
import json
import os
from typing import Dict, Any, Optional
from datetime import datetime
import redis.asyncio as aioredis
from shared.utils.logger import get_logger
logger = get_logger(__name__)
class RedisTaskManager:
"""基于 Redis 的任务状态管理"""
_instance: Optional["RedisTaskManager"] = None
def __init__(self):
self._redis: Optional[aioredis.Redis] = None
self._prefix = "moldinsight:task:"
self._ttl = 86400 * 7 # 任务默认保留 7 天
self._connected = False
@classmethod
def get_instance(cls) -> "RedisTaskManager":
if cls._instance is None:
cls._instance = RedisTaskManager()
return cls._instance
async def connect(self):
"""连接 Redis"""
if self._connected and self._redis:
return
host = os.getenv("REDIS_HOST", "szcjw")
port = int(os.getenv("REDIS_PORT", "6379"))
password = os.getenv("REDIS_PASSWORD", "")
db = int(os.getenv("REDIS_DB", "0"))
try:
self._redis = aioredis.Redis(
host=host,
port=port,
password=password if password else None,
db=db,
decode_responses=True,
socket_connect_timeout=5,
socket_timeout=5,
retry_on_timeout=True,
)
# 测试连接
await self._redis.ping()
self._connected = True
logger.info(f"Redis 连接成功: {host}:{port}")
except Exception as e:
logger.error(f"Redis 连接失败: {e},任务状态将使用内存回退")
self._redis = None
self._connected = False
async def disconnect(self):
"""断开 Redis 连接"""
if self._redis:
await self._redis.aclose()
self._redis = None
self._connected = False
logger.info("Redis 连接已断开")
@property
def is_connected(self) -> bool:
return self._connected and self._redis is not None
# ---- 内存回退 ----
_fallback_tasks: Dict[str, Dict[str, Any]] = {}
def _fallback_set(self, task_id: str, data: Dict[str, Any]):
self._fallback_tasks[task_id] = data
def _fallback_get(self, task_id: str) -> Optional[Dict[str, Any]]:
return self._fallback_tasks.get(task_id)
def _fallback_delete(self, task_id: str):
self._fallback_tasks.pop(task_id, None)
def _fallback_all(self) -> Dict[str, Dict[str, Any]]:
return dict(self._fallback_tasks)
def _fallback_count(self) -> int:
return len(self._fallback_tasks)
# ---- 公共接口 ----
async def set_task(self, task_id: str, data: Dict[str, Any], ttl: Optional[int] = None):
"""设置任务数据"""
effective_ttl = ttl or self._ttl
# 确保数据可序列化
serializable = self._make_serializable(data)
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
await self._redis.setex(key, effective_ttl, json.dumps(serializable, ensure_ascii=False))
return
except Exception as e:
logger.warning(f"Redis 写入失败,回退到内存: {e}")
self._fallback_set(task_id, serializable)
async def get_task(self, task_id: str) -> Optional[Dict[str, Any]]:
"""获取任务数据"""
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
raw = await self._redis.get(key)
if raw:
return json.loads(raw)
return None
except Exception as e:
logger.warning(f"Redis 读取失败,回退到内存: {e}")
return self._fallback_get(task_id)
async def update_task(self, task_id: str, updates: Dict[str, Any]):
"""更新任务的部分字段"""
current = await self.get_task(task_id)
if current is None:
logger.warning(f"任务 {task_id} 不存在,无法更新")
return
current.update(self._make_serializable(updates))
await self.set_task(task_id, current)
async def delete_task(self, task_id: str):
"""删除任务"""
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
await self._redis.delete(key)
return
except Exception as e:
logger.warning(f"Redis 删除失败,回退到内存: {e}")
self._fallback_delete(task_id)
async def get_all_tasks(self) -> Dict[str, Dict[str, Any]]:
"""获取所有任务"""
if self.is_connected:
try:
pattern = f"{self._prefix}*"
keys = []
async for key in self._redis.scan_iter(match=pattern):
keys.append(key)
result = {}
for key in keys:
task_id = key.replace(self._prefix, "")
raw = await self._redis.get(key)
if raw:
result[task_id] = json.loads(raw)
return result
except Exception as e:
logger.warning(f"Redis 扫描失败,回退到内存: {e}")
return self._fallback_all()
async def get_task_count(self) -> int:
"""获取任务总数"""
if self.is_connected:
try:
pattern = f"{self._prefix}*"
count = 0
async for _ in self._redis.scan_iter(match=pattern):
count += 1
return count
except Exception as e:
logger.warning(f"Redis 计数失败,回退到内存: {e}")
return self._fallback_count()
async def cleanup_old_tasks(self, max_age_seconds: int = 86400 * 7):
"""清理过期任务(Redis 由 TTL 自动管理,内存回退需手动清理)"""
now = datetime.now()
to_delete = []
for task_id, task in self._fallback_tasks.items():
completed_at = task.get("completed_at")
if completed_at:
try:
completed_dt = datetime.fromisoformat(completed_at)
if (now - completed_dt).total_seconds() > max_age_seconds:
to_delete.append(task_id)
except (ValueError, TypeError):
pass
for task_id in to_delete:
del self._fallback_tasks[task_id]
if to_delete:
logger.info(f"清理了 {len(to_delete)} 个过期内存任务")
# ---- 工具方法 ----
@staticmethod
def _make_serializable(obj: Any) -> Any:
"""确保对象可 JSON 序列化"""
if isinstance(obj, dict):
return {k: RedisTaskManager._make_serializable(v) for k, v in obj.items()}
if isinstance(obj, (list, tuple)):
return [RedisTaskManager._make_serializable(v) for v in obj]
if isinstance(obj, datetime):
return obj.isoformat()
if hasattr(obj, "value"):
# Enum 类型
return obj.value
if isinstance(obj, (int, float, str, bool, type(None))):
return obj
return str(obj)
# 全局单例
redis_task_manager = RedisTaskManager.get_instance()