479 lines
18 KiB
Python
479 lines
18 KiB
Python
"""驱动因子预算 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)}
|