from __future__ import annotations import copy import unittest from pathlib import Path import tempfile from datahub.adapters.base import AdapterError from datahub.adapters.tushare import TushareAdapter from datahub.crypto import SecretVault from datahub.db import HubDB from datahub.pipeline import Pipeline from datahub.scheduler import Scheduler from datahub.serving import V1API from datahub.settings import Settings from datahub.timeutil import SHANGHAI from tests.fixtures import RAW, TRADE_DATE, fake_transport from tests.test_eod_retry import clock_at from tests.test_quality_gates import FIELD_GATES SAMPLE_DAY = "20260907" NEXT_DAY = "20260908" SAMPLE_CODE = "003021.SZ" def _dated(row: dict, day: str) -> dict: item = dict(row) if "trade_date" in item: item["trade_date"] = day return item class RevisingTransport: """Fixture transport that can rewrite daily_basic after the first publish.""" DATE_APIS = {"daily", "daily_basic", "adj_factor", "moneyflow", "stk_auction", "index_daily"} def __init__(self, extra_calendar: list[dict] | None = None) -> None: self.calls: list[str] = [] self.fail_daily_basic = False self.empty_daily_basic = False self.null_volume_ratio = False self.turnover_by_code: dict[str, float] = {} self.extra_calendar = extra_calendar or [] def __call__(self, api_name: str, params: dict, fields: str): self.calls.append(api_name) day = str(params.get("trade_date") or "") if api_name == "trade_cal": rows = fake_transport(api_name, params, fields) extra = [ row for row in self.extra_calendar if str(params.get("start_date") or "") <= row["cal_date"] <= str(params.get("end_date") or "99999999") ] return rows + extra if self.fail_daily_basic and api_name == "daily_basic": raise AdapterError("tushare daily_basic unavailable") if self.empty_daily_basic and api_name == "daily_basic": return [] if api_name == "index_daily": code = params.get("ts_code") rows = [row for row in RAW["index_daily"] if row["ts_code"] == code] if day: rows = [_dated(row, day) for row in rows] return rows rows = fake_transport(api_name, params, fields) if api_name == "stock_basic": rows = list(rows) rows.append({ "ts_code": SAMPLE_CODE, "symbol": "003021", "name": "兆威机电", "area": "广东", "industry": "元器件", "market": "主板", "list_status": "L", "list_date": "20201202", }) return rows if api_name in self.DATE_APIS: template = RAW.get(api_name) or [] if not day: return [_dated(row, TRADE_DATE) for row in template] out = [_dated(row, day) for row in template] extra = copy.deepcopy(template[0]) extra["ts_code"] = SAMPLE_CODE extra["trade_date"] = day if api_name == "daily_basic": extra["turnover_rate"] = self.turnover_by_code.get(SAMPLE_CODE, extra.get("turnover_rate")) if self.null_volume_ratio: extra["volume_ratio"] = None for row in out: row["volume_ratio"] = None out.append(extra) if api_name == "daily_basic": for row in out: code = str(row.get("ts_code") or "") if code in self.turnover_by_code: row["turnover_rate"] = self.turnover_by_code[code] return out return rows def make_revision_env(quality_extra: dict | None = None, extra_calendar: list[dict] | None = None): tmp = tempfile.TemporaryDirectory() db = HubDB(Path(tmp.name) / "hub.db") transport = RevisingTransport(extra_calendar=extra_calendar) adapter = TushareAdapter("x", transport=transport) quality = { "daily_row_ratio": 0.5, "null_rate_max": 0.5, "max_publish_attempts": 2, "publication_generations": 3, "field_gates": FIELD_GATES, "revision_review_start": "20:00", "revision_review_interval_minutes": 30, "revision_review_cutoff": "23:20", "revision_review_datasets": ["valuation"], } if quality_extra: quality.update(quality_extra) settings = Settings( encryption_key=SecretVault.generate_key(), api_token="t" * 32, db_path=db.path, backup_dir=Path(tmp.name) / "backups", quality=quality, scheduler_enabled=False, ) pipe = Pipeline(db, adapter, settings) sched = Scheduler(db, pipe) return tmp, db, transport, pipe, sched SAMPLE_CALENDAR = [ {"exchange": "SSE", "cal_date": SAMPLE_DAY, "is_open": 1, "pretrade_date": "20260906"}, {"exchange": "SSE", "cal_date": NEXT_DAY, "is_open": 1, "pretrade_date": SAMPLE_DAY}, ] class RevisionReviewTests(unittest.TestCase): def _publish(self, pipe: Pipeline, day: str) -> None: pipe.ingest_reference(day) pipe.run_eod_batch_a(day) pipe.run_eod_batch_b(day) def _turnover(self, db: HubDB, day: str, code: str = SAMPLE_CODE) -> float | None: pub = db.fetchone( "SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?", (day,), ) row = db.fetchone( "SELECT turnover_rate FROM eod_valuation WHERE batch_id=? AND ts_code=?", (pub["active_batch"], code), ) return None if row is None else row["turnover_rate"] def _batch_ids(self, db: HubDB, day: str) -> set[str]: return {str(row["batch_id"]) for row in db.fetchall("SELECT batch_id FROM batches WHERE trade_date=?", (day,))} def test_no_change_does_not_create_a_new_batch(self) -> None: tmp, db, transport, pipe, sched = make_revision_env() self.addCleanup(tmp.cleanup) self._publish(pipe, TRADE_DATE) before = self._batch_ids(db, TRADE_DATE) sched.tick(clock_at(TRADE_DATE, 20, 0)) self.assertEqual(self._batch_ids(db, TRADE_DATE), before) progress = db.fetchone("SELECT * FROM revision_progress WHERE trade_date=?", (TRADE_DATE,)) self.assertEqual(progress["state"], "aligned") self.assertIn("无变化", progress["detail"]) status = sched.revision_status(TRADE_DATE, clock=clock_at(TRADE_DATE, 20, 0)) self.assertEqual(status["state"], "aligned") def test_hel423_20260907_single_field_revision_is_caught_up(self) -> None: tmp, db, transport, pipe, sched = make_revision_env(extra_calendar=SAMPLE_CALENDAR) self.addCleanup(tmp.cleanup) transport.turnover_by_code[SAMPLE_CODE] = 1.3565 self._publish(pipe, SAMPLE_DAY) self.assertEqual(self._turnover(db, SAMPLE_DAY), 1.3565) first = db.fetchone( "SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?", (SAMPLE_DAY,), )["active_batch"] transport.turnover_by_code[SAMPLE_CODE] = 1.3572 seen: list[float | None] = [] def watch() -> None: seen.append(self._turnover(db, SAMPLE_DAY)) pipe.before_commit = watch ran = sched.tick(clock_at(SAMPLE_DAY, 20, 0)) self.assertIn("eod_revise", ran) self.assertEqual(seen, [1.3565]) # readers still see the previous complete version mid-switch self.assertEqual(self._turnover(db, SAMPLE_DAY), 1.3572) second = db.fetchone( "SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?", (SAMPLE_DAY,), )["active_batch"] self.assertNotEqual(second, first) api = V1API(db, pipe, pipe.settings) payload = api.valuation({"date": SAMPLE_DAY, "code": SAMPLE_CODE}) row = next(item for item in payload["data"] if item["ts_code"] == SAMPLE_CODE) self.assertEqual(row["turnover_rate"], 1.3572) progress = db.fetchone("SELECT * FROM revision_progress WHERE trade_date=?", (SAMPLE_DAY,)) self.assertEqual(progress["state"], "aligned") self.assertEqual(progress["detail"], "已追平") audit = db.fetchone( "SELECT * FROM audit_log WHERE action='revision-review' ORDER BY id DESC" ) self.assertIn("1.3572", str(audit["detail"])) self.assertIn(SAMPLE_CODE, str(audit["detail"])) def test_empty_or_failed_upstream_keeps_previous_version(self) -> None: tmp, db, transport, pipe, sched = make_revision_env() self.addCleanup(tmp.cleanup) self._publish(pipe, TRADE_DATE) active = db.fetchone( "SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?", (TRADE_DATE,), )["active_batch"] batches = self._batch_ids(db, TRADE_DATE) transport.empty_daily_basic = True sched.tick(clock_at(TRADE_DATE, 20, 0)) self.assertEqual( db.fetchone( "SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?", (TRADE_DATE,), )["active_batch"], active, ) self.assertEqual( db.fetchone("SELECT state FROM revision_progress WHERE trade_date=?", (TRADE_DATE,))["state"], "review_failed", ) transport.empty_daily_basic = False transport.fail_daily_basic = True sched.tick(clock_at(TRADE_DATE, 20, 30)) self.assertEqual( db.fetchone( "SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?", (TRADE_DATE,), )["active_batch"], active, ) self.assertEqual(self._batch_ids(db, TRADE_DATE), batches) def test_quality_gate_rejects_catchup_and_keeps_previous(self) -> None: tmp, db, transport, pipe, sched = make_revision_env() self.addCleanup(tmp.cleanup) self._publish(pipe, TRADE_DATE) active = db.fetchone( "SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?", (TRADE_DATE,), )["active_batch"] transport.turnover_by_code[SAMPLE_CODE] = 9.9999 transport.null_volume_ratio = True sched.tick(clock_at(TRADE_DATE, 20, 0)) self.assertEqual( db.fetchone( "SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?", (TRADE_DATE,), )["active_batch"], active, ) self.assertEqual( db.fetchone("SELECT state FROM revision_progress WHERE trade_date=?", (TRADE_DATE,))["state"], "review_failed", ) def test_repeat_ticks_after_align_do_not_republish(self) -> None: tmp, db, transport, pipe, sched = make_revision_env() self.addCleanup(tmp.cleanup) transport.turnover_by_code[SAMPLE_CODE] = 1.3565 self._publish(pipe, TRADE_DATE) transport.turnover_by_code[SAMPLE_CODE] = 1.3572 sched.tick(clock_at(TRADE_DATE, 20, 0)) after_fix = self._batch_ids(db, TRADE_DATE) sched.tick(clock_at(TRADE_DATE, 20, 10)) # inside interval self.assertEqual(len(db.fetchall("SELECT * FROM job_runs WHERE job_id='eod_revise'")), 1) sched.tick(clock_at(TRADE_DATE, 20, 30)) # next light compare, no change self.assertEqual(self._batch_ids(db, TRADE_DATE), after_fix) self.assertEqual(self._turnover(db, TRADE_DATE), 1.3572) def test_restart_catches_up_inside_window(self) -> None: tmp, db, transport, pipe, sched = make_revision_env() self.addCleanup(tmp.cleanup) transport.turnover_by_code[SAMPLE_CODE] = 1.3565 self._publish(pipe, TRADE_DATE) transport.turnover_by_code[SAMPLE_CODE] = 1.3572 sched2 = Scheduler(db, pipe) ran = sched2.tick(clock_at(TRADE_DATE, 21, 0)) self.assertIn("eod_revise", ran) self.assertEqual(self._turnover(db, TRADE_DATE), 1.3572) def test_cutoff_stops_evening_reviews_and_morning_catchup_runs(self) -> None: tmp, db, transport, pipe, sched = make_revision_env(extra_calendar=SAMPLE_CALENDAR) self.addCleanup(tmp.cleanup) transport.turnover_by_code[SAMPLE_CODE] = 1.3565 self._publish(pipe, SAMPLE_DAY) sched.tick(clock_at(SAMPLE_DAY, 23, 25)) # past 23:20 cutoff, no review yet cutoff = db.fetchone("SELECT * FROM revision_progress WHERE trade_date=?", (SAMPLE_DAY,)) self.assertEqual(cutoff["state"], "cutoff") self.assertEqual(self._turnover(db, SAMPLE_DAY), 1.3565) transport.turnover_by_code[SAMPLE_CODE] = 1.3572 sched.tick(clock_at(SAMPLE_DAY, 23, 50)) # still same calendar day, no catch-up self.assertEqual(self._turnover(db, SAMPLE_DAY), 1.3565) ran = sched.tick(clock_at(NEXT_DAY, 8, 45)) self.assertIn("eod_revise", ran) self.assertEqual(self._turnover(db, SAMPLE_DAY), 1.3572) progress = db.fetchone("SELECT * FROM revision_progress WHERE trade_date=?", (SAMPLE_DAY,)) self.assertEqual(progress["state"], "aligned") self.assertEqual(int(progress["catchup_done"]), 1) batches = self._batch_ids(db, SAMPLE_DAY) sched.tick(clock_at(NEXT_DAY, 8, 50)) self.assertEqual(self._batch_ids(db, SAMPLE_DAY), batches) def test_only_valuation_is_light_fetched(self) -> None: tmp, db, transport, pipe, sched = make_revision_env() self.addCleanup(tmp.cleanup) self._publish(pipe, TRADE_DATE) before = [name for name in transport.calls] sched.tick(clock_at(TRADE_DATE, 20, 0)) extra = transport.calls[len(before):] self.assertIn("daily_basic", extra) self.assertNotIn("daily", extra) self.assertNotIn("moneyflow", extra) self.assertNotIn("stk_auction", extra) self.assertNotIn("index_daily", extra) if __name__ == "__main__": unittest.main()