Files
xiaobai-review/xiaobai-datahub/tests/test_history_backfill.py
T
2026-09-02 22:16:05 +08:00

237 lines
10 KiB
Python

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()