from contextvars import ContextVar from typing import Optional from fastapi import Depends, Request from sqlalchemy.ext.asyncio import AsyncSession from app.core.database import get_db from app.models.company import Company from app.services.company import company_service # 当前租户的上下文变量 _current_company_id: ContextVar[Optional[int]] = ContextVar("current_company_id", default=None) def get_current_company_id( request: Request, ) -> Optional[int]: """ 获取当前租户 ID 从请求头 X-Company-ID 获取租户 ID Args: request: FastAPI 请求对象 Returns: 当前租户 ID,未设置则返回 None """ company_id_str = request.headers.get("X-Company-ID") if not company_id_str: return None try: return int(company_id_str) except ValueError: return None def set_current_company_id(company_id: Optional[int]) -> None: """ 设置当前租户 ID Args: company_id: 租户 ID """ _current_company_id.set(company_id) async def get_current_company( request: Request, db: AsyncSession = Depends(get_db), ) -> Optional[Company]: """ 从请求头获取当前企业 支持两种方式: 1. HTTP Header: X-Company-ID 2. JWT Token 中的 company_id(后续实现认证后) Args: request: FastAPI 请求对象 db: 数据库会话 Returns: 当前企业对象,不存在则返回 None """ # 从请求头获取 company_id company_id_str = request.headers.get("X-Company-ID") if not company_id_str: return None try: company_id = int(company_id_str) except ValueError: return None # 设置上下文变量 set_current_company_id(company_id) # 查询企业 company = await company_service.get_company(db, company_id) return company def require_company( company: Optional[Company] = Depends(get_current_company), ) -> Company: """ 要求必须有企业上下文的依赖 Args: company: 当前企业 Returns: 企业对象 Raises: HTTPException: 未提供企业 ID 或企业不存在 """ if not company: from fastapi import HTTPException, status raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail="未提供企业 ID 或企业不存在", ) return company