from datetime import datetime, timedelta from typing import Optional from jose import JWTError, ExpiredSignatureError, jwt import bcrypt from fastapi import Depends, HTTPException, status, Request from fastapi.security import OAuth2PasswordBearer from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy import select from sqlalchemy.orm import selectinload from config.settings import settings from database.database import get_db_session from models.database import User, UserRole from utils.logger import get_logger logger = get_logger(__name__) pwd_context = bcrypt oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login", auto_error=False) def verify_password(plain_password: str, hashed_password: str) -> bool: return pwd_context.checkpw(plain_password.encode('utf-8'), hashed_password.encode('utf-8')) def get_password_hash(password: str) -> str: if len(password.encode('utf-8')) > 72: password = password[:72] return pwd_context.hashpw(password.encode('utf-8'), pwd_context.gensalt()).decode('utf-8') def create_access_token(data: dict, expires_delta: Optional[timedelta] = None) -> str: to_encode = data.copy() if expires_delta: expire = datetime.utcnow() + expires_delta else: expire = datetime.utcnow() + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) to_encode.update({"exp": expire}) encoded_jwt = jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM) return encoded_jwt async def get_current_user( token: Optional[str] = Depends(oauth2_scheme), db_session: AsyncSession = Depends(get_db_session) ) -> Optional[User]: if not token: return None try: payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) username: str = payload.get("sub") if username is None: logger.warning(f"[AUTH] Token 中缺少 sub 字段") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Token 格式无效:缺少用户标识", headers={"WWW-Authenticate": "Bearer"}, ) except ExpiredSignatureError: logger.warning(f"[AUTH] Token 已过期") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="登录已过期,请重新登录", headers={"WWW-Authenticate": "Bearer"}, ) except JWTError as e: logger.warning(f"[AUTH] Token 验证失败: {type(e).__name__}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Token 无效,请重新登录", headers={"WWW-Authenticate": "Bearer"}, ) result = await db_session.execute( select(User).options(selectinload(User.user_roles).selectinload(UserRole.role)).where(User.username == username) ) user = result.scalar_one_or_none() if user is None: logger.warning(f"[AUTH] Token 有效但用户不存在: {username}") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="用户账户不存在,请重新登录", headers={"WWW-Authenticate": "Bearer"}, ) if not user.is_active: logger.warning(f"[AUTH] 用户已被禁用: {username}") raise HTTPException(status_code=400, detail="用户已被禁用") return user async def get_current_active_user( current_user: Optional[User] = Depends(get_current_user) ) -> User: if not current_user: logger.warning("[AUTH] 未提供认证信息,拒绝访问") raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="请先登录", headers={"WWW-Authenticate": "Bearer"}, ) return current_user async def get_current_admin_user( current_user: User = Depends(get_current_active_user) ) -> User: if not current_user.is_superuser: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="需要管理员权限" ) return current_user async def authenticate_user(db_session: AsyncSession, username: str, password: str) -> Optional[User]: result = await db_session.execute( select(User).options(selectinload(User.user_roles).selectinload(UserRole.role)).where(User.username == username) ) user = result.scalar_one_or_none() if not user: return None if not verify_password(password, user.hashed_password): return None user.last_login = datetime.utcnow() await db_session.commit() return user async def create_user( db_session: AsyncSession, username: str, email: str, password: str, full_name: Optional[str] = None ) -> User: hashed_password = get_password_hash(password) user = User( username=username, email=email, hashed_password=hashed_password, full_name=full_name, is_active=True ) db_session.add(user) await db_session.commit() await db_session.refresh(user) return user async def get_user_by_username(db_session: AsyncSession, username: str) -> Optional[User]: result = await db_session.execute( select(User).where(User.username == username) ) return result.scalar_one_or_none() async def get_user_by_email(db_session: AsyncSession, email: str) -> Optional[User]: result = await db_session.execute( select(User).where(User.email == email) ) return result.scalar_one_or_none()