Files
xiaobai-review/xiaobai-datahub/tests/test_quality_gates.py
T
75c2e33b68 fix(HEL-461): CLI/后台强制重发改为整组边界切换
eod-refresh --force 与管理后台补数不再单数据集发布,
统一走 force_republish_boundary,避免绕过 A/B 完整边界。

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-05 11:26:27 +08:00

263 lines
11 KiB
Python

from __future__ import annotations
import unittest
from pathlib import Path
import tempfile
from datahub.adapters.tushare import TushareAdapter
from datahub.crypto import SecretVault
from datahub.db import HubDB
from datahub.pipeline import Pipeline, QualityError
from datahub.scheduler import Scheduler
from datahub.settings import Settings
from tests.fixtures import TRADE_DATE, fake_transport
from tests.test_eod_retry import clock_at
FIELD_GATES = {
"valuation": {
"fields": [
"turnover_rate", "volume_ratio", "total_mv", "circ_mv",
"pe_ttm", "pb", "ps_ttm", "dv_ttm",
],
"min_nonnull_rate": 0.9,
"min_nonnull_rate_by_field": {"pe_ttm": 0.5, "dv_ttm": 0.3},
"max_nonnull_drop_vs_prev": 0.15,
"max_nonfinite_rate": 0.01,
},
}
class ValuationTransport:
"""fake_transport with switchable daily_basic degradation modes."""
def __init__(self) -> None:
self.mode = "ok"
def __call__(self, api_name: str, params: dict, fields: str):
rows = fake_transport(api_name, params, fields)
if api_name != "daily_basic":
return rows
trade_date = str(params.get("trade_date") or "")
if trade_date:
rows = [{**row, "trade_date": trade_date} for row in rows]
if self.mode == "ok":
return rows
patched = []
for row in rows:
item = dict(row)
if self.mode == "fields_all_null":
item["volume_ratio"] = None
item["dv_ttm"] = None
elif self.mode == "vr_all_null":
item["volume_ratio"] = None
elif self.mode == "dv_all_null":
item["dv_ttm"] = None
elif self.mode == "nonfinite":
item["volume_ratio"] = float("inf")
patched.append(item)
return patched
def make_pipe(transport, quality_extra=None, clock=None):
tmp = tempfile.TemporaryDirectory()
db = HubDB(Path(tmp.name) / "hub.db")
adapter = TushareAdapter("test-token", transport=transport)
quality = {
"daily_row_ratio": 0.98,
"null_rate_max": 0.01,
"max_publish_attempts": 3,
"publication_generations": 3,
"job_run_retain_days": 90,
"staging_retain_days": 14,
"field_gates": FIELD_GATES,
}
if quality_extra:
quality.update(quality_extra)
settings = Settings(
encryption_key=SecretVault.generate_key(),
api_token="t" * 32,
admin_password="admin-pass",
tushare_token="test-token",
db_path=db.path,
quality=quality,
scheduler_enabled=False,
)
pipe = Pipeline(db, adapter, settings, clock=clock)
pipe._tmp = tmp
return pipe, db
class ValuationFieldGateTests(unittest.TestCase):
def setUp(self) -> None:
self.transport = ValuationTransport()
self.pipe, self.db = make_pipe(self.transport)
self.pipe.ingest_reference(TRADE_DATE)
def _active(self) -> str | None:
row = self.db.fetchone(
"SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?",
(TRADE_DATE,),
)
return str(row["active_batch"]) if row else None
def test_normal_batch_with_legit_dv_nulls_passes(self) -> None:
result = self.pipe.run_dataset("valuation", TRADE_DATE)
self.assertEqual(result["state"], "published")
fields = result["quality"]["fields"]
# fixture: 1 of 2 stocks has null dv_ttm → 0.5 non-null ≥ 0.3 floor
self.assertEqual(fields["dv_ttm"]["nonnull_rate"], 0.5)
self.assertEqual(fields["volume_ratio"]["nonnull_rate"], 1.0)
self.assertFalse(result["quality"]["errors"])
def test_all_null_fields_rejected_and_prev_batch_kept(self) -> None:
first = self.pipe.run_dataset("valuation", TRADE_DATE)
self.transport.mode = "fields_all_null"
with self.assertRaises(QualityError) as ctx:
self.pipe.run_dataset("valuation", TRADE_DATE)
errors = "; ".join(ctx.exception.report["errors"])
self.assertIn("field gate: valuation.volume_ratio non-null rate 0.0000 < 0.9", errors)
self.assertIn("field gate: valuation.dv_ttm non-null rate 0.0000 < 0.3", errors)
# previous good publication stays active
self.assertEqual(self._active(), first["batch_id"])
# rejected batch left staged with readable error + field stats
rejected = self.db.fetchone(
"SELECT * FROM batches WHERE state='staged' AND dataset='valuation' ORDER BY started_at DESC",
)
self.assertIsNotNone(rejected)
self.assertIn("field gate: valuation.volume_ratio", str(rejected["error"]))
import json
quality = json.loads(rejected["quality_json"])
self.assertEqual(quality["fields"]["volume_ratio"]["nonnull"], 0)
self.assertEqual(quality["fields"]["dv_ttm"]["nonnull"], 0)
def test_volume_ratio_all_null_alone_rejected(self) -> None:
self.pipe.run_dataset("valuation", TRADE_DATE)
self.transport.mode = "vr_all_null"
with self.assertRaises(QualityError):
self.pipe.run_dataset("valuation", TRADE_DATE)
self.assertEqual(
self.db.fetchone(
"SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?",
(TRADE_DATE,),
)["active_batch"],
"20240902-valuation-001",
)
def test_dv_ttm_all_null_rejected_by_floor_and_collapse(self) -> None:
prev_day = "20240830"
prev = self.pipe.run_dataset("valuation", prev_day) # prev dv nonnull 0.5
self.transport.mode = "dv_all_null"
with self.assertRaises(QualityError) as ctx:
self.pipe.run_dataset("valuation", TRADE_DATE)
errors = "; ".join(ctx.exception.report["errors"])
self.assertIn("field gate: valuation.dv_ttm non-null rate 0.0000 < 0.3", errors)
self.assertIn(f"dropped > 0.15 vs prev batch {prev['batch_id']}", errors)
def test_nonfinite_values_rejected(self) -> None:
self.pipe.run_dataset("valuation", TRADE_DATE)
rows = self.pipe.adapter.normalize(
"valuation", self.pipe._guarded_fetch("valuation", {"trade_date": TRADE_DATE})
)
for row in rows:
row["volume_ratio"] = float("inf")
with self.assertRaises(QualityError) as ctx:
self.pipe.run_dataset("valuation", TRADE_DATE, prepared_rows=rows)
errors = "; ".join(ctx.exception.report["errors"])
self.assertIn("field gate: valuation.volume_ratio non-finite rate 1.0000 > 0.01", errors)
def test_gate_off_when_not_configured(self) -> None:
pipe, _db = make_pipe(ValuationTransport(), quality_extra={"field_gates": {}})
pipe.ingest_reference(TRADE_DATE)
pipe.adapter._transport.mode = "fields_all_null"
result = pipe.run_dataset("valuation", TRADE_DATE)
self.assertEqual(result["state"], "published") # legacy behavior when unconfigured
def test_gate_applies_to_any_configured_dataset(self) -> None:
gates = {"daily": {"fields": ["volume"], "min_nonnull_rate": 0.9, "max_nonfinite_rate": 0.01}}
pipe, _db = make_pipe(ValuationTransport(), quality_extra={"field_gates": gates})
pipe.ingest_reference(TRADE_DATE)
def null_volume(api_name, params, fields):
if api_name != "daily":
return fake_transport(api_name, params, fields)
rows = fake_transport(api_name, params, fields)
for row in rows:
row["vol"] = None
return rows
pipe.adapter._transport = null_volume
with self.assertRaises(QualityError) as ctx:
pipe.run_dataset("daily", TRADE_DATE)
errors = "; ".join(ctx.exception.report["errors"])
self.assertIn("field gate: daily.volume non-null rate 0.0000 < 0.9", errors)
class GateRetryInterplayTests(unittest.TestCase):
def test_rejected_valuation_stays_missing_and_retry_publishes_later(self) -> None:
transport = ValuationTransport()
transport.mode = "fields_all_null"
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
db = HubDB(Path(tmp.name) / "hub.db")
adapter = TushareAdapter("x", transport=transport)
settings = Settings(
encryption_key=SecretVault.generate_key(),
db_path=db.path,
backup_dir=Path(tmp.name) / "backups",
quality={"field_gates": FIELD_GATES, "max_publish_attempts": 2},
)
pipe = Pipeline(db, adapter, settings)
pipe.ingest_reference("20240902")
sched = Scheduler(db, pipe)
sched.tick(clock_at("20240902", 15, 5)) # valuation rejected by field gate
sched.tick(clock_at("20240902", 15, 10))
self.assertIn("valuation", pipe.missing_official_datasets("20240902"))
self.assertEqual(
pipe.active_batch("valuation", "20240902"),
None,
)
transport.mode = "ok"
sched.tick(clock_at("20240902", 15, 45)) # retry passes the gate
self.assertNotIn("valuation", pipe.missing_official_datasets("20240902"))
rows = db.fetchall("SELECT * FROM eod_valuation WHERE trade_date='20240902'")
self.assertTrue(rows)
self.assertTrue(all(row["volume_ratio"] is not None for row in rows))
class ForceRepublishTests(unittest.TestCase):
def test_force_boundary_republish_keeps_prev_for_rollback(self) -> None:
transport = ValuationTransport()
pipe, db = make_pipe(transport)
pipe.ingest_reference(TRADE_DATE)
first = pipe.run_eod_batch_a(TRADE_DATE)
first_val = first["valuation"]["batch_id"]
first_daily = first["daily"]["batch_id"]
transport.mode = "vr_all_null"
blocked = pipe.force_republish_boundary("valuation", TRADE_DATE)
self.assertEqual(blocked["valuation"]["state"], "failed")
self.assertEqual(pipe.active_batch("valuation", TRADE_DATE), first_val)
self.assertEqual(pipe.active_batch("daily", TRADE_DATE), first_daily)
transport.mode = "ok"
second = pipe.force_republish_boundary("valuation", TRADE_DATE)
self.assertEqual(second["valuation"]["state"], "published")
self.assertNotEqual(second["valuation"]["batch_id"], first_val)
self.assertNotEqual(second["daily"]["batch_id"], first_daily)
pubs = db.fetchall(
"SELECT dataset, active_batch, prev_batch, published_at FROM publications WHERE trade_date=?",
(TRADE_DATE,),
)
by_ds = {str(row["dataset"]): row for row in pubs}
a_times = {by_ds[name]["published_at"] for name in ("daily", "valuation", "moneyflow", "auction", "stocks")}
self.assertEqual(len(a_times), 1)
self.assertEqual(by_ds["valuation"]["active_batch"], second["valuation"]["batch_id"])
self.assertEqual(by_ds["valuation"]["prev_batch"], first_val)
rolled = pipe.rollback("valuation", TRADE_DATE, actor="cli")
self.assertEqual(rolled["active_batch"], first_val)
if __name__ == "__main__":
unittest.main()