114 lines
3.9 KiB
Python
114 lines
3.9 KiB
Python
"""
|
|
安全工具函数
|
|
"""
|
|
|
|
from fastapi import Depends, HTTPException, status
|
|
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
|
|
from typing import Optional
|
|
import jwt
|
|
import logging
|
|
from datetime import datetime, timedelta
|
|
from config.settings import settings
|
|
from models.database import User
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from sqlalchemy import select
|
|
from utils.database import get_db
|
|
from typing import Dict, Any
|
|
from jwt.exceptions import ExpiredSignatureError, InvalidSignatureError, PyJWTError
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
security = HTTPBearer()
|
|
|
|
def create_access_token(data: dict, expires_delta: timedelta = None):
|
|
"""创建访问令牌
|
|
|
|
业务要求:普通登录至少 7 天内不掉线。
|
|
因此在未显式传入 expires_delta 时,强制最小有效期为 7 天。
|
|
"""
|
|
to_encode = data.copy()
|
|
|
|
if expires_delta:
|
|
expire = datetime.utcnow() + expires_delta
|
|
else:
|
|
min_expire_minutes = 60 * 24 * 7 # 7 days
|
|
configured_expire_minutes = int(getattr(settings, "ACCESS_TOKEN_EXPIRE_MINUTES", min_expire_minutes) or min_expire_minutes)
|
|
effective_expire_minutes = max(configured_expire_minutes, min_expire_minutes)
|
|
expire = datetime.utcnow() + timedelta(minutes=effective_expire_minutes)
|
|
|
|
to_encode.update({"exp": expire})
|
|
encoded_jwt = jwt.encode(to_encode, settings.JWT_SECRET_KEY, algorithm=settings.JWT_ALGORITHM)
|
|
return encoded_jwt
|
|
|
|
async def verify_token(token: str) -> Dict[str, Any]:
|
|
"""验证JWT令牌并返回用户信息"""
|
|
logger.info(f"Verifying token: {token[:50]}..." if len(token) > 50 else f"Verifying token: {token}")
|
|
|
|
if not token or token == 'undefined':
|
|
logger.warning("Invalid token: token is empty or 'undefined'")
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Invalid token"
|
|
)
|
|
|
|
try:
|
|
payload = jwt.decode(token, settings.JWT_SECRET_KEY, algorithms=[settings.JWT_ALGORITHM])
|
|
user_id: int = payload.get("sub")
|
|
if user_id is None:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Could not validate credentials"
|
|
)
|
|
|
|
logger.info(f"Token verified successfully for user_id: {user_id}")
|
|
|
|
user_info = {
|
|
"id": user_id,
|
|
"phone": payload.get("phone", ""),
|
|
"nickname": payload.get("nickname", ""),
|
|
"is_admin": payload.get("is_admin", False)
|
|
}
|
|
return user_info
|
|
except ExpiredSignatureError:
|
|
logger.warning("Token has expired")
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Token has expired"
|
|
)
|
|
except InvalidSignatureError:
|
|
logger.warning("Invalid token signature")
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Invalid token signature"
|
|
)
|
|
except PyJWTError as e:
|
|
logger.warning(f"JWT validation failed: {str(e)}")
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="Could not validate credentials"
|
|
)
|
|
|
|
async def get_current_user(credentials: HTTPAuthorizationCredentials = Depends(security)) -> dict:
|
|
"""
|
|
获取当前用户信息
|
|
验证JWT token并返回用户信息
|
|
"""
|
|
try:
|
|
token = credentials.credentials
|
|
user_info = await verify_token(token)
|
|
return user_info
|
|
except HTTPException:
|
|
raise
|
|
|
|
# 管理员认证函数
|
|
async def get_current_admin_user(current_user: dict = Depends(get_current_user)) -> dict:
|
|
"""
|
|
获取当前管理员用户信息
|
|
"""
|
|
if not current_user.get("is_admin", False):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="Not enough permissions"
|
|
)
|
|
return current_user
|