149 lines
4.4 KiB
Python
149 lines
4.4 KiB
Python
from fastapi import Depends, HTTPException, status
|
|
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.models.base import async_session
|
|
from app.models.user import User
|
|
from app.services.auth import decode_access_token, user_must_set_password
|
|
|
|
security = HTTPBearer(auto_error=False)
|
|
|
|
|
|
async def get_db():
|
|
async with async_session() as session:
|
|
try:
|
|
yield session
|
|
await session.commit()
|
|
except Exception:
|
|
await session.rollback()
|
|
raise
|
|
finally:
|
|
await session.close()
|
|
|
|
|
|
async def get_current_user_allow_password_pending(
|
|
credentials: HTTPAuthorizationCredentials | None = Depends(security),
|
|
db: AsyncSession = Depends(get_db),
|
|
) -> User:
|
|
if not credentials:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="未登录",
|
|
)
|
|
|
|
user_id = decode_access_token(credentials.credentials)
|
|
if not user_id:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="登录已过期",
|
|
)
|
|
|
|
# Skip captcha tokens
|
|
if user_id.startswith("captcha:"):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="无效的凭证",
|
|
)
|
|
|
|
result = await db.execute(select(User).where(User.id == user_id).limit(1))
|
|
user = result.scalar_one_or_none()
|
|
if not user or not user.is_active:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
|
detail="账号不存在或已禁用",
|
|
)
|
|
return user
|
|
|
|
|
|
async def get_current_user(
|
|
current_user: User = Depends(get_current_user_allow_password_pending),
|
|
) -> User:
|
|
if user_must_set_password(current_user):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail={
|
|
"code": "PASSWORD_REQUIRED",
|
|
"message": "请先设置登录密码",
|
|
},
|
|
)
|
|
return current_user
|
|
|
|
|
|
async def get_optional_current_user(
|
|
credentials: HTTPAuthorizationCredentials | None = Depends(security),
|
|
db: AsyncSession = Depends(get_db),
|
|
) -> User | None:
|
|
if not credentials:
|
|
return None
|
|
|
|
user_id = decode_access_token(credentials.credentials)
|
|
if not user_id:
|
|
return None
|
|
|
|
if user_id.startswith("captcha:"):
|
|
return None
|
|
|
|
result = await db.execute(select(User).where(User.id == user_id).limit(1))
|
|
user = result.scalar_one_or_none()
|
|
if not user or not user.is_active:
|
|
return None
|
|
|
|
if user_must_set_password(user):
|
|
return None
|
|
|
|
return user
|
|
|
|
|
|
async def get_admin_user(
|
|
current_user: User = Depends(get_current_user_allow_password_pending),
|
|
) -> User:
|
|
if not current_user.is_admin or current_user.user_type != "admin":
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="需要管理员权限",
|
|
)
|
|
return current_user
|
|
|
|
|
|
async def get_backend_user(
|
|
current_user: User = Depends(get_current_user_allow_password_pending),
|
|
) -> User:
|
|
if current_user.user_type != "admin":
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="需要后台用户权限",
|
|
)
|
|
return current_user
|
|
|
|
|
|
def require_menu_access(menu_path: str):
|
|
"""检查后台用户是否有指定菜单权限。
|
|
|
|
- 超级管理员 (is_admin=True): 直接放行
|
|
- 非管理员后台用户: 检查 allowed_menus 是否包含 menu_path
|
|
- 非后台用户: 403 拒绝
|
|
|
|
用法:
|
|
@router.get("/stats")
|
|
async def get_stats(admin: User = Depends(require_menu_access("/")), ...):
|
|
"""
|
|
async def _dependency(
|
|
current_user: User = Depends(get_current_user_allow_password_pending),
|
|
) -> User:
|
|
if current_user.user_type != "admin":
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="需要后台用户权限",
|
|
)
|
|
if current_user.is_admin:
|
|
return current_user
|
|
allowed = set(current_user.allowed_menus or [])
|
|
if menu_path not in allowed:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_403_FORBIDDEN,
|
|
detail="需要管理员权限",
|
|
)
|
|
return current_user
|
|
return _dependency
|