Files

479 lines
18 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""驱动因子预算 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)}