Compare commits

..
Author SHA1 Message Date
16ba83ec01 fix(HEL-461): 切换事务失败写入 release-group 审计日志
整组切换中断时除回滚与废弃批次外,同步记录
action=release-group 的失败审计,便于后台追踪。

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-05 16:02:11 +08:00
1c740a9d48 fix(HEL-461): 后台整组切换异常统一为 FAILED_PRECONDITION
管理后台补数在切换事务中断时不再抛出原始异常,
统一映射为 ApiError FAILED_PRECONDITION,并保留旧完整版本。

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-05 15:56:28 +08:00
75c2e33b68 fix(HEL-461): CLI/后台强制重发改为整组边界切换
eod-refresh --force 与管理后台补数不再单数据集发布,
统一走 force_republish_boundary,避免绕过 A/B 完整边界。

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-05 11:26:27 +08:00
32f565ecb9 fix(HEL-461): 整批发布按完整边界重暂存,主档与快照同事务
边界内任有缺失则整组重暂存后统一切换,避免旧新批次混发;
refresh_stocks 失败时主档保持旧值,并补齐回归测试。

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-05 08:52:50 +08:00
16841e9ae3 fix(HEL-459): 影子比较按请求字段投影,盘后整批原子发布
比较侧只对网站本次请求字段计业务差异,忽略数据中枢额外列;
盘后 A/B/重发改为先整批暂存与交叉校验,再单事务切换公开版本。

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-05 08:39:31 +08:00
multica-agentandmultica-agent bed6450992 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>
2026-09-04 21:36:20 +08:00
总工andmultica-agent c9892050c3 feat(HEL-435): 盘后未出数时晚间自动重试并提供安全补跑
Co-authored-by: multica-agent <github@multica.ai>
2026-09-03 22:30:15 +08:00
a836cda1b2 feat(HEL-421): 回补历史日历和指数并标记区间不完整
Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-02 22:16:05 +08:00
总工 25ff6bbe06 fix(HEL-417): 配置根日志让影子对比报告落入容器日志 2026-09-02 21:53:19 +08:00
28 changed files with 3345 additions and 142 deletions
+12
View File
@@ -1,11 +1,23 @@
from __future__ import annotations
import argparse
import logging
from http.server import ThreadingHTTPServer
from typing import Any
def configure_logging() -> None:
"""让 INFO 级结构化日志(含 datahub 影子对比报告)落到容器日志。"""
if logging.getLogger().handlers:
return
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s %(message)s",
)
def main(handler_class: type[Any] | None = None, service: Any | None = None) -> None:
configure_logging()
if handler_class is None or service is None:
from backend.application import RequestHandler, SERVICE
+8 -2
View File
@@ -122,10 +122,12 @@ class DatahubBridge:
legacy_rows = legacy_query(api_name, params, fields)
except Exception as exc:
if flags.read and hub_rows is not None and hub_error is None:
self._emit_shadow(compare_rows(dataset, [], hub_canonical, hub_meta, self._error_text(exc)))
self._emit_shadow(
compare_rows(dataset, [], hub_canonical, hub_meta, self._error_text(exc), fields)
)
return project_fields(hub_rows, fields)
raise
self._emit_shadow(compare_rows(dataset, legacy_rows, hub_canonical, hub_meta, hub_error))
self._emit_shadow(compare_rows(dataset, legacy_rows, hub_canonical, hub_meta, hub_error, fields))
if flags.read and hub_rows is not None and hub_error is None:
return project_fields(hub_rows, fields)
return legacy_rows
@@ -208,6 +210,10 @@ class DatahubBridge:
raise DatahubError("STALE", f"{dataset} data is stale")
if dataset in EMPTY_FAIL_DATASETS and not rows:
raise DatahubError("EMPTY", f"{dataset} returned no rows")
coverage = meta.get("coverage") if isinstance(meta.get("coverage"), dict) else {}
if meta.get("incomplete") is True or coverage.get("complete") is False:
missing = coverage.get("missing_count")
raise DatahubError("INCOMPLETE", f"{dataset} range is incomplete missing={missing}")
def _require_fresh(self, response: DatahubResponse, dataset: str) -> DatahubResponse:
self._validate_usable(dataset, list(response.data or []) if isinstance(response.data, list) else [], response)
+29 -2
View File
@@ -5,6 +5,7 @@ from typing import Any
from backend.data.datahub.native import SCALE_FIELDS, row_key, to_canonical_row, yyyymmdd
NUMERIC_TOLERANCE = 1e-4
CANONICAL_ALIASES = {"volume": "vol"}
def compare_rows(
@@ -13,8 +14,10 @@ def compare_rows(
hub_rows: list[dict[str, Any]] | None,
hub_meta: dict[str, Any] | None = None,
hub_error: str | None = None,
fields: str = "",
) -> dict[str, Any]:
hub = hub_rows or []
requested = _requested_fields(fields)
legacy_map = {row_key(dataset, row): row for row in legacy_rows}
hub_map = {row_key(dataset, _align_hub_row(row)): row for row in hub}
missing_hub = sorted(key for key in legacy_map if key not in hub_map)
@@ -26,7 +29,7 @@ def compare_rows(
hub_row = hub_map.get(key)
if hub_row is None:
continue
field_report = _compare_fields(dataset, legacy, hub_row)
field_report = _compare_fields(dataset, legacy, hub_row, requested)
if field_report["unit_conversion"]:
unit_conversion.append({"key": list(key), "fields": field_report["unit_conversion"]})
if field_report["value_diff"]:
@@ -53,6 +56,7 @@ def compare_rows(
"published_at": (hub_meta or {}).get("published_at"),
"trade_date": yyyymmdd((hub_meta or {}).get("trade_date")),
"hub_error": hub_error,
"fields_compared": sorted(requested) if requested is not None else None,
"equal": (
not hub_error
and not missing_hub
@@ -71,13 +75,36 @@ def _align_hub_row(row: dict[str, Any]) -> dict[str, Any]:
return aligned
def _compare_fields(dataset: str, legacy: dict[str, Any], hub: dict[str, Any]) -> dict[str, list[dict[str, Any]]]:
def _requested_fields(fields: str) -> list[str] | None:
"""Fields the website actually asked for; None means "no projection"."""
keys = [item.strip() for item in str(fields or "").split(",") if item.strip()]
if not keys:
return None
seen: list[str] = []
for key in keys:
canonical = CANONICAL_ALIASES.get(key, key)
if canonical not in seen:
seen.append(canonical)
return seen
def _compare_fields(
dataset: str,
legacy: dict[str, Any],
hub: dict[str, Any],
requested: list[str] | None = None,
) -> dict[str, list[dict[str, Any]]]:
canonical_legacy = to_canonical_row(dataset, legacy)
hub_canonical = _hub_canonical(dataset, hub)
native_hub = _align_hub_row(hub)
value_diff: list[dict[str, Any]] = []
unit_conversion: list[dict[str, Any]] = []
keys = (set(canonical_legacy) | set(hub_canonical)) - {"batch_id", "updated_at", "volume"}
if requested is not None:
# Compare only what the website asked for. Extra hub columns are
# transport detail, not business differences; a requested field still
# alarms when it is missing or holds a different value.
keys = set(requested) - {"batch_id", "updated_at", "volume"}
scales = SCALE_FIELDS.get(dataset) or {}
for field in sorted(keys):
left = canonical_legacy.get(field)
+34
View File
@@ -0,0 +1,34 @@
import logging
import unittest
from backend.bootstrap.runtime import configure_logging
class ConfigureLoggingTest(unittest.TestCase):
def setUp(self) -> None:
self._saved_handlers = logging.getLogger().handlers[:]
self._saved_level = logging.getLogger().level
logging.getLogger().handlers.clear()
def tearDown(self) -> None:
logging.getLogger().handlers[:] = self._saved_handlers
logging.getLogger().setLevel(self._saved_level)
def test_configures_root_logger_at_info(self) -> None:
configure_logging()
root = logging.getLogger()
self.assertTrue(root.handlers)
self.assertEqual(root.level, logging.INFO)
with self.assertLogs("xiaobai.datahub", level="INFO") as captured:
logging.getLogger("xiaobai.datahub").info("datahub shadow %s", {"dataset": "daily"})
self.assertIn("datahub shadow", captured.output[0])
def test_keeps_existing_configuration(self) -> None:
handler = logging.NullHandler()
logging.getLogger().addHandler(handler)
configure_logging()
self.assertEqual(logging.getLogger().handlers, [handler])
if __name__ == "__main__":
unittest.main()
+114 -1
View File
@@ -126,7 +126,7 @@ class DatahubBridgeTests(unittest.TestCase):
self.assertEqual(calendar[0]["is_open"], 1)
self.assertEqual(calendar_client.paths, [])
def test_fallback_on_down_401_timeout_empty_unpublished_and_stale(self) -> None:
def test_fallback_on_down_401_timeout_empty_unpublished_stale_and_incomplete(self) -> None:
cases = [
DatahubError("UNAVAILABLE", "down"),
DatahubError("UNAUTHORIZED", "401"),
@@ -134,6 +134,7 @@ class DatahubBridgeTests(unittest.TestCase):
DatahubError("EMPTY", "no rows"),
DatahubError("DATASET_NOT_PUBLISHED", "not ready"),
DatahubError("STALE", "old"),
DatahubError("INCOMPLETE", "truncated"),
]
for error in cases:
with self.subTest(error=error.code):
@@ -144,6 +145,16 @@ class DatahubBridgeTests(unittest.TestCase):
data=[dict(HUB_DAILY)],
meta={"stale": True, "staleness_seconds": 999999},
))
elif error.code == "INCOMPLETE":
client = FakeClient(response=DatahubResponse(
data=[dict(HUB_DAILY)],
meta={
"stale": False,
"staleness_seconds": 0,
"incomplete": True,
"coverage": {"complete": False, "missing_count": 80},
},
))
else:
client = FakeClient(error=error)
legacy = FakeLegacy([LEGACY_DAILY])
@@ -190,6 +201,87 @@ class DatahubBridgeTests(unittest.TestCase):
skew = compare_rows("daily", [LEGACY_DAILY], [HUB_DAILY], {"stale": False, "staleness_seconds": 12})
self.assertTrue(skew["time_skew"])
def test_shadow_extra_hub_columns_are_not_false_diffs_when_projected(self) -> None:
hub_full = {**HUB_DAILY, "adj_factor": 1.1}
legacy_close_only = {k: LEGACY_DAILY[k] for k in ("ts_code", "trade_date", "close")}
report = compare_rows(
"daily", [legacy_close_only], [hub_full],
{"stale": False, "staleness_seconds": 0},
fields="ts_code,trade_date,close",
)
self.assertTrue(report["equal"])
self.assertEqual(report["value_diff_count"], 0)
self.assertEqual(report["fields_compared"], ["close", "trade_date", "ts_code"])
# without projection the same pair shows the historic false diff
unprojected = compare_rows("daily", [legacy_close_only], [hub_full])
self.assertFalse(unprojected["equal"])
legacy_stocks = {"ts_code": "600000.SH", "name": "浦发银行"}
hub_stocks = {
"ts_code": "600000.SH", "symbol": "600000", "name": "浦发银行", "area": "上海",
"industry": "银行", "market": "主板", "list_status": "L", "list_date": "19991110",
}
stocks = compare_rows("stocks", [legacy_stocks], [hub_stocks], {}, fields="ts_code,name")
self.assertTrue(stocks["equal"])
legacy_cal = {"cal_date": "20240902", "is_open": 1}
hub_cal = {
"cal_date": "20240902", "is_open": True,
"pretrade_date": "20240830", "prev_open": "20240830",
}
calendar = compare_rows(
"calendar", [legacy_cal], [hub_cal], {}, fields="cal_date,is_open"
)
self.assertTrue(calendar["equal"])
def test_shadow_projection_still_alarms_on_requested_field_problems(self) -> None:
hub_missing_field = {k: v for k, v in HUB_DAILY.items() if k != "close"}
legacy_close_only = {k: LEGACY_DAILY[k] for k in ("ts_code", "trade_date", "close")}
lost = compare_rows(
"daily", [legacy_close_only], [hub_missing_field], fields="ts_code,trade_date,close"
)
self.assertFalse(lost["equal"])
self.assertEqual(lost["value_diff_count"], 1)
changed = compare_rows(
"daily", [legacy_close_only], [{**HUB_DAILY, "close": 99.0}],
fields="ts_code,trade_date,close",
)
self.assertFalse(changed["equal"])
self.assertEqual(changed["value_diff_count"], 1)
self.assertEqual(changed["value_diffs"][0]["fields"][0]["field"], "close")
gone = compare_rows("daily", [LEGACY_DAILY], [], fields="ts_code,trade_date,close")
self.assertEqual(gone["missing_hub_count"], 1)
self.assertFalse(gone["equal"])
unit = compare_rows(
"daily", [LEGACY_DAILY], [{**HUB_DAILY, "amount": 2000.0, "volume": 1000.0}],
fields="ts_code,trade_date,vol,amount",
)
self.assertGreater(unit["unit_conversion_count"], 0)
self.assertFalse(unit["equal"])
def test_bridge_shadow_report_uses_website_request_fields(self) -> None:
hub_full = {**HUB_DAILY, "adj_factor": 1.1}
legacy_close_only = {k: LEGACY_DAILY[k] for k in ("ts_code", "trade_date", "close", "vol", "amount")}
reports: list[dict[str, Any]] = []
client = FakeClient(
response=DatahubResponse(
data=[hub_full],
meta={"tier": "official", "trade_date": "20240902", "stale": False, "staleness_seconds": 0},
)
)
wrapped = DatahubAwareTushareClient(
FakeLegacy([legacy_close_only]),
DatahubBridge(flags(daily=(False, True)), client, shadow_sink=reports.append),
)
rows = wrapped.query("daily", {"trade_date": "20240902"}, "ts_code,trade_date,close,vol,amount")
self.assertEqual(rows[0]["close"], 10.20)
self.assertEqual(rows[0]["vol"], 1000.0)
self.assertTrue(reports[0]["equal"])
self.assertEqual(reports[0]["matched"], 1)
def test_native_roundtrip_matches_known_scales(self) -> None:
native = to_native_row("daily", HUB_DAILY)
self.assertEqual(native["vol"], 1000.0)
@@ -235,6 +327,27 @@ class DatahubBridgeTests(unittest.TestCase):
self.assertIsInstance(client, DatahubAwareTushareClient)
self.assertFalse(gateway.datahub.settings.any_enabled())
def test_stock_detail_range_query_is_not_silently_accepted_when_incomplete(self) -> None:
source = (ROOT / "backend" / "data" / "providers" / "tushare_stocks.py").read_text(encoding="utf-8")
self.assertIn('"daily"', source)
self.assertIn("start_date", source)
self.assertIn("end_date", source)
client = FakeClient(
response=DatahubResponse(
data=[dict(HUB_DAILY)],
meta={"stale": False, "staleness_seconds": 0, "incomplete": True, "coverage": {"complete": False, "missing_count": 89}},
)
)
legacy = FakeLegacy([LEGACY_DAILY])
wrapped = DatahubAwareTushareClient(legacy, DatahubBridge(flags(daily=(True, False)), client))
rows = wrapped.query(
"daily",
{"ts_code": "600000.SH", "start_date": "20240301", "end_date": "20240902"},
"ts_code,amount",
)
self.assertEqual(rows[0]["amount"], 2000.0)
self.assertEqual(len(legacy.calls), 1)
def test_features_do_not_import_datahub_client(self) -> None:
violations = []
for path in (ROOT / "backend" / "features").rglob("*.py"):
+60
View File
@@ -63,6 +63,66 @@ python -m unittest discover -s tests -v
不调用真实 Tushare;用内存/临时库和假适配器。
## 历史回补
交易日历默认从 `20160101` 拉到今天后 30 天;盘前 `precheck` 与手动回补都走同一 UPSERT,可重复执行。
网站实际使用的指数(上证、深成、创业板、沪深300)按交易日增量发布,默认覆盖 260 个交易日(大于现有 90 天窗口,并覆盖智能选股基准回看)。已发布日期默认跳过。
```bash
cd xiaobai-datahub
python -m datahub history-backfill
# 可选:--calendar-start 20160101 --index-days 260 --force
```
管理后台也可手动跑 `history_backfill` 任务,或 `POST /admin/api/backfill``dataset=history`、确认词 `history:full`
区间接口在 `meta.coverage` / `meta.incomplete` 标明覆盖是否完整;网站只读接入把不完整区间视为不可用并回旧链路。个股日 K 的 90 天区间查询依赖已核实,本阶段不回补全市场历史。
## 估值字段级质量门
`hub-quality.config.json``field_gates` 按数据集配置关键字段:非空率下限(支持按字段覆盖,如 `dv_ttm` 合法高空值)、非有限值比例上限、以及相对上一已发布批次的非空率塌陷保护。字段大面积为空的批次会被拒绝发布、保留上一份正常正式数据,失败原因逐字段写入 `batches.error` / `quality_json`。被拒后数据集仍视为缺失,盘后自动重试(HEL-435 机制)会继续尝试直到成功或截止。配置对任意数据集生效,不写死单日或单字段。
## 整批原子发布(release group
盘后发布/重发(eod_a、eod_retry、`eod-refresh`、跨数据集重发)不再逐数据集各自切换,而是走整批原子可见机制:
- 一致性边界:日 K、估值、资金流、竞价同属 A 组整批;指数日 K 为 B 组;当日股票主档快照随 A 组一同切换(主档 `stock_master` 的 UPSERT 与快照发布同一事务,不会出现主档先行/滞后)。
- 流程:组内全部成员先在暂存表完成拉取、字段质量门、覆盖检查和跨数据集交叉校验(`cross_gates` 配置 ts_code 覆盖重叠率下限),全部达标后才在**一个 SQLite 事务**里复制正式表并翻转全部 `publications` 指针。
- 任一成员失败(拉取失败、质量门拒绝、交叉校验不过、切换事务中断)→ 整批不切换,对外继续提供上一份完整正式版本,失败原因写入 `batches.error``audit_log``action=release-group`),等待晚间自动重试。
- 读取侧任何时刻只会看到"旧完整版本"或"新完整版本":发布指针在单事务内统一翻转,容器重启/事务中断自动回滚,不暴露字段残缺或跨数据集混合版本。
- 幂等:仅当一致性边界内全部成员都已发布时才整组跳过;边界内任有缺失则整组重暂存后统一切换,避免旧批次与新批次混在同一次重发中。重复执行、并发重试不会在完整边界已就绪时生成重复批次(调度器另有 EOD 互斥锁)。
## 股票主档每日刷新与发布
交易日 20:00 与 23:10`stocks_refresh_times` 可配)自动刷新股票主档并发布版本化快照(`eod_stocks` + `publications.dataset='stocks'`),覆盖当日新上市、证券简称变化和上市首日 N/C 前缀摘除;无变化则跳过,重复执行幂等。`/v1/stocks` 从最新已发布快照提供数据并带 `batch_id` / `published_at``/v1/datasets/status` 同步展示 stocks 状态。
```bash
cd xiaobai-datahub
python -m datahub stocks-refresh # 手动触发;--force 无变化也重发
```
## 资金流历史回补
网站会沿真实调用链查最近若干交易日的 moneyflow(个股详情任意日期点查 + 智能选股最近 5 个交易日),默认回补最近 60 个交易日(`moneyflow_history_trading_days` 可配,已发布日期自动跳过)。点查未覆盖的历史日期返回 `DATASET_NOT_PUBLISHED` 并附 `available_from` / `available_to`(低于下界时 `reason=history_not_backfilled`),网站据此明确回退旧链路,不会静默拿到半截数据。
```bash
cd xiaobai-datahub
python -m datahub moneyflow-backfill # --trading-days 60 --end-date --force 可选
```
## 盘后补跑与强制重发
```bash
cd xiaobai-datahub
python -m datahub eod-refresh --trade-date 20260904 # 补不完整的 A/B 边界
python -m datahub eod-refresh --trade-date 20260904 --force --dataset valuation
# --force 按一致性边界整组重发:valuation/daily/moneyflow/auction/stocks → A 组;
# index_daily → B 组。不可再单独切换某一个正式数据集。
```
管理后台「补数」对盘后正式数据集同样走 `force_republish_boundary`,不会绕过 A/B 整批边界。
## 备份
每日 00:40 任务把 `datahub.db` 备份到 `data/backups/`(保留 14 份)。也可手动:
+19 -1
View File
@@ -104,11 +104,29 @@ async function render() {
if (state.page === "overview") {
const data = await api("/admin/api/overview");
$("phase").textContent = data.session_phase;
const eod = data.eod_status || {};
const eodLabels = {
pending_first_attempt: "等待首次尝试",
waiting_upstream: "等待上游",
done: "已成功",
cutoff_failed: "已截止失败",
closed_day: "休市",
};
const eodExtra = [];
if (eod.state === "waiting_upstream") {
eodExtra.push(`已试 ${eod.attempts}`);
if (eod.next_retry_at) eodExtra.push(`下次重试 ${esc(String(eod.next_retry_at).replace("T", " ").slice(11, 16))}`);
if (eod.missing_datasets && eod.missing_datasets.length) eodExtra.push(`${esc(eod.missing_datasets.join(","))}`);
}
if (eod.state === "cutoff_failed" && eod.missing_datasets) {
eodExtra.push(`${esc(eod.missing_datasets.join(","))}`);
}
page.innerHTML = `
<div class="cards">
<div class="card"><div class="muted">交易日</div><strong>${esc(data.trade_date)}</strong></div>
<div class="card"><div class="muted">阶段</div><strong>${esc(data.session_phase)}</strong></div>
<div class="card"><div class="muted">今日发布</div><strong>${data.publications.length}</strong></div>
<div class="card"><div class="muted">盘后补跑</div><strong>${esc(eodLabels[eod.state] || eod.state || "-")}</strong><div class="muted">${eodExtra.join(" · ")}</div></div>
<div class="card"><div class="muted">异常批次</div><strong class="${data.anomalies.length ? "fail" : "ok"}">${data.anomalies.length}</strong></div>
</div>
<h2>最近调用</h2>
@@ -251,7 +269,7 @@ function renderRelease(data) {
async function dangerous(kind, dataset) {
const date = ($("rel-date") && $("rel-date").value) || "";
const ds = dataset || prompt("数据集(daily / valuation / moneyflow / auction / index_daily / reference", "daily");
const ds = dataset || prompt("数据集(daily/valuation/moneyflow/auction/stocks→A组整批;index_daily→B组;或 reference", "daily");
if (!ds) return;
const password = prompt("二次确认:输入管理密码");
if (!password) return;
+44 -1
View File
@@ -11,5 +11,48 @@
"publication_generations": 3,
"tushare_rate_per_minute": 300,
"list_limit_default": 5000,
"list_limit_max": 5000
"list_limit_max": 5000,
"calendar_start": "20160101",
"index_history_trading_days": 260,
"eod_retry_start": "15:15",
"eod_retry_interval_minutes": 30,
"eod_retry_cutoff": "23:30",
"moneyflow_history_trading_days": 60,
"stocks_refresh_times": [
"20:00",
"23:10"
],
"cross_gates": [
{
"left": "daily",
"right": "valuation",
"min_key_overlap": 0.98
},
{
"left": "daily",
"right": "moneyflow",
"min_key_overlap": 0.98
}
],
"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
}
}
}
+4
View File
@@ -0,0 +1,4 @@
from datahub.cli import main
if __name__ == "__main__":
raise SystemExit(main())
+4 -1
View File
@@ -44,7 +44,10 @@ DATASET_API = {
"auction": "stk_auction",
}
DEFAULT_INDEX_CODES = ("000001.SH", "399001.SZ", "399006.SZ", "000300.SH")
# Website actual index usage: market cards / 90-day charts (SH/SZ/CYB) plus
# screener 沪深300 benchmark (lookback up to 260 trading days).
WEBSITE_INDEX_CODES = ("000001.SH", "399001.SZ", "399006.SZ", "000300.SH")
DEFAULT_INDEX_CODES = WEBSITE_INDEX_CODES
class TushareAdapter(MarketAdapter):
+30 -6
View File
@@ -6,7 +6,7 @@ from typing import Any
from datahub.adapters import RESERVED
from datahub.auth import AuthService
from datahub.db import HubDB
from datahub.pipeline import Pipeline
from datahub.pipeline import OFFICIAL_DATASETS, STOCKS_DATASET, Pipeline
from datahub.scheduler import Scheduler
from datahub.serving import ApiError
from datahub.timeutil import isoformat, now_shanghai, session_phase, yyyymmdd
@@ -38,6 +38,7 @@ class AdminAPI:
"trade_date": today,
"session_phase": session_phase(now_shanghai(), is_open),
"is_open_day": is_open,
"eod_status": self.scheduler.eod_status(today),
"publications": pubs,
"anomalies": failed,
"recent_calls": _public_calls(calls),
@@ -83,11 +84,15 @@ class AdminAPI:
def jobs(self) -> dict[str, Any]:
runs = self.db.fetchall("SELECT * FROM job_runs ORDER BY id DESC LIMIT 100")
stocks_times = "/".join(self.pipeline.settings.stocks_refresh_times) or "20:00"
return {
"jobs": [
{"id": "precheck", "at": "08:45", "title": "盘前预检"},
{"id": "eod_a", "at": "15:05", "title": "盘后批 A daily/valuation/moneyflow/auction"},
{"id": "eod_b", "at": "15:10", "title": "盘后批 B index_daily"},
{"id": "eod_retry", "at": "15:15-23:30", "title": "盘后未出数自动重试(每 30 分钟,成功即停)"},
{"id": "stocks_refresh", "at": stocks_times, "title": "股票主档刷新与正式发布(新上市/更名,无变化跳过)"},
{"id": "history_backfill", "at": "manual", "title": "回补历史日历与指数日 K"},
{"id": "cleanup", "at": "00:30", "title": "清理 staging / 日志"},
{"id": "backup", "at": "00:40", "title": "SQLite 备份"},
],
@@ -130,12 +135,31 @@ class AdminAPI:
return result
def backfill(self, dataset: str, trade_date: str, password: str, confirm: str, actor: str) -> dict[str, Any]:
self._dangerous(password, confirm, f"{dataset}:{trade_date}")
if dataset == "reference":
result = self.pipeline.ingest_reference(trade_date)
day = yyyymmdd(trade_date or now_shanghai())
if dataset == "history":
self._dangerous(password, confirm, "history:full")
result = self.pipeline.backfill_history(day)
else:
result = self.pipeline.run_dataset(dataset, trade_date)
self.pipeline.audit(actor, "backfill", f"{dataset}:{trade_date}", json.dumps({"ok": True}))
self._dangerous(password, confirm, f"{dataset}:{day}")
if dataset == "reference":
result = self.pipeline.ingest_reference(day)
elif dataset in OFFICIAL_DATASETS or dataset == STOCKS_DATASET:
# Manual same-day republish must rebuild the full A/B boundary.
# Gate failures and mid-switch exceptions both surface as
# FAILED_PRECONDITION so the admin API never leaks raw
# transaction errors to the client.
try:
result = self.pipeline.force_republish_boundary(dataset, day)
failures = self.pipeline.eod_failures(result)
if failures:
raise ApiError("FAILED_PRECONDITION", "; ".join(failures))
except ApiError:
raise
except Exception as exc:
raise ApiError("FAILED_PRECONDITION", str(exc)) from exc
else:
raise ApiError("INVALID_ARGUMENT", f"unsupported backfill dataset: {dataset}")
self.pipeline.audit(actor, "backfill", f"{dataset}:{day}", json.dumps({"ok": True}))
return result
def _dangerous(self, password: str, confirm: str, expected: str) -> None:
+112
View File
@@ -0,0 +1,112 @@
"""Command-line entry for one-shot datahub operations."""
from __future__ import annotations
import argparse
import json
import sys
from datahub.hub import build_hub
from datahub.pipeline import EOD_A_DATASETS, OFFICIAL_DATASETS, STOCKS_DATASET
from datahub.settings import load_settings
from datahub.timeutil import yyyymmdd
def main(argv: list[str] | None = None) -> int:
parser = argparse.ArgumentParser(description="xiaobai-datahub CLI")
sub = parser.add_subparsers(dest="command", required=True)
history = sub.add_parser("history-backfill", help="回补 2016 年起交易日历和网站所用指数日 K")
history.add_argument("--calendar-start", default=None, help="日历起点,默认配置 calendar_start")
history.add_argument("--index-days", type=int, default=None, help="指数回补交易日数量,默认 260")
history.add_argument("--force", action="store_true", help="覆盖已发布的指数日期")
refresh = sub.add_parser("eod-refresh", help="对指定交易日补跑盘后正式数据(跳过已完整发布的一致性边界,仍走质量门禁)")
refresh.add_argument("--trade-date", default=None, help="交易日 YYYYMMDD,默认今天")
refresh.add_argument(
"--force", action="store_true",
help="强制重发 --dataset 所属的完整一致性边界(A 组或 B 组),生成新批次并保留上一批次可回滚",
)
refresh.add_argument(
"--dataset", default=None,
help="配合 --force:指定边界内任一成员(如 valuation→整组 Aindex_daily→整组 B",
)
stocks_refresh = sub.add_parser("stocks-refresh", help="刷新股票主档并发布正式快照(幂等:无变化则跳过)")
stocks_refresh.add_argument("--trade-date", default=None, help="交易日 YYYYMMDD,默认今天")
stocks_refresh.add_argument("--force", action="store_true", help="即使快照无变化也重新发布")
moneyflow_backfill = sub.add_parser(
"moneyflow-backfill", help="回补资金流历史(默认覆盖网站所需的最近 N 个交易日,跳过已发布日期)",
)
moneyflow_backfill.add_argument("--end-date", default=None, help="截止交易日 YYYYMMDD,默认今天")
moneyflow_backfill.add_argument("--trading-days", type=int, default=None, help="回补交易日数量,默认配置 moneyflow_history_trading_days")
moneyflow_backfill.add_argument("--force", action="store_true", help="覆盖已发布的资金流日期")
args = parser.parse_args(argv)
settings = load_settings()
hub = build_hub(settings)
if args.command == "history-backfill":
result = hub.pipeline.backfill_history(
calendar_start=args.calendar_start,
index_days=args.index_days,
force=args.force,
)
json.dump(result, sys.stdout, ensure_ascii=False, indent=2, default=str)
sys.stdout.write("\n")
return 0 if result.get("ok") else 1
if args.command == "eod-refresh":
day = yyyymmdd(args.trade_date) if args.trade_date else yyyymmdd()
if args.force:
allowed = set(OFFICIAL_DATASETS) | {STOCKS_DATASET}
if not args.dataset:
parser.error("--force requires --dataset (e.g. --dataset valuation)")
if args.dataset not in allowed:
parser.error(f"unknown dataset: {args.dataset}")
result = hub.pipeline.force_republish_boundary(args.dataset, day)
boundary = "A" if args.dataset in EOD_A_DATASETS or args.dataset == STOCKS_DATASET else "B"
else:
result = hub.pipeline.run_eod_missing(day)
boundary = None
hub.pipeline.audit("cli", "eod-refresh", f"eod:{day}", json.dumps(
{"force": bool(args.force), "dataset": args.dataset, "boundary": boundary,
**{name: item.get("state") for name, item in result.items() if isinstance(item, dict)}},
ensure_ascii=False,
))
if args.force:
failures = hub.pipeline.eod_failures(result)
payload = {"trade_date": day, "boundary": boundary, "datasets": result}
json.dump(payload, sys.stdout, ensure_ascii=False, indent=2, default=str)
sys.stdout.write("\n")
return 0 if not failures else 1
missing = hub.pipeline.missing_official_datasets(day)
payload = {"trade_date": day, "datasets": result, "missing_after": missing}
json.dump(payload, sys.stdout, ensure_ascii=False, indent=2, default=str)
sys.stdout.write("\n")
return 0 if not missing else 1
if args.command == "stocks-refresh":
day = yyyymmdd(args.trade_date) if args.trade_date else yyyymmdd()
result = hub.pipeline.refresh_stocks(day, force=args.force)
hub.pipeline.audit("cli", "stocks-refresh", f"stocks:{day}", json.dumps(
{"force": bool(args.force), "state": result.get("state"), "batch_id": result.get("batch_id")},
ensure_ascii=False,
))
json.dump(result, sys.stdout, ensure_ascii=False, indent=2, default=str)
sys.stdout.write("\n")
return 0 if result.get("state") != "failed" else 1
if args.command == "moneyflow-backfill":
result = hub.pipeline.backfill_moneyflow_history(
end_date=args.end_date,
trading_days=args.trading_days,
force=args.force,
)
hub.pipeline.audit("cli", "moneyflow-backfill", f"moneyflow:{result.get('end')}", json.dumps(
{"published": len(result.get("published") or []), "skipped": len(result.get("skipped") or []),
"failed": len(result.get("failed") or [])},
ensure_ascii=False,
))
json.dump(result, sys.stdout, ensure_ascii=False, indent=2, default=str)
sys.stdout.write("\n")
return 0 if result.get("ok") else 1
parser.error(f"unknown command: {args.command}")
return 2
if __name__ == "__main__":
raise SystemExit(main())
+130
View File
@@ -0,0 +1,130 @@
from __future__ import annotations
from typing import Any, Iterable
from datahub.db import HubDB
from datahub.timeutil import iter_yyyymmdd, yyyymmdd
MISSING_SAMPLE_LIMIT = 10
def coverage_payload(
*,
kind: str,
start: str,
end: str,
expected: Iterable[str],
available: Iterable[str],
extra: dict[str, Any] | None = None,
) -> dict[str, Any]:
start = yyyymmdd(start)
end = yyyymmdd(end)
expected_list = sorted({yyyymmdd(item) for item in expected if item})
available_set = {yyyymmdd(item) for item in available if item}
missing = [item for item in expected_list if item not in available_set]
payload: dict[str, Any] = {
"kind": kind,
"complete": not missing,
"requested_from": start,
"requested_to": end,
"available_from": min(available_set) if available_set else None,
"available_to": max(available_set) if available_set else None,
"expected_count": len(expected_list),
"available_count": len(available_set),
"missing_count": len(missing),
"missing_sample": missing[:MISSING_SAMPLE_LIMIT],
}
if extra:
payload.update(extra)
return payload
def calendar_coverage(db: HubDB, start: str, end: str, exchange: str = "SSE") -> dict[str, Any]:
start = yyyymmdd(start)
end = yyyymmdd(end)
expected = list(iter_yyyymmdd(start, end))
rows = db.fetchall(
"SELECT cal_date FROM trade_calendar WHERE exchange = ? AND cal_date >= ? AND cal_date <= ?",
(exchange, start, end),
)
return coverage_payload(
kind="calendar",
start=start,
end=end,
expected=expected,
available=(row["cal_date"] for row in rows),
extra={"exchange": exchange},
)
def published_range_coverage(
db: HubDB,
dataset: str,
start: str,
end: str,
ts_code: str = "",
table: str = "",
) -> dict[str, Any]:
start = yyyymmdd(start)
end = yyyymmdd(end)
calendar = calendar_coverage(db, start, end)
open_rows = db.fetchall(
"""
SELECT cal_date FROM trade_calendar
WHERE exchange = 'SSE' AND is_open = 1 AND cal_date >= ? AND cal_date <= ?
ORDER BY cal_date
""",
(start, end),
)
expected_open = [row["cal_date"] for row in open_rows]
pubs = db.fetchall(
"""
SELECT trade_date, active_batch FROM publications
WHERE dataset = ? AND trade_date >= ? AND trade_date <= ?
ORDER BY trade_date
""",
(dataset, start, end),
)
published_dates = [row["trade_date"] for row in pubs]
available = list(published_dates)
extra: dict[str, Any] = {
"dataset": dataset,
"calendar_complete": calendar["complete"],
"calendar_missing_count": calendar["missing_count"],
}
if ts_code and table and pubs:
present_code: list[str] = []
for pub in pubs:
hit = db.fetchone(
f"SELECT 1 AS ok FROM {table} WHERE trade_date = ? AND batch_id = ? AND ts_code = ? LIMIT 1",
(pub["trade_date"], pub["active_batch"], ts_code),
)
if hit:
present_code.append(pub["trade_date"])
available = present_code
extra["code"] = ts_code
payload = coverage_payload(
kind="published_range",
start=start,
end=end,
expected=expected_open,
available=available,
extra=extra,
)
if not calendar["complete"]:
payload["complete"] = False
payload["calendar_missing_sample"] = calendar["missing_sample"]
return payload
def point_coverage(trade_date: str, dataset: str = "") -> dict[str, Any]:
day = yyyymmdd(trade_date)
payload = coverage_payload(
kind="point",
start=day,
end=day,
expected=[day],
available=[day],
extra={"dataset": dataset} if dataset else None,
)
return payload
+27
View File
@@ -114,6 +114,21 @@ CREATE TABLE IF NOT EXISTS eod_index_bars (
PRIMARY KEY (ts_code, trade_date, batch_id)
) WITHOUT ROWID;
CREATE TABLE IF NOT EXISTS eod_stocks (
ts_code TEXT NOT NULL, trade_date TEXT NOT NULL,
symbol TEXT, name TEXT, area TEXT, industry TEXT, market TEXT,
list_status TEXT, list_date TEXT,
batch_id TEXT NOT NULL,
PRIMARY KEY (ts_code, trade_date, batch_id)
) WITHOUT ROWID;
CREATE TABLE IF NOT EXISTS staging_stocks (
ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, batch_id TEXT NOT NULL,
symbol TEXT, name TEXT, area TEXT, industry TEXT, market TEXT,
list_status TEXT, list_date TEXT,
PRIMARY KEY (batch_id, ts_code, trade_date)
);
CREATE TABLE IF NOT EXISTS staging_bars (
ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, batch_id TEXT NOT NULL,
open REAL, high REAL, low REAL, close REAL, pct_chg REAL,
@@ -212,6 +227,17 @@ CREATE TABLE IF NOT EXISTS job_runs (
detail TEXT
);
CREATE TABLE IF NOT EXISTS eod_progress (
trade_date TEXT PRIMARY KEY,
state TEXT NOT NULL,
attempts INTEGER NOT NULL DEFAULT 0,
last_attempt_at TEXT,
next_retry_at TEXT,
finished_at TEXT,
detail TEXT,
updated_at TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS audit_log (
id INTEGER PRIMARY KEY AUTOINCREMENT,
actor TEXT NOT NULL,
@@ -262,6 +288,7 @@ DATASET_TABLES = {
"moneyflow": ("eod_moneyflow", "staging_moneyflow"),
"auction": ("eod_auction", "staging_auction"),
"index_daily": ("eod_index_bars", "staging_index_bars"),
"stocks": ("eod_stocks", "staging_stocks"),
}
File diff suppressed because it is too large Load Diff
+207 -6
View File
@@ -2,7 +2,7 @@ from __future__ import annotations
import threading
from collections.abc import Callable
from datetime import datetime, time
from datetime import datetime, time, timedelta
from typing import Any
from datahub.db import HubDB
@@ -14,6 +14,8 @@ LOGGER = get_logger()
JobFn = Callable[[str], Any]
EOD_JOB_IDS = {"eod_a", "eod_b", "eod_retry"}
def is_open_day(db: HubDB, day: str) -> bool:
row = db.fetchone(
@@ -25,8 +27,19 @@ def is_open_day(db: HubDB, day: str) -> bool:
return int(row["is_open"]) == 1
def _hhmm(value: str) -> time:
return datetime.strptime(value, "%H:%M").time()
class Scheduler:
"""Calendar-driven in-process scheduler. Non-trading days skip EOD fetches."""
"""Calendar-driven in-process scheduler. Non-trading days skip EOD fetches.
EOD datasets that failed to publish (e.g. upstream not ready at 15:05)
are retried automatically every ``eod_retry_interval_minutes`` between
``eod_retry_start`` and ``eod_retry_cutoff``. Progress is persisted in
``eod_progress`` so a container restart catches up instead of waiting
for the next day, and completed days are never re-fetched.
"""
def __init__(self, db: HubDB, pipeline: Pipeline, jobs: dict[str, JobFn] | None = None) -> None:
self.db = db
@@ -35,12 +48,16 @@ class Scheduler:
"precheck": self._precheck,
"eod_a": self._eod_a,
"eod_b": self._eod_b,
"eod_retry": self._eod_retry,
"stocks_refresh": self._stocks_refresh,
"cleanup": self._cleanup,
"backup": self._backup,
"history_backfill": self._history_backfill,
}
self._stop = threading.Event()
self._thread: threading.Thread | None = None
self._fired: set[tuple[str, str, str]] = set()
self._eod_lock = threading.Lock()
def start(self, interval_seconds: float = 30.0) -> None:
if self._thread and self._thread.is_alive():
@@ -62,9 +79,9 @@ class Scheduler:
self._thread.join(timeout)
def tick(self, clock: datetime | None = None) -> list[str]:
now = clock or now_shanghai()
now = now_shanghai(clock)
day = yyyymmdd(now)
current = now.timetz() if False else now.time()
current = now.time()
ran: list[str] = []
plan = [
("precheck", time(8, 45)),
@@ -73,6 +90,8 @@ class Scheduler:
("cleanup", time(0, 30)),
("backup", time(0, 40)),
]
for refresh_at in self.pipeline.settings.stocks_refresh_times:
plan.append(("stocks_refresh", _hhmm(refresh_at)))
open_day = is_open_day(self.db, day)
for job_id, at in plan:
if current < at:
@@ -80,18 +99,187 @@ class Scheduler:
key = (job_id, day, at.strftime("%H%M"))
if key in self._fired:
continue
if job_id in {"eod_a", "eod_b"} and not open_day:
if job_id in {"eod_a", "eod_b", "stocks_refresh"} and not open_day:
self._fired.add(key)
continue
self._fired.add(key)
self.run_job(job_id, day)
if job_id in {"eod_a", "eod_b"}:
# Record the attempt before running: even a crash must not
# hide that today's first EOD try already happened.
self._record_eod_attempt(day, now)
try:
self.run_job(job_id, day)
except Exception:
if job_id not in {"eod_a", "eod_b", "stocks_refresh"}:
raise
# Keep the tick alive; evening retries take over.
LOGGER.exception("scheduled job %s failed for %s", job_id, day)
ran.append(job_id)
if job_id in {"eod_a", "eod_b"}:
self._settle_eod(day)
ran.extend(self._eod_retry_tick(now, day, open_day))
return ran
# ------------------------------------------------------------------
# EOD retry window
# ------------------------------------------------------------------
def _eod_retry_tick(self, now: datetime, day: str, open_day: bool) -> list[str]:
if not open_day:
return []
settings = self.pipeline.settings
current = now.time()
start = _hhmm(settings.eod_retry_start)
cutoff = _hhmm(settings.eod_retry_cutoff)
interval = timedelta(minutes=settings.eod_retry_interval_minutes)
missing = self.pipeline.missing_official_datasets(day)
row = self.eod_progress(day)
if not missing:
if row is None or row["state"] != "done":
self._save_eod_progress(day, state="done", finished_at=isoformat(now))
return []
if current < start:
return []
if row and row["state"] == "cutoff_failed":
return []
if current >= cutoff:
detail = "截止时间已到,缺失数据集: " + ",".join(missing)
self._save_eod_progress(day, state="cutoff_failed", finished_at=isoformat(now), detail=detail)
with self.db.write() as connection:
connection.execute(
"INSERT INTO job_runs(job_id, state, started_at, finished_at, error, attempt, detail)"
" VALUES ('eod_retry','failed',?,?,?,?,?)",
(isoformat(now), isoformat(now), detail, int((row or {}).get("attempts") or 0), "eod cutoff reached"),
)
LOGGER.warning(
"eod retry window closed without data",
extra={"hub": {"trade_date": day, "missing": missing, "reason": "eod_cutoff"}},
)
return []
last = None
if row and row["last_attempt_at"]:
try:
last = datetime.fromisoformat(str(row["last_attempt_at"]))
except ValueError:
last = None
if last is not None and now_shanghai(last).replace(tzinfo=None) + interval > now.replace(tzinfo=None):
return []
if "eod_retry" not in self.jobs:
return []
self._record_eod_attempt(day, now)
ran = []
try:
self.run_job("eod_retry", day)
except Exception:
# job_runs already carries the failure; the window keeps retrying.
LOGGER.warning("eod retry failed for %s", day, exc_info=True)
ran.append("eod_retry")
self._settle_eod(day)
return ran
def _settle_eod(self, day: str) -> None:
"""Flip the day to done as soon as every official dataset is published."""
if not self.pipeline.missing_official_datasets(day):
row = self.eod_progress(day)
if row is None or row["state"] != "done":
self._save_eod_progress(day, state="done", finished_at=isoformat())
def eod_progress(self, day: str) -> dict[str, Any] | None:
return self.db.fetchone("SELECT * FROM eod_progress WHERE trade_date = ?", (day,))
def eod_status(self, trade_date: str | None = None, clock: datetime | None = None) -> dict[str, Any]:
"""Human/admin facing view: 等待上游 / 下次重试 / 已成功 / 已截止失败."""
day = yyyymmdd(trade_date or now_shanghai(clock))
now = now_shanghai(clock)
row = self.eod_progress(day)
open_day = is_open_day(self.db, day)
missing = self.pipeline.missing_official_datasets(day)
if row and row["state"] == "done":
state = "done"
elif not open_day:
state = "closed_day"
elif not missing:
state = "done"
elif row and row["state"] == "cutoff_failed":
state = "cutoff_failed"
elif now.time() < _hhmm("15:05"):
state = "pending_first_attempt"
else:
state = "waiting_upstream"
return {
"trade_date": day,
"is_open_day": open_day,
"state": state,
"missing_datasets": missing,
"attempts": int((row or {}).get("attempts") or 0),
"last_attempt_at": (row or {}).get("last_attempt_at"),
"next_retry_at": (row or {}).get("next_retry_at") if state == "waiting_upstream" else None,
"finished_at": (row or {}).get("finished_at"),
"detail": (row or {}).get("detail"),
}
def _record_eod_attempt(self, day: str, now: datetime) -> None:
row = self.eod_progress(day)
attempts = int((row or {}).get("attempts") or 0) + 1
interval = self.pipeline.settings.eod_retry_interval_minutes
self._save_eod_progress(
day,
state="waiting_upstream",
attempts=attempts,
last_attempt_at=isoformat(now),
next_retry_at=isoformat(now + timedelta(minutes=interval)),
)
def _save_eod_progress(self, day: str, **fields: Any) -> None:
columns = [
"trade_date", "state", "attempts", "last_attempt_at",
"next_retry_at", "finished_at", "detail", "updated_at",
]
with self.db.write() as connection:
existing = connection.execute(
"SELECT trade_date FROM eod_progress WHERE trade_date = ?",
(day,),
).fetchone()
if existing is None:
payload = {name: None for name in columns}
payload.update({"trade_date": day, "state": "waiting_upstream", "attempts": 0})
payload.update(fields)
payload["updated_at"] = isoformat()
placeholders = ",".join("?" for _ in columns)
connection.execute(
f"INSERT INTO eod_progress({','.join(columns)}) VALUES ({placeholders})",
tuple(payload[name] for name in columns),
)
else:
assignments = ", ".join(f"{name} = ?" for name in fields)
connection.execute(
f"UPDATE eod_progress SET {assignments}, updated_at = ? WHERE trade_date = ?",
(*fields.values(), isoformat(), day),
)
# ------------------------------------------------------------------
# Job execution
# ------------------------------------------------------------------
def run_job(self, job_id: str, trade_date: str) -> dict[str, Any]:
fn = self.jobs.get(job_id)
if fn is None:
raise KeyError(job_id)
if job_id in EOD_JOB_IDS:
if not self._eod_lock.acquire(blocking=False):
return {
"job_id": job_id,
"state": "skipped",
"detail": "another EOD job is already running",
}
try:
return self._run_job(fn, job_id, trade_date)
finally:
self._eod_lock.release()
return self._run_job(fn, job_id, trade_date)
def _run_job(self, fn: JobFn, job_id: str, trade_date: str) -> dict[str, Any]:
started = isoformat()
run_id = None
with self.db.write() as connection:
@@ -102,6 +290,10 @@ class Scheduler:
run_id = cur.lastrowid
try:
result = fn(trade_date) or {}
if isinstance(result, dict):
failures = self.pipeline.eod_failures(result) if job_id in EOD_JOB_IDS else []
if failures:
raise RuntimeError("; ".join(failures))
with self.db.write() as connection:
connection.execute(
"UPDATE job_runs SET state=?, finished_at=?, rows_out=?, detail=? WHERE id=?",
@@ -125,6 +317,15 @@ class Scheduler:
def _eod_b(self, trade_date: str) -> dict[str, Any]:
return self.pipeline.run_eod_batch_b(trade_date)
def _eod_retry(self, trade_date: str) -> dict[str, Any]:
return self.pipeline.run_eod_missing(trade_date)
def _stocks_refresh(self, trade_date: str) -> dict[str, Any]:
return self.pipeline.refresh_stocks(trade_date)
def _history_backfill(self, trade_date: str) -> dict[str, Any]:
return self.pipeline.backfill_history(trade_date)
def _cleanup(self, trade_date: str) -> dict[str, Any]:
result = self.pipeline.cleanup()
if now_shanghai().weekday() == 6:
+75 -14
View File
@@ -6,6 +6,7 @@ from urllib.parse import parse_qs
from datahub import SCHEMA_VERSION
from datahub.codes import resolve_code
from datahub.coverage import calendar_coverage, point_coverage, published_range_coverage
from datahub.db import HubDB
from datahub.normalize import qfq_bar
from datahub.numbers import finite_number
@@ -129,10 +130,34 @@ class V1API:
}
for row in rows
]
return envelope(items, self._official_meta("calendar", end if items else start, source="tushare:trade_cal"))
meta = self._official_meta("calendar", end if items else start, source="tushare:trade_cal")
return envelope(items, attach_coverage(meta, calendar_coverage(self.db, start, end)))
def stocks(self, updated_since: str, q: dict[str, str]) -> dict[str, Any]:
limit, offset = self._page(q)
today = yyyymmdd(now_shanghai())
batch_id, snapshot = self.pipeline.published_stock_snapshot(today)
if batch_id:
# Formal view: the latest published stock snapshot, with batch
# metadata. Filters are applied in-memory on the snapshot.
pub = self.pipeline.latest_stocks_publication(today) or {}
rows = snapshot
if updated_since:
rows = []
rows = rows[offset: offset + limit]
return envelope(
rows,
{
"tier": "official",
"trade_date": pub.get("trade_date"),
"published_at": pub.get("published_at"),
"source": "tushare:stock_basic",
"batch_id": batch_id,
"stale": False,
"staleness_seconds": 0,
"state": pub.get("state"),
},
)
if updated_since:
rows = self.db.fetchall(
"SELECT * FROM stock_master WHERE updated_at >= ? ORDER BY ts_code LIMIT ? OFFSET ?",
@@ -174,7 +199,7 @@ class V1API:
def dataset_status(self, date: str) -> dict[str, Any]:
trade_date = yyyymmdd(date or now_shanghai())
datasets = ("daily", "valuation", "moneyflow", "auction", "index_daily")
datasets = ("daily", "valuation", "moneyflow", "auction", "index_daily", "stocks")
items = []
for dataset in datasets:
pub = self.db.fetchone(
@@ -249,7 +274,7 @@ class V1API:
raise ApiError(
"DATASET_NOT_PUBLISHED",
f"{dataset} {start} 尚未发布",
extra={"expected_at": "15:05+08:00"},
extra=self._unpublished_extra(dataset, start),
)
limit, offset = self._page(q)
sql = f"SELECT * FROM {table} WHERE trade_date = ? AND batch_id = ?"
@@ -272,14 +297,18 @@ class V1API:
"staleness_seconds": 0,
"state": pub["state"],
}
return envelope(rows, meta)
return envelope(rows, attach_coverage(meta, point_coverage(start, dataset)))
# multi-day: walk published dates
pubs = self.db.fetchall(
"SELECT * FROM publications WHERE dataset = ? AND trade_date >= ? AND trade_date <= ? ORDER BY trade_date",
(dataset, start, end),
)
if not pubs:
raise ApiError("DATASET_NOT_PUBLISHED", f"{dataset} {start}-{end} 尚未发布")
raise ApiError(
"DATASET_NOT_PUBLISHED",
f"{dataset} {start}-{end} 尚未发布",
extra=self._unpublished_extra(dataset, end),
)
rows: list[dict[str, Any]] = []
limit, offset = self._page(q)
for pub in pubs:
@@ -294,17 +323,28 @@ class V1API:
if adjust == "qfq" and dataset == "daily":
sliced = self._apply_qfq(sliced)
last = pubs[-1]
coverage = published_range_coverage(
self.db,
dataset,
start,
end,
ts_code=ts_code,
table=table,
)
return envelope(
sliced,
{
"tier": "official",
"trade_date": last["trade_date"],
"published_at": last["published_at"],
"source": source,
"batch_id": last["active_batch"],
"stale": False,
"staleness_seconds": 0,
},
attach_coverage(
{
"tier": "official",
"trade_date": last["trade_date"],
"published_at": last["published_at"],
"source": source,
"batch_id": last["active_batch"],
"stale": False,
"staleness_seconds": 0,
},
coverage,
),
)
def _apply_qfq(self, rows: list[dict[str, Any]]) -> list[dict[str, Any]]:
@@ -337,6 +377,20 @@ class V1API:
offset = max(0, offset)
return limit, offset
def _unpublished_extra(self, dataset: str, trade_date: str) -> dict[str, Any]:
"""Identifiable coverage info: is this a history gap or today-not-yet?"""
extra: dict[str, Any] = {"expected_at": "15:05+08:00"}
row = self.db.fetchone(
"SELECT MIN(trade_date) AS a, MAX(trade_date) AS b FROM publications WHERE dataset = ?",
(dataset,),
)
if row and row.get("a"):
extra["available_from"] = row["a"]
extra["available_to"] = row["b"]
if str(trade_date) < str(row["a"]):
extra["reason"] = "history_not_backfilled"
return extra
def _official_meta(self, dataset: str, trade_date: str, source: str) -> dict[str, Any]:
pub = self.db.fetchone(
"SELECT * FROM publications WHERE dataset = ? AND trade_date = ?",
@@ -359,6 +413,13 @@ def add_default(days: int) -> str:
return (now_shanghai() + timedelta(days=days)).strftime("%Y%m%d")
def attach_coverage(meta: dict[str, Any], coverage: dict[str, Any]) -> dict[str, Any]:
merged = dict(meta)
merged["coverage"] = coverage
merged["incomplete"] = not bool(coverage.get("complete"))
return merged
def parse_query(raw: str) -> dict[str, list[str]]:
return parse_qs(raw, keep_blank_values=True)
+31
View File
@@ -48,6 +48,37 @@ class Settings:
def list_limit_max(self) -> int:
return int(self.quality.get("list_limit_max") or 5000)
@property
def calendar_start(self) -> str:
return str(self.quality.get("calendar_start") or "20160101")
@property
def index_history_trading_days(self) -> int:
return int(self.quality.get("index_history_trading_days") or 260)
@property
def moneyflow_history_trading_days(self) -> int:
return int(self.quality.get("moneyflow_history_trading_days") or 60)
@property
def stocks_refresh_times(self) -> tuple[str, ...]:
raw = self.quality.get("stocks_refresh_times") or ["20:00", "23:10"]
if isinstance(raw, str):
raw = [raw]
return tuple(str(item) for item in raw)
@property
def eod_retry_start(self) -> str:
return str(self.quality.get("eod_retry_start") or "15:15")
@property
def eod_retry_interval_minutes(self) -> int:
return int(self.quality.get("eod_retry_interval_minutes") or 30)
@property
def eod_retry_cutoff(self) -> str:
return str(self.quality.get("eod_retry_cutoff") or "23:30")
def load_settings(
env: dict[str, str] | None = None,
+10
View File
@@ -58,6 +58,16 @@ def add_days(trade_date: str, days: int) -> str:
return (parse_trade_date(trade_date) + timedelta(days=days)).strftime("%Y%m%d")
def iter_yyyymmdd(start: str, end: str):
cursor = parse_trade_date(start)
last = parse_trade_date(end)
if cursor > last:
return
while cursor <= last:
yield cursor.strftime("%Y%m%d")
cursor += timedelta(days=1)
def utc_timestamp(value: Any) -> str:
if isinstance(value, datetime):
return isoformat(value)
+12 -2
View File
@@ -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},
@@ -51,7 +51,17 @@ RAW = {
def fake_transport(api_name: str, params: dict, fields: str):
if api_name == "index_daily":
code = params.get("ts_code")
return [row for row in RAW["index_daily"] if row["ts_code"] == code]
rows = [row for row in RAW["index_daily"] 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
if api_name == "trade_cal":
start = str(params.get("start_date") or "")
end = str(params.get("end_date") or "99999999")
+3
View File
@@ -107,6 +107,9 @@ class ApiContractTests(unittest.TestCase):
self.assertIn("data", body)
self.assertIn("meta", body)
self.assertIn("tier", body["meta"])
if "calendar" in path or "bars" in path or "indexes" in path or "valuation" in path or "moneyflow" in path or "auction" in path:
self.assertIn("coverage", body["meta"])
self.assertIn("incomplete", body["meta"])
def test_qfq_matches_formula(self) -> None:
_, none = self._get(f"/v1/bars/daily?date={TRADE_DATE}&code=600000.SH&adjust=none", token=self.token)
@@ -0,0 +1,408 @@
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.settings import Settings
from datahub.serving import V1API
from tests.fixtures import TRADE_DATE, fake_transport
GROUP_A = ("daily", "valuation", "moneyflow", "auction")
class GroupTransport:
"""fake_transport with per-API degradation switches for release-group tests."""
def __init__(self) -> None:
self.empty: set[str] = set()
self.keep_rows: dict[str, int] = {}
self.stocks: list[dict] | None = None
self.calls: list[str] = []
def __call__(self, api_name: str, params: dict, fields: str):
self.calls.append(api_name)
if api_name in self.empty:
return []
if api_name == "stock_basic" and self.stocks is not None:
return [dict(row) for row in self.stocks]
rows = fake_transport(api_name, params, fields)
keep = self.keep_rows.get(api_name)
if keep is not None:
return rows[:keep]
return rows
def make_pipe(transport: GroupTransport, quality_extra: dict | None = 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": 2,
"publication_generations": 3,
}
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)
pipe._tmp = tmp
return pipe, db
def publications_map(db: HubDB, day: str) -> dict[str, str]:
rows = db.fetchall("SELECT dataset, active_batch FROM publications WHERE trade_date = ?", (day,))
return {str(row["dataset"]): str(row["active_batch"]) for row in rows}
class ReleaseGroupSwitchTests(unittest.TestCase):
def setUp(self) -> None:
self.transport = GroupTransport()
self.pipe, self.db = make_pipe(self.transport)
self.pipe.ingest_reference(TRADE_DATE)
def test_whole_group_switches_in_one_publish_instant(self) -> None:
results = self.pipe.run_eod_batch_a(TRADE_DATE)
self.assertEqual(set(results), {*GROUP_A, "stocks"})
self.assertEqual({item["state"] for item in results.values()}, {"published"})
pubs = self.db.fetchall("SELECT * FROM publications WHERE trade_date = ?", (TRADE_DATE,))
self.assertEqual(len(pubs), 5)
self.assertEqual(len({row["published_at"] for row in pubs}), 1)
# official rows copied and serving resolves the new batches
api = V1API(self.db, self.pipe, self.pipe.settings)
payload = api.handle("/v1/bars/daily", {"date": [TRADE_DATE], "code": ["600000.SH"]})
self.assertEqual(payload["meta"]["batch_id"], results["daily"]["batch_id"])
stocks = api.handle("/v1/stocks", {})
self.assertEqual(stocks["meta"]["batch_id"], results["stocks"]["batch_id"])
def test_any_member_failure_blocks_entire_group(self) -> None:
self.transport.empty = {"daily_basic"} # valuation upstream returns nothing
results = self.pipe.run_eod_batch_a(TRADE_DATE)
self.assertEqual(results["valuation"]["state"], "failed")
self.assertEqual(results["moneyflow"]["state"], "aborted")
self.assertEqual(results["auction"]["state"], "aborted")
self.assertEqual(results["daily"]["state"], "failed") # staged fine, then abandoned
# nothing became visible, and the reason is recorded
self.assertEqual(publications_map(self.db, TRADE_DATE), {})
abandoned = self.db.fetchall(
"SELECT * FROM batches WHERE trade_date = ? AND state = 'failed'",
(TRADE_DATE,),
)
self.assertTrue(any("release group not switched" in str(row["error"] or "") for row in abandoned))
audit = self.db.fetchone(
"SELECT * FROM audit_log WHERE action = 'release-group' ORDER BY id DESC"
)
self.assertIn("valuation", str(audit["detail"]))
# still missing → evening retries keep trying
self.assertIn("daily", self.pipe.missing_official_datasets(TRADE_DATE))
def test_failure_keeps_previous_complete_version_serving(self) -> None:
first = self.pipe.run_dataset("daily", TRADE_DATE)
self.transport.empty = {"daily_basic"}
results = self.pipe.run_eod_missing(TRADE_DATE)
# incomplete A-group restages daily with the others; valuation fails → no A switch
self.assertEqual(results["daily"]["state"], "failed")
self.assertEqual(results["valuation"]["state"], "failed")
# the already-published daily batch is untouched and keeps serving
self.assertEqual(self.pipe.active_batch("daily", TRADE_DATE), first["batch_id"])
pubs = publications_map(self.db, TRADE_DATE)
self.assertEqual(pubs["daily"], first["batch_id"])
self.assertNotIn("valuation", pubs)
self.assertNotIn("moneyflow", pubs)
self.assertNotIn("auction", pubs)
# B-group is an independent boundary and may still publish
self.assertEqual(results["index_daily"]["state"], "published")
payload = V1API(self.db, self.pipe, self.pipe.settings).handle(
"/v1/bars/daily", {"date": [TRADE_DATE], "code": ["600000.SH"]}
)
self.assertEqual(payload["meta"]["batch_id"], first["batch_id"])
def test_partial_group_retry_does_not_mix_batches(self) -> None:
"""Already-published A members must be restaged with missing ones."""
first_daily = self.pipe.run_dataset("daily", TRADE_DATE)
first_moneyflow = self.pipe.run_dataset("moneyflow", TRADE_DATE)
results = self.pipe.run_eod_missing(TRADE_DATE)
# A-group switched as one boundary; B-group (index) also published
for name in (*GROUP_A, "stocks"):
self.assertEqual(results[name]["state"], "published", name)
self.assertEqual(results["index_daily"]["state"], "published")
pubs = self.db.fetchall(
"SELECT dataset, active_batch, published_at FROM publications WHERE trade_date = ?",
(TRADE_DATE,),
)
by_ds = {str(row["dataset"]): row for row in pubs}
# old partial batches replaced — no cross-batch mix of the first wave
self.assertNotEqual(by_ds["daily"]["active_batch"], first_daily["batch_id"])
self.assertNotEqual(by_ds["moneyflow"]["active_batch"], first_moneyflow["batch_id"])
a_times = {by_ds[name]["published_at"] for name in (*GROUP_A, "stocks")}
self.assertEqual(len(a_times), 1)
# serving resolves the new complete A-group batches
api = V1API(self.db, self.pipe, self.pipe.settings)
daily = api.handle("/v1/bars/daily", {"date": [TRADE_DATE], "code": ["600000.SH"]})
self.assertEqual(daily["meta"]["batch_id"], results["daily"]["batch_id"])
self.assertEqual(daily["meta"]["batch_id"], by_ds["daily"]["active_batch"])
def test_reads_during_switch_see_old_state_until_commit(self) -> None:
snapshots: list[dict] = []
def watcher() -> None:
with self.db.connect() as connection:
rows = connection.execute(
"SELECT dataset, active_batch FROM publications WHERE trade_date = ?",
(TRADE_DATE,),
).fetchall()
snapshots.append({str(row["dataset"]): row["active_batch"] for row in rows})
self.pipe.before_commit = watcher
self.pipe.run_eod_batch_a(TRADE_DATE)
# inside the switch transaction the group was still invisible
self.assertEqual(snapshots[0], {})
after = publications_map(self.db, TRADE_DATE)
self.assertEqual(set(after), {*GROUP_A, "stocks"})
def test_switch_crash_rolls_back_whole_group(self) -> None:
def explode() -> None:
raise RuntimeError("killed mid-switch")
self.pipe.before_commit = explode
with self.assertRaises(RuntimeError):
self.pipe.run_eod_batch_a(TRADE_DATE)
self.assertEqual(publications_map(self.db, TRADE_DATE), {})
for table in ("eod_bars", "eod_valuation", "eod_moneyflow", "eod_auction", "eod_stocks"):
rows = self.db.fetchall(f"SELECT * FROM {table} WHERE trade_date = ?", (TRADE_DATE,))
self.assertEqual(rows, [], table)
audit = self.db.fetchone(
"SELECT * FROM audit_log WHERE action = 'release-group' ORDER BY id DESC"
)
self.assertIsNotNone(audit)
detail = str(audit["detail"])
self.assertIn("killed mid-switch", detail)
self.assertIn("failed", detail)
def test_duplicate_runs_are_idempotent(self) -> None:
self.pipe.run_eod_batch_a(TRADE_DATE)
self.pipe.run_eod_batch_b(TRADE_DATE)
batches_before = {
str(row["batch_id"])
for row in self.db.fetchall("SELECT batch_id FROM batches WHERE trade_date = ?", (TRADE_DATE,))
}
calls_before = len(self.transport.calls)
again = self.pipe.run_eod_missing(TRADE_DATE)
self.assertEqual({item["state"] for item in again.values()}, {"skipped"})
self.assertEqual({item["reason"] for item in again.values()}, {"already_published"})
batches_after = {
str(row["batch_id"])
for row in self.db.fetchall("SELECT batch_id FROM batches WHERE trade_date = ?", (TRADE_DATE,))
}
self.assertEqual(batches_after, batches_before)
self.assertEqual(len(self.transport.calls), calls_before)
self.assertEqual(self.pipe.missing_official_datasets(TRADE_DATE), [])
def test_cross_gate_failure_blocks_switch(self) -> None:
transport = GroupTransport()
pipe, db = make_pipe(
transport,
quality_extra={"cross_gates": [
{"left": "daily", "right": "moneyflow", "min_key_overlap": 1.0},
]},
)
pipe.ingest_reference(TRADE_DATE)
transport.keep_rows["moneyflow"] = 1 # moneyflow covers only half the market
results = pipe.run_eod_batch_a(TRADE_DATE)
self.assertEqual(results["moneyflow"]["state"], "failed")
self.assertIn("cross gate", str(results["moneyflow"]["error"]))
self.assertEqual(publications_map(db, TRADE_DATE), {})
def test_stocks_master_and_snapshot_switch_together_or_not_at_all(self) -> None:
original = [
{"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"},
]
renamed = [dict(original[0]), {**original[1], "name": "金钛股份"}]
self.transport.stocks = renamed
self.pipe.run_eod_batch_a(TRADE_DATE)
master = self.db.fetchone("SELECT name FROM stock_master WHERE ts_code = '920071.BJ'")
self.assertEqual(master["name"], "金钛股份")
stocks_pub = self.db.fetchone(
"SELECT active_batch FROM publications WHERE dataset = 'stocks' AND trade_date = ?",
(TRADE_DATE,),
)
self.assertIsNotNone(stocks_pub)
# failure path: rename staged but the group is blocked → master stays untouched
transport = GroupTransport()
transport.stocks = original
pipe, db = make_pipe(
transport,
quality_extra={"cross_gates": [
{"left": "daily", "right": "moneyflow", "min_key_overlap": 1.0},
]},
)
pipe.ingest_reference(TRADE_DATE) # master seeded with "N金钛"
transport.stocks = renamed
transport.keep_rows["moneyflow"] = 1
results = pipe.run_eod_batch_a(TRADE_DATE)
self.assertEqual(results["stocks"]["state"], "failed")
master = db.fetchone("SELECT name FROM stock_master WHERE ts_code = '920071.BJ'")
self.assertEqual(master["name"], "N金钛") # rename not applied
stocks_pub = db.fetchone(
"SELECT active_batch FROM publications WHERE dataset = 'stocks' AND trade_date = ?",
(TRADE_DATE,),
)
self.assertIsNone(stocks_pub)
class StocksRefreshAtomicTests(unittest.TestCase):
def setUp(self) -> None:
self.transport = GroupTransport()
self.pipe, self.db = make_pipe(self.transport)
self.pipe.ingest_reference(TRADE_DATE)
self.transport.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"},
]
first = self.pipe.refresh_stocks(TRADE_DATE)
self.assertEqual(first["state"], "published")
self.first_batch = first["batch_id"]
def test_refresh_keeps_master_when_snapshot_publish_fails(self) -> None:
self.transport.stocks = [
{"ts_code": "600000.SH", "symbol": "600000", "name": "浦发银行", "area": "上海",
"industry": "银行", "market": "主板", "list_status": "L", "list_date": "19991110"},
{"ts_code": "920071.BJ", "symbol": "920071", "name": "金钛股份", "area": "辽宁",
"industry": "小金属", "market": "北交所", "list_status": "L", "list_date": "20240901"},
]
def explode() -> None:
raise RuntimeError("snapshot switch killed")
self.pipe.before_commit = explode
with self.assertRaises(RuntimeError):
self.pipe.refresh_stocks(TRADE_DATE)
master = self.db.fetchone("SELECT name FROM stock_master WHERE ts_code = '920071.BJ'")
self.assertEqual(master["name"], "N金钛") # rename not applied
self.assertEqual(self.pipe.active_batch("stocks", TRADE_DATE), self.first_batch)
audit = self.db.fetchone(
"SELECT * FROM audit_log WHERE action = 'stocks-refresh' ORDER BY id DESC"
)
self.assertIn("failed", str(audit["detail"]))
self.assertIn("snapshot switch killed", str(audit["detail"]))
def test_refresh_keeps_master_when_quality_gate_rejects(self) -> None:
self.transport.stocks = [] # empty → hard fail before publish
with self.assertRaises(Exception):
self.pipe.refresh_stocks(TRADE_DATE)
master = self.db.fetchone("SELECT name FROM stock_master WHERE ts_code = '920071.BJ'")
self.assertEqual(master["name"], "N金钛")
self.assertEqual(self.pipe.active_batch("stocks", TRADE_DATE), self.first_batch)
audit = self.db.fetchone(
"SELECT * FROM audit_log WHERE action = 'stocks-refresh' ORDER BY id DESC"
)
self.assertIn("failed", str(audit["detail"]))
class ForceBoundaryEntryTests(unittest.TestCase):
"""CLI force / admin backfill must rebuild the full A/B boundary."""
def setUp(self) -> None:
self.transport = GroupTransport()
self.pipe, self.db = make_pipe(self.transport)
self.pipe.ingest_reference(TRADE_DATE)
self.first = self.pipe.run_eod_batch_a(TRADE_DATE)
self.pipe.run_eod_batch_b(TRADE_DATE)
def test_force_republish_valuation_rebuilds_whole_a_group(self) -> None:
before = publications_map(self.db, TRADE_DATE)
results = self.pipe.force_republish_boundary("valuation", TRADE_DATE)
self.assertEqual({item["state"] for item in results.values()}, {"published"})
after = publications_map(self.db, TRADE_DATE)
for name in (*GROUP_A, "stocks"):
self.assertNotEqual(after[name], before[name], name)
self.assertEqual(after[name], results[name]["batch_id"], name)
# B-group left alone
self.assertEqual(after["index_daily"], before["index_daily"])
pubs = self.db.fetchall(
"SELECT dataset, published_at FROM publications WHERE trade_date = ?",
(TRADE_DATE,),
)
a_times = {row["published_at"] for row in pubs if row["dataset"] in {*GROUP_A, "stocks"}}
self.assertEqual(len(a_times), 1)
def test_force_republish_index_rebuilds_only_b_group(self) -> None:
before = publications_map(self.db, TRADE_DATE)
results = self.pipe.force_republish_boundary("index_daily", TRADE_DATE)
self.assertEqual(results["index_daily"]["state"], "published")
after = publications_map(self.db, TRADE_DATE)
self.assertNotEqual(after["index_daily"], before["index_daily"])
for name in GROUP_A:
self.assertEqual(after[name], before[name], name)
def test_admin_backfill_official_dataset_uses_boundary(self) -> None:
from datahub.admin_api import AdminAPI
from datahub.auth import AuthService
from datahub.crypto import SecretVault
from datahub.scheduler import Scheduler
from datahub.serving import ApiError
vault = SecretVault(self.pipe.settings.encryption_key)
auth = AuthService(self.db, vault, self.pipe.settings.api_token, "StartPass1")
admin = AdminAPI(self.db, self.pipe, Scheduler(self.db, self.pipe), auth)
before = publications_map(self.db, TRADE_DATE)
result = admin.backfill("moneyflow", TRADE_DATE, "StartPass1", f"moneyflow:{TRADE_DATE}", "tester")
self.assertEqual(result["moneyflow"]["state"], "published")
after = publications_map(self.db, TRADE_DATE)
for name in (*GROUP_A, "stocks"):
self.assertNotEqual(after[name], before[name], name)
# bad password / wrong confirm still rejected
with self.assertRaises(ApiError):
admin.backfill("daily", TRADE_DATE, "wrong", f"daily:{TRADE_DATE}", "tester")
def test_admin_backfill_switch_crash_is_failed_precondition(self) -> None:
from datahub.admin_api import AdminAPI
from datahub.auth import AuthService
from datahub.crypto import SecretVault
from datahub.scheduler import Scheduler
from datahub.serving import ApiError
vault = SecretVault(self.pipe.settings.encryption_key)
auth = AuthService(self.db, vault, self.pipe.settings.api_token, "StartPass1")
admin = AdminAPI(self.db, self.pipe, Scheduler(self.db, self.pipe), auth)
before = publications_map(self.db, TRADE_DATE)
def explode() -> None:
raise RuntimeError("killed mid-switch")
self.pipe.before_commit = explode
with self.assertRaises(ApiError) as ctx:
admin.backfill("valuation", TRADE_DATE, "StartPass1", f"valuation:{TRADE_DATE}", "tester")
self.assertEqual(ctx.exception.code, "FAILED_PRECONDITION")
self.assertIn("killed mid-switch", ctx.exception.message)
# previous complete A/B versions keep serving
self.assertEqual(publications_map(self.db, TRADE_DATE), before)
audit = self.db.fetchone(
"SELECT * FROM audit_log WHERE action = 'release-group' ORDER BY id DESC"
)
self.assertIsNotNone(audit)
self.assertIn("failed", str(audit["detail"]))
self.assertIn("killed mid-switch", str(audit["detail"]))
if __name__ == "__main__":
unittest.main()
+241
View File
@@ -0,0 +1,241 @@
from __future__ import annotations
import unittest
from datetime import datetime
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.settings import Settings
from datahub.timeutil import SHANGHAI
from tests.fixtures import fake_transport
OFFICIAL = {"daily", "valuation", "moneyflow", "auction", "index_daily"}
class DelayedTransport:
"""Upstream that only returns rows for dates it has "published" yet."""
DATE_APIS = {"daily", "daily_basic", "adj_factor", "moneyflow", "stk_auction", "index_daily"}
def __init__(self, ready_dates: set[str]) -> None:
self.ready = set(ready_dates)
self.calls: list[str] = []
def __call__(self, api_name: str, params: dict, fields: str):
self.calls.append(api_name)
if api_name in self.DATE_APIS:
trade_date = str(params.get("trade_date") or "")
if trade_date and trade_date not in self.ready:
return []
return fake_transport(api_name, params, fields)
def clock_at(day: str, hh: int, mm: int) -> datetime:
return datetime(int(day[:4]), int(day[4:6]), int(day[6:8]), hh, mm, tzinfo=SHANGHAI)
class EodRetryTests(unittest.TestCase):
def _make(self, ready_dates: set[str]):
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
db = HubDB(Path(tmp.name) / "hub.db")
transport = DelayedTransport(ready_dates)
adapter = TushareAdapter("x", transport=transport)
settings = Settings(
encryption_key=SecretVault.generate_key(),
db_path=db.path,
backup_dir=Path(tmp.name) / "backups",
)
pipe = Pipeline(db, adapter, settings)
pipe.ingest_reference("20240902")
sched = Scheduler(db, pipe)
return db, transport, pipe, sched
def _job_runs(self, db: HubDB, job_id: str) -> list[dict]:
return db.fetchall("SELECT * FROM job_runs WHERE job_id = ? ORDER BY id", (job_id,))
def _batches(self, db: HubDB, day: str) -> list[dict]:
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]:
return {str(row["batch_id"]) for row in db.fetchall("SELECT batch_id FROM batches WHERE trade_date = ?", (day,))}
@staticmethod
def _eod_calls(transport: DelayedTransport) -> list[str]:
return [name for name in transport.calls if name in DelayedTransport.DATE_APIS]
def _published(self, db: HubDB, day: str) -> set[str]:
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:
day = "20240902"
db, transport, pipe, sched = self._make(set())
sched.tick(clock_at(day, 15, 5)) # eod_a: upstream empty -> failed
sched.tick(clock_at(day, 15, 10)) # eod_b: upstream empty -> failed
self.assertEqual(self._published(db, day), set()) # quality gate held
sched.tick(clock_at(day, 15, 20)) # inside window, but <30min since 15:10
self.assertEqual(self._job_runs(db, "eod_retry"), [])
status = sched.eod_status(day, clock=clock_at(day, 15, 20))
self.assertEqual(status["state"], "waiting_upstream")
self.assertTrue(status["next_retry_at"])
self.assertEqual(status["missing_datasets"], sorted(OFFICIAL))
sched.tick(clock_at(day, 15, 40)) # retry #1, still empty
runs = self._job_runs(db, "eod_retry")
self.assertEqual(len(runs), 1)
self.assertEqual(runs[0]["state"], "failed")
self.assertEqual(self._published(db, day), set())
transport.ready.add(day)
sched.tick(clock_at(day, 16, 10)) # retry #2 succeeds
self.assertEqual(self._published(db, day), OFFICIAL)
self.assertEqual(sched.eod_status(day, clock=clock_at(day, 16, 10))["state"], "done")
progress = db.fetchone("SELECT * FROM eod_progress WHERE trade_date = ?", (day,))
self.assertEqual(progress["state"], "done")
self.assertEqual(progress["attempts"], 4) # eod_a + eod_b + 2 retries
# success stops all further same-day requests
batches_before = len(self._batches(db, day))
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(self._eod_calls(transport)), eod_calls_before)
def test_never_ready_marks_cutoff_failed_and_stops(self) -> None:
day = "20240902"
db, transport, pipe, sched = self._make(set())
sched.tick(clock_at(day, 15, 5))
sched.tick(clock_at(day, 15, 10))
sched.tick(clock_at(day, 15, 40))
sched.tick(clock_at(day, 16, 10))
sched.tick(clock_at(day, 23, 29))
self.assertEqual(len(self._job_runs(db, "eod_retry")), 3)
sched.tick(clock_at(day, 23, 35)) # past cutoff 23:30
status = sched.eod_status(day, clock=clock_at(day, 23, 35))
self.assertEqual(status["state"], "cutoff_failed")
cutoff_runs = [r for r in self._job_runs(db, "eod_retry") if "截止" in str(r["error"])]
self.assertEqual(len(cutoff_runs), 1)
self.assertEqual(self._published(db, day), set())
attempts = db.fetchone("SELECT attempts FROM eod_progress WHERE trade_date = ?", (day,))["attempts"]
sched.tick(clock_at(day, 23, 59))
self.assertEqual(
db.fetchone("SELECT attempts FROM eod_progress WHERE trade_date = ?", (day,))["attempts"],
attempts,
)
self.assertEqual(len(self._job_runs(db, "eod_retry")), 4) # 3 retries + 1 cutoff record
self.assertEqual(self._published(db, day), set())
def test_restart_catches_up_without_overwriting(self) -> None:
day = "20240902"
db, transport, pipe, sched = self._make({day})
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 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
sched2 = Scheduler(db, pipe)
ran = sched2.tick(clock_at(day, 21, 0))
self.assertIn("eod_a", ran)
self.assertIn("eod_b", ran)
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 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")
def test_restart_with_partial_publish_only_fetches_missing(self) -> None:
day = "20240902"
db, transport, pipe, sched = self._make({day})
sched.tick(clock_at(day, 15, 5)) # eod_a publishes 4; container "crashes" before eod_b
self.assertEqual(self._published(db, day), {"daily", "valuation", "moneyflow", "auction"})
batches_before = self._batch_ids(db, day)
sched2 = Scheduler(db, pipe)
ran = sched2.tick(clock_at(day, 15, 20)) # restart: eod_b catch-up, eod_a all skipped
self.assertIn("eod_b", ran)
self.assertEqual(self._published(db, day), OFFICIAL)
new_ids = self._batch_ids(db, day) - batches_before
new_datasets = {str(b["dataset"]) for b in self._batches(db, day) if str(b["batch_id"]) in new_ids}
self.assertEqual(new_datasets, {"index_daily"})
self.assertEqual(sched2.eod_status(day, clock=clock_at(day, 15, 20))["state"], "done")
def test_closed_day_skips_all_eod_work(self) -> None:
day = "20240907" # closed in fixture calendar
db, transport, pipe, sched = self._make(set())
for hh, mm in ((15, 5), (15, 10), (15, 40), (16, 10), (20, 0), (23, 40)):
ran = sched.tick(clock_at(day, hh, mm))
self.assertNotIn("eod_retry", ran)
eod_runs = db.fetchall("SELECT * FROM job_runs WHERE job_id LIKE 'eod%'")
self.assertEqual(eod_runs, [])
self.assertIsNone(db.fetchone("SELECT * FROM eod_progress WHERE trade_date = ?", (day,)))
self.assertEqual(self._published(db, day), set())
self.assertEqual(sched.eod_status(day, clock=clock_at(day, 20, 0))["state"], "closed_day")
def test_duplicate_and_concurrent_execution_are_safe(self) -> None:
day = "20240902"
db, transport, pipe, sched = self._make({day})
sched.tick(clock_at(day, 15, 5))
sched.tick(clock_at(day, 15, 10))
self.assertEqual(self._published(db, day), OFFICIAL)
batches_before = len(self._batches(db, day))
calls_before = len(transport.calls)
out = sched.run_job("eod_retry", day) # manual duplicate run
self.assertEqual(out["state"], "ok")
self.assertEqual(len(self._batches(db, day)), batches_before)
self.assertEqual(len(transport.calls), calls_before)
sched._eod_lock.acquire() # simulate an in-flight EOD job
try:
busy = sched.run_job("eod_retry", day)
self.assertEqual(busy["state"], "skipped")
busy_a = sched.run_job("eod_a", day)
self.assertEqual(busy_a["state"], "skipped")
finally:
sched._eod_lock.release()
self.assertEqual(len(self._batches(db, day)), batches_before)
if __name__ == "__main__":
unittest.main()
@@ -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()
@@ -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()
+1
View File
@@ -69,6 +69,7 @@ class PipelineTests(unittest.TestCase):
pipe, db = make_pipeline()
ref = pipe.ingest_reference(TRADE_DATE)
self.assertEqual(ref["stocks"], 2)
self.assertEqual(ref["calendar_from"], "20160101")
result = pipe.run_dataset("daily", TRADE_DATE)
self.assertEqual(result["state"], "published")
self.assertEqual(result["rows"], 2)
+262
View File
@@ -0,0 +1,262 @@
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_force_boundary_republish_keeps_prev_for_rollback(self) -> None:
transport = ValuationTransport()
pipe, db = make_pipe(transport)
pipe.ingest_reference(TRADE_DATE)
first = pipe.run_eod_batch_a(TRADE_DATE)
first_val = first["valuation"]["batch_id"]
first_daily = first["daily"]["batch_id"]
transport.mode = "vr_all_null"
blocked = pipe.force_republish_boundary("valuation", TRADE_DATE)
self.assertEqual(blocked["valuation"]["state"], "failed")
self.assertEqual(pipe.active_batch("valuation", TRADE_DATE), first_val)
self.assertEqual(pipe.active_batch("daily", TRADE_DATE), first_daily)
transport.mode = "ok"
second = pipe.force_republish_boundary("valuation", TRADE_DATE)
self.assertEqual(second["valuation"]["state"], "published")
self.assertNotEqual(second["valuation"]["batch_id"], first_val)
self.assertNotEqual(second["daily"]["batch_id"], first_daily)
pubs = db.fetchall(
"SELECT dataset, active_batch, prev_batch, published_at FROM publications WHERE trade_date=?",
(TRADE_DATE,),
)
by_ds = {str(row["dataset"]): row for row in pubs}
a_times = {by_ds[name]["published_at"] for name in ("daily", "valuation", "moneyflow", "auction", "stocks")}
self.assertEqual(len(a_times), 1)
self.assertEqual(by_ds["valuation"]["active_batch"], second["valuation"]["batch_id"])
self.assertEqual(by_ds["valuation"]["prev_batch"], first_val)
rolled = pipe.rollback("valuation", TRADE_DATE, actor="cli")
self.assertEqual(rolled["active_batch"], first_val)
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()