feat(HEL-457): 估值字段级质量门、股票主档每日发布和资金流历史回补
- field_gates 按数据集配置关键字段非空率下限/非有限比例/相对上一批次的塌陷保护, 字段大面积为空的批次拒发并保留上一正式批次,可读失败原因入 batches.error - 股票主档交易日 20:00/23:10 自动刷新并发布版本化快照(eod_stocks + publications), 覆盖新上市/简称变化/N前缀摘除;/v1/stocks 携带 batch_id/published_at,无变化跳过 - moneyflow 历史回补(默认 60 交易日,跳过已发布日期);未发布点查返回 available_from/available_to 与 history_not_backfilled 标记,缺失不再静默 - eod-refresh 新增 --force --dataset 安全重发(仍走全部质量门,上一批次可回滚) - 保持 HEL-435 盘后重试机制;新增 22 项测试覆盖字段拒发/正常通过/旧批保留/ 主档新增改名/资金流覆盖/重复执行幂等 Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
co-authored by
multica-agent
parent
c9892050c3
commit
bed6450992
@@ -25,7 +25,7 @@ RAW = {
|
||||
],
|
||||
"daily_basic": [
|
||||
{"ts_code": "600000.SH", "trade_date": "20240902", "turnover_rate": 1.2, "volume_ratio": 0.8, "total_mv": 1000.0, "circ_mv": 800.0, "pe_ttm": 5.1, "pb": 0.6, "ps_ttm": 1.1, "dv_ttm": 4.0},
|
||||
{"ts_code": "000001.SZ", "trade_date": "20240902", "turnover_rate": 2.2, "volume_ratio": 1.1, "total_mv": 2000.0, "circ_mv": 1500.0, "pe_ttm": 6.2, "pb": 0.7, "ps_ttm": 1.2, "dv_ttm": 3.0},
|
||||
{"ts_code": "000001.SZ", "trade_date": "20240902", "turnover_rate": 2.2, "volume_ratio": 1.1, "total_mv": 2000.0, "circ_mv": 1500.0, "pe_ttm": 6.2, "pb": 0.7, "ps_ttm": 1.2, "dv_ttm": None},
|
||||
],
|
||||
"adj_factor": [
|
||||
{"ts_code": "600000.SH", "trade_date": "20240902", "adj_factor": 1.1},
|
||||
|
||||
@@ -60,7 +60,11 @@ class EodRetryTests(unittest.TestCase):
|
||||
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,))
|
||||
placeholders = ",".join("?" for _ in OFFICIAL)
|
||||
return db.fetchall(
|
||||
f"SELECT * FROM batches WHERE trade_date = ? AND dataset IN ({placeholders})",
|
||||
(day, *sorted(OFFICIAL)),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _batch_ids(db: HubDB, day: str) -> set[str]:
|
||||
@@ -71,7 +75,11 @@ class EodRetryTests(unittest.TestCase):
|
||||
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,))
|
||||
placeholders = ",".join("?" for _ in OFFICIAL)
|
||||
rows = db.fetchall(
|
||||
f"SELECT dataset FROM publications WHERE trade_date = ? AND dataset IN ({placeholders})",
|
||||
(day, *sorted(OFFICIAL)),
|
||||
)
|
||||
return {str(row["dataset"]) for row in rows}
|
||||
|
||||
def test_first_empty_then_retry_succeeds(self) -> None:
|
||||
@@ -105,12 +113,12 @@ class EodRetryTests(unittest.TestCase):
|
||||
|
||||
# success stops all further same-day requests
|
||||
batches_before = len(self._batches(db, day))
|
||||
calls_before = len(transport.calls)
|
||||
eod_calls_before = len(self._eod_calls(transport))
|
||||
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)
|
||||
self.assertEqual(len(self._eod_calls(transport)), eod_calls_before)
|
||||
|
||||
def test_never_ready_marks_cutoff_failed_and_stops(self) -> None:
|
||||
day = "20240902"
|
||||
@@ -144,9 +152,20 @@ class EodRetryTests(unittest.TestCase):
|
||||
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)
|
||||
|
||||
def official_batches() -> list[str]:
|
||||
placeholders = ",".join("?" for _ in OFFICIAL)
|
||||
return [
|
||||
str(row["batch_id"])
|
||||
for row in db.fetchall(
|
||||
f"SELECT batch_id FROM batches WHERE trade_date = ? AND dataset IN ({placeholders})",
|
||||
(day, *sorted(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)
|
||||
active_map = {row["dataset"]: row["active_batch"] for row in active if row["dataset"] in OFFICIAL}
|
||||
batches_before = set(official_batches())
|
||||
calls_before = self._eod_calls(transport)
|
||||
|
||||
# container restart: fresh scheduler, missed-time catch-up fires eod_a/eod_b
|
||||
@@ -157,8 +176,11 @@ class EodRetryTests(unittest.TestCase):
|
||||
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(
|
||||
{row["dataset"]: row["active_batch"] for row in after if row["dataset"] in OFFICIAL},
|
||||
active_map,
|
||||
)
|
||||
self.assertEqual(set(official_batches()), 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")
|
||||
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from datetime import date, timedelta
|
||||
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.serving import ApiError, V1API
|
||||
from datahub.settings import Settings
|
||||
from tests.fixtures import fake_transport
|
||||
|
||||
OPEN_DATES = ["20240826", "20240827", "20240828", "20240829", "20240830", "20240902", "20240903"]
|
||||
EMPTY_UPSTREAM = {"20240828"} # one date the upstream cannot serve
|
||||
|
||||
|
||||
def build_calendar(open_dates: list[str], span_days: int = 16) -> list[dict]:
|
||||
start = date(int(open_dates[0][:4]), int(open_dates[0][4:6]), int(open_dates[0][6:8]))
|
||||
rows = []
|
||||
open_set = set(open_dates)
|
||||
for offset in range(span_days):
|
||||
cursor = start + timedelta(days=offset)
|
||||
compact = cursor.strftime("%Y%m%d")
|
||||
rows.append(
|
||||
{"exchange": "SSE", "cal_date": compact, "is_open": 1 if compact in open_set else 0, "pretrade_date": compact}
|
||||
)
|
||||
return rows
|
||||
|
||||
|
||||
def moneyflow_rows(day: str) -> list[dict]:
|
||||
return [
|
||||
{
|
||||
"ts_code": "600000.SH", "trade_date": day,
|
||||
"buy_sm_amount": 10 + int(day[-2:]), "sell_sm_amount": 8, "buy_md_amount": 20, "sell_md_amount": 15,
|
||||
"buy_lg_amount": 30, "sell_lg_amount": 25, "buy_elg_amount": 40, "sell_elg_amount": 35, "net_mf_amount": 17,
|
||||
},
|
||||
{
|
||||
"ts_code": "000001.SZ", "trade_date": day,
|
||||
"buy_sm_amount": 11, "sell_sm_amount": 9, "buy_md_amount": 21, "sell_md_amount": 16,
|
||||
"buy_lg_amount": 31, "sell_lg_amount": 26, "buy_elg_amount": 41, "sell_elg_amount": 36, "net_mf_amount": 18,
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
class MoneyflowHistoryTransport:
|
||||
def __init__(self) -> None:
|
||||
self.calendar = build_calendar(OPEN_DATES)
|
||||
self.moneyflow_fetches: list[str] = []
|
||||
|
||||
def __call__(self, api_name: str, params: dict, fields: str):
|
||||
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 self.calendar if start <= row["cal_date"] <= end]
|
||||
if api_name == "moneyflow":
|
||||
day = str(params.get("trade_date") or "")
|
||||
self.moneyflow_fetches.append(day)
|
||||
if day in EMPTY_UPSTREAM:
|
||||
return []
|
||||
return moneyflow_rows(day)
|
||||
return fake_transport(api_name, params, fields)
|
||||
|
||||
|
||||
class MoneyflowBackfillTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.transport = MoneyflowHistoryTransport()
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
self.db = HubDB(Path(tmp.name) / "hub.db")
|
||||
adapter = TushareAdapter("x", transport=self.transport)
|
||||
settings = Settings(
|
||||
encryption_key=SecretVault.generate_key(),
|
||||
db_path=self.db.path,
|
||||
backup_dir=Path(tmp.name) / "backups",
|
||||
quality={"max_publish_attempts": 2, "publication_generations": 3},
|
||||
)
|
||||
self.pipe = Pipeline(self.db, adapter, settings)
|
||||
self.pipe.ingest_reference("20240903")
|
||||
self.api = V1API(self.db, self.pipe, settings)
|
||||
|
||||
def test_backfill_publishes_window_and_reports_failures(self) -> None:
|
||||
result = self.pipe.backfill_moneyflow_history(end_date="20240903", trading_days=5)
|
||||
published = [item["trade_date"] for item in result["published"]]
|
||||
self.assertEqual(published, ["20240829", "20240830", "20240902", "20240903"])
|
||||
self.assertEqual(result["failed"][0]["trade_date"], "20240828")
|
||||
self.assertFalse(result["ok"])
|
||||
rows = self.db.fetchall("SELECT * FROM eod_moneyflow WHERE trade_date='20240830'")
|
||||
self.assertEqual(len(rows), 2)
|
||||
|
||||
def test_backfill_is_idempotent(self) -> None:
|
||||
self.pipe.backfill_moneyflow_history(end_date="20240903", trading_days=5)
|
||||
fetches_after_first = list(self.transport.moneyflow_fetches)
|
||||
second = self.pipe.backfill_moneyflow_history(end_date="20240903", trading_days=5)
|
||||
# only the still-missing date is re-fetched; published dates are skipped
|
||||
self.assertEqual(self.transport.moneyflow_fetches[len(fetches_after_first):], ["20240828"])
|
||||
self.assertEqual(len(second["skipped"]), 4)
|
||||
|
||||
def test_point_query_on_backfilled_date_serves_data(self) -> None:
|
||||
self.pipe.backfill_moneyflow_history(end_date="20240903", trading_days=5)
|
||||
payload = self.api.handle("/v1/moneyflow", {"date": ["20240830"]})
|
||||
self.assertEqual(len(payload["data"]), 2)
|
||||
self.assertEqual(payload["data"][0]["net_mf_amount"], 180000.0)
|
||||
|
||||
def test_unpublished_point_below_window_is_identifiable(self) -> None:
|
||||
self.pipe.backfill_moneyflow_history(end_date="20240903", trading_days=5)
|
||||
with self.assertRaises(ApiError) as ctx:
|
||||
self.api.handle("/v1/moneyflow", {"date": ["20240801"]})
|
||||
extra = ctx.exception.extra
|
||||
self.assertEqual(extra["available_from"], "20240829") # window starts at the first published date
|
||||
self.assertEqual(extra["available_to"], "20240903")
|
||||
self.assertEqual(extra["reason"], "history_not_backfilled")
|
||||
self.assertEqual(extra["expected_at"], "15:05+08:00")
|
||||
|
||||
def test_range_query_flags_missing_dates(self) -> None:
|
||||
self.pipe.backfill_moneyflow_history(end_date="20240903", trading_days=5)
|
||||
payload = self.api.handle("/v1/moneyflow", {"from": ["20240828"], "to": ["20240903"]})
|
||||
coverage = payload["meta"]["coverage"]
|
||||
self.assertFalse(coverage["complete"])
|
||||
self.assertEqual(coverage["missing_count"], 1)
|
||||
self.assertEqual(coverage["missing_sample"], ["20240828"])
|
||||
self.assertTrue(payload["meta"]["incomplete"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,253 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
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, QualityError
|
||||
from datahub.scheduler import Scheduler
|
||||
from datahub.settings import Settings
|
||||
from tests.fixtures import TRADE_DATE, fake_transport
|
||||
from tests.test_eod_retry import clock_at
|
||||
|
||||
FIELD_GATES = {
|
||||
"valuation": {
|
||||
"fields": [
|
||||
"turnover_rate", "volume_ratio", "total_mv", "circ_mv",
|
||||
"pe_ttm", "pb", "ps_ttm", "dv_ttm",
|
||||
],
|
||||
"min_nonnull_rate": 0.9,
|
||||
"min_nonnull_rate_by_field": {"pe_ttm": 0.5, "dv_ttm": 0.3},
|
||||
"max_nonnull_drop_vs_prev": 0.15,
|
||||
"max_nonfinite_rate": 0.01,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class ValuationTransport:
|
||||
"""fake_transport with switchable daily_basic degradation modes."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.mode = "ok"
|
||||
|
||||
def __call__(self, api_name: str, params: dict, fields: str):
|
||||
rows = fake_transport(api_name, params, fields)
|
||||
if api_name != "daily_basic":
|
||||
return rows
|
||||
trade_date = str(params.get("trade_date") or "")
|
||||
if trade_date:
|
||||
rows = [{**row, "trade_date": trade_date} for row in rows]
|
||||
if self.mode == "ok":
|
||||
return rows
|
||||
patched = []
|
||||
for row in rows:
|
||||
item = dict(row)
|
||||
if self.mode == "fields_all_null":
|
||||
item["volume_ratio"] = None
|
||||
item["dv_ttm"] = None
|
||||
elif self.mode == "vr_all_null":
|
||||
item["volume_ratio"] = None
|
||||
elif self.mode == "dv_all_null":
|
||||
item["dv_ttm"] = None
|
||||
elif self.mode == "nonfinite":
|
||||
item["volume_ratio"] = float("inf")
|
||||
patched.append(item)
|
||||
return patched
|
||||
|
||||
|
||||
def make_pipe(transport, quality_extra=None, clock=None):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
db = HubDB(Path(tmp.name) / "hub.db")
|
||||
adapter = TushareAdapter("test-token", transport=transport)
|
||||
quality = {
|
||||
"daily_row_ratio": 0.98,
|
||||
"null_rate_max": 0.01,
|
||||
"max_publish_attempts": 3,
|
||||
"publication_generations": 3,
|
||||
"job_run_retain_days": 90,
|
||||
"staging_retain_days": 14,
|
||||
"field_gates": FIELD_GATES,
|
||||
}
|
||||
if quality_extra:
|
||||
quality.update(quality_extra)
|
||||
settings = Settings(
|
||||
encryption_key=SecretVault.generate_key(),
|
||||
api_token="t" * 32,
|
||||
admin_password="admin-pass",
|
||||
tushare_token="test-token",
|
||||
db_path=db.path,
|
||||
quality=quality,
|
||||
scheduler_enabled=False,
|
||||
)
|
||||
pipe = Pipeline(db, adapter, settings, clock=clock)
|
||||
pipe._tmp = tmp
|
||||
return pipe, db
|
||||
|
||||
|
||||
class ValuationFieldGateTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.transport = ValuationTransport()
|
||||
self.pipe, self.db = make_pipe(self.transport)
|
||||
self.pipe.ingest_reference(TRADE_DATE)
|
||||
|
||||
def _active(self) -> str | None:
|
||||
row = self.db.fetchone(
|
||||
"SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?",
|
||||
(TRADE_DATE,),
|
||||
)
|
||||
return str(row["active_batch"]) if row else None
|
||||
|
||||
def test_normal_batch_with_legit_dv_nulls_passes(self) -> None:
|
||||
result = self.pipe.run_dataset("valuation", TRADE_DATE)
|
||||
self.assertEqual(result["state"], "published")
|
||||
fields = result["quality"]["fields"]
|
||||
# fixture: 1 of 2 stocks has null dv_ttm → 0.5 non-null ≥ 0.3 floor
|
||||
self.assertEqual(fields["dv_ttm"]["nonnull_rate"], 0.5)
|
||||
self.assertEqual(fields["volume_ratio"]["nonnull_rate"], 1.0)
|
||||
self.assertFalse(result["quality"]["errors"])
|
||||
|
||||
def test_all_null_fields_rejected_and_prev_batch_kept(self) -> None:
|
||||
first = self.pipe.run_dataset("valuation", TRADE_DATE)
|
||||
self.transport.mode = "fields_all_null"
|
||||
with self.assertRaises(QualityError) as ctx:
|
||||
self.pipe.run_dataset("valuation", TRADE_DATE)
|
||||
errors = "; ".join(ctx.exception.report["errors"])
|
||||
self.assertIn("field gate: valuation.volume_ratio non-null rate 0.0000 < 0.9", errors)
|
||||
self.assertIn("field gate: valuation.dv_ttm non-null rate 0.0000 < 0.3", errors)
|
||||
# previous good publication stays active
|
||||
self.assertEqual(self._active(), first["batch_id"])
|
||||
# rejected batch left staged with readable error + field stats
|
||||
rejected = self.db.fetchone(
|
||||
"SELECT * FROM batches WHERE state='staged' AND dataset='valuation' ORDER BY started_at DESC",
|
||||
)
|
||||
self.assertIsNotNone(rejected)
|
||||
self.assertIn("field gate: valuation.volume_ratio", str(rejected["error"]))
|
||||
import json
|
||||
|
||||
quality = json.loads(rejected["quality_json"])
|
||||
self.assertEqual(quality["fields"]["volume_ratio"]["nonnull"], 0)
|
||||
self.assertEqual(quality["fields"]["dv_ttm"]["nonnull"], 0)
|
||||
|
||||
def test_volume_ratio_all_null_alone_rejected(self) -> None:
|
||||
self.pipe.run_dataset("valuation", TRADE_DATE)
|
||||
self.transport.mode = "vr_all_null"
|
||||
with self.assertRaises(QualityError):
|
||||
self.pipe.run_dataset("valuation", TRADE_DATE)
|
||||
self.assertEqual(
|
||||
self.db.fetchone(
|
||||
"SELECT active_batch FROM publications WHERE dataset='valuation' AND trade_date=?",
|
||||
(TRADE_DATE,),
|
||||
)["active_batch"],
|
||||
"20240902-valuation-001",
|
||||
)
|
||||
|
||||
def test_dv_ttm_all_null_rejected_by_floor_and_collapse(self) -> None:
|
||||
prev_day = "20240830"
|
||||
prev = self.pipe.run_dataset("valuation", prev_day) # prev dv nonnull 0.5
|
||||
self.transport.mode = "dv_all_null"
|
||||
with self.assertRaises(QualityError) as ctx:
|
||||
self.pipe.run_dataset("valuation", TRADE_DATE)
|
||||
errors = "; ".join(ctx.exception.report["errors"])
|
||||
self.assertIn("field gate: valuation.dv_ttm non-null rate 0.0000 < 0.3", errors)
|
||||
self.assertIn(f"dropped > 0.15 vs prev batch {prev['batch_id']}", errors)
|
||||
|
||||
def test_nonfinite_values_rejected(self) -> None:
|
||||
self.pipe.run_dataset("valuation", TRADE_DATE)
|
||||
rows = self.pipe.adapter.normalize(
|
||||
"valuation", self.pipe._guarded_fetch("valuation", {"trade_date": TRADE_DATE})
|
||||
)
|
||||
for row in rows:
|
||||
row["volume_ratio"] = float("inf")
|
||||
with self.assertRaises(QualityError) as ctx:
|
||||
self.pipe.run_dataset("valuation", TRADE_DATE, prepared_rows=rows)
|
||||
errors = "; ".join(ctx.exception.report["errors"])
|
||||
self.assertIn("field gate: valuation.volume_ratio non-finite rate 1.0000 > 0.01", errors)
|
||||
|
||||
def test_gate_off_when_not_configured(self) -> None:
|
||||
pipe, _db = make_pipe(ValuationTransport(), quality_extra={"field_gates": {}})
|
||||
pipe.ingest_reference(TRADE_DATE)
|
||||
pipe.adapter._transport.mode = "fields_all_null"
|
||||
result = pipe.run_dataset("valuation", TRADE_DATE)
|
||||
self.assertEqual(result["state"], "published") # legacy behavior when unconfigured
|
||||
|
||||
def test_gate_applies_to_any_configured_dataset(self) -> None:
|
||||
gates = {"daily": {"fields": ["volume"], "min_nonnull_rate": 0.9, "max_nonfinite_rate": 0.01}}
|
||||
pipe, _db = make_pipe(ValuationTransport(), quality_extra={"field_gates": gates})
|
||||
pipe.ingest_reference(TRADE_DATE)
|
||||
|
||||
def null_volume(api_name, params, fields):
|
||||
if api_name != "daily":
|
||||
return fake_transport(api_name, params, fields)
|
||||
rows = fake_transport(api_name, params, fields)
|
||||
for row in rows:
|
||||
row["vol"] = None
|
||||
return rows
|
||||
|
||||
pipe.adapter._transport = null_volume
|
||||
with self.assertRaises(QualityError) as ctx:
|
||||
pipe.run_dataset("daily", TRADE_DATE)
|
||||
errors = "; ".join(ctx.exception.report["errors"])
|
||||
self.assertIn("field gate: daily.volume non-null rate 0.0000 < 0.9", errors)
|
||||
|
||||
|
||||
class GateRetryInterplayTests(unittest.TestCase):
|
||||
def test_rejected_valuation_stays_missing_and_retry_publishes_later(self) -> None:
|
||||
transport = ValuationTransport()
|
||||
transport.mode = "fields_all_null"
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
db = HubDB(Path(tmp.name) / "hub.db")
|
||||
adapter = TushareAdapter("x", transport=transport)
|
||||
settings = Settings(
|
||||
encryption_key=SecretVault.generate_key(),
|
||||
db_path=db.path,
|
||||
backup_dir=Path(tmp.name) / "backups",
|
||||
quality={"field_gates": FIELD_GATES, "max_publish_attempts": 2},
|
||||
)
|
||||
pipe = Pipeline(db, adapter, settings)
|
||||
pipe.ingest_reference("20240902")
|
||||
sched = Scheduler(db, pipe)
|
||||
|
||||
sched.tick(clock_at("20240902", 15, 5)) # valuation rejected by field gate
|
||||
sched.tick(clock_at("20240902", 15, 10))
|
||||
self.assertIn("valuation", pipe.missing_official_datasets("20240902"))
|
||||
self.assertEqual(
|
||||
pipe.active_batch("valuation", "20240902"),
|
||||
None,
|
||||
)
|
||||
|
||||
transport.mode = "ok"
|
||||
sched.tick(clock_at("20240902", 15, 45)) # retry passes the gate
|
||||
self.assertNotIn("valuation", pipe.missing_official_datasets("20240902"))
|
||||
rows = db.fetchall("SELECT * FROM eod_valuation WHERE trade_date='20240902'")
|
||||
self.assertTrue(rows)
|
||||
self.assertTrue(all(row["volume_ratio"] is not None for row in rows))
|
||||
|
||||
|
||||
class ForceRepublishTests(unittest.TestCase):
|
||||
def test_run_dataset_over_published_keeps_prev_for_rollback(self) -> None:
|
||||
transport = ValuationTransport()
|
||||
pipe, db = make_pipe(transport)
|
||||
pipe.ingest_reference(TRADE_DATE)
|
||||
first = pipe.run_dataset("valuation", TRADE_DATE)
|
||||
transport.mode = "vr_all_null"
|
||||
with self.assertRaises(QualityError):
|
||||
pipe.run_dataset("valuation", TRADE_DATE) # gate holds: bad re-publish refused
|
||||
transport.mode = "ok"
|
||||
second = pipe.run_dataset("valuation", TRADE_DATE) # CLI --force path
|
||||
self.assertNotEqual(first["batch_id"], second["batch_id"])
|
||||
pub = db.fetchone(
|
||||
"SELECT * FROM publications WHERE dataset='valuation' AND trade_date=?",
|
||||
(TRADE_DATE,),
|
||||
)
|
||||
self.assertEqual(pub["active_batch"], second["batch_id"])
|
||||
self.assertEqual(pub["prev_batch"], first["batch_id"])
|
||||
rolled = pipe.rollback("valuation", TRADE_DATE, actor="cli")
|
||||
self.assertEqual(rolled["active_batch"], first["batch_id"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,168 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
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.serving import V1API
|
||||
from datahub.settings import Settings
|
||||
from tests.fixtures import TRADE_DATE, fake_transport
|
||||
from tests.test_eod_retry import clock_at
|
||||
|
||||
|
||||
class StockMasterTransport:
|
||||
"""fake_transport with a mutable stock_basic list (new listings / renames)."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.stocks = [
|
||||
{"ts_code": "600000.SH", "symbol": "600000", "name": "浦发银行", "area": "上海", "industry": "银行", "market": "主板", "list_status": "L", "list_date": "19991110"},
|
||||
{"ts_code": "920071.BJ", "symbol": "920071", "name": "N金钛", "area": "辽宁", "industry": "小金属", "market": "北交所", "list_status": "L", "list_date": "20240901"},
|
||||
]
|
||||
|
||||
def __call__(self, api_name: str, params: dict, fields: str):
|
||||
if api_name == "stock_basic":
|
||||
return [dict(row) for row in self.stocks]
|
||||
return fake_transport(api_name, params, fields)
|
||||
|
||||
def rename_and_add(self) -> None:
|
||||
for row in self.stocks:
|
||||
if row["ts_code"] == "920071.BJ":
|
||||
row["name"] = "金钛股份" # N-prefix removed the day after listing
|
||||
self.stocks.append(
|
||||
{"ts_code": "920289.BJ", "symbol": "920289", "name": "N华汇", "area": "广东", "industry": "专用机械", "market": "北交所", "list_status": "L", "list_date": "20240902"}
|
||||
)
|
||||
|
||||
|
||||
def make_pipe(transport):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
db = HubDB(Path(tmp.name) / "hub.db")
|
||||
adapter = TushareAdapter("test-token", transport=transport)
|
||||
settings = Settings(
|
||||
encryption_key=SecretVault.generate_key(),
|
||||
api_token="t" * 32,
|
||||
admin_password="admin-pass",
|
||||
tushare_token="test-token",
|
||||
db_path=db.path,
|
||||
quality={"max_publish_attempts": 3, "publication_generations": 3},
|
||||
scheduler_enabled=False,
|
||||
)
|
||||
pipe = Pipeline(db, adapter, settings)
|
||||
pipe._tmp = tmp
|
||||
return pipe, db
|
||||
|
||||
|
||||
class StocksRefreshTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.transport = StockMasterTransport()
|
||||
self.pipe, self.db = make_pipe(self.transport)
|
||||
self.pipe.ingest_reference(TRADE_DATE)
|
||||
|
||||
def _stocks_api(self) -> dict:
|
||||
return V1API(self.db, self.pipe, self.pipe.settings).handle("/v1/stocks", {})
|
||||
|
||||
def test_first_refresh_publishes_snapshot_with_meta(self) -> None:
|
||||
result = self.pipe.refresh_stocks(TRADE_DATE)
|
||||
self.assertEqual(result["state"], "published")
|
||||
self.assertEqual(result["rows"], 2)
|
||||
self.assertTrue(result["batch_id"].startswith("20240902-stocks-"))
|
||||
payload = self._stocks_api()
|
||||
self.assertEqual(payload["meta"]["batch_id"], result["batch_id"])
|
||||
self.assertIsNotNone(payload["meta"]["published_at"])
|
||||
self.assertEqual(len(payload["data"]), 2)
|
||||
names = {row["ts_code"]: row["name"] for row in payload["data"]}
|
||||
self.assertEqual(names["920071.BJ"], "N金钛")
|
||||
self.assertNotIn("batch_id", payload["data"][0])
|
||||
|
||||
def test_new_listing_and_rename_publish_new_batch(self) -> None:
|
||||
first = self.pipe.refresh_stocks(TRADE_DATE)
|
||||
self.transport.rename_and_add()
|
||||
second = self.pipe.refresh_stocks(TRADE_DATE)
|
||||
self.assertEqual(second["state"], "published")
|
||||
self.assertNotEqual(second["batch_id"], first["batch_id"])
|
||||
payload = self._stocks_api()
|
||||
names = {row["ts_code"]: row["name"] for row in payload["data"]}
|
||||
self.assertEqual(names["920071.BJ"], "金钛股份")
|
||||
self.assertIn("920289.BJ", names)
|
||||
self.assertEqual(names["920289.BJ"], "N华汇")
|
||||
# stock_master is refreshed too (code resolution stays current)
|
||||
master = self.db.fetchone("SELECT name FROM stock_master WHERE ts_code='920289.BJ'")
|
||||
self.assertEqual(master["name"], "N华汇")
|
||||
|
||||
def test_unchanged_refresh_is_idempotent(self) -> None:
|
||||
first = self.pipe.refresh_stocks(TRADE_DATE)
|
||||
again = self.pipe.refresh_stocks(TRADE_DATE)
|
||||
self.assertEqual(again["state"], "skipped")
|
||||
self.assertEqual(again["reason"], "unchanged")
|
||||
self.assertEqual(again["batch_id"], first["batch_id"])
|
||||
count = self.db.fetchone(
|
||||
"SELECT COUNT(*) AS n FROM batches WHERE dataset='stocks' AND trade_date=?",
|
||||
(TRADE_DATE,),
|
||||
)["n"]
|
||||
self.assertEqual(count, 1)
|
||||
|
||||
def test_force_republishes_even_unchanged(self) -> None:
|
||||
first = self.pipe.refresh_stocks(TRADE_DATE)
|
||||
forced = self.pipe.refresh_stocks(TRADE_DATE, force=True)
|
||||
self.assertEqual(forced["state"], "published")
|
||||
self.assertNotEqual(forced["batch_id"], first["batch_id"])
|
||||
|
||||
def test_snapshot_pinned_until_next_publish(self) -> None:
|
||||
first = self.pipe.refresh_stocks(TRADE_DATE)
|
||||
self.transport.rename_and_add()
|
||||
# upstream changed but no refresh ran: published snapshot is untouched
|
||||
_, snapshot = self.pipe.published_stock_snapshot(TRADE_DATE)
|
||||
names = {row["ts_code"]: row["name"] for row in snapshot}
|
||||
self.assertEqual(names["920071.BJ"], "N金钛")
|
||||
self.assertNotIn("920289.BJ", names)
|
||||
self.assertEqual(len(snapshot), 2)
|
||||
|
||||
def test_dataset_status_includes_stocks(self) -> None:
|
||||
result = self.pipe.refresh_stocks(TRADE_DATE)
|
||||
payload = V1API(self.db, self.pipe, self.pipe.settings).handle(
|
||||
"/v1/datasets/status", {"date": [TRADE_DATE]}
|
||||
)
|
||||
by_name = {item["dataset"]: item for item in payload["data"]}
|
||||
self.assertIn("stocks", by_name)
|
||||
self.assertEqual(by_name["stocks"]["batch_id"], result["batch_id"])
|
||||
self.assertIsNotNone(by_name["stocks"]["published_at"])
|
||||
|
||||
|
||||
class StocksRefreshSchedulingTests(unittest.TestCase):
|
||||
def _make(self):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
db = HubDB(Path(tmp.name) / "hub.db")
|
||||
transport = StockMasterTransport()
|
||||
adapter = TushareAdapter("x", transport=transport)
|
||||
settings = Settings(
|
||||
encryption_key=SecretVault.generate_key(),
|
||||
db_path=db.path,
|
||||
backup_dir=Path(tmp.name) / "backups",
|
||||
quality={"stocks_refresh_times": ["20:00", "23:10"]},
|
||||
)
|
||||
pipe = Pipeline(db, adapter, settings)
|
||||
pipe.ingest_reference(TRADE_DATE)
|
||||
return db, pipe, Scheduler(db, pipe)
|
||||
|
||||
def test_scheduled_refresh_runs_on_open_day(self) -> None:
|
||||
db, pipe, sched = self._make()
|
||||
ran = sched.tick(clock_at(TRADE_DATE, 20, 0))
|
||||
self.assertIn("stocks_refresh", ran)
|
||||
ran = sched.tick(clock_at(TRADE_DATE, 23, 10))
|
||||
self.assertIn("stocks_refresh", ran) # second slot catches late renames
|
||||
self.assertIsNotNone(pipe.active_batch("stocks", TRADE_DATE))
|
||||
|
||||
def test_no_refresh_on_closed_day(self) -> None:
|
||||
db, _pipe, sched = self._make()
|
||||
sched.tick(clock_at("20240907", 20, 30)) # fixture: Saturday closed
|
||||
runs = db.fetchall("SELECT * FROM job_runs WHERE job_id='stocks_refresh'")
|
||||
self.assertEqual(runs, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user