Files
WonderQ-Project/WonderQ-Admin/app/auth.py
2026-08-25 22:47:59 +08:00

192 lines
7.3 KiB
Python

from datetime import datetime, timedelta, timezone
from uuid import uuid4
import bcrypt
import jwt
from fastapi import Depends, HTTPException, Request, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from sqlalchemy.orm import Session
from .config import get_settings
from .database import get_db
from .models import AdminUser, Customer
from .redis_session import (
AdminSessionStore,
RedisUnavailableError,
SessionRecord,
get_admin_session_store,
hash_refresh_token,
new_refresh_token,
)
bearer = HTTPBearer(auto_error=False)
def verify_password(password: str, password_hash: str) -> bool:
return bcrypt.checkpw(password.encode("utf-8"), password_hash.encode("utf-8"))
def hash_password(password: str, rounds: int = 12) -> str:
return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt(rounds)).decode("utf-8")
def create_token(user: AdminUser) -> str:
settings = get_settings()
now = datetime.now(timezone.utc)
payload = {
"typ": "admin",
"sub": user.id,
"email": user.email,
"role": user.role,
"iat": now,
"exp": now + timedelta(hours=settings.jwt_expires_hours),
}
return jwt.encode(payload, settings.jwt_secret, algorithm="HS256")
def create_access_token(user: AdminUser, *, session_id: str, access_jti: str) -> str:
settings = get_settings()
now = datetime.now(timezone.utc)
payload = {
"typ": "admin_access",
"aud": "admin",
"sub": user.id,
"email": user.email,
"role": user.role,
"sid": session_id,
"jti": access_jti,
"iat": now,
"exp": now + timedelta(minutes=settings.admin_access_expires_minutes),
}
return jwt.encode(payload, settings.jwt_secret, algorithm="HS256")
def decode_admin_access_token(token: str) -> dict:
return jwt.decode(token, get_settings().jwt_secret, algorithms=["HS256"], options={"verify_aud": False})
def issue_admin_session(user: AdminUser, store: AdminSessionStore) -> tuple[str, str, int]:
settings = get_settings()
now = datetime.now(timezone.utc)
session_id = str(uuid4())
access_jti = str(uuid4())
refresh_token = new_refresh_token()
access_expires_at = now + timedelta(minutes=settings.admin_access_expires_minutes)
refresh_expires_at = now + timedelta(days=settings.admin_refresh_expires_days)
store.create(
session_id=session_id,
user_id=user.id,
access_jti=access_jti,
refresh_hash=hash_refresh_token(refresh_token),
access_expires_at=access_expires_at,
refresh_expires_at=refresh_expires_at,
)
return (
create_access_token(user, session_id=session_id, access_jti=access_jti),
refresh_token,
settings.admin_access_expires_minutes * 60,
)
def rotate_admin_session(user: AdminUser, record: SessionRecord, store: AdminSessionStore) -> tuple[str, str, int] | None:
settings = get_settings()
now = datetime.now(timezone.utc)
access_jti = str(uuid4())
refresh_token = new_refresh_token()
if not store.rotate(
session_id=record.session_id,
refresh_hash=record.refresh_hash,
access_jti=access_jti,
refresh_hash_next=hash_refresh_token(refresh_token),
access_expires_at=now + timedelta(minutes=settings.admin_access_expires_minutes),
refresh_expires_at=now + timedelta(days=settings.admin_refresh_expires_days),
):
return None
return create_access_token(user, session_id=record.session_id, access_jti=access_jti), refresh_token, settings.admin_access_expires_minutes * 60
def create_customer_token(customer: Customer) -> str:
settings = get_settings()
now = datetime.now(timezone.utc)
payload = {
"typ": "customer",
"sub": customer.id,
"iat": now,
"exp": now + timedelta(hours=settings.customer_jwt_expires_hours),
}
return jwt.encode(payload, settings.jwt_secret, algorithm="HS256")
def require_admin(
request: Request,
credentials: HTTPAuthorizationCredentials | None = Depends(bearer),
db: Session = Depends(get_db),
store: AdminSessionStore = Depends(get_admin_session_store),
) -> AdminUser:
if credentials is None:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录后台")
try:
payload = decode_admin_access_token(credentials.credentials)
except jwt.PyJWTError as exc:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录后台") from exc
token_type = payload.get("typ", "admin")
if token_type not in {"admin", "admin_access"}:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录后台")
if token_type == "admin_access":
try:
if not payload.get("sid") or not payload.get("jti") or not store.is_access_active(payload["sid"], payload["jti"]):
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录后台")
except RedisUnavailableError as exc:
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="后台会话服务暂时不可用") from exc
user_id = payload.get("sub")
user = db.get(AdminUser, user_id) if user_id else None
if not user or not user.isActive:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录后台")
request.state.actor_id = user.id
return user
def require_customer(
request: Request,
credentials: HTTPAuthorizationCredentials | None = Depends(bearer),
db: Session = Depends(get_db),
) -> Customer:
if credentials is None:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录")
try:
payload = jwt.decode(credentials.credentials, get_settings().jwt_secret, algorithms=["HS256"])
except jwt.PyJWTError as exc:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录") from exc
if payload.get("typ") != "customer":
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录")
customer_id = payload.get("sub")
customer = db.get(Customer, customer_id) if customer_id else None
if not customer:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录")
request.state.actor_id = customer.id
return customer
def optional_customer(
request: Request,
credentials: HTTPAuthorizationCredentials | None = Depends(bearer),
db: Session = Depends(get_db),
) -> Customer | None:
if credentials is None:
return None
try:
payload = jwt.decode(credentials.credentials, get_settings().jwt_secret, algorithms=["HS256"])
except jwt.PyJWTError as exc:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录") from exc
if payload.get("typ") != "customer":
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录")
customer_id = payload.get("sub")
customer = db.get(Customer, customer_id) if customer_id else None
if not customer:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录")
request.state.actor_id = customer.id
return customer
def get_actor_id(request: Request) -> str | None:
return getattr(request.state, "actor_id", None)