Files
geMoldInsight/src/services/auth_service.py
T
2026-05-06 10:28:16 +08:00

169 lines
5.5 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 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()