84 lines
2.7 KiB
Python
84 lines
2.7 KiB
Python
"""FastAPI 依赖注入"""
|
|
|
|
from typing import Optional
|
|
from fastapi import Depends, HTTPException, Request, status, WebSocket
|
|
from fastapi.security import OAuth2PasswordBearer
|
|
from sqlalchemy.orm import Session
|
|
from jose import JWTError
|
|
|
|
from .database import get_db, SessionLocal
|
|
from .security import decode_access_token
|
|
from .models import User
|
|
from .permissions import get_user_permissions
|
|
|
|
# auto_error=False: 允许 token 为空(后续从 cookie 读取)
|
|
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login", auto_error=False)
|
|
|
|
|
|
def get_current_user(
|
|
request: Request,
|
|
token: Optional[str] = Depends(oauth2_scheme),
|
|
db: Session = Depends(get_db),
|
|
) -> User:
|
|
credentials_exc = HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="无效的认证凭据",
|
|
headers={"WWW-Authenticate": "Bearer"},
|
|
)
|
|
# 优先从 Authorization Bearer header 读取,回退到 httpOnly cookie
|
|
if not token:
|
|
token = request.cookies.get("access_token")
|
|
if not token:
|
|
raise credentials_exc
|
|
try:
|
|
payload = decode_access_token(token)
|
|
if payload is None:
|
|
raise credentials_exc
|
|
user_id_raw = payload.get("sub")
|
|
if user_id_raw is None:
|
|
raise credentials_exc
|
|
user_id = int(user_id_raw)
|
|
except JWTError:
|
|
raise credentials_exc
|
|
|
|
user = db.query(User).filter(User.id == int(user_id)).first()
|
|
if user is None or not user.is_active:
|
|
raise credentials_exc
|
|
return user
|
|
|
|
|
|
def require_permission(permission: str):
|
|
"""权限检查依赖工厂。用法: Depends(require_permission('user:create'))"""
|
|
|
|
def checker(current_user: User = Depends(get_current_user)) -> User:
|
|
if not current_user.is_active:
|
|
raise HTTPException(status_code=403, detail="账号已禁用")
|
|
perms = get_user_permissions(current_user)
|
|
if permission not in perms:
|
|
raise HTTPException(status_code=403, detail=f"无权限: {permission}")
|
|
return current_user
|
|
|
|
return checker
|
|
|
|
|
|
def authenticate_websocket(websocket: WebSocket) -> Optional[User]:
|
|
"""WebSocket 认证:从 cookie 或 query param token 中验证用户身份。
|
|
|
|
Returns:
|
|
User 如果认证成功,None 如果失败。
|
|
"""
|
|
token = websocket.cookies.get("access_token") or websocket.query_params.get("token")
|
|
if not token:
|
|
return None
|
|
payload = decode_access_token(token)
|
|
if not payload or not payload.get("sub"):
|
|
return None
|
|
db = SessionLocal()
|
|
try:
|
|
user = db.query(User).filter(User.id == int(payload["sub"])).first()
|
|
if user and user.is_active:
|
|
return user
|
|
return None
|
|
finally:
|
|
db.close()
|