feat: 预测性成本智能MVP — KPI趋势预测引擎(线性回归/移动平均)+API+前端Tab+pytest覆盖

This commit is contained in:
Hermes CI Fix
2026-08-25 00:10:46 +08:00
parent fb9eba38a8
commit 13aa153875
6 changed files with 863 additions and 1 deletions
+308
View File
@@ -0,0 +1,308 @@
"""KPI预测引擎测试 — 预测性成本智能 MVP
覆盖:线性回归预测 / 移动平均预测 / 置信区间 / 趋势方向判定 /
边界情况(历史<2条、无历史、未来期数=0)/ entity_id 多租户隔离。
"""
import pytest
from fastapi.testclient import TestClient
from sqlalchemy.orm import Session
from app.models import Entity, KPIValue
from app.utils import kpi_forecast_engine as engine
from tests.conftest import create_test_kpi, create_test_user, get_token_for_user, auth_header
BASE = "/api/cma/predict"
# ── 测试数据工厂 ──
def _seed_history(db: Session, kpi, values: list, entity_id: int = 1, start_period: str = "2026-01"):
"""为KPI写入连续期数的历史值"""
for i, v in enumerate(values):
month = int(start_period.split("-")[1]) + i
year = int(start_period.split("-")[0]) + (month - 1) // 12
month = (month - 1) % 12 + 1
db.add(KPIValue(
kpi_id=kpi.id, entity_id=entity_id,
period=f"{year:04d}-{month:02d}", actual_value=float(v),
))
db.commit()
def _seed_entity2(db: Session) -> Entity:
ent = db.query(Entity).filter(Entity.id == 2).first()
if not ent:
ent = Entity(id=2, name="博海网络科技", short_name="博海", status="active")
db.add(ent)
db.commit()
return ent
# ── 纯函数单元测试 ──
class TestNextPeriod:
def test_increment(self):
assert engine.next_period("2026-05", 1) == "2026-06"
def test_cross_year(self):
assert engine.next_period("2026-12", 1) == "2027-01"
assert engine.next_period("2026-01", 2) == "2026-03"
def test_malformed_fallback(self):
assert engine.next_period("未知期", 1) == "未知期+1"
class TestLinearForecast:
def test_known_line(self):
"""y = 2x + 1 (x=0..4) → 预测 x=5→11, x=6→13"""
values = [1.0, 3.0, 5.0, 7.0, 9.0]
res = engine.linear_forecast(values, periods=3)
assert len(res["forecast"]) == 3
assert res["forecast"][0]["predicted"] == pytest.approx(11.0, abs=1e-6)
assert res["forecast"][1]["predicted"] == pytest.approx(13.0, abs=1e-6)
assert res["forecast"][2]["predicted"] == pytest.approx(15.0, abs=1e-6)
# 完全拟合 → R²=1
assert res["r2"] == pytest.approx(1.0, abs=1e-6)
def test_confidence_interval_bounds(self):
"""每个预测点 lower ≤ predicted ≤ upper"""
import random
random.seed(42)
values = [random.uniform(50, 150) for _ in range(10)]
res = engine.linear_forecast(values, periods=4)
for f in res["forecast"]:
assert f["lower"] <= f["predicted"] <= f["upper"]
class TestMovingAverageForecast:
def test_window_mean(self):
"""最近3期均值作为预测值"""
values = [10.0, 20.0, 30.0, 40.0, 50.0]
res = engine.moving_average_forecast(values, periods=3)
assert len(res["forecast"]) == 3
assert res["mean"] == pytest.approx(40.0) # (30+40+50)/3
for f in res["forecast"]:
assert f["predicted"] == pytest.approx(40.0)
assert f["lower"] <= f["predicted"] <= f["upper"]
def test_window_smaller_than_history(self):
values = [5.0, 6.0]
res = engine.moving_average_forecast(values, periods=2)
assert res["window"] == 2
assert res["forecast"][0]["predicted"] == pytest.approx(5.5)
class TestTrendAndConfidence:
def test_trend_up(self):
assert engine.judge_trend(8.0) == "up"
def test_trend_down(self):
assert engine.judge_trend(-8.0) == "down"
def test_trend_flat(self):
assert engine.judge_trend(1.0) == "flat"
assert engine.judge_trend(-1.0) == "flat"
def test_confidence_high(self):
assert engine.compute_confidence(20, "linear", r2=0.9) == "high"
def test_confidence_medium(self):
assert engine.compute_confidence(8, "linear", r2=0.5) == "medium"
def test_confidence_low(self):
assert engine.compute_confidence(3, "linear", r2=0.1) == "low"
# 移动平均高波动 → low
assert engine.compute_confidence(10, "moving_average", cv=0.8) == "low"
# ── 引擎级测试(真实DB ──
class TestForecastKpi:
def test_forecast_kpi_structure(self, db: Session):
"""完整返回结构"""
kpi = create_test_kpi(db, kpi_code="F_TEST_REV", kpi_name="测试收入",
dimension="finance", unit="万元")
_seed_history(db, kpi, [100, 110, 120, 130, 140])
r = engine.forecast_kpi(1, "F_TEST_REV", db, periods=3)
assert r is not None
assert r["kpi"]["code"] == "F_TEST_REV"
assert r["kpi"]["name"] == "测试收入"
assert r["kpi"]["unit"] == "万元"
assert r["trend"] == "up"
assert r["confidence"] in ("high", "medium", "low")
assert r["history_count"] == 5
assert len(r["history"]) == 5
assert len(r["forecast"]) == 3
# 历史升序
periods = [h["period"] for h in r["history"]]
assert periods == sorted(periods)
# 预测期号在历史之后
assert r["forecast"][0]["period"] > periods[-1]
assert r["next_target"] == r["forecast"][0]["predicted"]
assert "测试收入" in r["summary"]
def test_forecast_kpi_moving_average_model(self, db: Session):
kpi = create_test_kpi(db, kpi_code="F_TEST_MA", kpi_name="测试费用", unit="")
_seed_history(db, kpi, [10, 12, 11, 13, 12])
r = engine.forecast_kpi(1, "F_TEST_MA", db, periods=2, model="moving_average")
assert r["model"] == "moving_average"
assert len(r["forecast"]) == 2
def test_forecast_kpi_insufficient_history(self, db: Session):
kpi = create_test_kpi(db, kpi_code="F_TEST_1PT", kpi_name="单点KPI")
_seed_history(db, kpi, [100.0])
assert engine.forecast_kpi(1, "F_TEST_1PT", db) is None
def test_forecast_kpi_no_history(self, db: Session):
create_test_kpi(db, kpi_code="F_TEST_NOHIST", kpi_name="无历史KPI")
assert engine.forecast_kpi(1, "F_TEST_NOHIST", db) is None
def test_forecast_kpi_unknown_code(self, db: Session):
assert engine.forecast_kpi(1, "F_NOT_EXIST", db) is None
def test_forecast_kpi_zero_periods(self, db: Session):
kpi = create_test_kpi(db, kpi_code="F_TEST_ZP", kpi_name="零期KPI")
_seed_history(db, kpi, [100, 110, 120])
r = engine.forecast_kpi(1, "F_TEST_ZP", db, periods=0)
assert r is not None
assert r["forecast"] == []
assert r["next_target"] is None
# ── API 端点测试 ──
class TestKpiForecastApi:
def test_single_forecast(self, client: TestClient, db: Session):
kpi = create_test_kpi(db, kpi_code="F_API_REV", kpi_name="API收入",
dimension="finance", unit="万元")
_seed_history(db, kpi, [100, 108, 116, 124, 132])
resp = client.get(f"{BASE}/kpi-forecast", params={
"entity_id": 1, "kpi_code": "F_API_REV", "periods": 3,
})
assert resp.status_code == 200
data = resp.json()
assert data["kpi"]["code"] == "F_API_REV"
assert data["trend"] in ("up", "down", "flat")
assert len(data["forecast"]) == 3
assert data["next_target"] is not None
assert data["summary"]
def test_moving_average_param(self, client: TestClient, db: Session):
kpi = create_test_kpi(db, kpi_code="F_API_MA", kpi_name="API费用", unit="")
_seed_history(db, kpi, [5, 6, 7, 6, 8])
resp = client.get(f"{BASE}/kpi-forecast", params={
"entity_id": 1, "kpi_code": "F_API_MA", "model": "moving_average",
})
assert resp.status_code == 200
assert resp.json()["model"] == "moving_average"
def test_insufficient_history_404(self, client: TestClient, db: Session):
kpi = create_test_kpi(db, kpi_code="F_API_LESS", kpi_name="数据不足")
_seed_history(db, kpi, [100.0])
resp = client.get(f"{BASE}/kpi-forecast", params={"entity_id": 1, "kpi_code": "F_API_LESS"})
assert resp.status_code == 404
def test_no_history_404(self, client: TestClient, db: Session):
create_test_kpi(db, kpi_code="F_API_EMPTY", kpi_name="空历史")
resp = client.get(f"{BASE}/kpi-forecast", params={"entity_id": 1, "kpi_code": "F_API_EMPTY"})
assert resp.status_code == 404
def test_zero_periods(self, client: TestClient, db: Session):
kpi = create_test_kpi(db, kpi_code="F_API_ZERO", kpi_name="零期")
_seed_history(db, kpi, [100, 110, 120])
resp = client.get(f"{BASE}/kpi-forecast", params={
"entity_id": 1, "kpi_code": "F_API_ZERO", "periods": 0,
})
assert resp.status_code == 200
data = resp.json()
assert data["forecast"] == []
assert data["next_target"] is None
def test_bad_model_400(self, client: TestClient, db: Session):
kpi = create_test_kpi(db, kpi_code="F_API_BADM", kpi_name="坏模型")
_seed_history(db, kpi, [1, 2, 3, 4])
resp = client.get(f"{BASE}/kpi-forecast", params={
"entity_id": 1, "kpi_code": "F_API_BADM", "model": "lstm",
})
assert resp.status_code == 400
def test_bad_periods_400(self, client: TestClient, db: Session):
resp = client.get(f"{BASE}/kpi-forecast", params={
"entity_id": 1, "kpi_code": "F_API_BADM", "periods": -1,
})
assert resp.status_code == 400
class TestKpiForecastFinanceApi:
def test_batch_finance(self, client: TestClient, db: Session):
"""只返回财务维度且历史≥3条的KPI,按可预测性排序"""
good1 = create_test_kpi(db, kpi_code="F_BATCH_REV", kpi_name="批量收入", dimension="finance")
_seed_history(db, good1, [100, 105, 110, 115, 120, 125, 130, 135]) # 8条, 完美上升 → high
good2 = create_test_kpi(db, kpi_code="F_BATCH_COST", kpi_name="批量成本", dimension="finance")
_seed_history(db, good2, [50, 52, 51, 53]) # 4条 → low/medium
short = create_test_kpi(db, kpi_code="F_BATCH_SHORT", kpi_name="数据不足", dimension="finance")
_seed_history(db, short, [10, 20]) # 2条 → 应被排除
non_fin = create_test_kpi(db, kpi_code="C_BATCH_SAT", kpi_name="客户满意", dimension="customer")
_seed_history(db, non_fin, [80, 82, 84, 86]) # 非财务 → 应被排除
resp = client.get(f"{BASE}/kpi-forecast/finance", params={"entity_id": 1})
assert resp.status_code == 200
data = resp.json()
assert data["total"] == 2
codes = [d["kpi"]["code"] for d in data["data"]]
assert "F_BATCH_REV" in codes and "F_BATCH_COST" in codes
assert "F_BATCH_SHORT" not in codes
assert "C_BATCH_SAT" not in codes
# 可预测性排序:F_BATCH_REV(8条 high) 应排前面
assert data["data"][0]["kpi"]["code"] == "F_BATCH_REV"
def test_batch_zero_periods(self, client: TestClient, db: Session):
kpi = create_test_kpi(db, kpi_code="F_BATCH_ZP", kpi_name="零期批量", dimension="finance")
_seed_history(db, kpi, [1, 2, 3, 4])
resp = client.get(f"{BASE}/kpi-forecast/finance", params={"entity_id": 1, "periods": 0})
assert resp.status_code == 200
assert resp.json()["data"][0]["forecast"] == []
class TestEntityIsolation:
def test_entity2_kpi_invisible_to_entity1(self, client: TestClient, db: Session):
"""entity2 的KPI在 entity1 下查询 → 404(不串数据)"""
_seed_entity2(db)
kpi2 = create_test_kpi(db, kpi_code="F_ISO_E2", kpi_name="博海专属KPI",
dimension="finance", entity_id=2)
_seed_history(db, kpi2, [100, 110, 120, 130], entity_id=2)
# 无token时 entity_id 走 query 参数 → entity1
resp = client.get(f"{BASE}/kpi-forecast", params={
"entity_id": 1, "kpi_code": "F_ISO_E2",
})
assert resp.status_code == 404
# entity2 自己能查到
resp = client.get(f"{BASE}/kpi-forecast", params={
"entity_id": 2, "kpi_code": "F_ISO_E2",
})
assert resp.status_code == 200
assert resp.json()["entity_id"] == 2
def test_batch_isolation(self, client: TestClient, db: Session):
"""批量预测:entity1 结果不含 entity2 的财务KPI"""
_seed_entity2(db)
kpi1 = create_test_kpi(db, kpi_code="F_ISO_E1", kpi_name="酣客收入", dimension="finance")
_seed_history(db, kpi1, [100, 110, 120, 130, 140], entity_id=1)
kpi2 = create_test_kpi(db, kpi_code="F_ISO_E2B", kpi_name="博海收入", dimension="finance", entity_id=2)
_seed_history(db, kpi2, [200, 210, 220, 230], entity_id=2)
resp = client.get(f"{BASE}/kpi-forecast/finance", params={"entity_id": 1})
codes = [d["kpi"]["code"] for d in resp.json()["data"]]
assert "F_ISO_E1" in codes
assert "F_ISO_E2B" not in codes
def test_token_binds_entity(self, client: TestClient, db: Session):
"""token绑定entity1时显式传entity_id=2 → 403 越权拦截"""
create_test_user(db)
token = get_token_for_user(client)
kpi = create_test_kpi(db, kpi_code="F_ISO_TK", kpi_name="TokenKPI")
_seed_history(db, kpi, [1, 2, 3, 4])
resp = client.get(f"{BASE}/kpi-forecast", params={
"entity_id": 2, "kpi_code": "F_ISO_TK",
}, headers=auth_header(token))
assert resp.status_code == 403