156 lines
5.1 KiB
Python
156 lines
5.1 KiB
Python
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
|