Compare commits

..
Author SHA1 Message Date
031eefab4d fix(HEL-386): 清理任务按 ISO 截止时间删除 job_runs/src_calls
YYYYMMDD 与 ISO 字符串比较会把同年保留期内记录全部误删。

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-02 12:26:09 +08:00
3498dd7a4b feat(HEL-382): 搭建 datahub 底座和盘后正式数据链路
新增独立 xiaobai-datahub 服务(SQLite WAL、Tushare 盘后发布、/v1 契约和管理后台),不改现站页面与数据链路。

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-09-02 12:05:26 +08:00
52 changed files with 4302 additions and 0 deletions
+3
View File
@@ -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
+45
View File
@@ -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"
+10
View File
@@ -0,0 +1,10 @@
.git
.gitignore
.env
.env.*
!.env.example
__pycache__/
*.py[cod]
*.log
data/
tests/
+13
View File
@@ -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
+36
View File
@@ -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"]
+78
View File
@@ -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`
- 回滚、补数需重新输入密码 + 确认词
- 容器非 rootuid 10002)、read_only、cap_drop ALL
+268
View File
@@ -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) => ({ "&": "&amp;", "<": "&lt;", ">": "&gt;", '"': "&quot;" }[ch]));
}
function table(headers, rows) {
const thead = headers.map((h) => `<th>${esc(h)}</th>`).join("");
const body = rows.length
? rows.map((cols) => `<tr>${cols.map((c) => `<td>${c}</td>`).join("")}</tr>`).join("")
: `<tr><td colspan="${headers.length}">暂无数据</td></tr>`;
return `<table><thead><tr>${thead}</tr></thead><tbody>${body}</tbody></table>`;
}
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 = `
<div class="cards">
<div class="card"><div class="muted">交易日</div><strong>${esc(data.trade_date)}</strong></div>
<div class="card"><div class="muted">阶段</div><strong>${esc(data.session_phase)}</strong></div>
<div class="card"><div class="muted">今日发布</div><strong>${data.publications.length}</strong></div>
<div class="card"><div class="muted">异常批次</div><strong class="${data.anomalies.length ? "fail" : "ok"}">${data.anomalies.length}</strong></div>
</div>
<h2>最近调用</h2>
${table(["时间", "源", "端点", "结果", "耗时"], data.recent_calls.map((row) => [
esc(row.created_at), esc(row.provider), esc(row.endpoint),
row.ok ? '<span class="ok">成功</span>' : `<span class="fail">${esc(row.error)}</span>`,
`${row.latency_ms ?? "-"} ms`,
]))}
`;
return;
}
if (state.page === "sources") {
const data = await api("/admin/api/sources");
page.innerHTML = `<h2>数据源</h2>` + 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,
`<button data-probe="${esc(item.provider)}">探测一次</button>`,
];
}),
);
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 = `
<h2>调度任务</h2>
${table(["任务", "时刻", "操作"], data.jobs.map((job) => [
`${esc(job.id)} · ${esc(job.title)}`, esc(job.at),
`<button data-run="${esc(job.id)}">手动触发</button>`,
]))}
<h3>最近运行</h3>
${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 = `
<h2>盘后发布 ${esc(data.trade_date)}</h2>
<div class="toolbar">
<label>日期 <input id="rel-date" value="${esc(data.trade_date)}" /></label>
<button type="button" id="rel-load">查看</button>
<button type="button" id="rel-backfill">补数</button>
</div>
<h3>当前映射</h3>
${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 ? `<button class="danger" data-rollback="${esc(row.dataset)}">回滚</button>` : "-",
]))}
<h3>批次</h3>
${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 = `
<h2>数据集 / 质量 ${esc(data.trade_date)}</h2>
${table(["数据集", "批次", "状态", "发布时间"], data.publications.map((row) => [
esc(row.dataset), esc(row.active_batch), esc(row.state), esc(row.published_at),
]))}
<h3>源间差异</h3>
${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 = `<h2>审计</h2>` + 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 `
<h2>盘后发布 ${esc(data.trade_date)}</h2>
<div class="toolbar">
<label>日期 <input id="rel-date" value="${esc(data.trade_date)}" /></label>
<button type="button" id="rel-load">查看</button>
<button type="button" id="rel-backfill">补数</button>
</div>
<h3>当前映射</h3>
${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 ? `<button class="danger" data-rollback="${esc(row.dataset)}">回滚</button>` : "-",
]))}
<h3>批次</h3>
${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();
+53
View File
@@ -0,0 +1,53 @@
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1" />
<title>xiaobai-datahub 管理后台</title>
<link rel="stylesheet" href="/admin/styles.css" />
</head>
<body>
<div id="app">
<section id="login-view" class="panel auth-panel">
<h1>数据中枢</h1>
<p class="muted">内网管理后台,用于查看源状态、调度和盘后发布批次。</p>
<form id="login-form">
<label>账号 <input name="username" value="hub_admin" autocomplete="username" /></label>
<label>密码 <input name="password" type="password" autocomplete="current-password" /></label>
<button type="submit">登录</button>
<p id="login-error" class="error" hidden></p>
</form>
</section>
<section id="change-view" class="panel auth-panel" hidden>
<h1>修改初始密码</h1>
<form id="change-form">
<label>当前密码 <input name="current" type="password" /></label>
<label>新密码(至少 8 位) <input name="new_password" type="password" /></label>
<button type="submit">保存并继续</button>
<p id="change-error" class="error" hidden></p>
</form>
</section>
<section id="shell" hidden>
<header class="top">
<strong>xiaobai-datahub</strong>
<span id="phase" class="pill"></span>
<span id="who" class="muted"></span>
<button type="button" id="theme-btn" class="ghost">夜间</button>
<button type="button" id="logout-btn" class="ghost">退出</button>
</header>
<nav>
<button data-page="overview" class="active">总览</button>
<button data-page="sources">数据源</button>
<button data-page="jobs">调度任务</button>
<button data-page="release">盘后发布</button>
<button data-page="datasets">数据集</button>
<button data-page="audit">审计</button>
</nav>
<main id="page"></main>
</section>
</div>
<script src="/admin/app.js"></script>
</body>
</html>
+51
View File
@@ -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; }
+39
View File
@@ -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"
@@ -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
}
+4
View File
@@ -0,0 +1,4 @@
"""xiaobai-datahub: independent market-data service for xiaobai-review."""
__version__ = "0.1.0"
SCHEMA_VERSION = 1
@@ -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,
}
@@ -0,0 +1,3 @@
from datahub.adapters.base import ReservedAdapter
ADAPTER = ReservedAdapter("akshare")
+47
View File
@@ -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 []
@@ -0,0 +1,3 @@
from datahub.adapters.base import ReservedAdapter
ADAPTER = ReservedAdapter("eastmoney")
@@ -0,0 +1,3 @@
from datahub.adapters.base import ReservedAdapter
ADAPTER = ReservedAdapter("ifind")
@@ -0,0 +1,3 @@
from datahub.adapters.base import ReservedAdapter
ADAPTER = ReservedAdapter("tencent")
+3
View File
@@ -0,0 +1,3 @@
from datahub.adapters.base import ReservedAdapter
ADAPTER = ReservedAdapter("ths")
+154
View File
@@ -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 []]
+3
View File
@@ -0,0 +1,3 @@
from datahub.adapters.base import ReservedAdapter
ADAPTER = ReservedAdapter("xgb")
+162
View File
@@ -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
+190
View File
@@ -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 "")
+26
View File
@@ -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}"
+42
View File
@@ -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:]
+340
View File
@@ -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")
@@ -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",
]
@@ -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,
)
+84
View File
@@ -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
@@ -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)
@@ -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)
+242
View File
@@ -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()
+54
View File
@@ -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())
+56
View File
@@ -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")
+201
View File
@@ -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
+23
View File
@@ -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)
+480
View File
@@ -0,0 +1,480 @@
from __future__ import annotations
import json
import time
from collections.abc import Callable
from datetime import timedelta
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)
now = now_shanghai(self.clock())
cutoff_staging = add_days(yyyymmdd(now), -staging_days)
cutoff_jobs = isoformat(now - timedelta(days=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,
),
)
+145
View File
@@ -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)}
+376
View File
@@ -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
+72
View File
@@ -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"},
)
+64
View File
@@ -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()
+1
View File
@@ -0,0 +1 @@
cryptography==49.0.0
+27
View File
@@ -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()
View File
+59
View File
@@ -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 [])
+96
View File
@@ -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()
+147
View File
@@ -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()
+56
View File
@@ -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()
+30
View File
@@ -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()
+78
View File
@@ -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()
+149
View File
@@ -0,0 +1,149 @@
from __future__ import annotations
import tempfile
import unittest
from datetime import datetime, timedelta
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 datahub.timeutil import SHANGHAI, isoformat
from tests.fixtures import TRADE_DATE, fake_transport
def make_pipeline(before_commit=None, clock=None, quality=None) -> tuple[Pipeline, HubDB]:
tmp = tempfile.TemporaryDirectory()
db = HubDB(Path(tmp.name) / "hub.db")
adapter = TushareAdapter("test-token", transport=fake_transport)
quality_cfg = {
"daily_row_ratio": 0.98,
"null_rate_max": 0.01,
"max_publish_attempts": 3,
"publication_generations": 3,
"job_run_retain_days": 90,
"staging_retain_days": 14,
}
if quality:
quality_cfg.update(quality)
settings = Settings(
encryption_key=SecretVault.generate_key(),
api_token="t" * 32,
admin_password="admin-pass",
tushare_token="test-token",
db_path=db.path,
quality=quality_cfg,
scheduler_enabled=False,
)
pipe = Pipeline(db, adapter, settings, before_commit=before_commit, clock=clock)
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")
def test_cleanup_iso_timestamps_respect_retention_on_job_and_src(self) -> None:
# job_runs.started_at / src_calls.created_at 存 ISO;旧实现用 YYYYMMDD 比较会误删同年记录。
frozen = datetime(2026, 9, 2, 0, 30, tzinfo=SHANGHAI)
retain_days = 90
pipe, db = make_pipeline(clock=lambda: frozen, quality={"job_run_retain_days": retain_days})
samples = {
"today": isoformat(frozen),
"within": isoformat(frozen - timedelta(days=retain_days - 1)),
"expired": isoformat(frozen - timedelta(days=retain_days + 1)),
}
with db.write() as connection:
for job_id, stamp in samples.items():
connection.execute(
"INSERT INTO job_runs(job_id, state, started_at, attempt) VALUES (?,?,?,1)",
(job_id, "ok", stamp),
)
connection.execute(
"INSERT INTO src_calls(provider, endpoint, ok, latency_ms, error, created_at) VALUES (?,?,?,?,?,?)",
("tushare", job_id, 1, 10, None, stamp),
)
pipe.cleanup()
jobs = {row["job_id"] for row in db.fetchall("SELECT job_id FROM job_runs")}
calls = {row["endpoint"] for row in db.fetchall("SELECT endpoint FROM src_calls")}
kept = {"today", "within"}
self.assertEqual(jobs, kept)
self.assertEqual(calls, kept)
if __name__ == "__main__":
unittest.main()
+62
View File
@@ -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()