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, *, remember_me: bool = False) -> 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, remember_me=remember_me, ) 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)