"""驱动因子预算 API — 科目/驱动因子模式切换 + 行业包 + 敏感性分析""" import json import math from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.orm import Session from typing import Optional, List from datetime import datetime from app.database import get_db from app.auth_middleware import require_auth, require_role from app.models import SystemConfig, OperationLog from app.models.driver_budget import DriverFactorTemplate, DriverFactorBudget router = APIRouter( prefix="/api/cma/budget/driver", tags=["驱动因子预算"], dependencies=[Depends(require_role("ceo", "finance", "it"))], ) # ────────────────────────────────────────────── # 模板预置数据 # ────────────────────────────────────────────── PRESET_TEMPLATES = [ # 通用收入模板 { "name": "收入预测(通用)", "industry": "general", "category": "revenue", "formula_desc": "收入 = 客户数 × 客单价", "formula_text": "客户数×客单价", "factors": [ {"key": "customer_count", "label": "客户数", "type": "number", "default": 100, "unit": "家"}, {"key": "avg_price", "label": "客单价", "type": "number", "default": 50000, "unit": "元/家"}, ], }, # 通用费用模板 { "name": "费用预测(通用)", "industry": "general", "category": "expense", "formula_desc": "费用 = 人数 × 人均成本", "formula_text": "人数×人均成本", "factors": [ {"key": "headcount", "label": "人数", "type": "number", "default": 50, "unit": "人"}, {"key": "avg_cost_per_head", "label": "人均成本", "type": "number", "default": 8000, "unit": "元/人"}, ], }, # 贸易经销版 — 渠补预算 { "name": "渠补预算(贸易经销版)", "industry": "trade", "category": "expense", "formula_desc": "渠补预算 = 计划维护渠道数 × 平均渠补率 × 平均渠道交易额", "formula_text": "计划维护渠道数×平均渠补率×平均渠道交易额", "factors": [ {"key": "channel_count", "label": "计划维护渠道数", "type": "number", "default": 80, "unit": "家"}, {"key": "channel_subsidy_rate", "label": "平均渠补率", "type": "percent", "default": 75, "unit": "%"}, {"key": "avg_transaction", "label": "平均渠道交易额", "type": "number", "default": 150000, "unit": "元/家"}, ], }, # 贸易经销版 — 收入预测 { "name": "收入预测(贸易经销版)", "industry": "trade", "category": "revenue", "formula_desc": "收入 = 渠道数 × 平均交易额", "formula_text": "渠道数×平均交易额", "factors": [ {"key": "channel_count", "label": "渠道数", "type": "number", "default": 80, "unit": "家"}, {"key": "avg_transaction", "label": "平均交易额", "type": "number", "default": 150000, "unit": "元/家"}, ], }, # IT服务版 — 收入 { "name": "收入预测(IT服务版)", "industry": "it", "category": "revenue", "formula_desc": "收入 = 计划新签客户数×平均合同额 + 续约客户数×续约率×平均合同额", "formula_text": "计划新签客户数×平均合同额+续约客户数×续约率×平均合同额", "factors": [ {"key": "new_customers", "label": "计划新签客户数", "type": "number", "default": 36, "unit": "家/年"}, {"key": "avg_contract_value", "label": "平均合同额", "type": "number", "default": 80000, "unit": "元/家"}, {"key": "renewal_customers", "label": "续约客户数", "type": "number", "default": 260, "unit": "家"}, {"key": "renewal_rate", "label": "续约率", "type": "percent", "default": 90, "unit": "%"}, ], }, # IT服务版 — 销售费用 { "name": "销售费用预算(IT服务版)", "industry": "it", "category": "expense", "formula_desc": "销售费用 = 新签客户数×平均获客成本 + 续约客户数×维护成本", "formula_text": "新签客户数×平均获客成本+续约客户数×维护成本", "factors": [ {"key": "new_customers", "label": "计划新签客户数", "type": "number", "default": 36, "unit": "家/年"}, {"key": "acquisition_cost", "label": "平均获客成本", "type": "number", "default": 3000, "unit": "元/家"}, {"key": "renewal_customers", "label": "续约客户数", "type": "number", "default": 260, "unit": "家"}, {"key": "maintenance_cost", "label": "维护成本", "type": "number", "default": 500, "unit": "元/家"}, ], }, ] # ────────────────────────────────────────────── # 模式切换 # ────────────────────────────────────────────── @router.get("/mode") def get_driver_mode(db: Session = Depends(get_db)): """获取当前预算编制模式: subject / driver""" cfg = db.query(SystemConfig).filter(SystemConfig.config_key == "budget_driver_mode").first() if not cfg: return {"mode": "subject", "label": "科目模式"} try: val = json.loads(cfg.config_value) except (json.JSONDecodeError, TypeError): val = {"mode": "subject"} return val @router.post("/mode") def set_driver_mode( data: dict, db: Session = Depends(get_db), current_user=Depends(require_auth), ): """切换预算编制模式: subject(科目模式) / driver(驱动因子模式)""" mode = data.get("mode", "subject") if mode not in ("subject", "driver"): raise HTTPException(400, "模式必须是 subject 或 driver") cfg = db.query(SystemConfig).filter(SystemConfig.config_key == "budget_driver_mode").first() val = json.dumps({"mode": mode, "label": "驱动因子模式" if mode == "driver" else "科目模式"}) if cfg: cfg.config_value = val else: cfg = SystemConfig( config_key="budget_driver_mode", config_value=val, description="预算编制模式: subject=科目模式, driver=驱动因子模式", ) db.add(cfg) db.commit() log = OperationLog( action="update", target_type="budget", target_id=0, detail=json.dumps({"mode": mode, "action": "切换编制模式"}, ensure_ascii=False), ) db.add(log) db.commit() return {"message": f"已切换为{'驱动因子模式' if mode == 'driver' else '科目模式'}", "mode": mode} # ────────────────────────────────────────────── # 模板管理 # ────────────────────────────────────────────── @router.get("/templates") def list_driver_templates( industry: Optional[str] = Query(None, description="行业: general/trade/it"), category: Optional[str] = Query(None, description="类别: revenue/expense"), db: Session = Depends(get_db), ): """获取驱动因子模板列表(含预置模板)""" # 先从数据库读 query = db.query(DriverFactorTemplate).filter(DriverFactorTemplate.is_active == 1) if industry: query = query.filter(DriverFactorTemplate.industry == industry) if category: query = query.filter(DriverFactorTemplate.category == category) db_templates = query.order_by(DriverFactorTemplate.id.asc()).all() # 合并预置模板(数据库没有则返回预置) if db_templates: result = [] for t in db_templates: result.append({ "id": t.id, "name": t.name, "industry": t.industry, "category": t.category, "formula_desc": t.formula_desc, "formula_text": t.formula_text, "factors": t.factors, "is_preset": False, }) return {"data": result, "total": len(result)} else: # 返回预置模板 filtered = PRESET_TEMPLATES if industry: filtered = [t for t in filtered if t["industry"] == industry] if category: filtered = [t for t in filtered if t["category"] == category] return {"data": filtered, "total": len(filtered)} @router.get("/industries") def list_driver_industries(db: Session = Depends(get_db)): """获取行业包列表""" industries = [ {"key": "general", "label": "通用模板", "icon": "📦"}, {"key": "trade", "label": "贸易经销版", "icon": "🏪"}, {"key": "it", "label": "IT服务版", "icon": "💻"}, ] return {"data": industries} # ────────────────────────────────────────────── # 驱动因子计算 # ────────────────────────────────────────────── def _calculate_formula(formula_text: str, factors: dict, factor_defs: list) -> float: """ 根据公式文本和驱动因子值计算预算。 支持的运算: + - × * / 公式文本中的因子名可以是key或label,函数会自动映射到值。 """ # 替换 × 为 * expr = formula_text.replace("×", "*") # 构建变量映射:key -> value, label -> value factor_map = {} for fd in factor_defs: key = fd["key"] val = factors.get(key, fd.get("default", 0)) if fd.get("type") == "percent": val = float(val) / 100.0 else: val = float(val) factor_map[key] = val # 也映射 label(中文名) factor_map[fd["label"]] = val # 按长度降序替换,避免短名被错误替换 sorted_names = sorted(factor_map.keys(), key=len, reverse=True) for name in sorted_names: expr = expr.replace(name, str(factor_map[name])) try: result = eval(expr, {"__builtins__": {}}, {}) return round(float(result), 2) except Exception as e: raise HTTPException(400, f"公式计算失败: {expr}, 错误: {str(e)}") @router.post("/calculate") def driver_calculate( data: dict, db: Session = Depends(get_db), current_user=Depends(require_auth), ): """ 驱动因子计算预算。 接收驱动因子值,自动按公式计算预算结果。 如果指定了 template_id,使用模板的公式和因子定义。 否则使用 data.name, data.formula_text, data.factors 中的因子定义。 """ template_id = data.get("template_id") factors_input = data.get("factors", {}) # {key: value} period = data.get("period", datetime.now().strftime("%Y-%m")) # 获取模板 formula_text = None factor_defs = [] if template_id: template = db.query(DriverFactorTemplate).filter( DriverFactorTemplate.id == template_id, DriverFactorTemplate.is_active == 1, ).first() if not template: # 检查预置模板 if 1 <= template_id <= len(PRESET_TEMPLATES): preset = PRESET_TEMPLATES[template_id - 1] formula_text = preset["formula_text"] factor_defs = preset["factors"] name = preset["name"] industry = preset["industry"] else: raise HTTPException(404, "模板不存在") else: formula_text = template.formula_text factor_defs = template.factors name = template.name industry = template.industry else: name = data.get("name", "自定义驱动因子预算") industry = data.get("industry", "general") formula_text = data.get("formula_text") factor_defs = data.get("factor_defs", []) if not formula_text: raise HTTPException(400, "缺少 formula_text (公式文本)") if not factor_defs: raise HTTPException(400, "缺少 factor_defs (因子定义)") # 填充默认值 complete_factors = {} for fd in factor_defs: key = fd["key"] val = factors_input.get(key) if val is None: val = fd.get("default", 0) complete_factors[key] = val # 计算结果 result = _calculate_formula(formula_text, complete_factors, factor_defs) # 记录计算结果 budget_record = DriverFactorBudget( name=name, industry=industry, template_id=template_id, factors=complete_factors, calculated_value=result, formula_text=formula_text, period=period, created_by=current_user.name if hasattr(current_user, "name") else "", ) db.add(budget_record) db.commit() db.refresh(budget_record) # 记录操作日志 log = OperationLog( action="calculate", target_type="budget", target_id=budget_record.id, detail=json.dumps({ "name": name, "formula": formula_text, "factors": complete_factors, "result": result, }, ensure_ascii=False), ) db.add(log) db.commit() # 计算原始因子信息(用于前端展示) factor_details = [] for fd in factor_defs: key = fd["key"] factor_details.append({ "key": key, "label": fd["label"], "value": complete_factors[key], "unit": fd.get("unit", ""), "type": fd.get("type", "number"), }) return { "id": budget_record.id, "name": name, "industry": industry, "template_id": template_id, "formula_text": formula_text, "formula_desc": data.get("formula_desc", ""), "factors": factor_details, "calculated_value": result, "period": period, } # ────────────────────────────────────────────── # 敏感性分析 # ────────────────────────────────────────────── @router.post("/sensitivity") def driver_sensitivity( data: dict, db: Session = Depends(get_db), current_user=Depends(require_auth), ): """ 敏感性分析:对指定驱动因子做 ±10%、±20% 变动,显示预算变动。 """ factors_input = data.get("factors", {}) formula_text = data.get("formula_text") factor_defs = data.get("factor_defs", []) target_factor = data.get("target_factor") # 要分析的因子key,不指定则分析所有因子 if not formula_text or not factor_defs: raise HTTPException(400, "缺少 formula_text 或 factor_defs") # 计算基准值 base_result = _calculate_formula(formula_text, factors_input, factor_defs) # 敏感性分析 sensitivity_results = [] factors_to_analyze = factor_defs if target_factor: factors_to_analyze = [fd for fd in factor_defs if fd["key"] == target_factor] if not factors_to_analyze: raise HTTPException(404, f"未找到因子: {target_factor}") for fd in factors_to_analyze: key = fd["key"] base_val = factors_input.get(key, fd.get("default", 0)) variations = [] for pct_change in [-20, -10, 10, 20]: if fd.get("type") == "percent": # 百分比因子:变动比例直接加在百分比上 changed_val = base_val * (1 + pct_change / 100.0) else: changed_val = base_val * (1 + pct_change / 100.0) test_factors = dict(factors_input) test_factors[key] = round(changed_val, 2) try: new_result = _calculate_formula(formula_text, test_factors, factor_defs) delta = round(new_result - base_result, 2) delta_pct = round((delta / base_result * 100) if base_result != 0 else 0, 2) variations.append({ "change_pct": pct_change, "factor_value": round(changed_val, 2), "budget_result": new_result, "delta": delta, "delta_pct": delta_pct, }) except Exception: variations.append({ "change_pct": pct_change, "factor_value": round(changed_val, 2), "budget_result": None, "delta": None, "delta_pct": None, "error": "计算失败", }) sensitivity_results.append({ "factor_key": key, "factor_label": fd["label"], "base_value": base_val, "unit": fd.get("unit", ""), "type": fd.get("type", "number"), "variations": variations, }) return { "base_value": base_result, "formula_text": formula_text, "sensitivity": sensitivity_results, } # ────────────────────────────────────────────── # 历史记录 # ────────────────────────────────────────────── @router.get("/history") def list_driver_history( limit: int = Query(50, ge=1, le=200), db: Session = Depends(get_db), ): """查看最近的驱动因子预算计算记录""" records = db.query(DriverFactorBudget).order_by( DriverFactorBudget.created_at.desc() ).limit(limit).all() result = [] for r in records: result.append({ "id": r.id, "name": r.name, "industry": r.industry, "factors": r.factors, "calculated_value": r.calculated_value, "formula_text": r.formula_text, "period": r.period, "created_at": r.created_at.isoformat() if r.created_at else None, }) return {"data": result, "total": len(result)}