63 lines
1.9 KiB
Python
63 lines
1.9 KiB
Python
"""认证与 RBAC 依赖。"""
|
|
from __future__ import annotations
|
|
|
|
from fastapi import Depends, HTTPException, status
|
|
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
|
from sqlalchemy.orm import Session
|
|
|
|
from database import get_db
|
|
from models import Permission, Role, RolePermission, User, UserRole
|
|
from services.security import decode_access_token
|
|
|
|
|
|
bearer = HTTPBearer(auto_error=False)
|
|
|
|
|
|
def get_current_user(
|
|
credentials: HTTPAuthorizationCredentials | None = Depends(bearer),
|
|
db: Session = Depends(get_db),
|
|
) -> User:
|
|
if credentials is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="请先登录",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
user_id = decode_access_token(credentials.credentials)
|
|
if user_id is None:
|
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="登录已失效")
|
|
user = db.get(User, user_id)
|
|
if user is None or not user.active:
|
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="用户不可用")
|
|
return user
|
|
|
|
|
|
def user_permissions(db: Session, user: User) -> set[str]:
|
|
rows = (
|
|
db.query(Permission.code)
|
|
.join(RolePermission, RolePermission.permission_id == Permission.id)
|
|
.join(Role, Role.id == RolePermission.role_id)
|
|
.join(UserRole, UserRole.role_id == Role.id)
|
|
.filter(UserRole.user_id == user.id)
|
|
.all()
|
|
)
|
|
return {code for (code,) in rows}
|
|
|
|
|
|
def require_permission(code: str):
|
|
def checker(
|
|
user: User = Depends(get_current_user),
|
|
db: Session = Depends(get_db),
|
|
) -> User:
|
|
if code not in user_permissions(db, user):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="没有权限执行该操作",
|
|
)
|
|
return user
|
|
|
|
return checker
|
|
|
|
|
|
require_learning_user = require_permission("learning:use")
|