feat(HEL-421): 回补历史日历和指数并标记区间不完整
Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
co-authored by
Cursor
multica-agent
parent
25ff6bbe06
commit
a836cda1b2
@@ -0,0 +1,236 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user