from fastapi import Request, HTTPException, status, Depends from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select import redis.asyncio as redis from typing import Callable from database import get_db from utils.jwt import decode_token from utils.redis_client import get_redis, check_token_in_redis from models import User security = HTTPBearer() async def get_current_user( request: Request, credentials: HTTPAuthorizationCredentials = Depends(security), db: AsyncSession = Depends(get_db), redis_conn: redis.Redis = Depends(get_redis), ): """获取当前用户""" token = credentials.credentials # 解码令牌 payload = decode_token(token) if not payload: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="无效的认证凭据", headers={"WWW-Authenticate": "Bearer"}, ) user_id = int(payload.get("sub")) # 检查令牌是否在Redis中 is_valid = await check_token_in_redis(redis_conn, user_id, token) if not is_valid: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="令牌已过期或已登出", headers={"WWW-Authenticate": "Bearer"}, ) result = await db.execute(select(User).filter(User.id == user_id)) user = result.scalar_one_or_none() if not user: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="用户不存在") return user def has_permission(required_permission_code: str) -> Callable: """ 权限验证依赖函数工厂 使用示例: @router.get("/protected", dependencies=[Depends(has_permission("user:read"))]) async def protected_route(): ... """ async def permission_checker( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): # 获取用户角色的所有权限 from models import Permission, Role result = await db.execute( select(Permission) .join(Role.permissions) .filter(Role.id == current_user.role_id) ) permissions = result.scalars().all() # 检查是否有 required_permission_code 权限 has_perm = any(perm.code == required_permission_code for perm in permissions) if not has_perm: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=f"没有访问权限,需要权限: {required_permission_code}", ) return current_user return permission_checker def has_any_permission(required_permission_codes: list[str]) -> Callable: """ 多权限验证依赖函数工厂(满足任一权限即可) 使用示例: @router.get("/protected", dependencies=[Depends(has_any_permission(["user:read", "admin:read"]))]) async def protected_route(): ... """ async def permission_checker( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): from models import Permission, Role result = await db.execute( select(Permission) .join(Role.permissions) .filter(Role.id == current_user.role_id) ) permissions = result.scalars().all() user_permission_codes = {perm.code for perm in permissions} # 检查是否有任一需要的权限 has_perm = any(code in user_permission_codes for code in required_permission_codes) if not has_perm: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=f"没有访问权限,需要以下任一权限: {', '.join(required_permission_codes)}", ) return current_user return permission_checker def has_all_permissions(required_permission_codes: list[str]) -> Callable: """ 多权限验证依赖函数工厂(需要满足所有权限) 使用示例: @router.get("/protected", dependencies=[Depends(has_all_permissions(["user:read", "user:write"]))]) async def protected_route(): ... """ async def permission_checker( current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): from models import Permission, Role result = await db.execute( select(Permission) .join(Role.permissions) .filter(Role.id == current_user.role_id) ) permissions = result.scalars().all() user_permission_codes = {perm.code for perm in permissions} # 检查是否满足所有需要的权限 missing_permissions = [code for code in required_permission_codes if code not in user_permission_codes] if missing_permissions: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=f"缺少必要权限: {', '.join(missing_permissions)}", ) return current_user return permission_checker