后端模块拆分
This commit is contained in:
@@ -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": "权限已删除"}
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user