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 shared.config.settings import settings from shared.database.database import get_db_session from shared.models.identity import User, UserRole from shared.utils.logger import get_logger logger = get_logger(__name__) pwd_context = bcrypt oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login", auto_error=False) def _require_secret_key() -> str: """SECRET_KEY 惰性校验:未配置时给出明确错误,而不是让 jwt.encode/decode 报晦涩 TypeError。""" if not settings.SECRET_KEY: raise RuntimeError("SECRET_KEY 未配置:请在 .env 中设置后重启服务(认证功能不可用)") return settings.SECRET_KEY def verify_password(plain_password: str, hashed_password: str) -> bool: # 比较侧按 bcrypt 语义截断到 72 字节:兼容历史上被截断存储的口令, # 且避免 checkpw 对超长输入直接抛 ValueError(登录会变 500); # 新口令的超长拒绝在 get_password_hash 中完成 password_bytes = plain_password.encode('utf-8')[:72] try: return pwd_context.checkpw(password_bytes, hashed_password.encode('utf-8')) except ValueError: return False def get_password_hash(password: str) -> str: # bcrypt 算法上限 72 字节:超长密码必须显式拒绝,静默截断会改变有效密码 if len(password.encode('utf-8')) > 72: raise ValueError("密码长度超过 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, _require_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, _require_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()