from __future__ import annotations import unittest from datetime import datetime 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.scheduler import Scheduler from datahub.settings import Settings from datahub.timeutil import SHANGHAI from tests.fixtures import fake_transport OFFICIAL = {"daily", "valuation", "moneyflow", "auction", "index_daily"} class DelayedTransport: """Upstream that only returns rows for dates it has "published" yet.""" DATE_APIS = {"daily", "daily_basic", "adj_factor", "moneyflow", "stk_auction", "index_daily"} def __init__(self, ready_dates: set[str]) -> None: self.ready = set(ready_dates) self.calls: list[str] = [] def __call__(self, api_name: str, params: dict, fields: str): self.calls.append(api_name) if api_name in self.DATE_APIS: trade_date = str(params.get("trade_date") or "") if trade_date and trade_date not in self.ready: return [] return fake_transport(api_name, params, fields) def clock_at(day: str, hh: int, mm: int) -> datetime: return datetime(int(day[:4]), int(day[4:6]), int(day[6:8]), hh, mm, tzinfo=SHANGHAI) class EodRetryTests(unittest.TestCase): def _make(self, ready_dates: set[str]): tmp = tempfile.TemporaryDirectory() self.addCleanup(tmp.cleanup) db = HubDB(Path(tmp.name) / "hub.db") transport = DelayedTransport(ready_dates) adapter = TushareAdapter("x", transport=transport) settings = Settings( encryption_key=SecretVault.generate_key(), db_path=db.path, backup_dir=Path(tmp.name) / "backups", ) pipe = Pipeline(db, adapter, settings) pipe.ingest_reference("20240902") sched = Scheduler(db, pipe) return db, transport, pipe, sched def _job_runs(self, db: HubDB, job_id: str) -> list[dict]: return db.fetchall("SELECT * FROM job_runs WHERE job_id = ? ORDER BY id", (job_id,)) def _batches(self, db: HubDB, day: str) -> list[dict]: return db.fetchall("SELECT * FROM batches WHERE trade_date = ?", (day,)) @staticmethod def _batch_ids(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,))} @staticmethod def _eod_calls(transport: DelayedTransport) -> list[str]: return [name for name in transport.calls if name in DelayedTransport.DATE_APIS] def _published(self, db: HubDB, day: str) -> set[str]: rows = db.fetchall("SELECT dataset FROM publications WHERE trade_date = ?", (day,)) return {str(row["dataset"]) for row in rows} def test_first_empty_then_retry_succeeds(self) -> None: day = "20240902" db, transport, pipe, sched = self._make(set()) sched.tick(clock_at(day, 15, 5)) # eod_a: upstream empty -> failed sched.tick(clock_at(day, 15, 10)) # eod_b: upstream empty -> failed self.assertEqual(self._published(db, day), set()) # quality gate held sched.tick(clock_at(day, 15, 20)) # inside window, but <30min since 15:10 self.assertEqual(self._job_runs(db, "eod_retry"), []) status = sched.eod_status(day, clock=clock_at(day, 15, 20)) self.assertEqual(status["state"], "waiting_upstream") self.assertTrue(status["next_retry_at"]) self.assertEqual(status["missing_datasets"], sorted(OFFICIAL)) sched.tick(clock_at(day, 15, 40)) # retry #1, still empty runs = self._job_runs(db, "eod_retry") self.assertEqual(len(runs), 1) self.assertEqual(runs[0]["state"], "failed") self.assertEqual(self._published(db, day), set()) transport.ready.add(day) sched.tick(clock_at(day, 16, 10)) # retry #2 succeeds self.assertEqual(self._published(db, day), OFFICIAL) self.assertEqual(sched.eod_status(day, clock=clock_at(day, 16, 10))["state"], "done") progress = db.fetchone("SELECT * FROM eod_progress WHERE trade_date = ?", (day,)) self.assertEqual(progress["state"], "done") self.assertEqual(progress["attempts"], 4) # eod_a + eod_b + 2 retries # success stops all further same-day requests batches_before = len(self._batches(db, day)) calls_before = len(transport.calls) sched.tick(clock_at(day, 17, 0)) sched.tick(clock_at(day, 23, 0)) self.assertEqual(len(self._job_runs(db, "eod_retry")), 2) self.assertEqual(len(self._batches(db, day)), batches_before) self.assertEqual(len(transport.calls), calls_before) def test_never_ready_marks_cutoff_failed_and_stops(self) -> None: day = "20240902" db, transport, pipe, sched = self._make(set()) sched.tick(clock_at(day, 15, 5)) sched.tick(clock_at(day, 15, 10)) sched.tick(clock_at(day, 15, 40)) sched.tick(clock_at(day, 16, 10)) sched.tick(clock_at(day, 23, 29)) self.assertEqual(len(self._job_runs(db, "eod_retry")), 3) sched.tick(clock_at(day, 23, 35)) # past cutoff 23:30 status = sched.eod_status(day, clock=clock_at(day, 23, 35)) self.assertEqual(status["state"], "cutoff_failed") cutoff_runs = [r for r in self._job_runs(db, "eod_retry") if "截止" in str(r["error"])] self.assertEqual(len(cutoff_runs), 1) self.assertEqual(self._published(db, day), set()) attempts = db.fetchone("SELECT attempts FROM eod_progress WHERE trade_date = ?", (day,))["attempts"] sched.tick(clock_at(day, 23, 59)) self.assertEqual( db.fetchone("SELECT attempts FROM eod_progress WHERE trade_date = ?", (day,))["attempts"], attempts, ) self.assertEqual(len(self._job_runs(db, "eod_retry")), 4) # 3 retries + 1 cutoff record self.assertEqual(self._published(db, day), set()) def test_restart_catches_up_without_overwriting(self) -> None: day = "20240902" db, transport, pipe, sched = self._make({day}) sched.tick(clock_at(day, 15, 5)) # eod_a publishes 4 datasets sched.tick(clock_at(day, 15, 10)) # eod_b publishes index self.assertEqual(self._published(db, day), OFFICIAL) active = db.fetchall("SELECT dataset, active_batch FROM publications WHERE trade_date = ?", (day,)) active_map = {row["dataset"]: row["active_batch"] for row in active} batches_before = self._batch_ids(db, day) calls_before = self._eod_calls(transport) # container restart: fresh scheduler, missed-time catch-up fires eod_a/eod_b sched2 = Scheduler(db, pipe) ran = sched2.tick(clock_at(day, 21, 0)) self.assertIn("eod_a", ran) self.assertIn("eod_b", ran) self.assertNotIn("eod_retry", ran) self.assertEqual(self._published(db, day), OFFICIAL) after = db.fetchall("SELECT dataset, active_batch FROM publications WHERE trade_date = ?", (day,)) self.assertEqual({row["dataset"]: row["active_batch"] for row in after}, active_map) self.assertEqual(self._batch_ids(db, day), batches_before) # no duplicate batches self.assertEqual(self._eod_calls(transport), calls_before) # no duplicate upstream EOD calls self.assertEqual(sched2.eod_status(day, clock=clock_at(day, 21, 0))["state"], "done") def test_restart_with_partial_publish_only_fetches_missing(self) -> None: day = "20240902" db, transport, pipe, sched = self._make({day}) sched.tick(clock_at(day, 15, 5)) # eod_a publishes 4; container "crashes" before eod_b self.assertEqual(self._published(db, day), {"daily", "valuation", "moneyflow", "auction"}) batches_before = self._batch_ids(db, day) sched2 = Scheduler(db, pipe) ran = sched2.tick(clock_at(day, 15, 20)) # restart: eod_b catch-up, eod_a all skipped self.assertIn("eod_b", ran) self.assertEqual(self._published(db, day), OFFICIAL) new_ids = self._batch_ids(db, day) - batches_before new_datasets = {str(b["dataset"]) for b in self._batches(db, day) if str(b["batch_id"]) in new_ids} self.assertEqual(new_datasets, {"index_daily"}) self.assertEqual(sched2.eod_status(day, clock=clock_at(day, 15, 20))["state"], "done") def test_closed_day_skips_all_eod_work(self) -> None: day = "20240907" # closed in fixture calendar db, transport, pipe, sched = self._make(set()) for hh, mm in ((15, 5), (15, 10), (15, 40), (16, 10), (20, 0), (23, 40)): ran = sched.tick(clock_at(day, hh, mm)) self.assertNotIn("eod_retry", ran) eod_runs = db.fetchall("SELECT * FROM job_runs WHERE job_id LIKE 'eod%'") self.assertEqual(eod_runs, []) self.assertIsNone(db.fetchone("SELECT * FROM eod_progress WHERE trade_date = ?", (day,))) self.assertEqual(self._published(db, day), set()) self.assertEqual(sched.eod_status(day, clock=clock_at(day, 20, 0))["state"], "closed_day") def test_duplicate_and_concurrent_execution_are_safe(self) -> None: day = "20240902" db, transport, pipe, sched = self._make({day}) sched.tick(clock_at(day, 15, 5)) sched.tick(clock_at(day, 15, 10)) self.assertEqual(self._published(db, day), OFFICIAL) batches_before = len(self._batches(db, day)) calls_before = len(transport.calls) out = sched.run_job("eod_retry", day) # manual duplicate run self.assertEqual(out["state"], "ok") self.assertEqual(len(self._batches(db, day)), batches_before) self.assertEqual(len(transport.calls), calls_before) sched._eod_lock.acquire() # simulate an in-flight EOD job try: busy = sched.run_job("eod_retry", day) self.assertEqual(busy["state"], "skipped") busy_a = sched.run_job("eod_a", day) self.assertEqual(busy_a["state"], "skipped") finally: sched._eod_lock.release() self.assertEqual(len(self._batches(db, day)), batches_before) if __name__ == "__main__": unittest.main()