diff --git a/backend/app/api/dashboard.py b/backend/app/api/dashboard.py index 233be174..ab3d9b74 100644 --- a/backend/app/api/dashboard.py +++ b/backend/app/api/dashboard.py @@ -6,6 +6,7 @@ from datetime import datetime, timedelta from typing import Optional from app.database import get_db from app.auth_middleware import require_auth, require_role +from app.deps import get_entity_id from app.models import KPIDefinition, KPIValue, KPIAlert, User from app.utils.cache import get as cache_get, set as cache_set import json @@ -52,15 +53,18 @@ def period_prefix(period_type: str): return None @router.get("/summary") -def get_dashboard_summary(role: str = Query("ceo"), period: str = Query("month"), db: Session = Depends(get_db)): - cache_key = f"summary:{role}:{period}" +def get_dashboard_summary(role: str = Query("ceo"), period: str = Query("month"), + db: Session = Depends(get_db), entity_id: int = Depends(get_entity_id)): + cache_key = f"summary:{role}:{period}:{entity_id}" cached = cache_get("dashboard", cache_key) if cached: return cached - kpi_total = db.query(func.count(KPIDefinition.id)).filter(KPIDefinition.status == "active").scalar() + kpi_total = db.query(func.count(KPIDefinition.id)).filter( + KPIDefinition.status == "active", KPIDefinition.entity_id == entity_id).scalar() alert_count = db.query(func.count(KPIAlert.id)).filter(KPIAlert.status == "pending").scalar() dims = db.query(KPIDefinition.dimension, func.count(KPIDefinition.id)).filter( - KPIDefinition.status == "active").group_by(KPIDefinition.dimension).all() + KPIDefinition.status == "active", KPIDefinition.entity_id == entity_id + ).group_by(KPIDefinition.dimension).all() # 读取最近一次同步状态(从日志文件最后一行) sync_status = {"last_sync": None, "status": "unknown", "detail": ""} @@ -94,11 +98,14 @@ def get_dashboard_summary(role: str = Query("ceo"), period: str = Query("month") @router.get("/kpis") def get_dashboard_kpis(role: str = Query("ceo"), period: str = Query("month"), start_date: str = Query(None), end_date: str = Query(None), - db: Session = Depends(get_db)): + db: Session = Depends(get_db), entity_id: int = Depends(get_entity_id)): start, end = parse_period(period, start_date, end_date) period_str = start.strftime("%Y-%m") - kpis = db.query(KPIDefinition).filter(KPIDefinition.status == "active").all() + kpis = db.query(KPIDefinition).filter( + KPIDefinition.status == "active", + KPIDefinition.entity_id == entity_id + ).all() result = [] for k in kpis: @@ -213,6 +220,7 @@ def get_finance_analysis( current_user: User = Depends(require_auth), period: str = Query("month"), db: Session = Depends(get_db), + entity_id: int = Depends(get_entity_id), ): """财务工作台分析数据""" period_str = datetime.now().strftime("%Y-%m") @@ -220,6 +228,7 @@ def get_finance_analysis( finance_kpis = db.query(KPIDefinition).filter( KPIDefinition.status == "active", KPIDefinition.dimension == "finance", + KPIDefinition.entity_id == entity_id, ).all() kpi_data = [] @@ -269,7 +278,7 @@ def get_finance_analysis( @router.get("/predict") -def predict_kpis(db: Session = Depends(get_db)): +def predict_kpis(db: Session = Depends(get_db), entity_id: int = Depends(get_entity_id)): """基于历史趋势预测下月KPI值(简单线性回归)""" from datetime import datetime, timedelta @@ -281,7 +290,10 @@ def predict_kpis(db: Session = Depends(get_db)): next_year += 1 next_period = f"{next_year}-{next_month:02d}" - kpis = db.query(KPIDefinition).filter(KPIDefinition.status == "active").all() + kpis = db.query(KPIDefinition).filter( + KPIDefinition.status == "active", + KPIDefinition.entity_id == entity_id + ).all() predictions = [] for k in kpis: diff --git a/backend/app/deps.py b/backend/app/deps.py new file mode 100644 index 00000000..210de261 --- /dev/null +++ b/backend/app/deps.py @@ -0,0 +1,26 @@ +"""多租户公共依赖 — 从请求头/参数读取当前企业ID""" +from fastapi import Request, Header, Query, Depends +from typing import Optional + + +def get_entity_id( + request: Request, + x_entity_id: Optional[str] = Header(None, alias="X-Entity-Id"), + entity_id: Optional[int] = Query(None, ge=1), +) -> int: + """解析当前企业ID:优先query参数 → header → 默认1(酣客) + 前端拦截器已统一附加 X-Entity-Id header 和 entity_id query参数 + """ + if entity_id is not None: + return entity_id + if x_entity_id and x_entity_id.isdigit(): + return int(x_entity_id) + # 兼容 body 中的 entity_id(POST场景) + if request.method in ("POST", "PUT", "PATCH"): + try: + body = request.state.body_json or {} + if body.get("entity_id"): + return int(body["entity_id"]) + except Exception: + pass + return 1 # 默认酣客 diff --git a/frontend/src/api/index.ts b/frontend/src/api/index.ts index 6d527320..daad802f 100644 --- a/frontend/src/api/index.ts +++ b/frontend/src/api/index.ts @@ -8,6 +8,12 @@ const api = axios.create({ api.interceptors.request.use((config) => { const token = localStorage.getItem('cma_token') if (token) config.headers.Authorization = `Bearer ${token}` + // 多租户实时切换:统一附加当前企业ID(解决29个页面未传entity_id的问题) + const entityId = Number(localStorage.getItem('cma_entity_id') || 1) + config.headers['X-Entity-Id'] = String(entityId) + if (config.method === 'get' || config.method === 'GET') { + config.params = { ...(config.params || {}), entity_id: entityId } + } return config })