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