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 from datahub.settings import Settings from datahub.serving import V1API from tests.fixtures import TRADE_DATE, fake_transport GROUP_A = ("daily", "valuation", "moneyflow", "auction") class GroupTransport: """fake_transport with per-API degradation switches for release-group tests.""" def __init__(self) -> None: self.empty: set[str] = set() self.keep_rows: dict[str, int] = {} self.stocks: list[dict] = [] self.calls: list[str] = [] def __call__(self, api_name: str, params: dict, fields: str): self.calls.append(api_name) if api_name in self.empty: return [] if api_name == "stock_basic" and self.stocks: return [dict(row) for row in self.stocks] rows = fake_transport(api_name, params, fields) keep = self.keep_rows.get(api_name) if keep is not None: return rows[:keep] return rows def make_pipe(transport: GroupTransport, quality_extra: dict | None = 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": 2, "publication_generations": 3, } 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) pipe._tmp = tmp return pipe, db def publications_map(db: HubDB, day: str) -> dict[str, str]: rows = db.fetchall("SELECT dataset, active_batch FROM publications WHERE trade_date = ?", (day,)) return {str(row["dataset"]): str(row["active_batch"]) for row in rows} class ReleaseGroupSwitchTests(unittest.TestCase): def setUp(self) -> None: self.transport = GroupTransport() self.pipe, self.db = make_pipe(self.transport) self.pipe.ingest_reference(TRADE_DATE) def test_whole_group_switches_in_one_publish_instant(self) -> None: results = self.pipe.run_eod_batch_a(TRADE_DATE) self.assertEqual(set(results), {*GROUP_A, "stocks"}) self.assertEqual({item["state"] for item in results.values()}, {"published"}) pubs = self.db.fetchall("SELECT * FROM publications WHERE trade_date = ?", (TRADE_DATE,)) self.assertEqual(len(pubs), 5) self.assertEqual(len({row["published_at"] for row in pubs}), 1) # official rows copied and serving resolves the new batches api = V1API(self.db, self.pipe, self.pipe.settings) payload = api.handle("/v1/bars/daily", {"date": [TRADE_DATE], "code": ["600000.SH"]}) self.assertEqual(payload["meta"]["batch_id"], results["daily"]["batch_id"]) stocks = api.handle("/v1/stocks", {}) self.assertEqual(stocks["meta"]["batch_id"], results["stocks"]["batch_id"]) def test_any_member_failure_blocks_entire_group(self) -> None: self.transport.empty = {"daily_basic"} # valuation upstream returns nothing results = self.pipe.run_eod_batch_a(TRADE_DATE) self.assertEqual(results["valuation"]["state"], "failed") self.assertEqual(results["moneyflow"]["state"], "aborted") self.assertEqual(results["auction"]["state"], "aborted") self.assertEqual(results["daily"]["state"], "failed") # staged fine, then abandoned # nothing became visible, and the reason is recorded self.assertEqual(publications_map(self.db, TRADE_DATE), {}) abandoned = self.db.fetchall( "SELECT * FROM batches WHERE trade_date = ? AND state = 'failed'", (TRADE_DATE,), ) self.assertTrue(any("release group not switched" in str(row["error"] or "") for row in abandoned)) audit = self.db.fetchone( "SELECT * FROM audit_log WHERE action = 'release-group' ORDER BY id DESC" ) self.assertIn("valuation", str(audit["detail"])) # still missing → evening retries keep trying self.assertIn("daily", self.pipe.missing_official_datasets(TRADE_DATE)) def test_failure_keeps_previous_complete_version_serving(self) -> None: first = self.pipe.run_dataset("daily", TRADE_DATE) self.transport.empty = {"daily_basic"} results = self.pipe.run_eod_missing(TRADE_DATE) self.assertEqual(results["daily"]["state"], "skipped") self.assertEqual(results["valuation"]["state"], "failed") # the already-published complete batch is untouched and keeps serving self.assertEqual(self.pipe.active_batch("daily", TRADE_DATE), first["batch_id"]) self.assertEqual( publications_map(self.db, TRADE_DATE), {"daily": first["batch_id"]}, ) payload = V1API(self.db, self.pipe, self.pipe.settings).handle( "/v1/bars/daily", {"date": [TRADE_DATE], "code": ["600000.SH"]} ) self.assertEqual(payload["meta"]["batch_id"], first["batch_id"]) def test_reads_during_switch_see_old_state_until_commit(self) -> None: snapshots: list[dict] = [] def watcher() -> None: with self.db.connect() as connection: rows = connection.execute( "SELECT dataset, active_batch FROM publications WHERE trade_date = ?", (TRADE_DATE,), ).fetchall() snapshots.append({str(row["dataset"]): row["active_batch"] for row in rows}) self.pipe.before_commit = watcher self.pipe.run_eod_batch_a(TRADE_DATE) # inside the switch transaction the group was still invisible self.assertEqual(snapshots[0], {}) after = publications_map(self.db, TRADE_DATE) self.assertEqual(set(after), {*GROUP_A, "stocks"}) def test_switch_crash_rolls_back_whole_group(self) -> None: def explode() -> None: raise RuntimeError("killed mid-switch") self.pipe.before_commit = explode with self.assertRaises(RuntimeError): self.pipe.run_eod_batch_a(TRADE_DATE) self.assertEqual(publications_map(self.db, TRADE_DATE), {}) for table in ("eod_bars", "eod_valuation", "eod_moneyflow", "eod_auction", "eod_stocks"): rows = self.db.fetchall(f"SELECT * FROM {table} WHERE trade_date = ?", (TRADE_DATE,)) self.assertEqual(rows, [], table) def test_duplicate_runs_are_idempotent(self) -> None: self.pipe.run_eod_batch_a(TRADE_DATE) self.pipe.run_eod_batch_b(TRADE_DATE) batches_before = { str(row["batch_id"]) for row in self.db.fetchall("SELECT batch_id FROM batches WHERE trade_date = ?", (TRADE_DATE,)) } calls_before = len(self.transport.calls) again = self.pipe.run_eod_missing(TRADE_DATE) self.assertEqual({item["state"] for item in again.values()}, {"skipped"}) self.assertEqual({item["reason"] for item in again.values()}, {"already_published"}) batches_after = { str(row["batch_id"]) for row in self.db.fetchall("SELECT batch_id FROM batches WHERE trade_date = ?", (TRADE_DATE,)) } self.assertEqual(batches_after, batches_before) self.assertEqual(len(self.transport.calls), calls_before) self.assertEqual(self.pipe.missing_official_datasets(TRADE_DATE), []) def test_cross_gate_failure_blocks_switch(self) -> None: transport = GroupTransport() pipe, db = make_pipe( transport, quality_extra={"cross_gates": [ {"left": "daily", "right": "moneyflow", "min_key_overlap": 1.0}, ]}, ) pipe.ingest_reference(TRADE_DATE) transport.keep_rows["moneyflow"] = 1 # moneyflow covers only half the market results = pipe.run_eod_batch_a(TRADE_DATE) self.assertEqual(results["moneyflow"]["state"], "failed") self.assertIn("cross gate", str(results["moneyflow"]["error"])) self.assertEqual(publications_map(db, TRADE_DATE), {}) def test_stocks_master_and_snapshot_switch_together_or_not_at_all(self) -> None: original = [ {"ts_code": "600000.SH", "symbol": "600000", "name": "浦发银行", "area": "上海", "industry": "银行", "market": "主板", "list_status": "L", "list_date": "19991110"}, {"ts_code": "920071.BJ", "symbol": "920071", "name": "N金钛", "area": "辽宁", "industry": "小金属", "market": "北交所", "list_status": "L", "list_date": "20240901"}, ] renamed = [dict(original[0]), {**original[1], "name": "金钛股份"}] self.transport.stocks = renamed self.pipe.run_eod_batch_a(TRADE_DATE) master = self.db.fetchone("SELECT name FROM stock_master WHERE ts_code = '920071.BJ'") self.assertEqual(master["name"], "金钛股份") stocks_pub = self.db.fetchone( "SELECT active_batch FROM publications WHERE dataset = 'stocks' AND trade_date = ?", (TRADE_DATE,), ) self.assertIsNotNone(stocks_pub) # failure path: rename staged but the group is blocked → master stays untouched transport = GroupTransport() transport.stocks = original pipe, db = make_pipe( transport, quality_extra={"cross_gates": [ {"left": "daily", "right": "moneyflow", "min_key_overlap": 1.0}, ]}, ) pipe.ingest_reference(TRADE_DATE) # master seeded with "N金钛" transport.stocks = renamed transport.keep_rows["moneyflow"] = 1 results = pipe.run_eod_batch_a(TRADE_DATE) self.assertEqual(results["stocks"]["state"], "failed") master = db.fetchone("SELECT name FROM stock_master WHERE ts_code = '920071.BJ'") self.assertEqual(master["name"], "N金钛") # rename not applied stocks_pub = db.fetchone( "SELECT active_batch FROM publications WHERE dataset = 'stocks' AND trade_date = ?", (TRADE_DATE,), ) self.assertIsNone(stocks_pub) if __name__ == "__main__": unittest.main()