157 lines
5.7 KiB
Python
157 lines
5.7 KiB
Python
"""
|
||
API 依赖注入模块
|
||
|
||
职责:
|
||
提供 FastAPI 的 Depends 依赖函数,供各路由文件使用,包括:
|
||
- get_current_user:从请求头提取并验证 JWT 令牌,返回当前登录用户。
|
||
- require_roles:角色鉴权,仅允许指定角色访问接口。
|
||
- require_permissions:权限鉴权,仅拥有指定权限码的用户可访问。
|
||
- 各业务服务的工厂函数(get_order_service 等),用于解耦路由与服务实例化。
|
||
"""
|
||
from fastapi import Depends, Header, Response
|
||
from sqlalchemy.orm import Session
|
||
|
||
from backend.app.core.error_codes import ErrorCode
|
||
from backend.app.core.exceptions import AppException
|
||
from backend.app.core.security import should_refresh_token, create_access_token
|
||
from backend.app.db import get_db_session
|
||
from backend.app.services.audit_service import audit_service
|
||
from backend.app.services.auth_service import auth_service
|
||
from backend.app.services.config_service import config_service
|
||
from backend.app.services.customer_service import customer_service
|
||
from backend.app.services.logistics_service import logistics_service
|
||
from backend.app.services.order_service import order_service
|
||
from backend.app.services.product_service import product_service
|
||
from backend.app.services.reminder_service import reminder_service
|
||
from backend.app.services.report_service import report_service
|
||
from backend.app.services.supplier_service import supplier_service
|
||
from backend.app.services.system_service import system_service
|
||
from backend.app.services.ai_service import ai_service
|
||
|
||
|
||
# 统一依赖入口,后续如果切换到真实容器或数据库实现,只需要改这里。
|
||
def get_auth_service():
|
||
"""获取认证服务实例"""
|
||
return auth_service
|
||
|
||
|
||
def get_current_user(
|
||
response: Response, # FastAPI Response 对象,用于设置响应头
|
||
authorization: str | None = Header(default=None), # 请求头中的 Authorization 字段,格式为 "Bearer <token>"
|
||
auth_service=Depends(get_auth_service), # 注入认证服务实例
|
||
session: Session = Depends(get_db_session), # 注入数据库会话
|
||
) -> dict:
|
||
"""从请求头解析 JWT 令牌并获取当前登录用户
|
||
|
||
用途:所有需要认证的接口通过 Depends(get_current_user) 调用。
|
||
请求参数:通过请求头 Authorization 传递 Bearer Token。
|
||
返回值:当前用户信息字典,包含 user_id、role_code、permissions 等。
|
||
权限要求:必须携带有效的 JWT 令牌。
|
||
特性:当 Token 剩余有效期不足 30 天时,自动续期并通过响应头 X-New-Token 返回新 Token。
|
||
"""
|
||
token = (authorization or "").removeprefix("Bearer").strip()
|
||
user = auth_service.get_me(token, session)
|
||
if not user:
|
||
raise AppException(code=ErrorCode.UNAUTHORIZED, message="未登录或登录失效", status_code=401)
|
||
|
||
# 滑动过期:Token 剩余不足 30 天时自动续期
|
||
if should_refresh_token(token, threshold_days=30):
|
||
new_token = create_access_token({
|
||
"user_id": user["user_id"],
|
||
"role_code": user["role_code"],
|
||
})
|
||
response.headers["X-New-Token"] = new_token
|
||
|
||
return user
|
||
|
||
|
||
def require_roles(*role_codes: str):
|
||
"""角色鉴权工厂函数
|
||
|
||
用途:返回一个依赖函数,校验当前用户的角色是否在允许列表中。
|
||
参数:*role_codes - 允许访问的角色编码列表(如 "admin"、"manager"、"salesman")。
|
||
返回值:当前用户信息字典(角色校验通过时)。
|
||
权限要求:当前用户的角色编码必须属于 role_codes 之一。
|
||
"""
|
||
def _require_roles(current_user: dict = Depends(get_current_user)) -> dict:
|
||
if role_codes and current_user.get("role_code") not in role_codes:
|
||
raise AppException(code=ErrorCode.FORBIDDEN, message="无权限访问", status_code=403)
|
||
return current_user
|
||
|
||
return _require_roles
|
||
|
||
|
||
def require_permissions(*permission_codes: str):
|
||
"""权限鉴权工厂函数
|
||
|
||
用途:返回一个依赖函数,校验当前用户是否拥有指定权限码。
|
||
参数:*permission_codes - 所需权限码列表(如 "order:list"、"order:create")。
|
||
返回值:当前用户信息字典(权限校验通过时)。
|
||
权限要求:当前用户的 permissions 列表必须包含所有指定的权限码。
|
||
"""
|
||
def _require_permissions(current_user: dict = Depends(get_current_user)) -> dict:
|
||
if not permission_codes:
|
||
return current_user
|
||
granted = set(current_user.get("permissions") or [])
|
||
missing = [code for code in permission_codes if code not in granted]
|
||
if missing:
|
||
raise AppException(code=ErrorCode.FORBIDDEN, message="缺少必要权限", status_code=403)
|
||
return current_user
|
||
|
||
return _require_permissions
|
||
|
||
|
||
def get_order_service():
|
||
"""获取订单服务实例"""
|
||
return order_service
|
||
|
||
|
||
def get_customer_service():
|
||
"""获取客户服务实例"""
|
||
return customer_service
|
||
|
||
|
||
def get_product_service():
|
||
"""获取产品服务实例"""
|
||
return product_service
|
||
|
||
|
||
def get_logistics_service():
|
||
"""获取物流服务实例"""
|
||
return logistics_service
|
||
|
||
|
||
def get_supplier_service():
|
||
"""获取供应商服务实例"""
|
||
return supplier_service
|
||
|
||
|
||
def get_reminder_service():
|
||
"""获取提醒服务实例"""
|
||
return reminder_service
|
||
|
||
|
||
def get_report_service():
|
||
"""获取报表服务实例"""
|
||
return report_service
|
||
|
||
|
||
def get_config_service():
|
||
"""获取配置服务实例"""
|
||
return config_service
|
||
|
||
|
||
def get_audit_service():
|
||
"""获取审计日志服务实例"""
|
||
return audit_service
|
||
|
||
|
||
def get_system_service():
|
||
"""获取系统服务实例"""
|
||
return system_service
|
||
|
||
|
||
def get_ai_service():
|
||
"""获取 AI 服务实例"""
|
||
return ai_service
|