From 3498dd7a4b53d1a73082ee6a39817128ce5f882e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=80=BB=E5=B7=A5?= Date: Wed, 2 Sep 2026 12:05:26 +0800 Subject: [PATCH] =?UTF-8?q?feat(HEL-382):=20=E6=90=AD=E5=BB=BA=20datahub?= =?UTF-8?q?=20=E5=BA=95=E5=BA=A7=E5=92=8C=E7=9B=98=E5=90=8E=E6=AD=A3?= =?UTF-8?q?=E5=BC=8F=E6=95=B0=E6=8D=AE=E9=93=BE=E8=B7=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新增独立 xiaobai-datahub 服务(SQLite WAL、Tushare 盘后发布、/v1 契约和管理后台),不改现站页面与数据链路。 Co-authored-by: Cursor Co-authored-by: multica-agent --- .gitignore | 3 + compose.datahub.yaml | 45 ++ xiaobai-datahub/.dockerignore | 10 + xiaobai-datahub/.env.example | 13 + xiaobai-datahub/Dockerfile | 36 ++ xiaobai-datahub/README.md | 78 +++ xiaobai-datahub/admin/app.js | 268 ++++++++++ xiaobai-datahub/admin/index.html | 53 ++ xiaobai-datahub/admin/styles.css | 51 ++ xiaobai-datahub/compose.yaml | 39 ++ .../config/hub-quality.config.json | 15 + xiaobai-datahub/datahub/__init__.py | 4 + xiaobai-datahub/datahub/adapters/__init__.py | 15 + xiaobai-datahub/datahub/adapters/akshare.py | 3 + xiaobai-datahub/datahub/adapters/base.py | 47 ++ xiaobai-datahub/datahub/adapters/eastmoney.py | 3 + xiaobai-datahub/datahub/adapters/ifind.py | 3 + xiaobai-datahub/datahub/adapters/tencent.py | 3 + xiaobai-datahub/datahub/adapters/ths.py | 3 + xiaobai-datahub/datahub/adapters/tushare.py | 154 ++++++ xiaobai-datahub/datahub/adapters/xgb.py | 3 + xiaobai-datahub/datahub/admin_api.py | 162 ++++++ xiaobai-datahub/datahub/auth.py | 190 +++++++ xiaobai-datahub/datahub/codes.py | 26 + xiaobai-datahub/datahub/crypto.py | 42 ++ xiaobai-datahub/datahub/db.py | 340 +++++++++++++ .../datahub/governance/__init__.py | 13 + xiaobai-datahub/datahub/governance/circuit.py | 107 ++++ xiaobai-datahub/datahub/governance/lkg.py | 84 +++ .../datahub/governance/ratelimit.py | 36 ++ xiaobai-datahub/datahub/governance/retry.py | 35 ++ xiaobai-datahub/datahub/httpapp.py | 242 +++++++++ xiaobai-datahub/datahub/hub.py | 54 ++ xiaobai-datahub/datahub/logutil.py | 56 ++ xiaobai-datahub/datahub/normalize.py | 201 ++++++++ xiaobai-datahub/datahub/numbers.py | 23 + xiaobai-datahub/datahub/pipeline.py | 478 ++++++++++++++++++ xiaobai-datahub/datahub/scheduler.py | 145 ++++++ xiaobai-datahub/datahub/serving.py | 376 ++++++++++++++ xiaobai-datahub/datahub/settings.py | 72 +++ xiaobai-datahub/datahub/timeutil.py | 64 +++ xiaobai-datahub/requirements.txt | 1 + xiaobai-datahub/server.py | 27 + xiaobai-datahub/tests/__init__.py | 0 xiaobai-datahub/tests/fixtures.py | 59 +++ xiaobai-datahub/tests/test_admin.py | 96 ++++ xiaobai-datahub/tests/test_api.py | 147 ++++++ xiaobai-datahub/tests/test_governance.py | 56 ++ xiaobai-datahub/tests/test_layout.py | 30 ++ xiaobai-datahub/tests/test_normalize.py | 78 +++ xiaobai-datahub/tests/test_pipeline.py | 108 ++++ xiaobai-datahub/tests/test_scheduler.py | 62 +++ 52 files changed, 4259 insertions(+) create mode 100644 compose.datahub.yaml create mode 100644 xiaobai-datahub/.dockerignore create mode 100644 xiaobai-datahub/.env.example create mode 100644 xiaobai-datahub/Dockerfile create mode 100644 xiaobai-datahub/README.md create mode 100644 xiaobai-datahub/admin/app.js create mode 100644 xiaobai-datahub/admin/index.html create mode 100644 xiaobai-datahub/admin/styles.css create mode 100644 xiaobai-datahub/compose.yaml create mode 100644 xiaobai-datahub/config/hub-quality.config.json create mode 100644 xiaobai-datahub/datahub/__init__.py create mode 100644 xiaobai-datahub/datahub/adapters/__init__.py create mode 100644 xiaobai-datahub/datahub/adapters/akshare.py create mode 100644 xiaobai-datahub/datahub/adapters/base.py create mode 100644 xiaobai-datahub/datahub/adapters/eastmoney.py create mode 100644 xiaobai-datahub/datahub/adapters/ifind.py create mode 100644 xiaobai-datahub/datahub/adapters/tencent.py create mode 100644 xiaobai-datahub/datahub/adapters/ths.py create mode 100644 xiaobai-datahub/datahub/adapters/tushare.py create mode 100644 xiaobai-datahub/datahub/adapters/xgb.py create mode 100644 xiaobai-datahub/datahub/admin_api.py create mode 100644 xiaobai-datahub/datahub/auth.py create mode 100644 xiaobai-datahub/datahub/codes.py create mode 100644 xiaobai-datahub/datahub/crypto.py create mode 100644 xiaobai-datahub/datahub/db.py create mode 100644 xiaobai-datahub/datahub/governance/__init__.py create mode 100644 xiaobai-datahub/datahub/governance/circuit.py create mode 100644 xiaobai-datahub/datahub/governance/lkg.py create mode 100644 xiaobai-datahub/datahub/governance/ratelimit.py create mode 100644 xiaobai-datahub/datahub/governance/retry.py create mode 100644 xiaobai-datahub/datahub/httpapp.py create mode 100644 xiaobai-datahub/datahub/hub.py create mode 100644 xiaobai-datahub/datahub/logutil.py create mode 100644 xiaobai-datahub/datahub/normalize.py create mode 100644 xiaobai-datahub/datahub/numbers.py create mode 100644 xiaobai-datahub/datahub/pipeline.py create mode 100644 xiaobai-datahub/datahub/scheduler.py create mode 100644 xiaobai-datahub/datahub/serving.py create mode 100644 xiaobai-datahub/datahub/settings.py create mode 100644 xiaobai-datahub/datahub/timeutil.py create mode 100644 xiaobai-datahub/requirements.txt create mode 100644 xiaobai-datahub/server.py create mode 100644 xiaobai-datahub/tests/__init__.py create mode 100644 xiaobai-datahub/tests/fixtures.py create mode 100644 xiaobai-datahub/tests/test_admin.py create mode 100644 xiaobai-datahub/tests/test_api.py create mode 100644 xiaobai-datahub/tests/test_governance.py create mode 100644 xiaobai-datahub/tests/test_layout.py create mode 100644 xiaobai-datahub/tests/test_normalize.py create mode 100644 xiaobai-datahub/tests/test_pipeline.py create mode 100644 xiaobai-datahub/tests/test_scheduler.py diff --git a/.gitignore b/.gitignore index b62287d..a4c36b3 100644 --- a/.gitignore +++ b/.gitignore @@ -8,6 +8,9 @@ data/*.db data/*.db-shm data/*.db-wal data/backups/ +datahub-data/ +xiaobai-datahub/data/ +xiaobai-datahub/.venv/ data/*.bak data/*.backup *.log diff --git a/compose.datahub.yaml b/compose.datahub.yaml new file mode 100644 index 0000000..786e78c --- /dev/null +++ b/compose.datahub.yaml @@ -0,0 +1,45 @@ +# Optional overlay. Does not replace the existing xiaobai-review service. +# Start later (总工部署时) with: +# docker compose -f compose.yaml -f compose.datahub.yaml up -d +# +# Required .env keys: DATAHUB_ENCRYPTION_KEY, DATAHUB_TOKEN, DATAHUB_ADMIN_PASSWORD, TUSHARE_TOKEN + +services: + xiaobai-datahub: + build: + context: ./xiaobai-datahub + dockerfile: Dockerfile + image: xiaobai-datahub:local + container_name: xiaobai-datahub + ports: + - "0.0.0.0:8766:8766/tcp" + env_file: + - ./xiaobai-datahub/.env + environment: + DATAHUB_ENCRYPTION_KEY: "${DATAHUB_ENCRYPTION_KEY:?DATAHUB_ENCRYPTION_KEY must be set}" + DATAHUB_TOKEN: "${DATAHUB_TOKEN:?DATAHUB_TOKEN must be set}" + DATAHUB_ADMIN_PASSWORD: "${DATAHUB_ADMIN_PASSWORD:?DATAHUB_ADMIN_PASSWORD must be set}" + TUSHARE_TOKEN: "${TUSHARE_TOKEN:-}" + DATAHUB_DB_PATH: /app/data/datahub.db + DATAHUB_BACKUP_DIR: /app/data/backups + TZ: Asia/Shanghai + PYTHONUTF8: "1" + volumes: + - type: bind + source: ./datahub-data + target: /app/data + restart: unless-stopped + init: true + read_only: true + tmpfs: + - /tmp:size=64m,mode=1777 + security_opt: + - no-new-privileges:true + cap_drop: + - ALL + stop_grace_period: 30s + logging: + driver: json-file + options: + max-size: "10m" + max-file: "3" diff --git a/xiaobai-datahub/.dockerignore b/xiaobai-datahub/.dockerignore new file mode 100644 index 0000000..db68c82 --- /dev/null +++ b/xiaobai-datahub/.dockerignore @@ -0,0 +1,10 @@ +.git +.gitignore +.env +.env.* +!.env.example +__pycache__/ +*.py[cod] +*.log +data/ +tests/ diff --git a/xiaobai-datahub/.env.example b/xiaobai-datahub/.env.example new file mode 100644 index 0000000..e5cd65b --- /dev/null +++ b/xiaobai-datahub/.env.example @@ -0,0 +1,13 @@ +# Fernet key. Generate with: python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())" +DATAHUB_ENCRYPTION_KEY= + +# Consumer API token for /v1 (32+ random bytes, shown once). Never log this value. +DATAHUB_TOKEN= + +# Initial admin password for /admin. Forced change on first login. +DATAHUB_ADMIN_PASSWORD= + +# Tushare Pro token. Stored encrypted after first launch; never returned by API or admin pages. +TUSHARE_TOKEN= + +TZ=Asia/Shanghai diff --git a/xiaobai-datahub/Dockerfile b/xiaobai-datahub/Dockerfile new file mode 100644 index 0000000..a594bbb --- /dev/null +++ b/xiaobai-datahub/Dockerfile @@ -0,0 +1,36 @@ +FROM python:3.12-slim-bookworm + +ARG APP_UID=10002 +ARG APP_GID=10002 + +ENV PYTHONDONTWRITEBYTECODE=1 \ + PYTHONUNBUFFERED=1 \ + PYTHONUTF8=1 \ + PIP_DISABLE_PIP_VERSION_CHECK=1 \ + TZ=Asia/Shanghai + +WORKDIR /app + +RUN apt-get update \ + && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends \ + ca-certificates \ + tzdata \ + && groupadd --gid "${APP_GID}" datahub \ + && useradd --uid "${APP_UID}" --gid "${APP_GID}" --create-home --shell /usr/sbin/nologin datahub \ + && rm -rf /var/lib/apt/lists/* + +COPY requirements.txt ./ +RUN python -m pip install --no-cache-dir -r requirements.txt + +COPY --chown=datahub:datahub . . +RUN mkdir -p /app/data /app/data/backups && chown -R datahub:datahub /app/data + +USER datahub + +EXPOSE 8766 +STOPSIGNAL SIGINT + +HEALTHCHECK --interval=30s --timeout=5s --start-period=20s --retries=3 \ + CMD ["python", "-c", "import urllib.request; urllib.request.urlopen('http://127.0.0.1:8766/livez', timeout=4).read()"] + +CMD ["python", "-u", "server.py", "--host", "0.0.0.0", "--port", "8766"] diff --git a/xiaobai-datahub/README.md b/xiaobai-datahub/README.md new file mode 100644 index 0000000..bd59a44 --- /dev/null +++ b/xiaobai-datahub/README.md @@ -0,0 +1,78 @@ +# xiaobai-datahub + +独立行情数据中枢(HEL-382 / P0)。与 `xiaobai-review` 同仓库、不同容器、不共享数据库文件。 +本阶段不部署现网;只提供可本地运行、可自测的底座和盘后正式数据链路。 + +## 做什么 + +- SQLite WAL `datahub.db`,容器名 `xiaobai-datahub`,端口 `8766` +- Tushare 盘后正式数据:交易日历、股票主档、daily、daily_basic、adj_factor、index_daily、moneyflow、stk_auction +- 暂存 → 校验 → 整批原子发布 → 可回滚 +- `/v1` 稳定接口(`X-Datahub-Token`) +- `/admin/` 最小管理后台(总览 / 数据源 / 调度 / 发布 / 数据集 / 审计) +- 东财/腾讯/同花顺/选股宝/AKShare/iFinD 适配器位已预留,本阶段不拉实时源 + +## 单位口径(相对现站) + +现站 `xiaobai-review` 按 Tushare 原始单位入库、展示时再换算。中枢在归一化层一次换算: + +| 字段 | Tushare / 现站 | 中枢 canonical | +|---|---|---| +| `daily.amount` / `index_daily.amount` | 千元 | 元(×1000) | +| `daily.vol` / `index_daily.vol` | 手 | 股(×100) | +| `moneyflow.*_amount` | 万元 | 元(×1e4) | +| `daily_basic.total_mv` / `circ_mv` | 万元 | 元(×1e4) | +| `stk_auction.amount` | 元 | 元 | + +差异为口径升级,golden 测试按上表对照,不为 0 的字段都有说明。 + +## 本地启动(不走 Docker) + +```bash +cd xiaobai-datahub +python -m venv .venv && .venv/bin/pip install -r requirements.txt +cp .env.example .env +# 填入 DATAHUB_ENCRYPTION_KEY / DATAHUB_TOKEN / DATAHUB_ADMIN_PASSWORD / TUSHARE_TOKEN +# 生成 Fernet 密钥: +# python -c "from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())" +.venv/bin/python server.py --host 127.0.0.1 --port 8766 +``` + +- 管理后台:http://127.0.0.1:8766/admin/ +- 存活检查:http://127.0.0.1:8766/livez (无需 token) +- `/v1/*` 必须带请求头 `X-Datahub-Token` + +## Docker(独立 compose,不改现网 review 服务) + +```bash +cd xiaobai-datahub +cp .env.example .env # 填密钥 +mkdir -p data +docker compose build +docker compose up -d +``` + +仓库根目录另有 `compose.datahub.yaml`,供总工以后与现有 `compose.yaml` 叠加部署,本卡不执行现网 `up`。 + +## 自测 + +```bash +cd xiaobai-datahub +python -m unittest discover -s tests -v +``` + +不调用真实 Tushare;用内存/临时库和假适配器。 + +## 备份 + +每日 00:40 任务把 `datahub.db` 备份到 `data/backups/`(保留 14 份)。也可手动: + +```bash +python -c "from pathlib import Path; from datahub.db import HubDB; HubDB(Path('data/datahub.db')).backup_to(Path('data/backups/manual.db'))" +``` + +## 安全 + +- 密钥只以 `configured / 末4位 / 更新时间` 出现在后台,不进日志、不进 `/v1` +- 回滚、补数需重新输入密码 + 确认词 +- 容器非 root(uid 10002)、read_only、cap_drop ALL diff --git a/xiaobai-datahub/admin/app.js b/xiaobai-datahub/admin/app.js new file mode 100644 index 0000000..8f29b7f --- /dev/null +++ b/xiaobai-datahub/admin/app.js @@ -0,0 +1,268 @@ +const state = { csrf: "", page: "overview" }; + +function $(id) { return document.getElementById(id); } + +async function api(path, options = {}) { + const headers = Object.assign({ "Content-Type": "application/json" }, options.headers || {}); + if (state.csrf && (options.method || "GET") !== "GET") headers["X-CSRF-Token"] = state.csrf; + const res = await fetch(path, Object.assign({}, options, { headers, credentials: "same-origin" })); + const body = await res.json(); + if (!res.ok) { + const msg = (body.error && body.error.message) || body.error || res.statusText; + throw new Error(msg); + } + return body; +} + +function show(id) { + ["login-view", "change-view", "shell"].forEach((key) => { $(key).hidden = key !== id; }); +} + +function esc(value) { + return String(value ?? "").replace(/[&<>"]/g, (ch) => ({ "&": "&", "<": "<", ">": ">", '"': """ }[ch])); +} + +function table(headers, rows) { + const thead = headers.map((h) => `${esc(h)}`).join(""); + const body = rows.length + ? rows.map((cols) => `${cols.map((c) => `${c}`).join("")}`).join("") + : `暂无数据`; + return `${thead}${body}
`; +} + +async function boot() { + try { + const session = await api("/admin/api/session"); + state.csrf = session.csrf; + $("who").textContent = session.username; + if (session.must_change) { show("change-view"); return; } + show("shell"); + await render(); + } catch { + show("login-view"); + } +} + +$("login-form").addEventListener("submit", async (event) => { + event.preventDefault(); + const form = new FormData(event.target); + $("login-error").hidden = true; + try { + const result = await api("/admin/api/login", { + method: "POST", + body: JSON.stringify({ username: form.get("username"), password: form.get("password") }), + }); + state.csrf = result.csrf; + if (result.must_change) show("change-view"); + else { show("shell"); await render(); } + } catch (err) { + $("login-error").hidden = false; + $("login-error").textContent = err.message; + } +}); + +$("change-form").addEventListener("submit", async (event) => { + event.preventDefault(); + const form = new FormData(event.target); + try { + await api("/admin/api/change-password", { + method: "POST", + body: JSON.stringify({ current: form.get("current"), new_password: form.get("new_password") }), + }); + show("shell"); + await render(); + } catch (err) { + $("change-error").hidden = false; + $("change-error").textContent = err.message; + } +}); + +$("logout-btn").addEventListener("click", async () => { + await api("/admin/api/logout", { method: "POST", body: "{}" }); + show("login-view"); +}); + +$("theme-btn").addEventListener("click", () => { + const root = document.documentElement; + const next = root.getAttribute("data-theme") === "night" ? "" : "night"; + if (next) root.setAttribute("data-theme", next); + else root.removeAttribute("data-theme"); + $("theme-btn").textContent = next ? "日间" : "夜间"; +}); + +document.querySelectorAll("nav button").forEach((btn) => { + btn.addEventListener("click", () => { + document.querySelectorAll("nav button").forEach((item) => item.classList.remove("active")); + btn.classList.add("active"); + state.page = btn.dataset.page; + render(); + }); +}); + +async function render() { + const page = $("page"); + if (state.page === "overview") { + const data = await api("/admin/api/overview"); + $("phase").textContent = data.session_phase; + page.innerHTML = ` +
+
交易日
${esc(data.trade_date)}
+
阶段
${esc(data.session_phase)}
+
今日发布
${data.publications.length}
+
异常批次
${data.anomalies.length}
+
+

最近调用

+ ${table(["时间", "源", "端点", "结果", "耗时"], data.recent_calls.map((row) => [ + esc(row.created_at), esc(row.provider), esc(row.endpoint), + row.ok ? '成功' : `${esc(row.error)}`, + `${row.latency_ms ?? "-"} ms`, + ]))} + `; + return; + } + if (state.page === "sources") { + const data = await api("/admin/api/sources"); + page.innerHTML = `

数据源

` + table( + ["源", "角色", "状态", "凭据", "操作"], + data.items.map((item) => { + const cred = item.credential || {}; + const credText = cred.configured ? `已配置 · ${esc(cred.last4 || "****")}` : "未配置"; + return [ + esc(item.provider), + esc(item.role), + esc((item.health && (item.health.state || item.health.status)) || "-"), + credText, + ``, + ]; + }), + ); + page.querySelectorAll("[data-probe]").forEach((btn) => { + btn.addEventListener("click", async () => { + const result = await api(`/admin/api/sources/${btn.dataset.probe}/probe`, { method: "POST", body: "{}" }); + alert(JSON.stringify(result)); + render(); + }); + }); + return; + } + if (state.page === "jobs") { + const data = await api("/admin/api/jobs"); + page.innerHTML = ` +

调度任务

+ ${table(["任务", "时刻", "操作"], data.jobs.map((job) => [ + `${esc(job.id)} · ${esc(job.title)}`, esc(job.at), + ``, + ]))} +

最近运行

+ ${table(["ID", "任务", "状态", "开始", "结束", "错误"], data.runs.map((row) => [ + row.id, esc(row.job_id), esc(row.state), esc(row.started_at), esc(row.finished_at), esc(row.error), + ]))} + `; + page.querySelectorAll("[data-run]").forEach((btn) => { + btn.addEventListener("click", async () => { + const date = prompt("交易日 YYYYMMDD(可留空=今天)", "") || ""; + await api(`/admin/api/jobs/${btn.dataset.run}/run`, { method: "POST", body: JSON.stringify({ trade_date: date }) }); + render(); + }); + }); + return; + } + if (state.page === "release") { + const date = new Date().toISOString().slice(0, 10).replace(/-/g, ""); + const data = await api(`/admin/api/batches?date=${date}`); + page.innerHTML = ` +

盘后发布 ${esc(data.trade_date)}

+
+ + + +
+

当前映射

+ ${table(["数据集", "活跃批次", "上一批次", "状态", "发布时间", "操作"], data.publications.map((row) => [ + esc(row.dataset), esc(row.active_batch), esc(row.prev_batch), esc(row.state), esc(row.published_at), + row.prev_batch ? `` : "-", + ]))} +

批次

+ ${table(["batch_id", "数据集", "状态", "行数", "错误"], data.batches.map((row) => [ + esc(row.batch_id), esc(row.dataset), esc(row.state), row.rows_out ?? "", esc(row.error), + ]))} + `; + $bindRelease(page); + return; + } + if (state.page === "datasets") { + const data = await api("/admin/api/datasets?date="); + page.innerHTML = ` +

数据集 / 质量 ${esc(data.trade_date)}

+ ${table(["数据集", "批次", "状态", "发布时间"], data.publications.map((row) => [ + esc(row.dataset), esc(row.active_batch), esc(row.state), esc(row.published_at), + ]))} +

源间差异

+ ${table(["指标", "左", "右", "偏差", "样本"], data.diff_reports.map((row) => [ + esc(row.metric), esc(row.left_value), esc(row.right_value), esc(row.deviation), row.sample_count ?? "", + ]))} + `; + return; + } + if (state.page === "audit") { + const data = await api("/admin/api/audit"); + page.innerHTML = `

审计

` + table( + ["时间", "操作者", "动作", "对象", "详情"], + data.items.map((row) => [esc(row.created_at), esc(row.actor), esc(row.action), esc(row.target), esc(row.detail)]), + ); + } +} + +function $bindRelease(page) { + page.querySelector("#rel-load").addEventListener("click", async () => { + const date = page.querySelector("#rel-date").value; + const data = await api(`/admin/api/batches?date=${encodeURIComponent(date)}`); + state.page = "release"; + // re-render with fetched date by writing location hash + history.replaceState(null, "", `#release-${date}`); + $("page").innerHTML = renderRelease(data); + $bindRelease($("page")); + }); + page.querySelector("#rel-backfill").addEventListener("click", () => dangerous("backfill")); + page.querySelectorAll("[data-rollback]").forEach((btn) => { + btn.addEventListener("click", () => dangerous("rollback", btn.dataset.rollback)); + }); +} + +function renderRelease(data) { + return ` +

盘后发布 ${esc(data.trade_date)}

+
+ + + +
+

当前映射

+ ${table(["数据集", "活跃批次", "上一批次", "状态", "发布时间", "操作"], data.publications.map((row) => [ + esc(row.dataset), esc(row.active_batch), esc(row.prev_batch), esc(row.state), esc(row.published_at), + row.prev_batch ? `` : "-", + ]))} +

批次

+ ${table(["batch_id", "数据集", "状态", "行数", "错误"], data.batches.map((row) => [ + esc(row.batch_id), esc(row.dataset), esc(row.state), row.rows_out ?? "", esc(row.error), + ]))} + `; +} + +async function dangerous(kind, dataset) { + const date = ($("rel-date") && $("rel-date").value) || ""; + const ds = dataset || prompt("数据集(daily / valuation / moneyflow / auction / index_daily / reference)", "daily"); + if (!ds) return; + const password = prompt("二次确认:输入管理密码"); + if (!password) return; + const confirmWord = `${ds}:${date}`; + const typed = prompt(`请输入确认词:${confirmWord}`); + const path = kind === "rollback" ? "/admin/api/rollback" : "/admin/api/backfill"; + await api(path, { + method: "POST", + body: JSON.stringify({ dataset: ds, trade_date: date, password, confirm: typed }), + }); + render(); +} + +boot(); diff --git a/xiaobai-datahub/admin/index.html b/xiaobai-datahub/admin/index.html new file mode 100644 index 0000000..b7eb893 --- /dev/null +++ b/xiaobai-datahub/admin/index.html @@ -0,0 +1,53 @@ + + + + + + xiaobai-datahub 管理后台 + + + +
+
+

数据中枢

+

内网管理后台,用于查看源状态、调度和盘后发布批次。

+
+ + + + +
+
+ + + + +
+ + + diff --git a/xiaobai-datahub/admin/styles.css b/xiaobai-datahub/admin/styles.css new file mode 100644 index 0000000..1631ac1 --- /dev/null +++ b/xiaobai-datahub/admin/styles.css @@ -0,0 +1,51 @@ +:root { + color-scheme: light; + --bg: #f4f5f7; + --surface: #ffffff; + --text: #1f2329; + --muted: #646a73; + --line: #dee0e3; + --action: #3370ff; + --danger: #e04536; + --ok: #16a34a; + --warn: #b45309; + --radius: 8px; + --pad: 16px; + font-family: "Segoe UI", "PingFang SC", "Noto Sans SC", sans-serif; +} +:root[data-theme="night"] { + color-scheme: dark; + --bg: #111318; + --surface: #1b1e24; + --text: #e8eaed; + --muted: #9aa0a6; + --line: #2a2f38; + --action: #5b8cff; +} +* { box-sizing: border-box; } +body { margin: 0; background: var(--bg); color: var(--text); } +.panel, header.top, nav, main { background: var(--surface); } +.auth-panel { max-width: 420px; margin: 12vh auto; padding: 28px; border-radius: var(--radius); border: 1px solid var(--line); } +label { display: block; margin: 12px 0; } +input, select { width: 100%; padding: 8px 10px; border: 1px solid var(--line); border-radius: 4px; background: var(--bg); color: var(--text); } +button { background: var(--action); color: #fff; border: 0; border-radius: 4px; padding: 8px 14px; cursor: pointer; } +button.ghost { background: transparent; color: var(--text); border: 1px solid var(--line); } +button.danger { background: var(--danger); } +.muted { color: var(--muted); } +.error { color: var(--danger); } +.top { display: flex; gap: 12px; align-items: center; padding: 10px var(--pad); border-bottom: 1px solid var(--line); } +nav { display: flex; gap: 4px; padding: 8px var(--pad); border-bottom: 1px solid var(--line); } +nav button { background: transparent; color: var(--muted); } +nav button.active { color: var(--action); background: transparent; font-weight: 600; } +main { padding: var(--pad); min-height: calc(100vh - 96px); } +.cards { display: grid; grid-template-columns: repeat(auto-fit, minmax(180px, 1fr)); gap: 12px; margin-bottom: 16px; } +.card { border: 1px solid var(--line); border-radius: var(--radius); padding: 12px; } +table { width: 100%; border-collapse: collapse; font-size: 13px; } +th, td { text-align: left; padding: 8px; border-bottom: 1px solid var(--line); vertical-align: top; } +.pill { font-size: 12px; padding: 2px 8px; border-radius: 999px; border: 1px solid var(--line); } +.ok { color: var(--ok); } +.warn { color: var(--warn); } +.fail { color: var(--danger); } +.toolbar { display: flex; gap: 8px; flex-wrap: wrap; margin: 12px 0; align-items: end; } +.toolbar label { margin: 0; } +dialog { border: 1px solid var(--line); border-radius: var(--radius); background: var(--surface); color: var(--text); padding: 20px; } diff --git a/xiaobai-datahub/compose.yaml b/xiaobai-datahub/compose.yaml new file mode 100644 index 0000000..d2a8c08 --- /dev/null +++ b/xiaobai-datahub/compose.yaml @@ -0,0 +1,39 @@ +services: + xiaobai-datahub: + build: + context: . + dockerfile: Dockerfile + image: xiaobai-datahub:local + container_name: xiaobai-datahub + ports: + - "0.0.0.0:8766:8766/tcp" + env_file: + - ./.env + environment: + DATAHUB_ENCRYPTION_KEY: "${DATAHUB_ENCRYPTION_KEY:?DATAHUB_ENCRYPTION_KEY must be set}" + DATAHUB_TOKEN: "${DATAHUB_TOKEN:?DATAHUB_TOKEN must be set}" + DATAHUB_ADMIN_PASSWORD: "${DATAHUB_ADMIN_PASSWORD:?DATAHUB_ADMIN_PASSWORD must be set}" + TUSHARE_TOKEN: "${TUSHARE_TOKEN:-}" + DATAHUB_DB_PATH: /app/data/datahub.db + DATAHUB_BACKUP_DIR: /app/data/backups + TZ: Asia/Shanghai + PYTHONUTF8: "1" + volumes: + - type: bind + source: ./data + target: /app/data + restart: unless-stopped + init: true + read_only: true + tmpfs: + - /tmp:size=64m,mode=1777 + security_opt: + - no-new-privileges:true + cap_drop: + - ALL + stop_grace_period: 30s + logging: + driver: json-file + options: + max-size: "10m" + max-file: "3" diff --git a/xiaobai-datahub/config/hub-quality.config.json b/xiaobai-datahub/config/hub-quality.config.json new file mode 100644 index 0000000..3ae6a17 --- /dev/null +++ b/xiaobai-datahub/config/hub-quality.config.json @@ -0,0 +1,15 @@ +{ + "daily_row_ratio": 0.98, + "null_rate_max": 0.01, + "cross_check_price_deviation": 0.03, + "cross_check_outlier_ratio": 0.05, + "index_price_deviation": 0.005, + "max_publish_attempts": 5, + "staging_retain_days": 14, + "job_run_retain_days": 90, + "backup_retain": 14, + "publication_generations": 3, + "tushare_rate_per_minute": 300, + "list_limit_default": 5000, + "list_limit_max": 5000 +} diff --git a/xiaobai-datahub/datahub/__init__.py b/xiaobai-datahub/datahub/__init__.py new file mode 100644 index 0000000..a7122bc --- /dev/null +++ b/xiaobai-datahub/datahub/__init__.py @@ -0,0 +1,4 @@ +"""xiaobai-datahub: independent market-data service for xiaobai-review.""" + +__version__ = "0.1.0" +SCHEMA_VERSION = 1 diff --git a/xiaobai-datahub/datahub/adapters/__init__.py b/xiaobai-datahub/datahub/adapters/__init__.py new file mode 100644 index 0000000..a5dcfb8 --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/__init__.py @@ -0,0 +1,15 @@ +from datahub.adapters.akshare import ADAPTER as akshare +from datahub.adapters.eastmoney import ADAPTER as eastmoney +from datahub.adapters.ifind import ADAPTER as ifind +from datahub.adapters.tencent import ADAPTER as tencent +from datahub.adapters.ths import ADAPTER as ths +from datahub.adapters.xgb import ADAPTER as xgb + +RESERVED = { + "eastmoney": eastmoney, + "tencent": tencent, + "ths": ths, + "xgb": xgb, + "akshare": akshare, + "ifind": ifind, +} diff --git a/xiaobai-datahub/datahub/adapters/akshare.py b/xiaobai-datahub/datahub/adapters/akshare.py new file mode 100644 index 0000000..cc8cfe3 --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/akshare.py @@ -0,0 +1,3 @@ +from datahub.adapters.base import ReservedAdapter + +ADAPTER = ReservedAdapter("akshare") diff --git a/xiaobai-datahub/datahub/adapters/base.py b/xiaobai-datahub/datahub/adapters/base.py new file mode 100644 index 0000000..3bf148e --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/base.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +from abc import ABC, abstractmethod +from typing import Any + + +class AdapterError(RuntimeError): + pass + + +class MarketAdapter(ABC): + """Uniform adapter: probe / fetch / normalize. Realtime adapters may be stubs in P0.""" + + name: str = "base" + + @abstractmethod + def probe(self) -> dict[str, Any]: + """Liveness check. Must not leak credentials.""" + + @abstractmethod + def fetch(self, dataset: str, params: dict[str, Any]) -> list[dict[str, Any]]: + """Return provider-native rows (pre-canonical).""" + + @abstractmethod + def normalize(self, dataset: str, rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Map provider-native rows onto hub canonical fields.""" + + +class ReservedAdapter(MarketAdapter): + """Placeholder for a later free/licensed source. Does not pull data in P0.""" + + def __init__(self, name: str) -> None: + self.name = name + + def probe(self) -> dict[str, Any]: + return { + "provider": self.name, + "configured": False, + "state": "reserved", + "message": "适配器位已预留,本阶段不接入", + } + + def fetch(self, dataset: str, params: dict[str, Any]) -> list[dict[str, Any]]: + raise AdapterError(f"{self.name} 适配器本阶段未接入") + + def normalize(self, dataset: str, rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + return [] diff --git a/xiaobai-datahub/datahub/adapters/eastmoney.py b/xiaobai-datahub/datahub/adapters/eastmoney.py new file mode 100644 index 0000000..29d5034 --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/eastmoney.py @@ -0,0 +1,3 @@ +from datahub.adapters.base import ReservedAdapter + +ADAPTER = ReservedAdapter("eastmoney") diff --git a/xiaobai-datahub/datahub/adapters/ifind.py b/xiaobai-datahub/datahub/adapters/ifind.py new file mode 100644 index 0000000..759d48c --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/ifind.py @@ -0,0 +1,3 @@ +from datahub.adapters.base import ReservedAdapter + +ADAPTER = ReservedAdapter("ifind") diff --git a/xiaobai-datahub/datahub/adapters/tencent.py b/xiaobai-datahub/datahub/adapters/tencent.py new file mode 100644 index 0000000..015c701 --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/tencent.py @@ -0,0 +1,3 @@ +from datahub.adapters.base import ReservedAdapter + +ADAPTER = ReservedAdapter("tencent") diff --git a/xiaobai-datahub/datahub/adapters/ths.py b/xiaobai-datahub/datahub/adapters/ths.py new file mode 100644 index 0000000..ffae8a7 --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/ths.py @@ -0,0 +1,3 @@ +from datahub.adapters.base import ReservedAdapter + +ADAPTER = ReservedAdapter("ths") diff --git a/xiaobai-datahub/datahub/adapters/tushare.py b/xiaobai-datahub/datahub/adapters/tushare.py new file mode 100644 index 0000000..c16cd42 --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/tushare.py @@ -0,0 +1,154 @@ +from __future__ import annotations + +import json +import time +import urllib.error +import urllib.request +from typing import Any, Callable + +from datahub.adapters.base import AdapterError, MarketAdapter +from datahub.normalize import ( + normalize_auction, + normalize_calendar, + normalize_daily, + normalize_index_daily, + normalize_moneyflow, + normalize_stock, + normalize_valuation, +) + +TUSHARE_URL = "http://api.tushare.pro" + +TUSHARE_FIELDS = { + "trade_cal": "exchange,cal_date,is_open,pretrade_date", + "stock_basic": "ts_code,symbol,name,area,industry,market,list_status,list_date", + "daily": "ts_code,trade_date,open,high,low,close,pct_chg,vol,amount", + "daily_basic": "ts_code,trade_date,turnover_rate,volume_ratio,total_mv,circ_mv,pe_ttm,pb,ps_ttm,dv_ttm", + "adj_factor": "ts_code,trade_date,adj_factor", + "index_daily": "ts_code,trade_date,open,high,low,close,pct_chg,vol,amount", + "moneyflow": ( + "ts_code,trade_date,buy_sm_amount,sell_sm_amount,buy_md_amount,sell_md_amount," + "buy_lg_amount,sell_lg_amount,buy_elg_amount,sell_elg_amount,net_mf_amount" + ), + "stk_auction": "ts_code,trade_date,vol,price,amount,pre_close,turnover_rate,volume_ratio,float_share", +} + +DATASET_API = { + "calendar": "trade_cal", + "stocks": "stock_basic", + "daily": "daily", + "valuation": "daily_basic", + "adj_factor": "adj_factor", + "index_daily": "index_daily", + "moneyflow": "moneyflow", + "auction": "stk_auction", +} + +DEFAULT_INDEX_CODES = ("000001.SH", "399001.SZ", "399006.SZ", "000300.SH") + + +class TushareAdapter(MarketAdapter): + name = "tushare" + + def __init__( + self, + token: str, + timeout: int = 30, + transport: Callable[[str, dict[str, Any], str], list[dict[str, Any]]] | None = None, + ) -> None: + self.token = token + self.timeout = timeout + self._transport = transport + + def probe(self) -> dict[str, Any]: + if not self.token: + return {"provider": self.name, "configured": False, "state": "unconfigured"} + started = time.perf_counter() + try: + rows = self.fetch("calendar", {"exchange": "SSE", "start_date": "20200102", "end_date": "20200102"}) + except AdapterError as exc: + return { + "provider": self.name, + "configured": True, + "state": "error", + "message": str(exc), + "latency_ms": round((time.perf_counter() - started) * 1000), + } + return { + "provider": self.name, + "configured": True, + "state": "ok" if rows else "empty", + "latency_ms": round((time.perf_counter() - started) * 1000), + } + + def fetch(self, dataset: str, params: dict[str, Any]) -> list[dict[str, Any]]: + api_name = DATASET_API.get(dataset, dataset) + fields = TUSHARE_FIELDS.get(api_name, "") + query_params = dict(params) + if api_name == "stock_basic" and "list_status" not in query_params: + query_params["list_status"] = "L" + if api_name == "trade_cal" and "exchange" not in query_params: + query_params["exchange"] = "SSE" + if api_name == "index_daily" and "ts_code" not in query_params: + # Caller typically loops codes; a missing code would pull nothing useful. + query_params.setdefault("ts_code", DEFAULT_INDEX_CODES[0]) + return self._query(api_name, query_params, fields) + + def fetch_index_daily(self, trade_date: str, codes: tuple[str, ...] = DEFAULT_INDEX_CODES) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + for ts_code in codes: + rows.extend(self.fetch("index_daily", {"ts_code": ts_code, "trade_date": trade_date})) + return rows + + def normalize(self, dataset: str, rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + mapping = { + "calendar": normalize_calendar, + "trade_cal": normalize_calendar, + "stocks": normalize_stock, + "stock_basic": normalize_stock, + "daily": normalize_daily, + "valuation": normalize_valuation, + "daily_basic": normalize_valuation, + "moneyflow": normalize_moneyflow, + "auction": normalize_auction, + "stk_auction": normalize_auction, + "index_daily": normalize_index_daily, + } + fn = mapping.get(dataset) + if fn is None: + if dataset == "adj_factor": + return [ + { + "ts_code": str(row.get("ts_code") or "").upper(), + "trade_date": str(row.get("trade_date") or ""), + "adj_factor": row.get("adj_factor"), + } + for row in rows + ] + raise AdapterError(f"unsupported dataset: {dataset}") + return [fn(row) for row in rows] + + def _query(self, api_name: str, params: dict[str, Any], fields: str) -> list[dict[str, Any]]: + if self._transport is not None: + return self._transport(api_name, params, fields) + if not self.token: + raise AdapterError("Tushare token 未配置") + payload = json.dumps( + {"api_name": api_name, "token": self.token, "params": params, "fields": fields} + ).encode("utf-8") + request = urllib.request.Request( + TUSHARE_URL, + data=payload, + headers={"Content-Type": "application/json", "User-Agent": "XiaobaiDatahub/0.1"}, + method="POST", + ) + try: + with urllib.request.urlopen(request, timeout=self.timeout) as response: + result = json.loads(response.read().decode("utf-8")) + except (urllib.error.URLError, TimeoutError, json.JSONDecodeError) as exc: + raise AdapterError(f"Tushare request failed: {exc}") from exc + if result.get("code") != 0: + raise AdapterError(result.get("msg") or "Tushare returned an unknown error") + data = result.get("data") or {} + columns = data.get("fields") or [] + return [dict(zip(columns, item)) for item in data.get("items") or []] diff --git a/xiaobai-datahub/datahub/adapters/xgb.py b/xiaobai-datahub/datahub/adapters/xgb.py new file mode 100644 index 0000000..68cea88 --- /dev/null +++ b/xiaobai-datahub/datahub/adapters/xgb.py @@ -0,0 +1,3 @@ +from datahub.adapters.base import ReservedAdapter + +ADAPTER = ReservedAdapter("xgb") diff --git a/xiaobai-datahub/datahub/admin_api.py b/xiaobai-datahub/datahub/admin_api.py new file mode 100644 index 0000000..2c4899b --- /dev/null +++ b/xiaobai-datahub/datahub/admin_api.py @@ -0,0 +1,162 @@ +from __future__ import annotations + +import json +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.scheduler import Scheduler +from datahub.serving import ApiError +from datahub.timeutil import isoformat, now_shanghai, session_phase, yyyymmdd + + +class AdminAPI: + def __init__(self, db: HubDB, pipeline: Pipeline, scheduler: Scheduler, auth: AuthService) -> None: + self.db = db + self.pipeline = pipeline + self.scheduler = scheduler + self.auth = auth + + def overview(self) -> dict[str, Any]: + today = yyyymmdd(now_shanghai()) + cal = self.db.fetchone( + "SELECT is_open FROM trade_calendar WHERE exchange = 'SSE' AND cal_date = ?", + (today,), + ) + is_open = bool(cal and int(cal["is_open"]) == 1) + pubs = self.db.fetchall("SELECT * FROM publications WHERE trade_date = ?", (today,)) + failed = self.db.fetchall( + "SELECT * FROM batches WHERE trade_date = ? AND state IN ('failed','staged')", + (today,), + ) + calls = self.db.fetchall( + "SELECT * FROM src_calls ORDER BY id DESC LIMIT 20", + ) + return { + "trade_date": today, + "session_phase": session_phase(now_shanghai(), is_open), + "is_open_day": is_open, + "publications": pubs, + "anomalies": failed, + "recent_calls": _public_calls(calls), + "source_count": len(self.db.fetchall("SELECT provider FROM src_health")), + } + + def sources(self) -> dict[str, Any]: + health = {f"{row['provider']}:{row['endpoint_class']}": row for row in self.db.fetchall("SELECT * FROM src_health")} + items = [ + { + "provider": "tushare", + "role": "official", + "health": health.get("tushare:pro") or {"state": "unknown"}, + "credential": self.auth.credential_status("tushare_token") or {"configured": bool(self.pipeline.adapter.token)}, + } + ] + for name, adapter in RESERVED.items(): + items.append( + { + "provider": name, + "role": "reserved", + "health": adapter.probe(), + "credential": {"configured": False, "last4": "", "updated_at": ""}, + } + ) + # Prefer encrypted last4 if stored + cred = self.auth.credential_status("tushare_token") + if cred.get("configured"): + items[0]["credential"] = cred + elif self.pipeline.adapter.token: + from datahub.crypto import mask_secret + + items[0]["credential"] = {"configured": True, "last4": mask_secret(self.pipeline.adapter.token), "updated_at": ""} + return {"items": items} + + def probe(self, provider: str) -> dict[str, Any]: + if provider == "tushare": + return self.pipeline.adapter.probe() + adapter = RESERVED.get(provider) + if adapter is None: + raise ApiError("INVALID_ARGUMENT", f"unknown provider: {provider}") + return adapter.probe() + + def jobs(self) -> dict[str, Any]: + runs = self.db.fetchall("SELECT * FROM job_runs ORDER BY id DESC LIMIT 100") + 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": "cleanup", "at": "00:30", "title": "清理 staging / 日志"}, + {"id": "backup", "at": "00:40", "title": "SQLite 备份"}, + ], + "runs": runs, + } + + def run_job(self, job_id: str, trade_date: str) -> dict[str, Any]: + return self.scheduler.run_job(job_id, yyyymmdd(trade_date or now_shanghai())) + + def batches(self, date: str, dataset: str = "") -> dict[str, Any]: + trade_date = yyyymmdd(date or now_shanghai()) + if dataset: + rows = self.db.fetchall( + "SELECT * FROM batches WHERE trade_date = ? AND dataset = ? ORDER BY started_at", + (trade_date, dataset), + ) + else: + rows = self.db.fetchall( + "SELECT * FROM batches WHERE trade_date = ? ORDER BY started_at", + (trade_date,), + ) + pubs = self.db.fetchall("SELECT * FROM publications WHERE trade_date = ?", (trade_date,)) + return {"trade_date": trade_date, "batches": rows, "publications": pubs} + + def datasets(self, date: str) -> dict[str, Any]: + trade_date = yyyymmdd(date or now_shanghai()) + pubs = self.db.fetchall("SELECT * FROM publications WHERE trade_date = ?", (trade_date,)) + diffs = self.db.fetchall( + "SELECT * FROM diff_reports WHERE trade_date = ? ORDER BY id", + (trade_date,), + ) + return {"trade_date": trade_date, "publications": pubs, "diff_reports": diffs} + + def audit(self) -> dict[str, Any]: + return {"items": self.db.fetchall("SELECT * FROM audit_log ORDER BY id DESC LIMIT 200")} + + def rollback(self, dataset: str, trade_date: str, password: str, confirm: str, actor: str) -> dict[str, Any]: + self._dangerous(password, confirm, f"{dataset}:{trade_date}") + result = self.pipeline.rollback(dataset, trade_date, actor=actor) + 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) + else: + result = self.pipeline.run_dataset(dataset, trade_date) + self.pipeline.audit(actor, "backfill", f"{dataset}:{trade_date}", json.dumps({"ok": True})) + return result + + def _dangerous(self, password: str, confirm: str, expected: str) -> None: + if not self.auth.confirm_password(password): + raise ApiError("UNAUTHORIZED", "二次确认密码错误") + if confirm.strip() != expected: + raise ApiError("INVALID_ARGUMENT", f"确认词必须为 {expected}") + + +def _public_calls(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + out = [] + for row in rows: + out.append( + { + "id": row["id"], + "provider": row["provider"], + "endpoint": row["endpoint"], + "ok": bool(row["ok"]), + "latency_ms": row["latency_ms"], + "error": row["error"], + "created_at": row["created_at"], + } + ) + return out diff --git a/xiaobai-datahub/datahub/auth.py b/xiaobai-datahub/datahub/auth.py new file mode 100644 index 0000000..8c60de4 --- /dev/null +++ b/xiaobai-datahub/datahub/auth.py @@ -0,0 +1,190 @@ +from __future__ import annotations + +import base64 +import hashlib +import hmac +import os +import secrets +from datetime import timedelta +from typing import Any + +from datahub.crypto import SecretVault, mask_secret +from datahub.db import HubDB +from datahub.timeutil import isoformat, now_shanghai + +PBKDF2_ROUNDS = 200_000 +SESSION_HOURS = 12 +LOGIN_FAIL_LIMIT = 5 +LOCK_MINUTES = 10 + + +def hash_password(password: str, salt: bytes | None = None) -> tuple[str, str]: + raw_salt = salt or os.urandom(16) + digest = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), raw_salt, PBKDF2_ROUNDS, dklen=32) + return ( + base64.urlsafe_b64encode(raw_salt).decode("ascii"), + base64.urlsafe_b64encode(digest).decode("ascii"), + ) + + +def verify_password(password: str, salt_text: str, expected_hash: str) -> bool: + try: + salt = base64.urlsafe_b64decode(salt_text.encode("ascii")) + _, actual = hash_password(password, salt) + except (ValueError, TypeError): + return False + return hmac.compare_digest(actual, expected_hash) + + +def token_hash(token: str) -> str: + return hashlib.sha256(token.encode("utf-8")).hexdigest() + + +class AuthService: + def __init__(self, db: HubDB, vault: SecretVault, api_token: str, admin_password: str) -> None: + self.db = db + self.vault = vault + self._bootstrap(api_token, admin_password) + + def _bootstrap(self, api_token: str, admin_password: str) -> None: + if api_token: + existing = self.db.fetchone("SELECT token_hash FROM api_tokens WHERE name = ?", ("review",)) + hashed = token_hash(api_token) + last4 = mask_secret(api_token) + if existing is None: + self.db.execute( + "INSERT INTO api_tokens(token_hash, name, last4, created_at) VALUES (?,?,?,?)", + (hashed, "review", last4, isoformat()), + ) + elif existing["token_hash"] != hashed: + self.db.execute( + "UPDATE api_tokens SET token_hash = ?, last4 = ? WHERE name = ?", + (hashed, last4, "review"), + ) + admin = self.db.fetchone("SELECT id FROM hub_admin WHERE username = ?", ("hub_admin",)) + if admin is None and admin_password: + salt, hashed = hash_password(admin_password) + now = isoformat() + self.db.execute( + """ + INSERT INTO hub_admin(username, password_salt, password_hash, password_must_change, created_at, updated_at) + VALUES (?, ?, ?, 1, ?, ?) + """, + ("hub_admin", salt, hashed, now, now), + ) + + def check_api_token(self, supplied: str) -> bool: + if not supplied: + return False + row = self.db.fetchone( + "SELECT token_hash FROM api_tokens WHERE token_hash = ? AND revoked_at IS NULL", + (token_hash(supplied),), + ) + return row is not None + + def login(self, username: str, password: str) -> dict[str, Any]: + user = self.db.fetchone("SELECT * FROM hub_admin WHERE username = ?", (username,)) + if not user: + raise PermissionError("账号或密码错误") + now = now_shanghai() + locked_until = user.get("locked_until") + if locked_until: + try: + from datetime import datetime + + if datetime.fromisoformat(str(locked_until)) > now: + raise PermissionError("账号已锁定,请稍后再试") + except ValueError: + pass + if not verify_password(password, str(user["password_salt"]), str(user["password_hash"])): + fails = int(user["failed_attempts"] or 0) + 1 + lock = isoformat(now + timedelta(minutes=LOCK_MINUTES)) if fails >= LOGIN_FAIL_LIMIT else None + self.db.execute( + "UPDATE hub_admin SET failed_attempts = ?, locked_until = ? WHERE id = ?", + (fails, lock, user["id"]), + ) + raise PermissionError("账号或密码错误") + self.db.execute( + "UPDATE hub_admin SET failed_attempts = 0, locked_until = NULL WHERE id = ?", + (user["id"],), + ) + session = secrets.token_urlsafe(32) + csrf = secrets.token_urlsafe(24) + expires = isoformat(now + timedelta(hours=SESSION_HOURS)) + self.db.execute( + "INSERT INTO hub_sessions(token_hash, csrf_token, expires_at, created_at) VALUES (?,?,?,?)", + (token_hash(session), csrf, expires, isoformat(now)), + ) + return { + "session": session, + "csrf": csrf, + "must_change": bool(user["password_must_change"]), + "expires_at": expires, + } + + def session_user(self, raw_token: str) -> dict[str, Any] | None: + if not raw_token: + return None + row = self.db.fetchone( + "SELECT * FROM hub_sessions WHERE token_hash = ?", + (token_hash(raw_token),), + ) + if not row: + return None + if str(row["expires_at"]) < isoformat(): + self.db.execute("DELETE FROM hub_sessions WHERE token_hash = ?", (row["token_hash"],)) + return None + admin = self.db.fetchone("SELECT username, password_must_change FROM hub_admin WHERE username = ?", ("hub_admin",)) + return { + "username": (admin or {}).get("username") or "hub_admin", + "csrf_token": row["csrf_token"], + "must_change": bool((admin or {}).get("password_must_change")), + "token_hash": row["token_hash"], + } + + def logout(self, raw_token: str) -> None: + if raw_token: + self.db.execute("DELETE FROM hub_sessions WHERE token_hash = ?", (token_hash(raw_token),)) + + def change_password(self, current: str, new_password: str) -> None: + if len(new_password) < 8: + raise ValueError("新密码至少 8 位") + user = self.db.fetchone("SELECT * FROM hub_admin WHERE username = ?", ("hub_admin",)) + if not user or not verify_password(current, str(user["password_salt"]), str(user["password_hash"])): + raise PermissionError("当前密码错误") + salt, hashed = hash_password(new_password) + self.db.execute( + "UPDATE hub_admin SET password_salt=?, password_hash=?, password_must_change=0, updated_at=? WHERE id=?", + (salt, hashed, isoformat(), user["id"]), + ) + + def confirm_password(self, password: str) -> bool: + user = self.db.fetchone("SELECT * FROM hub_admin WHERE username = ?", ("hub_admin",)) + if not user: + return False + return verify_password(password, str(user["password_salt"]), str(user["password_hash"])) + + def credential_status(self, name: str) -> dict[str, Any]: + row = self.db.fetchone("SELECT last4, updated_at FROM credentials WHERE name = ?", (name,)) + if not row: + return {"configured": False, "last4": "", "updated_at": ""} + return {"configured": True, "last4": row["last4"], "updated_at": row["updated_at"]} + + def store_credential(self, name: str, secret: str) -> None: + payload = self.vault.encrypt_json({name: secret}) + self.db.execute( + """ + INSERT INTO credentials(name, encrypted_payload, last4, updated_at) + VALUES (?, ?, ?, ?) + ON CONFLICT(name) DO UPDATE SET + encrypted_payload=excluded.encrypted_payload, last4=excluded.last4, updated_at=excluded.updated_at + """, + (name, payload, mask_secret(secret), isoformat()), + ) + + def load_credential(self, name: str) -> str: + row = self.db.fetchone("SELECT encrypted_payload FROM credentials WHERE name = ?", (name,)) + if not row: + return "" + data = self.vault.decrypt_json(str(row["encrypted_payload"])) + return str(data.get(name) or "") diff --git a/xiaobai-datahub/datahub/codes.py b/xiaobai-datahub/datahub/codes.py new file mode 100644 index 0000000..a18d859 --- /dev/null +++ b/xiaobai-datahub/datahub/codes.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +from datahub.db import HubDB + + +def resolve_code(db: HubDB, raw: str) -> str | None: + text = str(raw or "").strip().upper() + if not text: + return None + if "." in text: + row = db.fetchone("SELECT ts_code FROM stock_master WHERE ts_code = ?", (text,)) + if row: + return row["ts_code"] + # indices are not always in stock_master + return text + matches = db.fetchall( + "SELECT ts_code FROM stock_master WHERE symbol = ? OR ts_code LIKE ?", + (text, f"{text}.%"), + ) + if len(matches) == 1: + return matches[0]["ts_code"] + if len(matches) > 1: + return None + # unique exchange guess for 6-digit codes + suffix = "SH" if text.startswith("6") or text.startswith("9") else "SZ" if text.startswith(("0", "3")) else "BJ" + return f"{text}.{suffix}" diff --git a/xiaobai-datahub/datahub/crypto.py b/xiaobai-datahub/datahub/crypto.py new file mode 100644 index 0000000..5d7882c --- /dev/null +++ b/xiaobai-datahub/datahub/crypto.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +import json +from typing import Any + +from cryptography.fernet import Fernet, InvalidToken + + +class SecretVault: + def __init__(self, key: str) -> None: + try: + self._fernet = Fernet(key.encode("ascii")) + except (ValueError, TypeError) as exc: + raise ValueError("DATAHUB_ENCRYPTION_KEY 格式无效。") from exc + + @staticmethod + def generate_key() -> str: + return Fernet.generate_key().decode("ascii") + + def encrypt_json(self, payload: dict[str, Any]) -> str: + raw = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + return self._fernet.encrypt(raw).decode("ascii") + + def decrypt_json(self, token: str) -> dict[str, Any]: + if not token: + return {} + try: + payload = json.loads(self._fernet.decrypt(token.encode("ascii")).decode("utf-8")) + except (InvalidToken, UnicodeDecodeError, json.JSONDecodeError) as exc: + raise ValueError("凭据无法解密,请检查 DATAHUB_ENCRYPTION_KEY。") from exc + if not isinstance(payload, dict): + raise ValueError("凭据格式无效。") + return payload + + +def mask_secret(value: str, last_n: int = 4) -> str: + text = str(value or "") + if not text: + return "" + if len(text) <= last_n: + return "*" * len(text) + return ("*" * max(4, len(text) - last_n)) + text[-last_n:] diff --git a/xiaobai-datahub/datahub/db.py b/xiaobai-datahub/datahub/db.py new file mode 100644 index 0000000..dec81f9 --- /dev/null +++ b/xiaobai-datahub/datahub/db.py @@ -0,0 +1,340 @@ +from __future__ import annotations + +import sqlite3 +import threading +from collections.abc import Iterator +from contextlib import contextmanager +from pathlib import Path +from typing import Any + +from datahub.timeutil import isoformat + +SCHEMA = """ +CREATE TABLE IF NOT EXISTS schema_migrations ( + version INTEGER PRIMARY KEY, + applied_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS credentials ( + name TEXT PRIMARY KEY, + encrypted_payload TEXT NOT NULL, + last4 TEXT, + updated_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS hub_admin ( + id INTEGER PRIMARY KEY, + username TEXT NOT NULL UNIQUE, + password_salt TEXT NOT NULL, + password_hash TEXT NOT NULL, + password_must_change INTEGER NOT NULL DEFAULT 1, + failed_attempts INTEGER NOT NULL DEFAULT 0, + locked_until TEXT, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS hub_sessions ( + token_hash TEXT PRIMARY KEY, + csrf_token TEXT NOT NULL, + expires_at TEXT NOT NULL, + created_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS api_tokens ( + token_hash TEXT PRIMARY KEY, + name TEXT NOT NULL, + last4 TEXT NOT NULL, + created_at TEXT NOT NULL, + revoked_at TEXT +); + +CREATE TABLE IF NOT EXISTS trade_calendar ( + exchange TEXT NOT NULL, + cal_date TEXT NOT NULL, + is_open INTEGER NOT NULL, + pretrade_date TEXT, + fetched_at TEXT NOT NULL, + PRIMARY KEY (exchange, cal_date) +); + +CREATE TABLE IF NOT EXISTS stock_master ( + ts_code TEXT PRIMARY KEY, + symbol TEXT, + name TEXT, + area TEXT, + industry TEXT, + market TEXT, + list_status TEXT, + list_date TEXT, + updated_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS eod_bars ( + ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, + open REAL, high REAL, low REAL, close REAL, pct_chg REAL, + volume REAL, amount REAL, adj_factor REAL, + batch_id TEXT NOT NULL, + PRIMARY KEY (ts_code, trade_date, batch_id) +) WITHOUT ROWID; + +CREATE TABLE IF NOT EXISTS eod_valuation ( + ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, + turnover_rate REAL, volume_ratio REAL, + total_mv REAL, circ_mv REAL, + pe_ttm REAL, pb REAL, ps_ttm REAL, dv_ttm REAL, + batch_id TEXT NOT NULL, + PRIMARY KEY (ts_code, trade_date, batch_id) +) WITHOUT ROWID; + +CREATE TABLE IF NOT EXISTS eod_moneyflow ( + ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, + buy_sm_amount REAL, sell_sm_amount REAL, + buy_md_amount REAL, sell_md_amount REAL, + buy_lg_amount REAL, sell_lg_amount REAL, + buy_elg_amount REAL, sell_elg_amount REAL, + net_mf_amount REAL, + batch_id TEXT NOT NULL, + PRIMARY KEY (ts_code, trade_date, batch_id) +) WITHOUT ROWID; + +CREATE TABLE IF NOT EXISTS eod_auction ( + ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, + volume REAL, price REAL, amount REAL, pre_close REAL, + turnover_rate REAL, volume_ratio REAL, float_share REAL, + batch_id TEXT NOT NULL, + PRIMARY KEY (ts_code, trade_date, batch_id) +) WITHOUT ROWID; + +CREATE TABLE IF NOT EXISTS eod_index_bars ( + ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, + open REAL, high REAL, low REAL, close REAL, pct_chg REAL, + volume REAL, amount REAL, + batch_id TEXT NOT NULL, + PRIMARY KEY (ts_code, trade_date, batch_id) +) WITHOUT ROWID; + +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, + volume REAL, amount REAL, adj_factor REAL, + PRIMARY KEY (batch_id, ts_code, trade_date) +); + +CREATE TABLE IF NOT EXISTS staging_valuation ( + ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, batch_id TEXT NOT NULL, + turnover_rate REAL, volume_ratio REAL, + total_mv REAL, circ_mv REAL, pe_ttm REAL, pb REAL, ps_ttm REAL, dv_ttm REAL, + PRIMARY KEY (batch_id, ts_code, trade_date) +); + +CREATE TABLE IF NOT EXISTS staging_moneyflow ( + ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, batch_id TEXT NOT NULL, + buy_sm_amount REAL, sell_sm_amount REAL, buy_md_amount REAL, sell_md_amount REAL, + buy_lg_amount REAL, sell_lg_amount REAL, buy_elg_amount REAL, sell_elg_amount REAL, + net_mf_amount REAL, + PRIMARY KEY (batch_id, ts_code, trade_date) +); + +CREATE TABLE IF NOT EXISTS staging_auction ( + ts_code TEXT NOT NULL, trade_date TEXT NOT NULL, batch_id TEXT NOT NULL, + volume REAL, price REAL, amount REAL, pre_close REAL, + turnover_rate REAL, volume_ratio REAL, float_share REAL, + PRIMARY KEY (batch_id, ts_code, trade_date) +); + +CREATE TABLE IF NOT EXISTS staging_index_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, + volume REAL, amount REAL, + PRIMARY KEY (batch_id, ts_code, trade_date) +); + +CREATE TABLE IF NOT EXISTS publications ( + dataset TEXT NOT NULL, trade_date TEXT NOT NULL, + active_batch TEXT NOT NULL, prev_batch TEXT, + state TEXT NOT NULL, + published_at TEXT NOT NULL, + PRIMARY KEY (dataset, trade_date) +); + +CREATE TABLE IF NOT EXISTS publication_history ( + dataset TEXT NOT NULL, trade_date TEXT NOT NULL, + batch_id TEXT NOT NULL, published_at TEXT NOT NULL, + generation INTEGER NOT NULL, + PRIMARY KEY (dataset, trade_date, batch_id) +); + +CREATE TABLE IF NOT EXISTS batches ( + batch_id TEXT PRIMARY KEY, + dataset TEXT NOT NULL, + trade_date TEXT NOT NULL, + state TEXT NOT NULL, + attempt INTEGER DEFAULT 0, + rows_in INTEGER, + rows_out INTEGER, + quality_json TEXT, + started_at TEXT, + finished_at TEXT, + error TEXT +); + +CREATE TABLE IF NOT EXISTS src_health ( + provider TEXT NOT NULL, endpoint_class TEXT NOT NULL, + state TEXT NOT NULL, + last_ok_at TEXT, last_error TEXT, + consec_failures INTEGER DEFAULT 0, + opened_at TEXT, + cooldown_until TEXT, + PRIMARY KEY (provider, endpoint_class) +); + +CREATE TABLE IF NOT EXISTS src_calls ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + provider TEXT NOT NULL, + endpoint TEXT NOT NULL, + ok INTEGER NOT NULL, + latency_ms INTEGER, + error TEXT, + created_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS job_runs ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + job_id TEXT NOT NULL, + state TEXT NOT NULL, + started_at TEXT, + finished_at TEXT, + rows_in INTEGER, + rows_out INTEGER, + error TEXT, + attempt INTEGER DEFAULT 1, + detail TEXT +); + +CREATE TABLE IF NOT EXISTS audit_log ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + actor TEXT NOT NULL, + action TEXT NOT NULL, + target TEXT, + detail TEXT, + created_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS rt_cache ( + cache_key TEXT PRIMARY KEY, + payload TEXT NOT NULL, + source TEXT NOT NULL, + stored_at TEXT NOT NULL, + expires_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS last_known_good ( + cache_key TEXT PRIMARY KEY, + payload TEXT NOT NULL, + source TEXT NOT NULL, + stored_at TEXT NOT NULL +); + +CREATE TABLE IF NOT EXISTS diff_reports ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + trade_date TEXT NOT NULL, + metric TEXT NOT NULL, + left_source TEXT, + right_source TEXT, + left_value REAL, + right_value REAL, + deviation REAL, + sample_count INTEGER, + created_at TEXT NOT NULL +); + +CREATE INDEX IF NOT EXISTS idx_batches_date ON batches(trade_date, dataset); +CREATE INDEX IF NOT EXISTS idx_job_runs_job ON job_runs(job_id, started_at); +CREATE INDEX IF NOT EXISTS idx_src_calls_created ON src_calls(created_at); +CREATE INDEX IF NOT EXISTS idx_eod_bars_date ON eod_bars(trade_date, batch_id); +CREATE INDEX IF NOT EXISTS idx_calendar_open ON trade_calendar(is_open, cal_date); +""" + +DATASET_TABLES = { + "daily": ("eod_bars", "staging_bars"), + "valuation": ("eod_valuation", "staging_valuation"), + "moneyflow": ("eod_moneyflow", "staging_moneyflow"), + "auction": ("eod_auction", "staging_auction"), + "index_daily": ("eod_index_bars", "staging_index_bars"), +} + + +class ManagedConnection(sqlite3.Connection): + def __exit__(self, exc_type, exc_value, traceback): + try: + return super().__exit__(exc_type, exc_value, traceback) + finally: + self.close() + + +class HubDB: + def __init__(self, path: Path, timeout_seconds: float = 20) -> None: + self.path = Path(path) + self.path.parent.mkdir(parents=True, exist_ok=True) + self.timeout_seconds = timeout_seconds + self._write_lock = threading.RLock() + self.initialize() + + def connect(self) -> sqlite3.Connection: + connection = sqlite3.connect( + self.path, + timeout=self.timeout_seconds, + factory=ManagedConnection, + ) + connection.row_factory = sqlite3.Row + connection.execute("PRAGMA journal_mode=WAL") + connection.execute("PRAGMA foreign_keys=ON") + connection.execute("PRAGMA busy_timeout=20000") + connection.execute("PRAGMA synchronous=NORMAL") + return connection + + def initialize(self) -> None: + with self.connect() as connection: + connection.executescript(SCHEMA) + row = connection.execute( + "SELECT version FROM schema_migrations ORDER BY version DESC LIMIT 1" + ).fetchone() + if row is None: + connection.execute( + "INSERT INTO schema_migrations(version, applied_at) VALUES (1, ?)", + (isoformat(),), + ) + + @contextmanager + def write(self) -> Iterator[sqlite3.Connection]: + with self._write_lock: + with self.connect() as connection: + yield connection + + def fetchall(self, sql: str, params: tuple[Any, ...] = ()) -> list[dict[str, Any]]: + with self.connect() as connection: + rows = connection.execute(sql, params).fetchall() + return [dict(row) for row in rows] + + def fetchone(self, sql: str, params: tuple[Any, ...] = ()) -> dict[str, Any] | None: + with self.connect() as connection: + row = connection.execute(sql, params).fetchone() + return dict(row) if row else None + + def execute(self, sql: str, params: tuple[Any, ...] = ()) -> None: + with self.write() as connection: + connection.execute(sql, params) + + def executemany(self, sql: str, rows: list[tuple[Any, ...]]) -> None: + with self.write() as connection: + connection.executemany(sql, rows) + + def backup_to(self, dest: Path) -> None: + dest.parent.mkdir(parents=True, exist_ok=True) + with self.connect() as source, sqlite3.connect(dest) as target: + source.backup(target) + + def vacuum(self) -> None: + with self.connect() as connection: + connection.execute("VACUUM") diff --git a/xiaobai-datahub/datahub/governance/__init__.py b/xiaobai-datahub/datahub/governance/__init__.py new file mode 100644 index 0000000..7e28faa --- /dev/null +++ b/xiaobai-datahub/datahub/governance/__init__.py @@ -0,0 +1,13 @@ +from datahub.governance.circuit import CircuitBreaker, CircuitState +from datahub.governance.lkg import LastKnownGood +from datahub.governance.ratelimit import TokenBucket +from datahub.governance.retry import RetryError, retry_call + +__all__ = [ + "CircuitBreaker", + "CircuitState", + "LastKnownGood", + "RetryError", + "TokenBucket", + "retry_call", +] diff --git a/xiaobai-datahub/datahub/governance/circuit.py b/xiaobai-datahub/datahub/governance/circuit.py new file mode 100644 index 0000000..54029a3 --- /dev/null +++ b/xiaobai-datahub/datahub/governance/circuit.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import threading +import time +from collections import deque +from dataclasses import dataclass + + +@dataclass +class CircuitState: + state: str = "closed" # closed | open | half_open + consec_failures: int = 0 + opened_at: float | None = None + cooldown_until: float = 0.0 + last_error: str = "" + last_ok_at: float | None = None + + +class CircuitBreaker: + """Sliding-window breaker: 5 consecutive failures or >50% of 60s window → open.""" + + def __init__( + self, + failure_threshold: int = 5, + window_seconds: float = 60.0, + open_seconds: float = 120.0, + max_open_seconds: float = 600.0, + clock=time.monotonic, + ) -> None: + self.failure_threshold = failure_threshold + self.window_seconds = window_seconds + self.open_seconds = open_seconds + self.max_open_seconds = max_open_seconds + self._clock = clock + self._lock = threading.Lock() + self._events: deque[tuple[float, bool]] = deque() + self.status = CircuitState() + self._open_stretch = open_seconds + + def allow(self) -> bool: + with self._lock: + self._refresh_locked() + if self.status.state == "open": + return False + if self.status.state == "half_open": + # single probe in flight: caller must record success/failure + return True + return True + + def record_success(self) -> CircuitState: + with self._lock: + now = self._clock() + self._events.append((now, True)) + self.status.last_ok_at = now + self.status.consec_failures = 0 + self.status.last_error = "" + self._open_stretch = self.open_seconds + self.status.state = "closed" + self.status.opened_at = None + self.status.cooldown_until = 0.0 + return self._copy() + + def record_failure(self, error: str = "") -> CircuitState: + with self._lock: + now = self._clock() + self._events.append((now, False)) + self.status.consec_failures += 1 + self.status.last_error = error + self._prune_locked(now) + failures = sum(1 for _, ok in self._events if not ok) + total = len(self._events) + rate = (failures / total) if total else 0.0 + trip = self.status.consec_failures >= self.failure_threshold or ( + total >= self.failure_threshold and rate > 0.5 + ) + if trip: + self.status.state = "open" + self.status.opened_at = now + self.status.cooldown_until = now + self._open_stretch + self._open_stretch = min(self.max_open_seconds, self._open_stretch * 2) + return self._copy() + + def snapshot(self) -> CircuitState: + with self._lock: + self._refresh_locked() + return self._copy() + + def _refresh_locked(self) -> None: + now = self._clock() + self._prune_locked(now) + if self.status.state == "open" and now >= self.status.cooldown_until: + self.status.state = "half_open" + + def _prune_locked(self, now: float) -> None: + cutoff = now - self.window_seconds + while self._events and self._events[0][0] < cutoff: + self._events.popleft() + + def _copy(self) -> CircuitState: + return CircuitState( + state=self.status.state, + consec_failures=self.status.consec_failures, + opened_at=self.status.opened_at, + cooldown_until=self.status.cooldown_until, + last_error=self.status.last_error, + last_ok_at=self.status.last_ok_at, + ) diff --git a/xiaobai-datahub/datahub/governance/lkg.py b/xiaobai-datahub/datahub/governance/lkg.py new file mode 100644 index 0000000..a75f569 --- /dev/null +++ b/xiaobai-datahub/datahub/governance/lkg.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +import json +from typing import Any + +from datahub.db import HubDB +from datahub.timeutil import isoformat, now_shanghai + + +class LastKnownGood: + def __init__(self, db: HubDB) -> None: + self.db = db + + def store(self, cache_key: str, payload: Any, source: str) -> None: + self.db.execute( + """ + INSERT INTO last_known_good(cache_key, payload, source, stored_at) + VALUES (?, ?, ?, ?) + ON CONFLICT(cache_key) DO UPDATE SET + payload=excluded.payload, source=excluded.source, stored_at=excluded.stored_at + """, + (cache_key, json.dumps(payload, ensure_ascii=False), source, isoformat()), + ) + + def load(self, cache_key: str) -> dict[str, Any] | None: + row = self.db.fetchone("SELECT * FROM last_known_good WHERE cache_key = ?", (cache_key,)) + if not row: + return None + return { + "payload": json.loads(row["payload"]), + "source": row["source"], + "stored_at": row["stored_at"], + } + + def put_rt(self, cache_key: str, payload: Any, source: str, ttl_seconds: int) -> None: + now = now_shanghai() + expires = isoformat(now.replace(microsecond=0)) + # expires_at stored as iso; compute by adding ttl via timestamp + from datetime import timedelta + + self.db.execute( + """ + INSERT INTO rt_cache(cache_key, payload, source, stored_at, expires_at) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(cache_key) DO UPDATE SET + payload=excluded.payload, source=excluded.source, + stored_at=excluded.stored_at, expires_at=excluded.expires_at + """, + ( + cache_key, + json.dumps(payload, ensure_ascii=False), + source, + isoformat(now), + isoformat(now + timedelta(seconds=ttl_seconds)), + ), + ) + self.store(cache_key, payload, source) + + def get_rt(self, cache_key: str, max_stale_seconds: int | None = None) -> dict[str, Any] | None: + row = self.db.fetchone("SELECT * FROM rt_cache WHERE cache_key = ?", (cache_key,)) + if not row: + lkg = self.load(cache_key) + if not lkg: + return None + return {**lkg, "stale": True} + stored_at = row["stored_at"] + expired = row["expires_at"] < isoformat() + result = { + "payload": json.loads(row["payload"]), + "source": row["source"], + "stored_at": stored_at, + "stale": expired, + } + if expired and max_stale_seconds is not None: + from datetime import datetime + + try: + stored = datetime.fromisoformat(stored_at) + age = (now_shanghai() - stored).total_seconds() + except ValueError: + age = max_stale_seconds + 1 + if age > max_stale_seconds: + return None + return result diff --git a/xiaobai-datahub/datahub/governance/ratelimit.py b/xiaobai-datahub/datahub/governance/ratelimit.py new file mode 100644 index 0000000..5a26ad2 --- /dev/null +++ b/xiaobai-datahub/datahub/governance/ratelimit.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +import threading +import time + + +class TokenBucket: + def __init__(self, rate_per_minute: float, capacity: float | None = None, clock=time.monotonic) -> None: + self.rate_per_second = max(0.001, rate_per_minute / 60.0) + self.capacity = float(capacity if capacity is not None else rate_per_minute) + self._tokens = self.capacity + self._updated = clock() + self._clock = clock + self._lock = threading.Lock() + + def acquire(self, tokens: float = 1.0, block: bool = True) -> bool: + while True: + with self._lock: + now = self._clock() + elapsed = max(0.0, now - self._updated) + self._tokens = min(self.capacity, self._tokens + elapsed * self.rate_per_second) + self._updated = now + if self._tokens >= tokens: + self._tokens -= tokens + return True + wait = (tokens - self._tokens) / self.rate_per_second + if not block: + return False + time.sleep(min(wait, 0.05)) + + @property + def remaining(self) -> float: + with self._lock: + now = self._clock() + elapsed = max(0.0, now - self._updated) + return min(self.capacity, self._tokens + elapsed * self.rate_per_second) diff --git a/xiaobai-datahub/datahub/governance/retry.py b/xiaobai-datahub/datahub/governance/retry.py new file mode 100644 index 0000000..4de2e6e --- /dev/null +++ b/xiaobai-datahub/datahub/governance/retry.py @@ -0,0 +1,35 @@ +from __future__ import annotations + +import time +from collections.abc import Callable +from typing import TypeVar + +T = TypeVar("T") + + +class RetryError(RuntimeError): + def __init__(self, message: str, attempts: int, last_error: BaseException | None = None) -> None: + super().__init__(message) + self.attempts = attempts + self.last_error = last_error + + +def retry_call( + fn: Callable[[], T], + attempts: int = 5, + base_delay: float = 0.2, + max_delay: float = 8.0, + sleeper: Callable[[float], None] = time.sleep, + retry_on: tuple[type[BaseException], ...] = (Exception,), +) -> T: + last: BaseException | None = None + for attempt in range(1, max(1, attempts) + 1): + try: + return fn() + except retry_on as exc: + last = exc + if attempt >= attempts: + break + delay = min(max_delay, base_delay * (2 ** (attempt - 1))) + sleeper(delay) + raise RetryError(f"retry exhausted after {attempts} attempts: {last}", attempts, last) diff --git a/xiaobai-datahub/datahub/httpapp.py b/xiaobai-datahub/datahub/httpapp.py new file mode 100644 index 0000000..de6328a --- /dev/null +++ b/xiaobai-datahub/datahub/httpapp.py @@ -0,0 +1,242 @@ +from __future__ import annotations + +import json +import mimetypes +import secrets +from http import HTTPStatus +from http.cookies import SimpleCookie +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any +from urllib.parse import unquote, urlparse + +from datahub.hub import Hub +from datahub.logutil import configure_logging, get_logger +from datahub.serving import ApiError, parse_query + +LOGGER = get_logger() +SESSION_COOKIE = "datahub_session" + + +class HubRequestHandler(BaseHTTPRequestHandler): + hub: Hub + + def log_message(self, format: str, *args: Any) -> None: + LOGGER.info(format % args) + + def do_GET(self) -> None: # noqa: N802 + self._dispatch("GET") + + def do_POST(self) -> None: # noqa: N802 + self._dispatch("POST") + + def do_OPTIONS(self) -> None: # noqa: N802 + self.send_response(HTTPStatus.NO_CONTENT) + self.send_header("Allow", "GET, POST, OPTIONS") + self.end_headers() + + def _dispatch(self, method: str) -> None: + parsed = urlparse(self.path) + path = unquote(parsed.path) + try: + if path in {"/livez", "/healthz"}: + self._json({"status": "ok"}, HTTPStatus.OK) + return + if path.startswith("/v1/"): + self._v1(path, parsed.query) + return + if path.startswith("/admin/api/"): + self._admin_api(method, path) + return + if path.startswith("/admin"): + self._admin_static(path) + return + if path == "/": + self.send_response(HTTPStatus.FOUND) + self.send_header("Location", "/admin/") + self.end_headers() + return + self._json({"error": {"code": "INVALID_ARGUMENT", "message": "Not found"}}, HTTPStatus.NOT_FOUND) + except ApiError as exc: + self._json(exc.payload(), exc.status) + except PermissionError as exc: + self._json({"error": {"code": "UNAUTHORIZED", "message": str(exc)}}, HTTPStatus.UNAUTHORIZED) + except ValueError as exc: + self._json({"error": {"code": "INVALID_ARGUMENT", "message": str(exc)}}, HTTPStatus.BAD_REQUEST) + except Exception: + LOGGER.exception("internal error") + self._json({"error": {"code": "INTERNAL", "message": "internal error"}}, HTTPStatus.INTERNAL_SERVER_ERROR) + + def _v1(self, path: str, query: str) -> None: + token = self.headers.get("X-Datahub-Token", "") + if not self.hub.auth.check_api_token(token): + self.hub.pipeline.audit("anonymous", "unauthorized", path, "") + raise ApiError("UNAUTHORIZED", "missing or invalid X-Datahub-Token") + payload = self.hub.api.handle(path, parse_query(query)) + self._json(payload, HTTPStatus.OK) + + def _admin_api(self, method: str, path: str) -> None: + if path == "/admin/api/login" and method == "POST": + body = self._read_json() + result = self.hub.auth.login(str(body.get("username") or "hub_admin"), str(body.get("password") or "")) + self._json( + {"ok": True, "must_change": result["must_change"], "csrf": result["csrf"]}, + HTTPStatus.OK, + extra_headers=[self._cookie(result["session"])], + ) + return + user = self.hub.auth.session_user(self._cookie_value(SESSION_COOKIE)) + if not user: + raise ApiError("UNAUTHORIZED", "请先登录") + if method == "POST" and path != "/admin/api/login": + csrf = self.headers.get("X-CSRF-Token", "") + if not csrf or not secrets.compare_digest(csrf, str(user["csrf_token"])): + raise ApiError("UNAUTHORIZED", "CSRF 校验失败") + if path == "/admin/api/logout" and method == "POST": + self.hub.auth.logout(self._cookie_value(SESSION_COOKIE)) + self._json({"ok": True}, HTTPStatus.OK, extra_headers=[self._cookie("", clear=True)]) + return + if path == "/admin/api/session" and method == "GET": + self._json({"username": user["username"], "must_change": user["must_change"], "csrf": user["csrf_token"]}, HTTPStatus.OK) + return + if path == "/admin/api/change-password" and method == "POST": + body = self._read_json() + self.hub.auth.change_password(str(body.get("current") or ""), str(body.get("new_password") or "")) + self.hub.pipeline.audit(user["username"], "change_password", "hub_admin", "") + self._json({"ok": True}, HTTPStatus.OK) + return + if user["must_change"] and path not in {"/admin/api/change-password", "/admin/api/session"}: + raise ApiError("UNAUTHORIZED", "请先修改初始密码") + if path == "/admin/api/overview" and method == "GET": + self._json(self.hub.admin.overview(), HTTPStatus.OK) + return + if path == "/admin/api/sources" and method == "GET": + self._json(self.hub.admin.sources(), HTTPStatus.OK) + return + if path.startswith("/admin/api/sources/") and path.endswith("/probe") and method == "POST": + provider = path.split("/")[4] + self._json(self.hub.admin.probe(provider), HTTPStatus.OK) + return + if path == "/admin/api/jobs" and method == "GET": + self._json(self.hub.admin.jobs(), HTTPStatus.OK) + return + if path.startswith("/admin/api/jobs/") and path.endswith("/run") and method == "POST": + job_id = path.split("/")[4] + body = self._read_json(allow_empty=True) + self._json(self.hub.admin.run_job(job_id, str(body.get("trade_date") or "")), HTTPStatus.OK) + return + if path == "/admin/api/batches" and method == "GET": + query = parse_query(urlparse(self.path).query) + date = (query.get("date") or [""])[0] + dataset = (query.get("dataset") or [""])[0] + self._json(self.hub.admin.batches(date, dataset), HTTPStatus.OK) + return + if path == "/admin/api/datasets" and method == "GET": + query = parse_query(urlparse(self.path).query) + self._json(self.hub.admin.datasets((query.get("date") or [""])[0]), HTTPStatus.OK) + return + if path == "/admin/api/audit" and method == "GET": + self._json(self.hub.admin.audit(), HTTPStatus.OK) + return + if path == "/admin/api/rollback" and method == "POST": + body = self._read_json() + result = self.hub.admin.rollback( + str(body.get("dataset") or ""), + str(body.get("trade_date") or ""), + str(body.get("password") or ""), + str(body.get("confirm") or ""), + user["username"], + ) + self._json(result, HTTPStatus.OK) + return + if path == "/admin/api/backfill" and method == "POST": + body = self._read_json() + result = self.hub.admin.backfill( + str(body.get("dataset") or ""), + str(body.get("trade_date") or ""), + str(body.get("password") or ""), + str(body.get("confirm") or ""), + user["username"], + ) + self._json(result, HTTPStatus.OK) + return + raise ApiError("INVALID_ARGUMENT", f"unknown admin endpoint: {path}") + + def _admin_static(self, path: str) -> None: + relative = path[len("/admin"):].lstrip("/") or "index.html" + candidate = (self.hub.static_dir / relative).resolve() + try: + candidate.relative_to(self.hub.static_dir.resolve()) + except ValueError: + self.send_error(HTTPStatus.FORBIDDEN) + return + if candidate.is_dir(): + candidate = candidate / "index.html" + if not candidate.is_file(): + candidate = self.hub.static_dir / "index.html" + content = candidate.read_bytes() + content_type = mimetypes.guess_type(candidate.name)[0] or "application/octet-stream" + if content_type.startswith("text/") or content_type in {"application/javascript", "application/json"}: + content_type += "; charset=utf-8" + self.send_response(HTTPStatus.OK) + self.send_header("Content-Type", content_type) + self.send_header("Content-Length", str(len(content))) + self.send_header("Cache-Control", "no-cache") + self.end_headers() + self.wfile.write(content) + + def _read_json(self, allow_empty: bool = False) -> dict[str, Any]: + length = int(self.headers.get("Content-Length", "0") or 0) + if length == 0 and allow_empty: + return {} + if length <= 0 or length > 65536: + raise ValueError("请求内容为空或过大") + return json.loads(self.rfile.read(length).decode("utf-8")) + + def _cookie_value(self, name: str) -> str: + cookie = SimpleCookie() + try: + cookie.load(self.headers.get("Cookie", "")) + except Exception: + return "" + morsel = cookie.get(name) + return morsel.value if morsel else "" + + def _cookie(self, value: str, clear: bool = False) -> str: + max_age = 0 if clear else 12 * 3600 + return f"{SESSION_COOKIE}={value}; Path=/; HttpOnly; SameSite=Strict; Max-Age={max_age}" + + def _json(self, payload: dict[str, Any], status: HTTPStatus, extra_headers: list[str] | None = None) -> None: + raw = json.dumps(payload, ensure_ascii=False).encode("utf-8") + self.send_response(status) + self.send_header("Content-Type", "application/json; charset=utf-8") + self.send_header("Content-Length", str(len(raw))) + self.send_header("Cache-Control", "no-store") + for header in extra_headers or []: + self.send_header("Set-Cookie", header) + self.end_headers() + self.wfile.write(raw) + + +def make_handler(hub: Hub) -> type[HubRequestHandler]: + class BoundHandler(HubRequestHandler): + pass + + BoundHandler.hub = hub + BoundHandler.protocol_version = "HTTP/1.1" + return BoundHandler + + +def serve(hub: Hub, host: str, port: int) -> None: + configure_logging(hub.settings.log_level) + handler = make_handler(hub) + server = ThreadingHTTPServer((host, port), handler) + hub.start() + LOGGER.info("xiaobai-datahub listening", extra={"hub": {"host": host, "port": port}}) + print(f"xiaobai-datahub is running at http://{host}:{port}/admin/") + try: + server.serve_forever() + except KeyboardInterrupt: + pass + finally: + hub.stop() + server.server_close() diff --git a/xiaobai-datahub/datahub/hub.py b/xiaobai-datahub/datahub/hub.py new file mode 100644 index 0000000..fbd25c8 --- /dev/null +++ b/xiaobai-datahub/datahub/hub.py @@ -0,0 +1,54 @@ +from __future__ import annotations + +from pathlib import Path + +from datahub.adapters.tushare import TushareAdapter +from datahub.admin_api import AdminAPI +from datahub.auth import AuthService +from datahub.crypto import SecretVault +from datahub.db import HubDB +from datahub.governance.circuit import CircuitBreaker +from datahub.governance.lkg import LastKnownGood +from datahub.governance.ratelimit import TokenBucket +from datahub.pipeline import Pipeline +from datahub.scheduler import Scheduler +from datahub.serving import V1API +from datahub.settings import Settings, load_settings + + +class Hub: + def __init__(self, settings: Settings, adapter: TushareAdapter | None = None) -> None: + if not settings.encryption_key: + raise SystemExit("DATAHUB_ENCRYPTION_KEY 未配置") + self.settings = settings + self.db = HubDB(settings.db_path) + self.vault = SecretVault(settings.encryption_key) + self.auth = AuthService(self.db, self.vault, settings.api_token, settings.admin_password) + token = settings.tushare_token or self.auth.load_credential("tushare_token") + if settings.tushare_token: + self.auth.store_credential("tushare_token", settings.tushare_token) + token = settings.tushare_token + self.adapter = adapter or TushareAdapter(token) + self.pipeline = Pipeline( + self.db, + self.adapter, + settings, + bucket=TokenBucket(settings.tushare_rate_per_minute), + breaker=CircuitBreaker(), + ) + self.lkg = LastKnownGood(self.db) + self.scheduler = Scheduler(self.db, self.pipeline) + self.api = V1API(self.db, self.pipeline, settings) + self.admin = AdminAPI(self.db, self.pipeline, self.scheduler, self.auth) + self.static_dir = Path(__file__).resolve().parents[1] / "admin" + + def start(self) -> None: + if self.settings.scheduler_enabled: + self.scheduler.start() + + def stop(self) -> None: + self.scheduler.stop() + + +def build_hub(settings: Settings | None = None) -> Hub: + return Hub(settings or load_settings()) diff --git a/xiaobai-datahub/datahub/logutil.py b/xiaobai-datahub/datahub/logutil.py new file mode 100644 index 0000000..8752c62 --- /dev/null +++ b/xiaobai-datahub/datahub/logutil.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +import json +import logging +import sys +from typing import Any + +from datahub.timeutil import isoformat + +_SECRET_KEYS = ( + "token", "password", "secret", "key", "authorization", "credential", + "tushare_token", "datahub_token", "encryption_key", "cookie", +) + + +def _redact(value: Any, key: str = "") -> Any: + lowered = key.lower() + if any(part in lowered for part in _SECRET_KEYS): + return "***" + if isinstance(value, dict): + return {str(item_key): _redact(item_value, str(item_key)) for item_key, item_value in value.items()} + if isinstance(value, list): + return [_redact(item) for item in value] + return value + + +class JsonFormatter(logging.Formatter): + def format(self, record: logging.LogRecord) -> str: + payload: dict[str, Any] = { + "ts": isoformat(), + "level": record.levelname, + "logger": record.name, + "message": record.getMessage(), + } + extra = getattr(record, "hub", None) + if isinstance(extra, dict): + payload.update(_redact(extra)) + if record.exc_info: + payload["exc"] = self.formatException(record.exc_info) + return json.dumps(payload, ensure_ascii=False, default=str) + + +def configure_logging(level: str = "INFO") -> logging.Logger: + logger = logging.getLogger("datahub") + if logger.handlers: + return logger + handler = logging.StreamHandler(sys.stdout) + handler.setFormatter(JsonFormatter()) + logger.addHandler(handler) + logger.setLevel(getattr(logging, level.upper(), logging.INFO)) + logger.propagate = False + return logger + + +def get_logger() -> logging.Logger: + return logging.getLogger("datahub") diff --git a/xiaobai-datahub/datahub/normalize.py b/xiaobai-datahub/datahub/normalize.py new file mode 100644 index 0000000..bf4c5f4 --- /dev/null +++ b/xiaobai-datahub/datahub/normalize.py @@ -0,0 +1,201 @@ +"""Canonical field normalization for Tushare-native rows. + +Units (architecture §7.1): +- price: 4 decimal REAL +- pct_chg: percent, 4 decimal REAL +- volume: shares (Tushare daily/index vol is 手 → ×100) +- amount: yuan (Tushare daily/index amount is 千元 → ×1000) +- moneyflow amounts: yuan (Tushare is 万元 → ×1e4) +- daily_basic total_mv / circ_mv: yuan (Tushare is 万元 → ×1e4) +- stk_auction.amount is already yuan in Tushare; volume 手 → ×100 + +Existing xiaobai-review stores Tushare native units and converts at display time. +Hub converts once at ingest. Golden tests compare hub output against applying +these same factors to review-native rows. +""" + +from __future__ import annotations + +from typing import Any + +from datahub.numbers import finite_number, round4 + +AMOUNT_THOUSAND_YUAN = 1000.0 +AMOUNT_WAN_YUAN = 10000.0 +VOLUME_LOT = 100.0 + +DAILY_FIELDS = ("ts_code", "trade_date", "open", "high", "low", "close", "pct_chg", "vol", "amount") +VALUATION_FIELDS = ( + "ts_code", "trade_date", "turnover_rate", "volume_ratio", + "total_mv", "circ_mv", "pe_ttm", "pb", "ps_ttm", "dv_ttm", +) +MONEYFLOW_FIELDS = ( + "ts_code", "trade_date", + "buy_sm_amount", "sell_sm_amount", "buy_md_amount", "sell_md_amount", + "buy_lg_amount", "sell_lg_amount", "buy_elg_amount", "sell_elg_amount", + "net_mf_amount", +) +AUCTION_FIELDS = ( + "ts_code", "trade_date", "vol", "price", "amount", "pre_close", + "turnover_rate", "volume_ratio", "float_share", +) +INDEX_FIELDS = ("ts_code", "trade_date", "open", "high", "low", "close", "pct_chg", "vol", "amount") +CALENDAR_FIELDS = ("exchange", "cal_date", "is_open", "pretrade_date") +STOCK_FIELDS = ("ts_code", "symbol", "name", "area", "industry", "market", "list_status", "list_date") + + +def _code(value: Any) -> str: + return str(value or "").strip().upper() + + +def _date(value: Any) -> str: + return str(value or "").replace("-", "")[:8] + + +def review_daily_to_canonical(row: dict[str, Any]) -> dict[str, Any]: + """Convert a review-stored daily row (Tushare native units) to hub canonical.""" + return normalize_daily(row) + + +def normalize_daily(row: dict[str, Any], adj_factor: float | None = None) -> dict[str, Any]: + return { + "ts_code": _code(row.get("ts_code")), + "trade_date": _date(row.get("trade_date")), + "open": round4(finite_number(row.get("open"))), + "high": round4(finite_number(row.get("high"))), + "low": round4(finite_number(row.get("low"))), + "close": round4(finite_number(row.get("close"))), + "pct_chg": round4(finite_number(row.get("pct_chg"))), + "volume": round4(_scale(row.get("vol"), VOLUME_LOT)), + "amount": round4(_scale(row.get("amount"), AMOUNT_THOUSAND_YUAN)), + "adj_factor": round4(finite_number(adj_factor if adj_factor is not None else row.get("adj_factor"))), + } + + +def normalize_valuation(row: dict[str, Any]) -> dict[str, Any]: + return { + "ts_code": _code(row.get("ts_code")), + "trade_date": _date(row.get("trade_date")), + "turnover_rate": round4(finite_number(row.get("turnover_rate"))), + "volume_ratio": round4(finite_number(row.get("volume_ratio"))), + "total_mv": round4(_scale(row.get("total_mv"), AMOUNT_WAN_YUAN)), + "circ_mv": round4(_scale(row.get("circ_mv"), AMOUNT_WAN_YUAN)), + "pe_ttm": round4(finite_number(row.get("pe_ttm"))), + "pb": round4(finite_number(row.get("pb"))), + "ps_ttm": round4(finite_number(row.get("ps_ttm"))), + "dv_ttm": round4(finite_number(row.get("dv_ttm"))), + } + + +def normalize_moneyflow(row: dict[str, Any]) -> dict[str, Any]: + converted = { + "ts_code": _code(row.get("ts_code")), + "trade_date": _date(row.get("trade_date")), + } + for field in MONEYFLOW_FIELDS[2:]: + converted[field] = round4(_scale(row.get(field), AMOUNT_WAN_YUAN)) + return converted + + +def normalize_auction(row: dict[str, Any]) -> dict[str, Any]: + return { + "ts_code": _code(row.get("ts_code")), + "trade_date": _date(row.get("trade_date")), + "volume": round4(_scale(row.get("vol") if row.get("vol") is not None else row.get("volume"), VOLUME_LOT)), + "price": round4(finite_number(row.get("price"))), + "amount": round4(finite_number(row.get("amount"))), + "pre_close": round4(finite_number(row.get("pre_close"))), + "turnover_rate": round4(finite_number(row.get("turnover_rate"))), + "volume_ratio": round4(finite_number(row.get("volume_ratio"))), + "float_share": round4(_scale(row.get("float_share"), AMOUNT_WAN_YUAN) if row.get("float_share") is not None else None), + } + + +def normalize_index_daily(row: dict[str, Any]) -> dict[str, Any]: + return { + "ts_code": _code(row.get("ts_code")), + "trade_date": _date(row.get("trade_date")), + "open": round4(finite_number(row.get("open"))), + "high": round4(finite_number(row.get("high"))), + "low": round4(finite_number(row.get("low"))), + "close": round4(finite_number(row.get("close"))), + "pct_chg": round4(finite_number(row.get("pct_chg"))), + "volume": round4(_scale(row.get("vol"), VOLUME_LOT)), + "amount": round4(_scale(row.get("amount"), AMOUNT_THOUSAND_YUAN)), + } + + +def normalize_calendar(row: dict[str, Any]) -> dict[str, Any]: + is_open = row.get("is_open") + if is_open in (True, "1", 1, "Y", "y"): + open_flag = 1 + elif is_open in (False, "0", 0, "N", "n", None, ""): + open_flag = 0 + else: + open_flag = int(is_open) + return { + "exchange": str(row.get("exchange") or "SSE"), + "cal_date": _date(row.get("cal_date") or row.get("calDate")), + "is_open": open_flag, + "pretrade_date": _date(row.get("pretrade_date")) or None, + } + + +def normalize_stock(row: dict[str, Any]) -> dict[str, Any]: + ts_code = _code(row.get("ts_code")) + symbol = str(row.get("symbol") or "").strip() or (ts_code.split(".")[0] if ts_code else "") + return { + "ts_code": ts_code, + "symbol": symbol, + "name": str(row.get("name") or "").strip(), + "area": str(row.get("area") or "").strip() or None, + "industry": str(row.get("industry") or "").strip() or None, + "market": str(row.get("market") or "").strip() or None, + "list_status": str(row.get("list_status") or "L").strip() or "L", + "list_date": _date(row.get("list_date")) or None, + } + + +def apply_qfq(price: float | None, factor: float | None, latest_factor: float | None) -> float | None: + if price is None: + return None + current = factor if factor not in (None, 0) else 1.0 + latest = latest_factor if latest_factor not in (None, 0) else current + return round4(price * current / latest) + + +def qfq_bar(row: dict[str, Any], latest_factor: float | None) -> dict[str, Any]: + factor = finite_number(row.get("adj_factor"), 1.0) or 1.0 + out = dict(row) + for field in ("open", "high", "low", "close"): + out[field] = apply_qfq(finite_number(row.get(field)), factor, latest_factor) + return out + + +NORMALIZERS = { + "daily": normalize_daily, + "valuation": normalize_valuation, + "daily_basic": normalize_valuation, + "moneyflow": normalize_moneyflow, + "auction": normalize_auction, + "stk_auction": normalize_auction, + "index_daily": normalize_index_daily, + "trade_cal": normalize_calendar, + "calendar": normalize_calendar, + "stock_basic": normalize_stock, + "stocks": normalize_stock, +} + + +def normalize_rows(dataset: str, rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + fn = NORMALIZERS.get(dataset) + if fn is None: + raise ValueError(f"unknown dataset: {dataset}") + return [fn(row) for row in rows] + + +def _scale(value: Any, factor: float) -> float | None: + number = finite_number(value) + if number is None: + return None + return number * factor diff --git a/xiaobai-datahub/datahub/numbers.py b/xiaobai-datahub/datahub/numbers.py new file mode 100644 index 0000000..509fdb8 --- /dev/null +++ b/xiaobai-datahub/datahub/numbers.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +import math +from typing import Any + + +def finite_number(value: Any, default: float | None = None) -> float | None: + """Return a finite float, or default (None means JSON null).""" + if value is None or value == "": + return default + try: + number = float(value) + except (TypeError, ValueError): + return default + if not math.isfinite(number): + return default + return number + + +def round4(value: float | None) -> float | None: + if value is None: + return None + return round(float(value), 4) diff --git a/xiaobai-datahub/datahub/pipeline.py b/xiaobai-datahub/datahub/pipeline.py new file mode 100644 index 0000000..356ac5f --- /dev/null +++ b/xiaobai-datahub/datahub/pipeline.py @@ -0,0 +1,478 @@ +from __future__ import annotations + +import json +import time +from collections.abc import Callable +from typing import Any + +from datahub.adapters.base import AdapterError +from datahub.adapters.tushare import DEFAULT_INDEX_CODES, TushareAdapter +from datahub.db import DATASET_TABLES, HubDB +from datahub.governance.circuit import CircuitBreaker +from datahub.governance.ratelimit import TokenBucket +from datahub.governance.retry import RetryError, retry_call +from datahub.logutil import get_logger +from datahub.normalize import finite_number, normalize_daily +from datahub.settings import Settings +from datahub.timeutil import add_days, isoformat, now_shanghai, yyyymmdd + +LOGGER = get_logger() + +HARD_DATASETS = {"daily", "valuation", "index_daily"} +SOFT_DATASETS = {"moneyflow", "auction"} + +STAGING_INSERT = { + "daily": ( + "INSERT INTO staging_bars(ts_code,trade_date,batch_id,open,high,low,close,pct_chg,volume,amount,adj_factor) " + "VALUES (?,?,?,?,?,?,?,?,?,?,?)", + lambda r, b: ( + r["ts_code"], r["trade_date"], b, r.get("open"), r.get("high"), r.get("low"), + r.get("close"), r.get("pct_chg"), r.get("volume"), r.get("amount"), r.get("adj_factor"), + ), + ), + "valuation": ( + "INSERT INTO staging_valuation(ts_code,trade_date,batch_id,turnover_rate,volume_ratio,total_mv,circ_mv,pe_ttm,pb,ps_ttm,dv_ttm) " + "VALUES (?,?,?,?,?,?,?,?,?,?,?)", + lambda r, b: ( + r["ts_code"], r["trade_date"], b, r.get("turnover_rate"), r.get("volume_ratio"), + r.get("total_mv"), r.get("circ_mv"), r.get("pe_ttm"), r.get("pb"), r.get("ps_ttm"), r.get("dv_ttm"), + ), + ), + "moneyflow": ( + "INSERT INTO staging_moneyflow(ts_code,trade_date,batch_id,buy_sm_amount,sell_sm_amount,buy_md_amount,sell_md_amount,buy_lg_amount,sell_lg_amount,buy_elg_amount,sell_elg_amount,net_mf_amount) " + "VALUES (?,?,?,?,?,?,?,?,?,?,?,?)", + lambda r, b: ( + r["ts_code"], r["trade_date"], b, + r.get("buy_sm_amount"), r.get("sell_sm_amount"), r.get("buy_md_amount"), r.get("sell_md_amount"), + r.get("buy_lg_amount"), r.get("sell_lg_amount"), r.get("buy_elg_amount"), r.get("sell_elg_amount"), + r.get("net_mf_amount"), + ), + ), + "auction": ( + "INSERT INTO staging_auction(ts_code,trade_date,batch_id,volume,price,amount,pre_close,turnover_rate,volume_ratio,float_share) " + "VALUES (?,?,?,?,?,?,?,?,?,?)", + lambda r, b: ( + r["ts_code"], r["trade_date"], b, r.get("volume"), r.get("price"), r.get("amount"), + r.get("pre_close"), r.get("turnover_rate"), r.get("volume_ratio"), r.get("float_share"), + ), + ), + "index_daily": ( + "INSERT INTO staging_index_bars(ts_code,trade_date,batch_id,open,high,low,close,pct_chg,volume,amount) " + "VALUES (?,?,?,?,?,?,?,?,?,?)", + lambda r, b: ( + r["ts_code"], r["trade_date"], b, r.get("open"), r.get("high"), r.get("low"), + r.get("close"), r.get("pct_chg"), r.get("volume"), r.get("amount"), + ), + ), +} + +EOD_COPY = { + "daily": ( + "INSERT OR REPLACE INTO eod_bars " + "SELECT ts_code,trade_date,open,high,low,close,pct_chg,volume,amount,adj_factor,batch_id " + "FROM staging_bars WHERE batch_id = ?" + ), + "valuation": ( + "INSERT OR REPLACE INTO eod_valuation " + "SELECT ts_code,trade_date,turnover_rate,volume_ratio,total_mv,circ_mv,pe_ttm,pb,ps_ttm,dv_ttm,batch_id " + "FROM staging_valuation WHERE batch_id = ?" + ), + "moneyflow": ( + "INSERT OR REPLACE INTO eod_moneyflow " + "SELECT ts_code,trade_date,buy_sm_amount,sell_sm_amount,buy_md_amount,sell_md_amount," + "buy_lg_amount,sell_lg_amount,buy_elg_amount,sell_elg_amount,net_mf_amount,batch_id " + "FROM staging_moneyflow WHERE batch_id = ?" + ), + "auction": ( + "INSERT OR REPLACE INTO eod_auction " + "SELECT ts_code,trade_date,volume,price,amount,pre_close,turnover_rate,volume_ratio,float_share,batch_id " + "FROM staging_auction WHERE batch_id = ?" + ), + "index_daily": ( + "INSERT OR REPLACE INTO eod_index_bars " + "SELECT ts_code,trade_date,open,high,low,close,pct_chg,volume,amount,batch_id " + "FROM staging_index_bars WHERE batch_id = ?" + ), +} + + +class QualityError(RuntimeError): + def __init__(self, message: str, report: dict[str, Any]) -> None: + super().__init__(message) + self.report = report + + +class Pipeline: + def __init__( + self, + db: HubDB, + adapter: TushareAdapter, + settings: Settings, + bucket: TokenBucket | None = None, + breaker: CircuitBreaker | None = None, + before_commit: Callable[[], None] | None = None, + clock=None, + ) -> None: + self.db = db + self.adapter = adapter + self.settings = settings + self.bucket = bucket or TokenBucket(settings.tushare_rate_per_minute) + self.breaker = breaker or CircuitBreaker() + self.before_commit = before_commit + self.clock = clock or now_shanghai + + def next_batch_id(self, dataset: str, trade_date: str) -> str: + row = self.db.fetchone( + "SELECT COUNT(*) AS n FROM batches WHERE dataset = ? AND trade_date = ?", + (dataset, trade_date), + ) + seq = int((row or {}).get("n") or 0) + 1 + return f"{trade_date}-{dataset}-{seq:03d}" + + def ingest_reference(self, trade_date: str | None = None) -> dict[str, Any]: + """Refresh trade calendar (window) and stock master. Not versioned by batch.""" + day = yyyymmdd(trade_date or self.clock()) + start = add_days(day, -400) + end = add_days(day, 30) + calendar = self.adapter.normalize( + "calendar", + self._guarded_fetch("calendar", {"exchange": "SSE", "start_date": start, "end_date": end}), + ) + stocks = self.adapter.normalize("stocks", self._guarded_fetch("stocks", {"list_status": "L"})) + fetched_at = isoformat(self.clock()) + with self.db.write() as connection: + for row in calendar: + connection.execute( + """ + INSERT INTO trade_calendar(exchange, cal_date, is_open, pretrade_date, fetched_at) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(exchange, cal_date) DO UPDATE SET + is_open=excluded.is_open, pretrade_date=excluded.pretrade_date, fetched_at=excluded.fetched_at + """, + (row["exchange"], row["cal_date"], row["is_open"], row.get("pretrade_date"), fetched_at), + ) + for row in stocks: + connection.execute( + """ + INSERT INTO stock_master(ts_code,symbol,name,area,industry,market,list_status,list_date,updated_at) + VALUES (?,?,?,?,?,?,?,?,?) + ON CONFLICT(ts_code) DO UPDATE SET + symbol=excluded.symbol, name=excluded.name, area=excluded.area, + industry=excluded.industry, market=excluded.market, + list_status=excluded.list_status, list_date=excluded.list_date, + updated_at=excluded.updated_at + """, + ( + row["ts_code"], row.get("symbol"), row.get("name"), row.get("area"), + row.get("industry"), row.get("market"), row.get("list_status"), + row.get("list_date"), fetched_at, + ), + ) + return {"calendar": len(calendar), "stocks": len(stocks), "trade_date": day} + + def run_dataset(self, dataset: str, trade_date: str, attempts: int | None = None) -> dict[str, Any]: + trade_date = yyyymmdd(trade_date) + batch_id = self.next_batch_id(dataset, trade_date) + max_attempts = attempts or self.settings.max_publish_attempts + self._set_batch(batch_id, dataset, trade_date, "scheduled", 0) + try: + self._set_batch(batch_id, dataset, trade_date, "fetching", 1) + rows = retry_call( + lambda: self._fetch_dataset(dataset, trade_date), + attempts=max_attempts, + base_delay=0.05, + sleeper=lambda _d: None if attempts == 1 else time.sleep(_d), + ) + self._stage(dataset, batch_id, rows) + self._set_batch(batch_id, dataset, trade_date, "staged", 1, rows_in=len(rows), rows_out=len(rows)) + self._set_batch(batch_id, dataset, trade_date, "validating", 1) + report = self.validate(dataset, batch_id, trade_date, rows) + if report["hard_fail"]: + self._set_batch( + batch_id, dataset, trade_date, "staged", 1, + rows_in=len(rows), rows_out=len(rows), + quality=report, error="; ".join(report["errors"]), + ) + raise QualityError("integrity gate failed", report) + self._set_batch(batch_id, dataset, trade_date, "deriving", 1, rows_in=len(rows), rows_out=len(rows), quality=report) + self._set_batch(batch_id, dataset, trade_date, "publishing", 1, rows_in=len(rows), rows_out=len(rows), quality=report) + state = "degraded" if report["soft_fail"] else "published" + self.publish(dataset, trade_date, batch_id, state=state) + self._set_batch( + batch_id, dataset, trade_date, "published", 1, + rows_in=len(rows), rows_out=len(rows), quality=report, finished=True, + ) + return {"batch_id": batch_id, "dataset": dataset, "trade_date": trade_date, "rows": len(rows), "state": state, "quality": report} + except RetryError as exc: + self._set_batch(batch_id, dataset, trade_date, "failed", max_attempts, error=str(exc), finished=True) + raise + except QualityError: + raise + except Exception as exc: + self._set_batch(batch_id, dataset, trade_date, "failed", 1, error=str(exc), finished=True) + raise + + def run_eod_batch_a(self, trade_date: str) -> dict[str, Any]: + results = {} + for dataset in ("daily", "valuation", "moneyflow", "auction"): + results[dataset] = self.run_dataset(dataset, trade_date) + return results + + def run_eod_batch_b(self, trade_date: str) -> dict[str, Any]: + return {"index_daily": self.run_dataset("index_daily", trade_date)} + + def validate(self, dataset: str, batch_id: str, trade_date: str, rows: list[dict[str, Any]]) -> dict[str, Any]: + quality = self.settings.quality + errors: list[str] = [] + warnings: list[str] = [] + listed = self.db.fetchone( + "SELECT COUNT(*) AS n FROM stock_master WHERE list_status = 'L'", + ) + listed_n = int((listed or {}).get("n") or 0) + row_n = len(rows) + keys = [(row.get("ts_code"), row.get("trade_date")) for row in rows] + dup = row_n - len(set(keys)) + if dup: + errors.append(f"duplicate keys: {dup}") + bad_date = sum(1 for row in rows if str(row.get("trade_date")) != trade_date) + if bad_date: + errors.append(f"date mismatch rows: {bad_date}") + ratio = (row_n / listed_n) if listed_n else 1.0 + if dataset == "daily" and listed_n and ratio < float(quality.get("daily_row_ratio") or 0.98): + errors.append(f"row ratio {ratio:.4f} < {quality.get('daily_row_ratio')}") + null_fields = ("open", "high", "low", "close", "amount") if dataset in {"daily", "index_daily"} else () + if null_fields and rows: + nulls = sum(1 for row in rows if any(row.get(field) is None for field in null_fields)) + null_rate = nulls / row_n + if null_rate >= float(quality.get("null_rate_max") or 0.01): + errors.append(f"null rate {null_rate:.4f}") + if dataset in SOFT_DATASETS and row_n == 0: + warnings.append("empty soft dataset") + hard_fail = bool(errors) and dataset in HARD_DATASETS.union({"daily", "valuation", "index_daily"}) + if dataset in SOFT_DATASETS: + hard_fail = bool(dup or bad_date) + return { + "rows": row_n, + "listed": listed_n, + "ratio": round(ratio, 4), + "errors": errors, + "warnings": warnings, + "hard_fail": hard_fail, + "soft_fail": bool(warnings) and not hard_fail, + "batch_id": batch_id, + } + + def publish(self, dataset: str, trade_date: str, batch_id: str, state: str = "published") -> None: + copy_sql = EOD_COPY[dataset] + published_at = isoformat(self.clock()) + with self.db.write() as connection: + current = connection.execute( + "SELECT active_batch FROM publications WHERE dataset = ? AND trade_date = ?", + (dataset, trade_date), + ).fetchone() + prev = str(current["active_batch"]) if current else None + connection.execute(copy_sql, (batch_id,)) + if self.before_commit: + self.before_commit() + connection.execute( + """ + INSERT INTO publications(dataset, trade_date, active_batch, prev_batch, state, published_at) + VALUES (?, ?, ?, ?, ?, ?) + ON CONFLICT(dataset, trade_date) DO UPDATE SET + prev_batch=excluded.prev_batch, + active_batch=excluded.active_batch, + state=excluded.state, + published_at=excluded.published_at + """, + (dataset, trade_date, batch_id, prev, state, published_at), + ) + max_gen = connection.execute( + "SELECT COALESCE(MAX(generation), 0) AS g FROM publication_history WHERE dataset = ? AND trade_date = ?", + (dataset, trade_date), + ).fetchone() + generation = int(max_gen["g"]) + 1 + connection.execute( + "INSERT OR REPLACE INTO publication_history(dataset, trade_date, batch_id, published_at, generation) VALUES (?,?,?,?,?)", + (dataset, trade_date, batch_id, published_at, generation), + ) + keep = int(self.settings.quality.get("publication_generations") or 3) + stale = connection.execute( + """ + SELECT batch_id FROM publication_history + WHERE dataset = ? AND trade_date = ? + ORDER BY generation DESC + """, + (dataset, trade_date), + ).fetchall() + for row in stale[keep:]: + connection.execute( + "DELETE FROM publication_history WHERE dataset = ? AND trade_date = ? AND batch_id = ?", + (dataset, trade_date, row["batch_id"]), + ) + + def rollback(self, dataset: str, trade_date: str, actor: str = "admin") -> dict[str, Any]: + trade_date = yyyymmdd(trade_date) + pub = self.db.fetchone( + "SELECT * FROM publications WHERE dataset = ? AND trade_date = ?", + (dataset, trade_date), + ) + if not pub or not pub.get("prev_batch"): + raise ValueError("没有可回滚的上一批次") + target = pub["prev_batch"] + published_at = isoformat(self.clock()) + with self.db.write() as connection: + connection.execute( + """ + UPDATE publications + SET prev_batch = active_batch, active_batch = ?, published_at = ?, state = 'published' + WHERE dataset = ? AND trade_date = ? + """, + (target, published_at, dataset, trade_date), + ) + self.audit(actor, "rollback", f"{dataset}:{trade_date}", json.dumps({"to": target, "from": pub["active_batch"]})) + return {"dataset": dataset, "trade_date": trade_date, "active_batch": target, "prev_batch": pub["active_batch"]} + + def active_batch(self, dataset: str, trade_date: str) -> str | None: + row = self.db.fetchone( + "SELECT active_batch FROM publications WHERE dataset = ? AND trade_date = ?", + (dataset, trade_date), + ) + return str(row["active_batch"]) if row else None + + def cleanup(self) -> dict[str, int]: + staging_days = int(self.settings.quality.get("staging_retain_days") or 14) + job_days = int(self.settings.quality.get("job_run_retain_days") or 90) + cutoff_staging = add_days(yyyymmdd(self.clock()), -staging_days) + cutoff_jobs = add_days(yyyymmdd(self.clock()), -job_days) + deleted = 0 + with self.db.write() as connection: + for dataset, (_eod, staging) in DATASET_TABLES.items(): + cur = connection.execute( + f"DELETE FROM {staging} WHERE trade_date < ?", + (cutoff_staging,), + ) + deleted += cur.rowcount + connection.execute("DELETE FROM job_runs WHERE started_at < ?", (cutoff_jobs,)) + connection.execute("DELETE FROM src_calls WHERE created_at < ?", (cutoff_jobs,)) + return {"staging_deleted": deleted} + + def audit(self, actor: str, action: str, target: str = "", detail: str = "") -> None: + self.db.execute( + "INSERT INTO audit_log(actor, action, target, detail, created_at) VALUES (?,?,?,?,?)", + (actor, action, target, detail, isoformat(self.clock())), + ) + + def _fetch_dataset(self, dataset: str, trade_date: str) -> list[dict[str, Any]]: + if dataset == "daily": + raw = self._guarded_fetch("daily", {"trade_date": trade_date}) + factors = { + (row["ts_code"], row["trade_date"]): finite_number(row.get("adj_factor")) + for row in self._guarded_fetch("adj_factor", {"trade_date": trade_date}) + } + return [ + normalize_daily(row, adj_factor=factors.get((str(row.get("ts_code") or "").upper(), str(row.get("trade_date") or "")))) + for row in raw + ] + if dataset == "index_daily": + rows: list[dict[str, Any]] = [] + for ts_code in DEFAULT_INDEX_CODES: + raw = self._guarded_fetch("index_daily", {"ts_code": ts_code, "trade_date": trade_date}) + rows.extend(self.adapter.normalize("index_daily", raw)) + return rows + api_dataset = dataset + raw = self._guarded_fetch(api_dataset, {"trade_date": trade_date}) + return self.adapter.normalize(api_dataset, raw) + + def _guarded_fetch(self, dataset: str, params: dict[str, Any]) -> list[dict[str, Any]]: + if not self.breaker.allow(): + raise AdapterError("Tushare circuit open") + self.bucket.acquire() + started = time.perf_counter() + try: + # For daily we want RAW tushare rows so adj_factor can be merged later. + rows = self.adapter.fetch(dataset, params) + latency = round((time.perf_counter() - started) * 1000) + self.breaker.record_success() + self._log_call(dataset, True, latency, "") + self._persist_health("ok") + return rows + except Exception as exc: + latency = round((time.perf_counter() - started) * 1000) + self.breaker.record_failure(str(exc)) + self._log_call(dataset, False, latency, str(exc)) + self._persist_health("error", str(exc)) + raise + + def _stage(self, dataset: str, batch_id: str, rows: list[dict[str, Any]]) -> None: + sql, mapper = STAGING_INSERT[dataset] + with self.db.write() as connection: + connection.execute( + f"DELETE FROM {DATASET_TABLES[dataset][1]} WHERE batch_id = ?", + (batch_id,), + ) + connection.executemany(sql, [mapper(row, batch_id) for row in rows]) + + def _set_batch( + self, + batch_id: str, + dataset: str, + trade_date: str, + state: str, + attempt: int, + rows_in: int | None = None, + rows_out: int | None = None, + quality: dict[str, Any] | None = None, + error: str | None = None, + finished: bool = False, + ) -> None: + now = isoformat(self.clock()) + existing = self.db.fetchone("SELECT batch_id FROM batches WHERE batch_id = ?", (batch_id,)) + payload = json.dumps(quality, ensure_ascii=False) if quality else None + with self.db.write() as connection: + if existing is None: + connection.execute( + """ + INSERT INTO batches(batch_id, dataset, trade_date, state, attempt, rows_in, rows_out, quality_json, started_at, finished_at, error) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + (batch_id, dataset, trade_date, state, attempt, rows_in, rows_out, payload, now, now if finished else None, error), + ) + else: + connection.execute( + """ + UPDATE batches SET state=?, attempt=?, + rows_in=COALESCE(?, rows_in), rows_out=COALESCE(?, rows_out), + quality_json=COALESCE(?, quality_json), + finished_at=CASE WHEN ? THEN ? ELSE finished_at END, + error=COALESCE(?, error) + WHERE batch_id = ? + """, + (state, attempt, rows_in, rows_out, payload, 1 if finished else 0, now, error, batch_id), + ) + + def _log_call(self, endpoint: str, ok: bool, latency_ms: int, error: str) -> None: + self.db.execute( + "INSERT INTO src_calls(provider, endpoint, ok, latency_ms, error, created_at) VALUES (?,?,?,?,?,?)", + ("tushare", endpoint, 1 if ok else 0, latency_ms, error, isoformat(self.clock())), + ) + + def _persist_health(self, state: str, error: str = "") -> None: + snap = self.breaker.snapshot() + self.db.execute( + """ + INSERT INTO src_health(provider, endpoint_class, state, last_ok_at, last_error, consec_failures, opened_at, cooldown_until) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(provider, endpoint_class) DO UPDATE SET + state=excluded.state, last_ok_at=excluded.last_ok_at, last_error=excluded.last_error, + consec_failures=excluded.consec_failures, opened_at=excluded.opened_at, cooldown_until=excluded.cooldown_until + """, + ( + "tushare", "pro", + snap.state, + isoformat(self.clock()) if state == "ok" else None, + error or snap.last_error, + snap.consec_failures, + isoformat(self.clock()) if snap.state == "open" else None, + None, + ), + ) diff --git a/xiaobai-datahub/datahub/scheduler.py b/xiaobai-datahub/datahub/scheduler.py new file mode 100644 index 0000000..dbefbb1 --- /dev/null +++ b/xiaobai-datahub/datahub/scheduler.py @@ -0,0 +1,145 @@ +from __future__ import annotations + +import threading +from collections.abc import Callable +from datetime import datetime, time +from typing import Any + +from datahub.db import HubDB +from datahub.logutil import get_logger +from datahub.pipeline import Pipeline +from datahub.timeutil import isoformat, now_shanghai, yyyymmdd + +LOGGER = get_logger() + +JobFn = Callable[[str], Any] + + +def is_open_day(db: HubDB, day: str) -> bool: + row = db.fetchone( + "SELECT is_open FROM trade_calendar WHERE exchange = 'SSE' AND cal_date = ?", + (day,), + ) + if row is None: + return True # unknown calendar: do not skip reference refresh + return int(row["is_open"]) == 1 + + +class Scheduler: + """Calendar-driven in-process scheduler. Non-trading days skip EOD fetches.""" + + def __init__(self, db: HubDB, pipeline: Pipeline, jobs: dict[str, JobFn] | None = None) -> None: + self.db = db + self.pipeline = pipeline + self.jobs = jobs or { + "precheck": self._precheck, + "eod_a": self._eod_a, + "eod_b": self._eod_b, + "cleanup": self._cleanup, + "backup": self._backup, + } + self._stop = threading.Event() + self._thread: threading.Thread | None = None + self._fired: set[tuple[str, str, str]] = set() + + def start(self, interval_seconds: float = 30.0) -> None: + if self._thread and self._thread.is_alive(): + return + + def loop() -> None: + while not self._stop.wait(interval_seconds): + try: + self.tick() + except Exception: + LOGGER.exception("scheduler tick failed") + + self._thread = threading.Thread(target=loop, name="datahub-scheduler", daemon=True) + self._thread.start() + + def stop(self, timeout: float = 5.0) -> None: + self._stop.set() + if self._thread and self._thread is not threading.current_thread(): + self._thread.join(timeout) + + def tick(self, clock: datetime | None = None) -> list[str]: + now = clock or now_shanghai() + day = yyyymmdd(now) + current = now.timetz() if False else now.time() + ran: list[str] = [] + plan = [ + ("precheck", time(8, 45)), + ("eod_a", time(15, 5)), + ("eod_b", time(15, 10)), + ("cleanup", time(0, 30)), + ("backup", time(0, 40)), + ] + open_day = is_open_day(self.db, day) + for job_id, at in plan: + if current < at: + continue + 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: + self._fired.add(key) + continue + self._fired.add(key) + self.run_job(job_id, day) + ran.append(job_id) + return ran + + 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) + started = isoformat() + run_id = None + with self.db.write() as connection: + cur = connection.execute( + "INSERT INTO job_runs(job_id, state, started_at, attempt) VALUES (?,?,?,1)", + (job_id, "running", started), + ) + run_id = cur.lastrowid + try: + result = fn(trade_date) or {} + with self.db.write() as connection: + connection.execute( + "UPDATE job_runs SET state=?, finished_at=?, rows_out=?, detail=? WHERE id=?", + ("ok", isoformat(), result.get("rows") if isinstance(result, dict) else None, str(result)[:2000], run_id), + ) + return {"job_id": job_id, "result": result, "state": "ok"} + except Exception as exc: + with self.db.write() as connection: + connection.execute( + "UPDATE job_runs SET state=?, finished_at=?, error=? WHERE id=?", + ("failed", isoformat(), str(exc), run_id), + ) + raise + + def _precheck(self, trade_date: str) -> dict[str, Any]: + return self.pipeline.ingest_reference(trade_date) + + def _eod_a(self, trade_date: str) -> dict[str, Any]: + return self.pipeline.run_eod_batch_a(trade_date) + + def _eod_b(self, trade_date: str) -> dict[str, Any]: + return self.pipeline.run_eod_batch_b(trade_date) + + def _cleanup(self, trade_date: str) -> dict[str, Any]: + result = self.pipeline.cleanup() + if now_shanghai().weekday() == 6: + self.pipeline.db.vacuum() + result["vacuum"] = True + return result + + def _backup(self, trade_date: str) -> dict[str, Any]: + from pathlib import Path + + dest_dir = Path(self.pipeline.settings.backup_dir) + dest = dest_dir / f"datahub-{trade_date}.db" + self.pipeline.db.backup_to(dest) + keep = int(self.pipeline.settings.quality.get("backup_retain") or 14) + backups = sorted(dest_dir.glob("datahub-*.db")) + for old in backups[:-keep]: + old.unlink(missing_ok=True) + return {"path": str(dest.name), "kept": min(len(backups), keep)} diff --git a/xiaobai-datahub/datahub/serving.py b/xiaobai-datahub/datahub/serving.py new file mode 100644 index 0000000..4c72374 --- /dev/null +++ b/xiaobai-datahub/datahub/serving.py @@ -0,0 +1,376 @@ +from __future__ import annotations + +from http import HTTPStatus +from typing import Any +from urllib.parse import parse_qs + +from datahub import SCHEMA_VERSION +from datahub.codes import resolve_code +from datahub.db import HubDB +from datahub.normalize import qfq_bar +from datahub.numbers import finite_number +from datahub.pipeline import Pipeline +from datahub.settings import Settings +from datahub.timeutil import isoformat, now_shanghai, session_phase, yyyymmdd + +ERROR_STATUS = { + "UNAUTHORIZED": HTTPStatus.UNAUTHORIZED, + "INVALID_ARGUMENT": HTTPStatus.BAD_REQUEST, + "RATE_LIMITED": HTTPStatus.TOO_MANY_REQUESTS, + "SOURCE_UNAVAILABLE": HTTPStatus.SERVICE_UNAVAILABLE, + "DATASET_NOT_PUBLISHED": HTTPStatus.NOT_FOUND, + "STALE_DATA": HTTPStatus.OK, + "INTERNAL": HTTPStatus.INTERNAL_SERVER_ERROR, +} + + +class ApiError(Exception): + def __init__(self, code: str, message: str, retry_after: int | None = None, extra: dict[str, Any] | None = None) -> None: + super().__init__(message) + self.code = code + self.message = message + self.retry_after = retry_after + self.extra = extra or {} + + def payload(self) -> dict[str, Any]: + body: dict[str, Any] = {"code": self.code, "message": self.message} + if self.retry_after is not None: + body["retry_after"] = self.retry_after + body.update(self.extra) + return {"error": body} + + @property + def status(self) -> HTTPStatus: + return ERROR_STATUS.get(self.code, HTTPStatus.INTERNAL_SERVER_ERROR) + + +def envelope(data: Any, meta: dict[str, Any]) -> dict[str, Any]: + return {"schema_version": SCHEMA_VERSION, "data": data, "meta": meta} + + +class V1API: + def __init__(self, db: HubDB, pipeline: Pipeline, settings: Settings) -> None: + self.db = db + self.pipeline = pipeline + self.settings = settings + + def handle(self, path: str, query: dict[str, list[str]]) -> dict[str, Any]: + q = {key: values[-1] if values else "" for key, values in query.items()} + if path == "/v1/health": + return self.health() + if path == "/v1/calendar": + return self.calendar(q.get("from") or "", q.get("to") or "") + if path == "/v1/stocks": + return self.stocks(q.get("updated_since") or "", q) + if path == "/v1/bars/daily": + return self.daily_bars(q) + if path == "/v1/indexes/bars": + return self.index_bars(q) + if path == "/v1/valuation": + return self.valuation(q) + if path == "/v1/moneyflow": + return self.moneyflow(q) + if path == "/v1/auction": + return self.auction(q) + if path == "/v1/datasets/status": + return self.dataset_status(q.get("date") or "") + if path == "/v1/batches": + return self.batches(q.get("date") or "", q.get("dataset") or "") + raise ApiError("INVALID_ARGUMENT", f"unknown endpoint: {path}") + + def health(self) -> dict[str, Any]: + today = yyyymmdd(now_shanghai()) + cal = self.db.fetchone( + "SELECT is_open FROM trade_calendar WHERE exchange = 'SSE' AND cal_date = ?", + (today,), + ) + is_open = bool(cal and cal["is_open"] == 1) + sources = self.db.fetchall("SELECT * FROM src_health") + return envelope( + { + "status": "ok", + "session_phase": session_phase(now_shanghai(), is_open), + "trade_date": today, + "is_open_day": is_open, + "sources": [ + { + "provider": row["provider"], + "endpoint_class": row["endpoint_class"], + "state": row["state"], + "last_ok_at": row["last_ok_at"], + "consec_failures": row["consec_failures"], + } + for row in sources + ], + }, + {"tier": "official", "trade_date": today, "source": "datahub", "stale": False, "staleness_seconds": 0}, + ) + + def calendar(self, start: str, end: str) -> dict[str, Any]: + start = yyyymmdd(start or add_default(-30)) + end = yyyymmdd(end or add_default(5)) + rows = self.db.fetchall( + """ + SELECT cal_date, is_open, pretrade_date, + (SELECT MAX(cal_date) FROM trade_calendar t2 + WHERE t2.exchange = 'SSE' AND t2.is_open = 1 AND t2.cal_date < t1.cal_date) AS prev_open + FROM trade_calendar t1 + WHERE exchange = 'SSE' AND cal_date >= ? AND cal_date <= ? + ORDER BY cal_date + """, + (start, end), + ) + items = [ + { + "cal_date": row["cal_date"], + "is_open": bool(row["is_open"]), + "pretrade_date": row["pretrade_date"], + "prev_open": row["prev_open"], + } + for row in rows + ] + return envelope(items, self._official_meta("calendar", end if items else start, source="tushare:trade_cal")) + + def stocks(self, updated_since: str, q: dict[str, str]) -> dict[str, Any]: + limit, offset = self._page(q) + if updated_since: + rows = self.db.fetchall( + "SELECT * FROM stock_master WHERE updated_at >= ? ORDER BY ts_code LIMIT ? OFFSET ?", + (updated_since, limit, offset), + ) + else: + rows = self.db.fetchall( + "SELECT * FROM stock_master ORDER BY ts_code LIMIT ? OFFSET ?", + (limit, offset), + ) + return envelope(rows, self._official_meta("stocks", yyyymmdd(), source="tushare:stock_basic")) + + def daily_bars(self, q: dict[str, str]) -> dict[str, Any]: + return self._published_rows( + dataset="daily", + table="eod_bars", + q=q, + source="tushare:daily", + adjust=q.get("adjust") or "none", + ) + + def index_bars(self, q: dict[str, str]) -> dict[str, Any]: + return self._published_rows( + dataset="index_daily", + table="eod_index_bars", + q=q, + source="tushare:index_daily", + default_code="000001.SH", + ) + + def valuation(self, q: dict[str, str]) -> dict[str, Any]: + return self._published_rows(dataset="valuation", table="eod_valuation", q=q, source="tushare:daily_basic") + + def moneyflow(self, q: dict[str, str]) -> dict[str, Any]: + return self._published_rows(dataset="moneyflow", table="eod_moneyflow", q=q, source="tushare:moneyflow") + + def auction(self, q: dict[str, str]) -> dict[str, Any]: + return self._published_rows(dataset="auction", table="eod_auction", q=q, source="tushare:stk_auction") + + def dataset_status(self, date: str) -> dict[str, Any]: + trade_date = yyyymmdd(date or now_shanghai()) + datasets = ("daily", "valuation", "moneyflow", "auction", "index_daily") + items = [] + for dataset in datasets: + pub = self.db.fetchone( + "SELECT * FROM publications WHERE dataset = ? AND trade_date = ?", + (dataset, trade_date), + ) + batch = None + if pub: + batch = self.db.fetchone("SELECT * FROM batches WHERE batch_id = ?", (pub["active_batch"],)) + items.append( + { + "dataset": dataset, + "trade_date": trade_date, + "state": (pub or {}).get("state") or "unpublished", + "batch_id": (pub or {}).get("active_batch"), + "published_at": (pub or {}).get("published_at"), + "rows_out": (batch or {}).get("rows_out"), + "quality": _parse_json((batch or {}).get("quality_json")), + } + ) + return envelope(items, self._official_meta("status", trade_date, source="datahub")) + + def batches(self, date: str, dataset: str) -> dict[str, Any]: + trade_date = yyyymmdd(date or now_shanghai()) + if dataset: + rows = self.db.fetchall( + "SELECT * FROM batches WHERE trade_date = ? AND dataset = ? ORDER BY started_at", + (trade_date, dataset), + ) + else: + rows = self.db.fetchall( + "SELECT * FROM batches WHERE trade_date = ? ORDER BY started_at", + (trade_date,), + ) + return envelope(rows, self._official_meta("batches", trade_date, source="datahub")) + + def _published_rows( + self, + dataset: str, + table: str, + q: dict[str, str], + source: str, + adjust: str = "none", + default_code: str = "", + ) -> dict[str, Any]: + trade_date = q.get("date") or q.get("trade_date") or "" + code = q.get("code") or default_code + start = q.get("from") or "" + end = q.get("to") or "" + if trade_date: + trade_date = yyyymmdd(trade_date) + start = end = trade_date + if not start or not end: + if not trade_date: + raise ApiError("INVALID_ARGUMENT", "date or from/to is required") + else: + start = yyyymmdd(start) + end = yyyymmdd(end) + ts_code = "" + if code: + resolved = resolve_code(self.db, code) + if resolved is None: + raise ApiError("INVALID_ARGUMENT", f"ambiguous code: {code}") + ts_code = resolved + # For a range, use per-date published batch. Single-date is the common path. + if start == end: + pub = self.db.fetchone( + "SELECT * FROM publications WHERE dataset = ? AND trade_date = ?", + (dataset, start), + ) + if not pub: + raise ApiError( + "DATASET_NOT_PUBLISHED", + f"{dataset} {start} 尚未发布", + extra={"expected_at": "15:05+08:00"}, + ) + limit, offset = self._page(q) + sql = f"SELECT * FROM {table} WHERE trade_date = ? AND batch_id = ?" + params: list[Any] = [start, pub["active_batch"]] + if ts_code: + sql += " AND ts_code = ?" + params.append(ts_code) + sql += " ORDER BY ts_code LIMIT ? OFFSET ?" + params.extend([limit, offset]) + rows = [dict(row) for row in self.db.fetchall(sql, tuple(params))] + if adjust == "qfq" and dataset == "daily": + rows = self._apply_qfq(rows) + meta = { + "tier": "official", + "trade_date": start, + "published_at": pub["published_at"], + "source": source, + "batch_id": pub["active_batch"], + "stale": False, + "staleness_seconds": 0, + "state": pub["state"], + } + return envelope(rows, meta) + # 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} 尚未发布") + rows: list[dict[str, Any]] = [] + limit, offset = self._page(q) + for pub in pubs: + sql = f"SELECT * FROM {table} WHERE trade_date = ? AND batch_id = ?" + params = [pub["trade_date"], pub["active_batch"]] + if ts_code: + sql += " AND ts_code = ?" + params.append(ts_code) + sql += " ORDER BY ts_code" + rows.extend(self.db.fetchall(sql, tuple(params))) + sliced = rows[offset: offset + limit] + if adjust == "qfq" and dataset == "daily": + sliced = self._apply_qfq(sliced) + last = pubs[-1] + 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, + }, + ) + + def _apply_qfq(self, rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + by_code: dict[str, list[dict[str, Any]]] = {} + for row in rows: + by_code.setdefault(str(row["ts_code"]), []).append(row) + out: list[dict[str, Any]] = [] + for code, group in by_code.items(): + latest = None + factors = [finite_number(item.get("adj_factor")) for item in group] + factors = [item for item in factors if item] + if factors: + latest = max(factors) + else: + extra = self.db.fetchone( + "SELECT MAX(adj_factor) AS f FROM eod_bars WHERE ts_code = ?", + (code,), + ) + latest = finite_number((extra or {}).get("f"), 1.0) + out.extend(qfq_bar(item, latest) for item in group) + return out + + def _page(self, q: dict[str, str]) -> tuple[int, int]: + try: + limit = int(q.get("limit") or self.settings.list_limit_default) + offset = int(q.get("offset") or 0) + except ValueError as exc: + raise ApiError("INVALID_ARGUMENT", "limit/offset must be integers") from exc + limit = max(1, min(limit, self.settings.list_limit_max)) + offset = max(0, offset) + return limit, offset + + 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 = ?", + (dataset, trade_date), + ) + return { + "tier": "official", + "trade_date": trade_date, + "published_at": (pub or {}).get("published_at"), + "source": source, + "batch_id": (pub or {}).get("active_batch"), + "stale": False, + "staleness_seconds": 0, + } + + +def add_default(days: int) -> str: + from datetime import timedelta + + return (now_shanghai() + timedelta(days=days)).strftime("%Y%m%d") + + +def parse_query(raw: str) -> dict[str, list[str]]: + return parse_qs(raw, keep_blank_values=True) + + +def _parse_json(raw: Any) -> Any: + if not raw: + return None + if isinstance(raw, dict): + return raw + import json + + try: + return json.loads(str(raw)) + except json.JSONDecodeError: + return None diff --git a/xiaobai-datahub/datahub/settings.py b/xiaobai-datahub/datahub/settings.py new file mode 100644 index 0000000..f1c65ef --- /dev/null +++ b/xiaobai-datahub/datahub/settings.py @@ -0,0 +1,72 @@ +from __future__ import annotations + +import json +import os +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +ROOT = Path(__file__).resolve().parents[1] +DEFAULT_DB_PATH = Path(os.environ.get("DATAHUB_DB_PATH") or (ROOT / "data" / "datahub.db")) +DEFAULT_BACKUP_DIR = Path(os.environ.get("DATAHUB_BACKUP_DIR") or (ROOT / "data" / "backups")) +DEFAULT_CONFIG_PATH = ROOT / "config" / "hub-quality.config.json" + + +def _load_quality(path: Path) -> dict[str, Any]: + if not path.is_file(): + return {} + return json.loads(path.read_text(encoding="utf-8")) + + +@dataclass +class Settings: + host: str = "127.0.0.1" + port: int = 8766 + encryption_key: str = "" + api_token: str = "" + admin_password: str = "" + tushare_token: str = "" + db_path: Path = DEFAULT_DB_PATH + backup_dir: Path = DEFAULT_BACKUP_DIR + quality: dict[str, Any] = field(default_factory=dict) + log_level: str = "INFO" + scheduler_enabled: bool = True + + @property + def tushare_rate_per_minute(self) -> int: + return int(self.quality.get("tushare_rate_per_minute") or 300) + + @property + def max_publish_attempts(self) -> int: + return int(self.quality.get("max_publish_attempts") or 5) + + @property + def list_limit_default(self) -> int: + return int(self.quality.get("list_limit_default") or 5000) + + @property + def list_limit_max(self) -> int: + return int(self.quality.get("list_limit_max") or 5000) + + +def load_settings( + env: dict[str, str] | None = None, + config_path: Path | None = None, +) -> Settings: + environ = env if env is not None else dict(os.environ) + quality_path = config_path or DEFAULT_CONFIG_PATH + db_path = Path(environ.get("DATAHUB_DB_PATH") or DEFAULT_DB_PATH) + backup_dir = Path(environ.get("DATAHUB_BACKUP_DIR") or DEFAULT_BACKUP_DIR) + return Settings( + host=environ.get("DATAHUB_HOST") or "127.0.0.1", + port=int(environ.get("DATAHUB_PORT") or 8766), + encryption_key=str(environ.get("DATAHUB_ENCRYPTION_KEY") or "").strip(), + api_token=str(environ.get("DATAHUB_TOKEN") or "").strip(), + admin_password=str(environ.get("DATAHUB_ADMIN_PASSWORD") or "").strip(), + tushare_token=str(environ.get("TUSHARE_TOKEN") or "").strip(), + db_path=db_path, + backup_dir=backup_dir, + quality=_load_quality(quality_path), + log_level=environ.get("DATAHUB_LOG_LEVEL") or "INFO", + scheduler_enabled=str(environ.get("DATAHUB_SCHEDULER") or "1") not in {"0", "false", "False"}, + ) diff --git a/xiaobai-datahub/datahub/timeutil.py b/xiaobai-datahub/datahub/timeutil.py new file mode 100644 index 0000000..58faebb --- /dev/null +++ b/xiaobai-datahub/datahub/timeutil.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +from datetime import date, datetime, time, timedelta, timezone +from typing import Any +from zoneinfo import ZoneInfo + +SHANGHAI = ZoneInfo("Asia/Shanghai") + + +def now_shanghai(clock: datetime | None = None) -> datetime: + if clock is not None: + if clock.tzinfo is None: + return clock.replace(tzinfo=SHANGHAI) + return clock.astimezone(SHANGHAI) + return datetime.now(SHANGHAI) + + +def isoformat(value: datetime | None = None) -> str: + current = now_shanghai(value) + return current.isoformat(timespec="seconds") + + +def yyyymmdd(value: date | datetime | str | None = None) -> str: + if value is None: + return now_shanghai().strftime("%Y%m%d") + if isinstance(value, str): + digits = value.replace("-", "")[:8] + if len(digits) != 8 or not digits.isdigit(): + raise ValueError(f"invalid trade_date: {value}") + return digits + if isinstance(value, datetime): + return value.astimezone(SHANGHAI).strftime("%Y%m%d") + return value.strftime("%Y%m%d") + + +def parse_trade_date(value: str) -> date: + text = yyyymmdd(value) + return date(int(text[:4]), int(text[4:6]), int(text[6:8])) + + +def session_phase(clock: datetime | None, is_open_day: bool) -> str: + """pre | intradaily | lunch | eod | closed""" + if not is_open_day: + return "closed" + current = now_shanghai(clock).time() + if current < time(9, 15): + return "pre" + if current < time(11, 30) or (time(13, 0) <= current <= time(15, 5)): + return "intraday" + if current < time(13, 0): + return "lunch" + if current <= time(23, 40): + return "eod" + return "closed" + + +def add_days(trade_date: str, days: int) -> str: + return (parse_trade_date(trade_date) + timedelta(days=days)).strftime("%Y%m%d") + + +def utc_timestamp(value: Any) -> str: + if isinstance(value, datetime): + return isoformat(value) + return isoformat() diff --git a/xiaobai-datahub/requirements.txt b/xiaobai-datahub/requirements.txt new file mode 100644 index 0000000..a2b39f4 --- /dev/null +++ b/xiaobai-datahub/requirements.txt @@ -0,0 +1 @@ +cryptography==49.0.0 diff --git a/xiaobai-datahub/server.py b/xiaobai-datahub/server.py new file mode 100644 index 0000000..eb72034 --- /dev/null +++ b/xiaobai-datahub/server.py @@ -0,0 +1,27 @@ +"""xiaobai-datahub process entry.""" + +from __future__ import annotations + +import argparse + +from datahub.hub import build_hub +from datahub.httpapp import serve +from datahub.settings import load_settings + + +def main() -> None: + parser = argparse.ArgumentParser(description="xiaobai-datahub") + parser.add_argument("--host", default=None) + parser.add_argument("--port", type=int, default=None) + args = parser.parse_args() + settings = load_settings() + if args.host: + settings.host = args.host + if args.port: + settings.port = args.port + hub = build_hub(settings) + serve(hub, settings.host, settings.port) + + +if __name__ == "__main__": + main() diff --git a/xiaobai-datahub/tests/__init__.py b/xiaobai-datahub/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/xiaobai-datahub/tests/fixtures.py b/xiaobai-datahub/tests/fixtures.py new file mode 100644 index 0000000..ad2bacc --- /dev/null +++ b/xiaobai-datahub/tests/fixtures.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +TRADE_DATE = "20240902" + +RAW = { + "trade_cal": [ + {"exchange": "SSE", "cal_date": "20240902", "is_open": 1, "pretrade_date": "20240830"}, + {"exchange": "SSE", "cal_date": "20240903", "is_open": 1, "pretrade_date": "20240902"}, + {"exchange": "SSE", "cal_date": "20240907", "is_open": 0, "pretrade_date": "20240906"}, + ], + "stock_basic": [ + {"ts_code": "600000.SH", "symbol": "600000", "name": "浦发银行", "area": "上海", "industry": "银行", "market": "主板", "list_status": "L", "list_date": "19991110"}, + {"ts_code": "000001.SZ", "symbol": "000001", "name": "平安银行", "area": "深圳", "industry": "银行", "market": "主板", "list_status": "L", "list_date": "19910403"}, + ], + "daily": [ + {"ts_code": "600000.SH", "trade_date": "20240902", "open": 10.11, "high": 10.25, "low": 10.01, "close": 10.20, "pct_chg": 1.2345, "vol": 1000.0, "amount": 2000.0}, + {"ts_code": "000001.SZ", "trade_date": "20240902", "open": 11.00, "high": 11.20, "low": 10.90, "close": 11.10, "pct_chg": -0.5, "vol": 2000.0, "amount": 4000.0}, + ], + "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}, + ], + "adj_factor": [ + {"ts_code": "600000.SH", "trade_date": "20240902", "adj_factor": 1.1}, + {"ts_code": "000001.SZ", "trade_date": "20240902", "adj_factor": 2.0}, + ], + "index_daily": [ + {"ts_code": "000001.SH", "trade_date": "20240902", "open": 2700, "high": 2750, "low": 2690, "close": 2740, "pct_chg": 0.5, "vol": 3000.0, "amount": 500000.0}, + {"ts_code": "399001.SZ", "trade_date": "20240902", "open": 8000, "high": 8100, "low": 7900, "close": 8050, "pct_chg": 0.4, "vol": 2000.0, "amount": 300000.0}, + {"ts_code": "399006.SZ", "trade_date": "20240902", "open": 1600, "high": 1620, "low": 1580, "close": 1610, "pct_chg": 0.3, "vol": 1000.0, "amount": 100000.0}, + {"ts_code": "000300.SH", "trade_date": "20240902", "open": 3500, "high": 3550, "low": 3480, "close": 3520, "pct_chg": 0.2, "vol": 1500.0, "amount": 200000.0}, + ], + "moneyflow": [ + {"ts_code": "600000.SH", "trade_date": "20240902", "buy_sm_amount": 10, "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": "20240902", "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}, + ], + "stk_auction": [ + {"ts_code": "600000.SH", "trade_date": "20240902", "vol": 100, "price": 10.15, "amount": 1500000, "pre_close": 10.00, "turnover_rate": 0.1, "volume_ratio": 1.2, "float_share": 2000}, + {"ts_code": "000001.SZ", "trade_date": "20240902", "vol": 80, "price": 11.05, "amount": 1200000, "pre_close": 11.10, "turnover_rate": 0.2, "volume_ratio": 0.9, "float_share": 1800}, + ], +} + + +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] + 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 RAW["trade_cal"] if start <= row["cal_date"] <= end] + return list(RAW.get(api_name) or []) diff --git a/xiaobai-datahub/tests/test_admin.py b/xiaobai-datahub/tests/test_admin.py new file mode 100644 index 0000000..450411e --- /dev/null +++ b/xiaobai-datahub/tests/test_admin.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import json +import tempfile +import threading +import unittest +from http.server import ThreadingHTTPServer +from pathlib import Path +from urllib.request import Request, urlopen + +from datahub.adapters.tushare import TushareAdapter +from datahub.crypto import SecretVault +from datahub.httpapp import make_handler +from datahub.hub import Hub +from datahub.settings import Settings +from tests.fixtures import fake_transport + + +class AdminTests(unittest.TestCase): + def setUp(self) -> None: + self.tmp = tempfile.TemporaryDirectory() + settings = Settings( + encryption_key=SecretVault.generate_key(), + api_token="z" * 32, + admin_password="StartPass1", + tushare_token="real-tushare-token-abcdef", + db_path=Path(self.tmp.name) / "hub.db", + scheduler_enabled=False, + ) + self.hub = Hub(settings, adapter=TushareAdapter("real-tushare-token-abcdef", transport=fake_transport)) + handler = make_handler(self.hub) + self.server = ThreadingHTTPServer(("127.0.0.1", 0), handler) + threading.Thread(target=self.server.serve_forever, daemon=True).start() + self.base = f"http://127.0.0.1:{self.server.server_address[1]}" + + def tearDown(self) -> None: + self.server.shutdown() + self.server.server_close() + self.tmp.cleanup() + + def _json(self, path, method="GET", body=None, cookie="", csrf=""): + data = None if body is None else json.dumps(body).encode() + headers = {"Content-Type": "application/json"} + if cookie: + headers["Cookie"] = cookie + if csrf: + headers["X-CSRF-Token"] = csrf + req = Request(self.base + path, data=data, headers=headers, method=method) + with urlopen(req, timeout=5) as resp: + set_cookie = resp.headers.get("Set-Cookie", "") + return resp.status, json.loads(resp.read().decode()), set_cookie + + def test_login_change_password_and_secret_masking(self) -> None: + status, body, cookie_header = self._json( + "/admin/api/login", "POST", {"username": "hub_admin", "password": "StartPass1"} + ) + self.assertEqual(status, 200) + self.assertTrue(body["must_change"]) + cookie = cookie_header.split(";")[0] + csrf = body["csrf"] + status, _, _ = self._json( + "/admin/api/change-password", + "POST", + {"current": "StartPass1", "new_password": "NewPass123"}, + cookie=cookie, + csrf=csrf, + ) + self.assertEqual(status, 200) + _, sources, _ = self._json("/admin/api/sources", cookie=cookie, csrf=csrf) + blob = json.dumps(sources) + self.assertNotIn("real-tushare-token-abcdef", blob) + self.assertTrue(sources["items"][0]["credential"]["configured"]) + self.assertTrue(str(sources["items"][0]["credential"]["last4"]).endswith("cdef") or "****" in str(sources["items"][0]["credential"]["last4"])) + + def test_rollback_requires_password_and_confirm(self) -> None: + _, body, cookie_header = self._json( + "/admin/api/login", "POST", {"username": "hub_admin", "password": "StartPass1"} + ) + cookie = cookie_header.split(";")[0] + csrf = body["csrf"] + self._json("/admin/api/change-password", "POST", {"current": "StartPass1", "new_password": "NewPass123"}, cookie, csrf) + from urllib.error import HTTPError + + with self.assertRaises(HTTPError) as ctx: + self._json( + "/admin/api/rollback", + "POST", + {"dataset": "daily", "trade_date": "20240902", "password": "wrong", "confirm": "daily:20240902"}, + cookie, + csrf, + ) + self.assertEqual(ctx.exception.code, 401) + + +if __name__ == "__main__": + unittest.main() diff --git a/xiaobai-datahub/tests/test_api.py b/xiaobai-datahub/tests/test_api.py new file mode 100644 index 0000000..8ec4f19 --- /dev/null +++ b/xiaobai-datahub/tests/test_api.py @@ -0,0 +1,147 @@ +from __future__ import annotations + +import json +import tempfile +import threading +import unittest +from http.server import ThreadingHTTPServer +from pathlib import Path +from urllib.error import HTTPError +from urllib.request import Request, urlopen + +from datahub.adapters.tushare import TushareAdapter +from datahub.crypto import SecretVault +from datahub.httpapp import make_handler +from datahub.hub import Hub +from datahub.settings import Settings +from tests.fixtures import TRADE_DATE, fake_transport + +ERROR_CODES = { + "UNAUTHORIZED", + "INVALID_ARGUMENT", + "RATE_LIMITED", + "SOURCE_UNAVAILABLE", + "DATASET_NOT_PUBLISHED", + "STALE_DATA", + "INTERNAL", +} + + +class ApiContractTests(unittest.TestCase): + def setUp(self) -> None: + self.tmp = tempfile.TemporaryDirectory() + key = SecretVault.generate_key() + self.token = "k" * 32 + settings = Settings( + host="127.0.0.1", + port=0, + encryption_key=key, + api_token=self.token, + admin_password="StartPass1", + tushare_token="tushare-secret-token-xyz", + db_path=Path(self.tmp.name) / "hub.db", + backup_dir=Path(self.tmp.name) / "backups", + scheduler_enabled=False, + quality={"daily_row_ratio": 0.5, "null_rate_max": 0.5, "list_limit_default": 5000, "list_limit_max": 5000}, + ) + adapter = TushareAdapter("tushare-secret-token-xyz", transport=fake_transport) + self.hub = Hub(settings, adapter=adapter) + self.hub.pipeline.ingest_reference(TRADE_DATE) + for dataset in ("daily", "valuation", "moneyflow", "auction", "index_daily"): + self.hub.pipeline.run_dataset(dataset, TRADE_DATE) + handler = make_handler(self.hub) + self.server = ThreadingHTTPServer(("127.0.0.1", 0), handler) + self.thread = threading.Thread(target=self.server.serve_forever, daemon=True) + self.thread.start() + self.base = f"http://127.0.0.1:{self.server.server_address[1]}" + + def tearDown(self) -> None: + self.server.shutdown() + self.server.server_close() + self.hub.stop() + self.tmp.cleanup() + + def _get(self, path: str, token: str | None = None) -> tuple[int, dict]: + headers = {} + if token is not None: + headers["X-Datahub-Token"] = token + req = Request(self.base + path, headers=headers) + try: + with urlopen(req, timeout=5) as resp: + return resp.status, json.loads(resp.read().decode()) + except HTTPError as exc: + return exc.code, json.loads(exc.read().decode()) + + def test_livez_no_token(self) -> None: + status, body = self._get("/livez", token=None) + self.assertEqual(status, 200) + self.assertEqual(body["status"], "ok") + + def test_missing_and_bad_token_401(self) -> None: + status, body = self._get("/v1/health", token=None) + self.assertEqual(status, 401) + self.assertEqual(body["error"]["code"], "UNAUTHORIZED") + status, body = self._get("/v1/health", token="wrong") + self.assertEqual(status, 401) + self.assertNotIn("tushare-secret-token-xyz", json.dumps(body)) + self.assertNotIn(self.token, json.dumps(body)) + + def test_core_endpoints_schema(self) -> None: + paths = [ + "/v1/health", + f"/v1/calendar?from=20240901&to=20240907", + "/v1/stocks", + f"/v1/bars/daily?date={TRADE_DATE}&code=600000.SH&adjust=none", + f"/v1/bars/daily?date={TRADE_DATE}&code=600000.SH&adjust=qfq", + f"/v1/indexes/bars?date={TRADE_DATE}&code=000001.SH", + f"/v1/valuation?date={TRADE_DATE}&code=600000.SH", + f"/v1/moneyflow?date={TRADE_DATE}&code=600000.SH", + f"/v1/auction?date={TRADE_DATE}", + f"/v1/datasets/status?date={TRADE_DATE}", + f"/v1/batches?date={TRADE_DATE}", + ] + for path in paths: + status, body = self._get(path, token=self.token) + self.assertEqual(status, 200, path) + self.assertEqual(body["schema_version"], 1) + self.assertIn("data", body) + self.assertIn("meta", body) + self.assertIn("tier", 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) + _, qfq = self._get(f"/v1/bars/daily?date={TRADE_DATE}&code=600000.SH&adjust=qfq", token=self.token) + raw = none["data"][0] + adj = qfq["data"][0] + expected = round(raw["close"] * raw["adj_factor"] / raw["adj_factor"], 4) + self.assertEqual(adj["close"], expected) + + def test_unpublished_code(self) -> None: + status, body = self._get("/v1/bars/daily?date=19990101", token=self.token) + self.assertEqual(status, 404) + self.assertEqual(body["error"]["code"], "DATASET_NOT_PUBLISHED") + self.assertIn("expected_at", body["error"]) + + def test_error_code_set_documented(self) -> None: + self.assertGreaterEqual(ERROR_CODES, {"UNAUTHORIZED", "DATASET_NOT_PUBLISHED", "INVALID_ARGUMENT"}) + + def test_six_digit_code(self) -> None: + status, body = self._get(f"/v1/bars/daily?date={TRADE_DATE}&code=600000", token=self.token) + self.assertEqual(status, 200) + self.assertEqual(body["data"][0]["ts_code"], "600000.SH") + + def test_amount_unit_is_yuan(self) -> None: + _, body = self._get(f"/v1/bars/daily?date={TRADE_DATE}&code=600000.SH", token=self.token) + self.assertEqual(body["data"][0]["amount"], 2_000_000.0) + _, flow = self._get(f"/v1/moneyflow?date={TRADE_DATE}&code=600000.SH", token=self.token) + self.assertEqual(flow["data"][0]["net_mf_amount"], 170000.0) + + def test_token_never_in_health_or_admin_sources(self) -> None: + _, health = self._get("/v1/health", token=self.token) + blob = json.dumps(health) + self.assertNotIn("tushare-secret-token-xyz", blob) + self.assertNotIn(self.token, blob) + + +if __name__ == "__main__": + unittest.main() diff --git a/xiaobai-datahub/tests/test_governance.py b/xiaobai-datahub/tests/test_governance.py new file mode 100644 index 0000000..79cec9f --- /dev/null +++ b/xiaobai-datahub/tests/test_governance.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +import unittest + +from datahub.governance.circuit import CircuitBreaker +from datahub.governance.ratelimit import TokenBucket +from datahub.governance.retry import RetryError, retry_call + + +class FakeClock: + def __init__(self) -> None: + self.value = 0.0 + + def __call__(self) -> float: + return self.value + + +class GovernanceTests(unittest.TestCase): + def test_token_bucket_caps_burst_at_capacity(self) -> None: + clock = FakeClock() + bucket = TokenBucket(rate_per_minute=300, capacity=300, clock=clock) + ok = 0 + for _ in range(400): + if bucket.acquire(block=False): + ok += 1 + self.assertEqual(ok, 300) + clock.value = 60 + self.assertTrue(bucket.acquire(block=False)) + + def test_circuit_opens_after_five_failures_and_half_opens(self) -> None: + clock = FakeClock() + breaker = CircuitBreaker(clock=clock, open_seconds=120) + for _ in range(5): + breaker.record_failure("boom") + self.assertEqual(breaker.snapshot().state, "open") + self.assertFalse(breaker.allow()) + clock.value = 120 + self.assertEqual(breaker.snapshot().state, "half_open") + self.assertTrue(breaker.allow()) + breaker.record_success() + self.assertEqual(breaker.snapshot().state, "closed") + + def test_retry_exhausts(self) -> None: + calls = {"n": 0} + + def fail(): + calls["n"] += 1 + raise RuntimeError("no") + + with self.assertRaises(RetryError): + retry_call(fail, attempts=3, sleeper=lambda _d: None) + self.assertEqual(calls["n"], 3) + + +if __name__ == "__main__": + unittest.main() diff --git a/xiaobai-datahub/tests/test_layout.py b/xiaobai-datahub/tests/test_layout.py new file mode 100644 index 0000000..7f70bf1 --- /dev/null +++ b/xiaobai-datahub/tests/test_layout.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +import unittest +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] + + +class LayoutTests(unittest.TestCase): + def test_dockerfile_and_compose_exist(self) -> None: + self.assertTrue((ROOT / "Dockerfile").is_file()) + self.assertTrue((ROOT / "compose.yaml").is_file()) + self.assertTrue((ROOT / "requirements.txt").read_text(encoding="utf-8").startswith("cryptography==")) + dockerfile = (ROOT / "Dockerfile").read_text(encoding="utf-8") + self.assertIn("10002", dockerfile) + self.assertIn("8766", dockerfile) + self.assertIn("livez", dockerfile) + + def test_reserved_adapters_present(self) -> None: + from datahub.adapters import RESERVED + + for name in ("eastmoney", "tencent", "ths", "xgb", "akshare", "ifind"): + self.assertIn(name, RESERVED) + probe = RESERVED[name].probe() + self.assertEqual(probe["state"], "reserved") + self.assertFalse(probe["configured"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/xiaobai-datahub/tests/test_normalize.py b/xiaobai-datahub/tests/test_normalize.py new file mode 100644 index 0000000..23b1101 --- /dev/null +++ b/xiaobai-datahub/tests/test_normalize.py @@ -0,0 +1,78 @@ +from __future__ import annotations + +import unittest + +from datahub.normalize import ( + AMOUNT_THOUSAND_YUAN, + AMOUNT_WAN_YUAN, + VOLUME_LOT, + apply_qfq, + normalize_auction, + normalize_daily, + normalize_index_daily, + normalize_moneyflow, + normalize_valuation, + review_daily_to_canonical, +) +from tests.fixtures import RAW + + +class NormalizeTests(unittest.TestCase): + def test_daily_matches_architecture_and_review_native_conversion(self) -> None: + raw = RAW["daily"][0] + hub = normalize_daily(raw, adj_factor=1.1) + # review stores Tushare native units; canonical = review * factor + review_native = dict(raw) + converted = review_daily_to_canonical(review_native) + self.assertEqual(hub["amount"], converted["amount"]) + self.assertEqual(hub["amount"], raw["amount"] * AMOUNT_THOUSAND_YUAN) + self.assertEqual(hub["volume"], raw["vol"] * VOLUME_LOT) + self.assertEqual(hub["close"], 10.2) + self.assertEqual(hub["adj_factor"], 1.1) + self.assertEqual(hub["ts_code"], "600000.SH") + + def test_moneyflow_wan_to_yuan(self) -> None: + raw = RAW["moneyflow"][0] + hub = normalize_moneyflow(raw) + self.assertEqual(hub["net_mf_amount"], raw["net_mf_amount"] * AMOUNT_WAN_YUAN) + self.assertEqual(hub["buy_lg_amount"], 30 * AMOUNT_WAN_YUAN) + + def test_valuation_mv_wan_to_yuan(self) -> None: + raw = RAW["daily_basic"][0] + hub = normalize_valuation(raw) + self.assertEqual(hub["total_mv"], 1000 * AMOUNT_WAN_YUAN) + self.assertEqual(hub["circ_mv"], 800 * AMOUNT_WAN_YUAN) + + def test_index_daily_amount_thousand_yuan(self) -> None: + raw = RAW["index_daily"][0] + hub = normalize_index_daily(raw) + self.assertEqual(hub["amount"], raw["amount"] * AMOUNT_THOUSAND_YUAN) + self.assertEqual(hub["volume"], raw["vol"] * VOLUME_LOT) + + def test_auction_amount_already_yuan(self) -> None: + raw = RAW["stk_auction"][0] + hub = normalize_auction(raw) + self.assertEqual(hub["amount"], raw["amount"]) + self.assertEqual(hub["volume"], raw["vol"] * VOLUME_LOT) + + def test_field_diff_against_review_native_is_explained(self) -> None: + """Golden: every non-zero diff vs review-native daily is a documented unit factor.""" + raw = RAW["daily"][0] + hub = normalize_daily(raw) + diffs = {} + for key in ("open", "high", "low", "close", "pct_chg"): + if hub[key] != raw[key]: + diffs[key] = (raw[key], hub[key]) + self.assertEqual(diffs, {}) + self.assertNotEqual(hub["amount"], raw["amount"]) + self.assertEqual(hub["amount"] / raw["amount"], AMOUNT_THOUSAND_YUAN) + self.assertEqual(hub["volume"] / raw["vol"], VOLUME_LOT) + + def test_qfq_formula(self) -> None: + self.assertEqual(apply_qfq(10.0, 1.1, 2.2), 5.0) + none_price = apply_qfq(None, 1.1, 2.2) + self.assertIsNone(none_price) + + +if __name__ == "__main__": + unittest.main() diff --git a/xiaobai-datahub/tests/test_pipeline.py b/xiaobai-datahub/tests/test_pipeline.py new file mode 100644 index 0000000..ea3d577 --- /dev/null +++ b/xiaobai-datahub/tests/test_pipeline.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +import tempfile +import unittest +from pathlib import Path + +from datahub.adapters.tushare import TushareAdapter +from datahub.crypto import SecretVault +from datahub.db import HubDB +from datahub.pipeline import Pipeline, QualityError +from datahub.settings import Settings +from tests.fixtures import TRADE_DATE, fake_transport + + +def make_pipeline(before_commit=None) -> tuple[Pipeline, HubDB]: + tmp = tempfile.TemporaryDirectory() + db = HubDB(Path(tmp.name) / "hub.db") + adapter = TushareAdapter("test-token", transport=fake_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={"daily_row_ratio": 0.98, "null_rate_max": 0.01, "max_publish_attempts": 3, "publication_generations": 3}, + scheduler_enabled=False, + ) + pipe = Pipeline(db, adapter, settings, before_commit=before_commit) + pipe._tmp = tmp # keep alive + return pipe, db + + +class PipelineTests(unittest.TestCase): + def test_reference_and_daily_publish(self) -> None: + pipe, db = make_pipeline() + ref = pipe.ingest_reference(TRADE_DATE) + self.assertEqual(ref["stocks"], 2) + result = pipe.run_dataset("daily", TRADE_DATE) + self.assertEqual(result["state"], "published") + self.assertEqual(result["rows"], 2) + pub = db.fetchone("SELECT * FROM publications WHERE dataset='daily' AND trade_date=?", (TRADE_DATE,)) + self.assertEqual(pub["active_batch"], result["batch_id"]) + rows = db.fetchall("SELECT * FROM eod_bars WHERE batch_id=?", (result["batch_id"],)) + self.assertEqual(len(rows), 2) + self.assertEqual(rows[0]["amount"] if rows[0]["ts_code"] == "600000.SH" else rows[1]["amount"], 2_000_000.0) + + def test_atomic_publish_abort_leaves_no_half_batch(self) -> None: + pipe, db = make_pipeline() + pipe.ingest_reference(TRADE_DATE) + first = pipe.run_dataset("daily", TRADE_DATE) + boom = {"n": 0} + + def explode() -> None: + boom["n"] += 1 + raise RuntimeError("killed") + + pipe.before_commit = explode + with self.assertRaises(RuntimeError): + pipe.run_dataset("daily", TRADE_DATE) + pub = db.fetchone("SELECT * FROM publications WHERE dataset='daily' AND trade_date=?", (TRADE_DATE,)) + self.assertEqual(pub["active_batch"], first["batch_id"]) + visible = db.fetchall( + "SELECT DISTINCT batch_id FROM eod_bars WHERE trade_date=? AND batch_id=?", + (TRADE_DATE, pub["active_batch"]), + ) + self.assertEqual(len(visible), 1) + + def test_rollback_switches_active_batch(self) -> None: + pipe, _db = make_pipeline() + pipe.ingest_reference(TRADE_DATE) + first = pipe.run_dataset("daily", TRADE_DATE) + second = pipe.run_dataset("daily", TRADE_DATE) + self.assertNotEqual(first["batch_id"], second["batch_id"]) + rolled = pipe.rollback("daily", TRADE_DATE, actor="test") + self.assertEqual(rolled["active_batch"], first["batch_id"]) + from datahub.serving import V1API + + api = V1API(pipe.db, pipe, pipe.settings) + payload = api.handle("/v1/bars/daily", {"date": [TRADE_DATE], "code": ["600000.SH"]}) + self.assertEqual(payload["meta"]["batch_id"], first["batch_id"]) + + def test_row_ratio_gate_rejects_short_batch(self) -> None: + pipe, _db = make_pipeline() + pipe.ingest_reference(TRADE_DATE) + original = fake_transport + + def short(api_name, params, fields): + rows = original(api_name, params, fields) + if api_name == "daily": + return rows[:1] + return rows + + pipe.adapter._transport = short + with self.assertRaises(QualityError) as ctx: + pipe.run_dataset("daily", TRADE_DATE) + self.assertTrue(ctx.exception.report["hard_fail"]) + pub = pipe.db.fetchone("SELECT * FROM publications WHERE dataset='daily'") + self.assertIsNone(pub) + + def test_wal_mode(self) -> None: + pipe, db = make_pipeline() + with db.connect() as connection: + mode = connection.execute("PRAGMA journal_mode").fetchone()[0] + self.assertEqual(str(mode).lower(), "wal") + + +if __name__ == "__main__": + unittest.main() diff --git a/xiaobai-datahub/tests/test_scheduler.py b/xiaobai-datahub/tests/test_scheduler.py new file mode 100644 index 0000000..da0de5b --- /dev/null +++ b/xiaobai-datahub/tests/test_scheduler.py @@ -0,0 +1,62 @@ +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 + + +class SchedulerTests(unittest.TestCase): + def test_skips_eod_on_closed_day(self) -> None: + tmp = tempfile.TemporaryDirectory() + db = HubDB(Path(tmp.name) / "hub.db") + adapter = TushareAdapter("x", transport=fake_transport) + settings = Settings(encryption_key=SecretVault.generate_key(), scheduler_enabled=False, db_path=db.path) + pipe = Pipeline(db, adapter, settings) + pipe.ingest_reference("20240902") + # 20240907 is closed in fixture + ran = {"eod_a": 0} + + def fake_eod(_date: str): + ran["eod_a"] += 1 + return {} + + sched = Scheduler(db, pipe, jobs={"precheck": lambda d: {}, "eod_a": fake_eod, "eod_b": lambda d: {}, "cleanup": lambda d: {}, "backup": lambda d: {}}) + clock = datetime(2024, 9, 7, 16, 0, tzinfo=SHANGHAI) + fired = sched.tick(clock) + self.assertNotIn("eod_a", fired) + self.assertEqual(ran["eod_a"], 0) + tmp.cleanup() + + def test_fires_eod_on_open_day(self) -> None: + tmp = tempfile.TemporaryDirectory() + db = HubDB(Path(tmp.name) / "hub.db") + adapter = TushareAdapter("x", transport=fake_transport) + settings = Settings(encryption_key=SecretVault.generate_key(), scheduler_enabled=False, db_path=db.path) + pipe = Pipeline(db, adapter, settings) + pipe.ingest_reference("20240902") + ran = {"eod_a": 0} + + def fake_eod(_date: str): + ran["eod_a"] += 1 + return {"rows": 1} + + sched = Scheduler(db, pipe, jobs={"precheck": lambda d: {}, "eod_a": fake_eod, "eod_b": lambda d: {}, "cleanup": lambda d: {}, "backup": lambda d: {}}) + clock = datetime(2024, 9, 2, 16, 0, tzinfo=SHANGHAI) + fired = sched.tick(clock) + self.assertIn("eod_a", fired) + self.assertEqual(ran["eod_a"], 1) + tmp.cleanup() + + +if __name__ == "__main__": + unittest.main()