from __future__ import annotations import unittest from datetime import date, timedelta from datahub.coverage import calendar_coverage, point_coverage, published_range_coverage from datahub.serving import V1API from tests.fixtures import TRADE_DATE, fake_transport from tests.test_pipeline import make_pipeline def history_transport(open_dates: list[str], extra_closed: list[str] | None = None): open_set = set(open_dates) start = date(int(open_dates[0][:4]), int(open_dates[0][4:6]), int(open_dates[0][6:8])) end = date(int(open_dates[-1][:4]), int(open_dates[-1][4:6]), int(open_dates[-1][6:8])) calendar = [] cursor = start while cursor <= end: compact = cursor.strftime("%Y%m%d") calendar.append( { "exchange": "SSE", "cal_date": compact, "is_open": 1 if compact in open_set else 0, "pretrade_date": compact, } ) cursor += timedelta(days=1) for day in extra_closed or []: calendar.append( {"exchange": "SSE", "cal_date": day, "is_open": 0, "pretrade_date": open_dates[0]} ) index_codes = ("000001.SH", "399001.SZ", "399006.SZ", "000300.SH") index_rows = [] for ts_code in index_codes: for day in open_dates: index_rows.append( { "ts_code": ts_code, "trade_date": day, "open": 100, "high": 101, "low": 99, "close": 100.5, "pct_chg": 0.1, "vol": 10.0, "amount": 20.0, } ) def transport(api_name, params, fields): if api_name == "trade_cal": start = str(params.get("start_date") or "") end = str(params.get("end_date") or "99999999") return [row for row in calendar if start <= row["cal_date"] <= end] if api_name == "index_daily": code = params.get("ts_code") rows = [row for row in index_rows if row["ts_code"] == code] trade_date = str(params.get("trade_date") or "") start = str(params.get("start_date") or "") end = str(params.get("end_date") or "") if trade_date: rows = [row for row in rows if row["trade_date"] == trade_date] if start: rows = [row for row in rows if row["trade_date"] >= start] if end: rows = [row for row in rows if row["trade_date"] <= end] return rows return fake_transport(api_name, params, fields) return transport def consecutive_open_days(end: str, count: int) -> list[str]: cursor = date(int(end[:4]), int(end[4:6]), int(end[6:8])) days: list[str] = [] while len(days) < count: if cursor.weekday() < 5: days.append(cursor.strftime("%Y%m%d")) cursor -= timedelta(days=1) return sorted(days) class CoverageApiTests(unittest.TestCase): def test_calendar_marks_holes_incomplete(self) -> None: pipe, _db = make_pipeline() pipe.ingest_reference(TRADE_DATE) api = V1API(pipe.db, pipe, pipe.settings) payload = api.handle("/v1/calendar", {"from": ["20240901"], "to": ["20240907"]}) self.assertTrue(payload["meta"]["incomplete"]) self.assertFalse(payload["meta"]["coverage"]["complete"]) self.assertGreater(payload["meta"]["coverage"]["missing_count"], 0) self.assertIn("20240901", payload["meta"]["coverage"]["missing_sample"]) def test_calendar_complete_when_every_day_present(self) -> None: pipe, _db = make_pipeline() pipe.ingest_reference(TRADE_DATE) api = V1API(pipe.db, pipe, pipe.settings) payload = api.handle("/v1/calendar", {"from": ["20240902"], "to": ["20240903"]}) self.assertFalse(payload["meta"]["incomplete"]) self.assertTrue(payload["meta"]["coverage"]["complete"]) self.assertEqual(payload["meta"]["coverage"]["expected_count"], 2) self.assertEqual(len(payload["data"]), 2) def test_index_range_incomplete_without_history(self) -> None: pipe, _db = make_pipeline() pipe.ingest_reference(TRADE_DATE) pipe.run_dataset("index_daily", TRADE_DATE) api = V1API(pipe.db, pipe, pipe.settings) payload = api.handle( "/v1/indexes/bars", {"from": ["20240902"], "to": ["20240903"], "code": ["000001.SH"]}, ) self.assertTrue(payload["meta"]["incomplete"]) self.assertFalse(payload["meta"]["coverage"]["complete"]) self.assertEqual(payload["meta"]["coverage"]["available_count"], 1) self.assertIn("20240903", payload["meta"]["coverage"]["missing_sample"]) def test_index_point_query_stays_complete(self) -> None: pipe, _db = make_pipeline() pipe.ingest_reference(TRADE_DATE) pipe.run_dataset("index_daily", TRADE_DATE) api = V1API(pipe.db, pipe, pipe.settings) payload = api.handle("/v1/indexes/bars", {"date": [TRADE_DATE], "code": ["000001.SH"]}) self.assertFalse(payload["meta"]["incomplete"]) self.assertTrue(payload["meta"]["coverage"]["complete"]) self.assertEqual(payload["meta"]["coverage"]["kind"], "point") def test_daily_range_incomplete_without_stock_history(self) -> None: pipe, _db = make_pipeline() pipe.ingest_reference(TRADE_DATE) pipe.run_dataset("daily", TRADE_DATE) api = V1API(pipe.db, pipe, pipe.settings) payload = api.handle( "/v1/bars/daily", {"from": ["20240902"], "to": ["20240903"], "code": ["600000.SH"]}, ) self.assertTrue(payload["meta"]["incomplete"]) self.assertFalse(payload["meta"]["coverage"]["complete"]) class HistoryBackfillTests(unittest.TestCase): def test_index_history_is_idempotent_and_covers_requested_days(self) -> None: open_dates = consecutive_open_days(TRADE_DATE, 5) pipe, db = make_pipeline(quality={"index_history_trading_days": 5, "calendar_start": open_dates[0]}) pipe.adapter._transport = history_transport(open_dates) first = pipe.backfill_history(TRADE_DATE, index_days=5) self.assertTrue(first["ok"]) self.assertEqual(first["calendar"]["calendar_from"], open_dates[0]) self.assertEqual(first["index_daily"]["requested_days"], 5) self.assertEqual(len(first["index_daily"]["published"]), 5) self.assertEqual(first["index_daily"]["skipped"], []) pubs = db.fetchall("SELECT trade_date FROM publications WHERE dataset='index_daily'") self.assertEqual(sorted(row["trade_date"] for row in pubs), open_dates) second = pipe.backfill_index_history(TRADE_DATE, trading_days=5) self.assertTrue(second["ok"]) self.assertEqual(second["published"], []) self.assertEqual(second["skipped"], open_dates) api = V1API(db, pipe, pipe.settings) payload = api.handle( "/v1/indexes/bars", {"from": [open_dates[0]], "to": [open_dates[-1]], "code": ["000001.SH"]}, ) self.assertFalse(payload["meta"]["incomplete"]) self.assertEqual(payload["meta"]["coverage"]["available_count"], 5) self.assertEqual(len(payload["data"]), 5) def test_index_history_retries_failed_dates_without_dropping_success(self) -> None: open_dates = consecutive_open_days(TRADE_DATE, 3) base = history_transport(open_dates) def missing_cyb(api_name, params, fields): if api_name == "index_daily" and params.get("ts_code") == "399006.SZ": raise RuntimeError("upstream down") return base(api_name, params, fields) pipe, db = make_pipeline(quality={"index_history_trading_days": 3, "max_publish_attempts": 1}) pipe.adapter._transport = missing_cyb first = pipe.backfill_history(TRADE_DATE, calendar_start=open_dates[0], index_days=3) self.assertFalse(first["ok"]) self.assertTrue(any(item.get("ts_code") == "399006.SZ" for item in first["index_daily"]["failed"])) published_first = { row["trade_date"] for row in db.fetchall("SELECT trade_date FROM publications WHERE dataset='index_daily'") } self.assertEqual(published_first, set(open_dates)) pipe.adapter._transport = base retry = pipe.backfill_index_history(TRADE_DATE, trading_days=3) self.assertTrue(retry["ok"]) self.assertEqual(len(retry["published"]), 3) for day in open_dates: rows = db.fetchall( """ SELECT DISTINCT ts_code FROM eod_index_bars WHERE trade_date = ? AND batch_id = ( SELECT active_batch FROM publications WHERE dataset='index_daily' AND trade_date = ? ) """, (day, day), ) self.assertEqual({row["ts_code"] for row in rows}, {"000001.SH", "399001.SZ", "399006.SZ", "000300.SH"}) def test_prepared_rows_skip_upstream_fetch(self) -> None: pipe, _db = make_pipeline() pipe.ingest_reference(TRADE_DATE) calls = {"n": 0} original = pipe.adapter._transport def counting(api_name, params, fields): calls["n"] += 1 return original(api_name, params, fields) pipe.adapter._transport = counting rows = pipe.adapter.normalize("index_daily", original("index_daily", {"ts_code": "000001.SH", "trade_date": TRADE_DATE}, "")) before = calls["n"] result = pipe.run_dataset("index_daily", TRADE_DATE, prepared_rows=rows) self.assertEqual(result["rows"], 1) self.assertEqual(calls["n"], before) def test_coverage_helpers_point_and_calendar(self) -> None: pipe, db = make_pipeline() pipe.ingest_reference(TRADE_DATE) point = point_coverage(TRADE_DATE, "index_daily") self.assertTrue(point["complete"]) cal = calendar_coverage(db, "20240902", "20240903") self.assertTrue(cal["complete"]) pub = published_range_coverage(db, "index_daily", "20240902", "20240903") self.assertFalse(pub["complete"]) if __name__ == "__main__": unittest.main()