Files
geMoldInsight/src/shared/services/auth_service.py
T
2026-09-16 17:55:04 +08:00

184 lines
6.3 KiB
Python

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.database 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()