fix: 多租户数据隔离P1a+P1b — 工作台/预警/行动计划/因果链按企业过滤

P1a(有entity表查询补齐):
- kpis get/update/delete/restore 跨企业404校验
- kpis create 强制entity=token企业, update禁止改归属
P1b(无entity表join隔离):
- alerts list/resolve join kpi_definitions 按企业过滤
- action_plans list join过滤 + create校验关联KPI归属
- kpi_causality full-network/list join过滤
- dashboard my_dashboard(用户发现) assigned/preset均按企业隔离
测试: 测试KPI种子entity对齐(2→1), pytest 451 passed
实证: 酣客token 6KPI(无博海id) vs 博海token 1KPI(414) 切换隔离正确
This commit is contained in:
Hermes CI Fix
2026-08-23 17:43:39 +08:00
parent 9c61b02d66
commit 50c15ddbf6
6 changed files with 52 additions and 22 deletions
+9 -2
View File
@@ -7,6 +7,7 @@ import re
import logging import logging
from calendar import monthrange from calendar import monthrange
from app.database import get_db from app.database import get_db
from app.deps import get_entity_id
from app.auth_middleware import require_role, require_auth from app.auth_middleware import require_role, require_auth
from app.models import ActionPlan, KPIAlert, KPIDefinition, User, Objective from app.models import ActionPlan, KPIAlert, KPIDefinition, User, Objective
@@ -88,9 +89,10 @@ def list_plans(
alert_id: Optional[int] = None, alert_id: Optional[int] = None,
db: Session = Depends(get_db), db: Session = Depends(get_db),
current_user: User = Depends(require_auth), current_user: User = Depends(require_auth),
entity_id: int = Depends(get_entity_id),
): ):
"""获取行动计划列表""" """获取行动计划列表(账套隔离: join KPI按企业过滤, 2026-08-23 P1b"""
query = db.query(ActionPlan).order_by(ActionPlan.created_at.desc()) query = db.query(ActionPlan).join(KPIDefinition, KPIDefinition.id == ActionPlan.kpi_id).filter(KPIDefinition.entity_id == entity_id).order_by(ActionPlan.created_at.desc())
if status: if status:
query = query.filter(ActionPlan.status == status) query = query.filter(ActionPlan.status == status)
@@ -123,12 +125,17 @@ def create_plan(
data: dict, data: dict,
db: Session = Depends(get_db), db: Session = Depends(get_db),
current_user: User = Depends(require_auth), current_user: User = Depends(require_auth),
entity_id: int = Depends(get_entity_id),
): ):
"""创建改善行动计划(也是OKR的KR""" """创建改善行动计划(也是OKR的KR"""
required = ["title", "kpi_id"] required = ["title", "kpi_id"]
for field in required: for field in required:
if field not in data: if field not in data:
raise HTTPException(400, f"缺少必填字段: {field}") raise HTTPException(400, f"缺少必填字段: {field}")
# 账套隔离: 关联KPI必须属于当前企业 (2026-08-23 P1b)
kpi_ent = db.query(KPIDefinition).filter(KPIDefinition.id == data["kpi_id"]).first()
if not kpi_ent or kpi_ent.entity_id != entity_id:
raise HTTPException(404, "关联KPI不存在")
due_date = datetime.fromisoformat(data["due_date"]) if data.get("due_date") else None due_date = datetime.fromisoformat(data["due_date"]) if data.get("due_date") else None
+8 -5
View File
@@ -2,8 +2,9 @@
from fastapi import APIRouter, Depends, Query, HTTPException from fastapi import APIRouter, Depends, Query, HTTPException
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.database import get_db from app.database import get_db
from app.deps import get_entity_id
from app.auth_middleware import require_auth, require_role from app.auth_middleware import require_auth, require_role
from app.models import KPIAlert, OperationLog, ActionPlan from app.models import KPIAlert, OperationLog, ActionPlan, KPIDefinition
import logging import logging
logger = logging.getLogger("cma.alerts") logger = logging.getLogger("cma.alerts")
@@ -13,8 +14,9 @@ router = APIRouter(prefix="/api/cma/alerts", tags=["预警"],
) )
@router.get("") @router.get("")
def list_alerts(status: str = None, page: int = Query(1, ge=1), db: Session = Depends(get_db)): def list_alerts(status: str = None, page: int = Query(1, ge=1), db: Session = Depends(get_db), entity_id: int = Depends(get_entity_id)):
query = db.query(KPIAlert) # 账套隔离: join kpi_definitions 按企业过滤 (2026-08-23 P1b)
query = db.query(KPIAlert).join(KPIDefinition, KPIDefinition.id == KPIAlert.kpi_id).filter(KPIDefinition.entity_id == entity_id)
if status: if status:
query = query.filter(KPIAlert.status == status) query = query.filter(KPIAlert.status == status)
total = query.count() total = query.count()
@@ -22,8 +24,9 @@ def list_alerts(status: str = None, page: int = Query(1, ge=1), db: Session = De
return {"total": total, "data": [{c.name: getattr(a, c.name) for c in KPIAlert.__table__.columns} for a in alerts]} return {"total": total, "data": [{c.name: getattr(a, c.name) for c in KPIAlert.__table__.columns} for a in alerts]}
@router.post("/{alert_id}/resolve") @router.post("/{alert_id}/resolve")
def resolve_alert(alert_id: int, data: dict, db: Session = Depends(get_db)): def resolve_alert(alert_id: int, data: dict, db: Session = Depends(get_db), entity_id: int = Depends(get_entity_id)):
alert = db.query(KPIAlert).filter(KPIAlert.id == alert_id).first() alert = db.query(KPIAlert).join(KPIDefinition, KPIDefinition.id == KPIAlert.kpi_id).filter(
KPIAlert.id == alert_id, KPIDefinition.entity_id == entity_id).first()
if not alert: if not alert:
raise HTTPException(404, "预警不存在") raise HTTPException(404, "预警不存在")
alert.status = "resolved" alert.status = "resolved"
+6 -3
View File
@@ -390,8 +390,9 @@ def predict_kpis(db: Session = Depends(get_db), entity_id: int = Depends(get_ent
def my_dashboard( def my_dashboard(
current_user: User = Depends(require_auth), current_user: User = Depends(require_auth),
db: Session = Depends(get_db), db: Session = Depends(get_db),
entity_id: int = Depends(get_entity_id),
): ):
"""个人工作台:返回我的KPI、改善行动、待办提醒""" """个人工作台:返回我的KPI、改善行动、待办提醒(账套隔离: 按token企业过滤)"""
username = current_user.username username = current_user.username
name = current_user.name name = current_user.name
role = current_user.role role = current_user.role
@@ -405,22 +406,24 @@ def my_dashboard(
} }
preset_codes = ROLE_PRESET_KPIS.get(role, []) preset_codes = ROLE_PRESET_KPIS.get(role, [])
# 1. 我的KPIresponsible_user匹配用户名或姓名)+ 角色预设 # 1. 我的KPIresponsible_user匹配用户名或姓名)+ 角色预设(均按企业隔离)
assigned_kpis = db.query(KPIDefinition).filter( assigned_kpis = db.query(KPIDefinition).filter(
or_( or_(
KPIDefinition.responsible_user == username, KPIDefinition.responsible_user == username,
KPIDefinition.responsible_user == name, KPIDefinition.responsible_user == name,
), ),
KPIDefinition.status == "active", KPIDefinition.status == "active",
KPIDefinition.entity_id == entity_id,
).all() ).all()
assigned_ids = {k.id for k in assigned_kpis} assigned_ids = {k.id for k in assigned_kpis}
# 补充角色预设KPI(去重) # 补充角色预设KPI(去重,按企业隔离
preset_kpis = [] preset_kpis = []
if preset_codes: if preset_codes:
q = db.query(KPIDefinition).filter( q = db.query(KPIDefinition).filter(
KPIDefinition.kpi_code.in_(preset_codes), KPIDefinition.kpi_code.in_(preset_codes),
KPIDefinition.status == "active", KPIDefinition.status == "active",
KPIDefinition.entity_id == entity_id,
) )
if assigned_ids: if assigned_ids:
q = q.filter(~KPIDefinition.id.in_(assigned_ids)) q = q.filter(~KPIDefinition.id.in_(assigned_ids))
+7 -5
View File
@@ -8,6 +8,7 @@ from typing import Optional
import logging import logging
from app.database import get_db from app.database import get_db
from app.deps import get_entity_id
from app.auth_middleware import require_auth, require_role from app.auth_middleware import require_auth, require_role
from app.models import KPIDefinition, KPICausality, KPIValue, OperationLog from app.models import KPIDefinition, KPICausality, KPIValue, OperationLog
@@ -28,9 +29,9 @@ def _to_dict(obj):
# ============================================================ # ============================================================
@router.get("/full-network") @router.get("/full-network")
def get_full_network(db: Session = Depends(get_db)): def get_full_network(db: Session = Depends(get_db), entity_id: int = Depends(get_entity_id)):
"""获取全局因果网络数据(用于力导向图)""" """获取全局因果网络数据(用于力导向图)— 账套隔离: 仅当前企业KPI的因果链 (2026-08-23 P1b)"""
edges = db.query(KPICausality).all() edges = db.query(KPICausality).join(KPIDefinition, KPIDefinition.id == KPICausality.source_kpi_id).filter(KPIDefinition.entity_id == entity_id).all()
node_ids = set() node_ids = set()
edge_list = [] edge_list = []
for e in edges: for e in edges:
@@ -212,9 +213,10 @@ def list_causalities(
source_kpi_id: Optional[int] = None, source_kpi_id: Optional[int] = None,
target_kpi_id: Optional[int] = None, target_kpi_id: Optional[int] = None,
db: Session = Depends(get_db), db: Session = Depends(get_db),
entity_id: int = Depends(get_entity_id),
): ):
"""获取因果链列表""" """获取因果链列表(账套隔离: 仅当前企业KPI, 2026-08-23 P1b"""
query = db.query(KPICausality) query = db.query(KPICausality).join(KPIDefinition, KPIDefinition.id == KPICausality.source_kpi_id).filter(KPIDefinition.entity_id == entity_id)
if source_kpi_id: if source_kpi_id:
query = query.filter(KPICausality.source_kpi_id == source_kpi_id) query = query.filter(KPICausality.source_kpi_id == source_kpi_id)
if target_kpi_id: if target_kpi_id:
+20 -6
View File
@@ -475,10 +475,13 @@ def get_kpi_causality_chain(
# ============================================================ # ============================================================
@router.get("/{kpi_id}") @router.get("/{kpi_id}")
def get_kpi(kpi_id: int, db: Session = Depends(get_db)): def get_kpi(kpi_id: int, db: Session = Depends(get_db), entity_id: int = Depends(get_entity_id)):
kpi = db.query(KPIDefinition).filter(KPIDefinition.id == kpi_id).first() kpi = db.query(KPIDefinition).filter(KPIDefinition.id == kpi_id).first()
if not kpi: if not kpi:
raise HTTPException(404, "KPI不存在") raise HTTPException(404, "KPI不存在")
# 账套隔离: 禁止跨企业读取 (2026-08-23 P1a)
if kpi.entity_id != entity_id:
raise HTTPException(404, "KPI不存在")
return kpi_to_dict(kpi) return kpi_to_dict(kpi)
@@ -586,15 +589,16 @@ def apply_calc_type_inference(data: dict, infer_missing: bool = True) -> dict:
@router.post("") @router.post("")
def create_kpi(data: dict, db: Session = Depends(get_db), user=WRITE_ROLES): def create_kpi(data: dict, db: Session = Depends(get_db), user=WRITE_ROLES, entity_id: int = Depends(get_entity_id)):
# 检查编码唯一性 # 检查编码唯一性
existing = db.query(KPIDefinition).filter(KPIDefinition.kpi_code == data.get("kpi_code", "")).first() existing = db.query(KPIDefinition).filter(KPIDefinition.kpi_code == data.get("kpi_code", ""), KPIDefinition.entity_id == entity_id).first()
if existing: if existing:
raise HTTPException(400, f"KPI编码 {data['kpi_code']} 已存在") raise HTTPException(400, f"KPI编码 {data['kpi_code']} 已存在")
# 数据治理校验(规则1强制拦截) # 数据治理校验(规则1强制拦截)
errs = _validate_kpi_data(data, db=db, is_update=False) errs = _validate_kpi_data(data, db=db, is_update=False)
if errs: if errs:
raise HTTPException(422, detail={"message": "数据校验不通过", "errors": errs}) raise HTTPException(422, detail={"message": "数据校验不通过", "errors": errs})
data["entity_id"] = entity_id # 账套隔离: 强制写入token企业 (2026-08-23 P1a)
data = apply_calc_type_inference(data) data = apply_calc_type_inference(data)
kpi = KPIDefinition(**data) kpi = KPIDefinition(**data)
db.add(kpi) db.add(kpi)
@@ -605,14 +609,18 @@ def create_kpi(data: dict, db: Session = Depends(get_db), user=WRITE_ROLES):
@router.put("/{kpi_id}") @router.put("/{kpi_id}")
def update_kpi(kpi_id: int, data: dict, db: Session = Depends(get_db), user=WRITE_ROLES): def update_kpi(kpi_id: int, data: dict, db: Session = Depends(get_db), user=WRITE_ROLES, entity_id: int = Depends(get_entity_id)):
kpi = db.query(KPIDefinition).filter(KPIDefinition.id == kpi_id).first() kpi = db.query(KPIDefinition).filter(KPIDefinition.id == kpi_id).first()
if not kpi: if not kpi:
raise HTTPException(404, "KPI不存在") raise HTTPException(404, "KPI不存在")
# 账套隔离: 禁止跨企业修改 (2026-08-23 P1a)
if kpi.entity_id != entity_id:
raise HTTPException(404, "KPI不存在")
# 数据治理校验(更新时只检查传了但为空的字段) # 数据治理校验(更新时只检查传了但为空的字段)
errs = _validate_kpi_data(data, db=db, current_kpi_id=kpi_id, is_update=True) errs = _validate_kpi_data(data, db=db, current_kpi_id=kpi_id, is_update=True)
if errs: if errs:
raise HTTPException(422, detail={"message": "数据校验不通过", "errors": errs}) raise HTTPException(422, detail={"message": "数据校验不通过", "errors": errs})
data.pop("entity_id", None) # 禁止通过update改企业归属
data = apply_calc_type_inference(data, infer_missing=False) data = apply_calc_type_inference(data, infer_missing=False)
for k, v in data.items(): for k, v in data.items():
if hasattr(kpi, k) and v is not None: if hasattr(kpi, k) and v is not None:
@@ -623,18 +631,24 @@ def update_kpi(kpi_id: int, data: dict, db: Session = Depends(get_db), user=WRIT
@router.delete("/{kpi_id}") @router.delete("/{kpi_id}")
def delete_kpi(kpi_id: int, db: Session = Depends(get_db), user=WRITE_ROLES): def delete_kpi(kpi_id: int, db: Session = Depends(get_db), user=WRITE_ROLES, entity_id: int = Depends(get_entity_id)):
kpi = db.query(KPIDefinition).filter(KPIDefinition.id == kpi_id).first() kpi = db.query(KPIDefinition).filter(KPIDefinition.id == kpi_id).first()
if kpi: if kpi:
# 账套隔离: 禁止跨企业删除 (2026-08-23 P1a)
if kpi.entity_id != entity_id:
raise HTTPException(404, "KPI不存在")
kpi.status = "disabled" kpi.status = "disabled"
db.commit() db.commit()
return {"message": "已删除"} return {"message": "已删除"}
@router.put("/{kpi_id}/restore") @router.put("/{kpi_id}/restore")
def restore_kpi(kpi_id: int, db: Session = Depends(get_db), user=WRITE_ROLES): def restore_kpi(kpi_id: int, db: Session = Depends(get_db), user=WRITE_ROLES, entity_id: int = Depends(get_entity_id)):
kpi = db.query(KPIDefinition).filter(KPIDefinition.id == kpi_id).first() kpi = db.query(KPIDefinition).filter(KPIDefinition.id == kpi_id).first()
if kpi: if kpi:
# 账套隔离: 禁止跨企业恢复 (2026-08-23 P1a)
if kpi.entity_id != entity_id:
raise HTTPException(404, "KPI不存在")
kpi.status = "active" kpi.status = "active"
db.commit() db.commit()
return {"message": "已恢复"} return {"message": "已恢复"}
+2 -1
View File
@@ -17,7 +17,8 @@ BASE = "/api/cma/kpi-causality"
def _seed_kpi(db: Session, code: str, name: str = None, dimension: str = "finance", def _seed_kpi(db: Session, code: str, name: str = None, dimension: str = "finance",
entity_id: int = 2) -> KPIDefinition: entity_id: int = 1) -> KPIDefinition:
# entity_id=1 与测试token(账套entity=1)对齐 — 2026-08-23 多租户隔离后必须匹配
kpi = KPIDefinition( kpi = KPIDefinition(
kpi_code=code, kpi_code=code,
kpi_name=name or code, kpi_name=name or code,