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>
This commit is contained in:
co-authored by
Cursor
multica-agent
parent
c2ebc0ab91
commit
3498dd7a4b
@@ -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")
|
||||
@@ -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")
|
||||
@@ -0,0 +1,3 @@
|
||||
from datahub.adapters.base import ReservedAdapter
|
||||
|
||||
ADAPTER = ReservedAdapter("ths")
|
||||
@@ -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 []]
|
||||
@@ -0,0 +1,3 @@
|
||||
from datahub.adapters.base import ReservedAdapter
|
||||
|
||||
ADAPTER = ReservedAdapter("xgb")
|
||||
@@ -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
|
||||
@@ -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 "")
|
||||
@@ -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}"
|
||||
@@ -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:]
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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())
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -0,0 +1,478 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from datahub.adapters.base import AdapterError
|
||||
from datahub.adapters.tushare import DEFAULT_INDEX_CODES, TushareAdapter
|
||||
from datahub.db import DATASET_TABLES, HubDB
|
||||
from datahub.governance.circuit import CircuitBreaker
|
||||
from datahub.governance.ratelimit import TokenBucket
|
||||
from datahub.governance.retry import RetryError, retry_call
|
||||
from datahub.logutil import get_logger
|
||||
from datahub.normalize import finite_number, normalize_daily
|
||||
from datahub.settings import Settings
|
||||
from datahub.timeutil import add_days, isoformat, now_shanghai, yyyymmdd
|
||||
|
||||
LOGGER = get_logger()
|
||||
|
||||
HARD_DATASETS = {"daily", "valuation", "index_daily"}
|
||||
SOFT_DATASETS = {"moneyflow", "auction"}
|
||||
|
||||
STAGING_INSERT = {
|
||||
"daily": (
|
||||
"INSERT INTO staging_bars(ts_code,trade_date,batch_id,open,high,low,close,pct_chg,volume,amount,adj_factor) "
|
||||
"VALUES (?,?,?,?,?,?,?,?,?,?,?)",
|
||||
lambda r, b: (
|
||||
r["ts_code"], r["trade_date"], b, r.get("open"), r.get("high"), r.get("low"),
|
||||
r.get("close"), r.get("pct_chg"), r.get("volume"), r.get("amount"), r.get("adj_factor"),
|
||||
),
|
||||
),
|
||||
"valuation": (
|
||||
"INSERT INTO staging_valuation(ts_code,trade_date,batch_id,turnover_rate,volume_ratio,total_mv,circ_mv,pe_ttm,pb,ps_ttm,dv_ttm) "
|
||||
"VALUES (?,?,?,?,?,?,?,?,?,?,?)",
|
||||
lambda r, b: (
|
||||
r["ts_code"], r["trade_date"], b, r.get("turnover_rate"), r.get("volume_ratio"),
|
||||
r.get("total_mv"), r.get("circ_mv"), r.get("pe_ttm"), r.get("pb"), r.get("ps_ttm"), r.get("dv_ttm"),
|
||||
),
|
||||
),
|
||||
"moneyflow": (
|
||||
"INSERT INTO staging_moneyflow(ts_code,trade_date,batch_id,buy_sm_amount,sell_sm_amount,buy_md_amount,sell_md_amount,buy_lg_amount,sell_lg_amount,buy_elg_amount,sell_elg_amount,net_mf_amount) "
|
||||
"VALUES (?,?,?,?,?,?,?,?,?,?,?,?)",
|
||||
lambda r, b: (
|
||||
r["ts_code"], r["trade_date"], b,
|
||||
r.get("buy_sm_amount"), r.get("sell_sm_amount"), r.get("buy_md_amount"), r.get("sell_md_amount"),
|
||||
r.get("buy_lg_amount"), r.get("sell_lg_amount"), r.get("buy_elg_amount"), r.get("sell_elg_amount"),
|
||||
r.get("net_mf_amount"),
|
||||
),
|
||||
),
|
||||
"auction": (
|
||||
"INSERT INTO staging_auction(ts_code,trade_date,batch_id,volume,price,amount,pre_close,turnover_rate,volume_ratio,float_share) "
|
||||
"VALUES (?,?,?,?,?,?,?,?,?,?)",
|
||||
lambda r, b: (
|
||||
r["ts_code"], r["trade_date"], b, r.get("volume"), r.get("price"), r.get("amount"),
|
||||
r.get("pre_close"), r.get("turnover_rate"), r.get("volume_ratio"), r.get("float_share"),
|
||||
),
|
||||
),
|
||||
"index_daily": (
|
||||
"INSERT INTO staging_index_bars(ts_code,trade_date,batch_id,open,high,low,close,pct_chg,volume,amount) "
|
||||
"VALUES (?,?,?,?,?,?,?,?,?,?)",
|
||||
lambda r, b: (
|
||||
r["ts_code"], r["trade_date"], b, r.get("open"), r.get("high"), r.get("low"),
|
||||
r.get("close"), r.get("pct_chg"), r.get("volume"), r.get("amount"),
|
||||
),
|
||||
),
|
||||
}
|
||||
|
||||
EOD_COPY = {
|
||||
"daily": (
|
||||
"INSERT OR REPLACE INTO eod_bars "
|
||||
"SELECT ts_code,trade_date,open,high,low,close,pct_chg,volume,amount,adj_factor,batch_id "
|
||||
"FROM staging_bars WHERE batch_id = ?"
|
||||
),
|
||||
"valuation": (
|
||||
"INSERT OR REPLACE INTO eod_valuation "
|
||||
"SELECT ts_code,trade_date,turnover_rate,volume_ratio,total_mv,circ_mv,pe_ttm,pb,ps_ttm,dv_ttm,batch_id "
|
||||
"FROM staging_valuation WHERE batch_id = ?"
|
||||
),
|
||||
"moneyflow": (
|
||||
"INSERT OR REPLACE INTO eod_moneyflow "
|
||||
"SELECT ts_code,trade_date,buy_sm_amount,sell_sm_amount,buy_md_amount,sell_md_amount,"
|
||||
"buy_lg_amount,sell_lg_amount,buy_elg_amount,sell_elg_amount,net_mf_amount,batch_id "
|
||||
"FROM staging_moneyflow WHERE batch_id = ?"
|
||||
),
|
||||
"auction": (
|
||||
"INSERT OR REPLACE INTO eod_auction "
|
||||
"SELECT ts_code,trade_date,volume,price,amount,pre_close,turnover_rate,volume_ratio,float_share,batch_id "
|
||||
"FROM staging_auction WHERE batch_id = ?"
|
||||
),
|
||||
"index_daily": (
|
||||
"INSERT OR REPLACE INTO eod_index_bars "
|
||||
"SELECT ts_code,trade_date,open,high,low,close,pct_chg,volume,amount,batch_id "
|
||||
"FROM staging_index_bars WHERE batch_id = ?"
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
class QualityError(RuntimeError):
|
||||
def __init__(self, message: str, report: dict[str, Any]) -> None:
|
||||
super().__init__(message)
|
||||
self.report = report
|
||||
|
||||
|
||||
class Pipeline:
|
||||
def __init__(
|
||||
self,
|
||||
db: HubDB,
|
||||
adapter: TushareAdapter,
|
||||
settings: Settings,
|
||||
bucket: TokenBucket | None = None,
|
||||
breaker: CircuitBreaker | None = None,
|
||||
before_commit: Callable[[], None] | None = None,
|
||||
clock=None,
|
||||
) -> None:
|
||||
self.db = db
|
||||
self.adapter = adapter
|
||||
self.settings = settings
|
||||
self.bucket = bucket or TokenBucket(settings.tushare_rate_per_minute)
|
||||
self.breaker = breaker or CircuitBreaker()
|
||||
self.before_commit = before_commit
|
||||
self.clock = clock or now_shanghai
|
||||
|
||||
def next_batch_id(self, dataset: str, trade_date: str) -> str:
|
||||
row = self.db.fetchone(
|
||||
"SELECT COUNT(*) AS n FROM batches WHERE dataset = ? AND trade_date = ?",
|
||||
(dataset, trade_date),
|
||||
)
|
||||
seq = int((row or {}).get("n") or 0) + 1
|
||||
return f"{trade_date}-{dataset}-{seq:03d}"
|
||||
|
||||
def ingest_reference(self, trade_date: str | None = None) -> dict[str, Any]:
|
||||
"""Refresh trade calendar (window) and stock master. Not versioned by batch."""
|
||||
day = yyyymmdd(trade_date or self.clock())
|
||||
start = add_days(day, -400)
|
||||
end = add_days(day, 30)
|
||||
calendar = self.adapter.normalize(
|
||||
"calendar",
|
||||
self._guarded_fetch("calendar", {"exchange": "SSE", "start_date": start, "end_date": end}),
|
||||
)
|
||||
stocks = self.adapter.normalize("stocks", self._guarded_fetch("stocks", {"list_status": "L"}))
|
||||
fetched_at = isoformat(self.clock())
|
||||
with self.db.write() as connection:
|
||||
for row in calendar:
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO trade_calendar(exchange, cal_date, is_open, pretrade_date, fetched_at)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
ON CONFLICT(exchange, cal_date) DO UPDATE SET
|
||||
is_open=excluded.is_open, pretrade_date=excluded.pretrade_date, fetched_at=excluded.fetched_at
|
||||
""",
|
||||
(row["exchange"], row["cal_date"], row["is_open"], row.get("pretrade_date"), fetched_at),
|
||||
)
|
||||
for row in stocks:
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO stock_master(ts_code,symbol,name,area,industry,market,list_status,list_date,updated_at)
|
||||
VALUES (?,?,?,?,?,?,?,?,?)
|
||||
ON CONFLICT(ts_code) DO UPDATE SET
|
||||
symbol=excluded.symbol, name=excluded.name, area=excluded.area,
|
||||
industry=excluded.industry, market=excluded.market,
|
||||
list_status=excluded.list_status, list_date=excluded.list_date,
|
||||
updated_at=excluded.updated_at
|
||||
""",
|
||||
(
|
||||
row["ts_code"], row.get("symbol"), row.get("name"), row.get("area"),
|
||||
row.get("industry"), row.get("market"), row.get("list_status"),
|
||||
row.get("list_date"), fetched_at,
|
||||
),
|
||||
)
|
||||
return {"calendar": len(calendar), "stocks": len(stocks), "trade_date": day}
|
||||
|
||||
def run_dataset(self, dataset: str, trade_date: str, attempts: int | None = None) -> dict[str, Any]:
|
||||
trade_date = yyyymmdd(trade_date)
|
||||
batch_id = self.next_batch_id(dataset, trade_date)
|
||||
max_attempts = attempts or self.settings.max_publish_attempts
|
||||
self._set_batch(batch_id, dataset, trade_date, "scheduled", 0)
|
||||
try:
|
||||
self._set_batch(batch_id, dataset, trade_date, "fetching", 1)
|
||||
rows = retry_call(
|
||||
lambda: self._fetch_dataset(dataset, trade_date),
|
||||
attempts=max_attempts,
|
||||
base_delay=0.05,
|
||||
sleeper=lambda _d: None if attempts == 1 else time.sleep(_d),
|
||||
)
|
||||
self._stage(dataset, batch_id, rows)
|
||||
self._set_batch(batch_id, dataset, trade_date, "staged", 1, rows_in=len(rows), rows_out=len(rows))
|
||||
self._set_batch(batch_id, dataset, trade_date, "validating", 1)
|
||||
report = self.validate(dataset, batch_id, trade_date, rows)
|
||||
if report["hard_fail"]:
|
||||
self._set_batch(
|
||||
batch_id, dataset, trade_date, "staged", 1,
|
||||
rows_in=len(rows), rows_out=len(rows),
|
||||
quality=report, error="; ".join(report["errors"]),
|
||||
)
|
||||
raise QualityError("integrity gate failed", report)
|
||||
self._set_batch(batch_id, dataset, trade_date, "deriving", 1, rows_in=len(rows), rows_out=len(rows), quality=report)
|
||||
self._set_batch(batch_id, dataset, trade_date, "publishing", 1, rows_in=len(rows), rows_out=len(rows), quality=report)
|
||||
state = "degraded" if report["soft_fail"] else "published"
|
||||
self.publish(dataset, trade_date, batch_id, state=state)
|
||||
self._set_batch(
|
||||
batch_id, dataset, trade_date, "published", 1,
|
||||
rows_in=len(rows), rows_out=len(rows), quality=report, finished=True,
|
||||
)
|
||||
return {"batch_id": batch_id, "dataset": dataset, "trade_date": trade_date, "rows": len(rows), "state": state, "quality": report}
|
||||
except RetryError as exc:
|
||||
self._set_batch(batch_id, dataset, trade_date, "failed", max_attempts, error=str(exc), finished=True)
|
||||
raise
|
||||
except QualityError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
self._set_batch(batch_id, dataset, trade_date, "failed", 1, error=str(exc), finished=True)
|
||||
raise
|
||||
|
||||
def run_eod_batch_a(self, trade_date: str) -> dict[str, Any]:
|
||||
results = {}
|
||||
for dataset in ("daily", "valuation", "moneyflow", "auction"):
|
||||
results[dataset] = self.run_dataset(dataset, trade_date)
|
||||
return results
|
||||
|
||||
def run_eod_batch_b(self, trade_date: str) -> dict[str, Any]:
|
||||
return {"index_daily": self.run_dataset("index_daily", trade_date)}
|
||||
|
||||
def validate(self, dataset: str, batch_id: str, trade_date: str, rows: list[dict[str, Any]]) -> dict[str, Any]:
|
||||
quality = self.settings.quality
|
||||
errors: list[str] = []
|
||||
warnings: list[str] = []
|
||||
listed = self.db.fetchone(
|
||||
"SELECT COUNT(*) AS n FROM stock_master WHERE list_status = 'L'",
|
||||
)
|
||||
listed_n = int((listed or {}).get("n") or 0)
|
||||
row_n = len(rows)
|
||||
keys = [(row.get("ts_code"), row.get("trade_date")) for row in rows]
|
||||
dup = row_n - len(set(keys))
|
||||
if dup:
|
||||
errors.append(f"duplicate keys: {dup}")
|
||||
bad_date = sum(1 for row in rows if str(row.get("trade_date")) != trade_date)
|
||||
if bad_date:
|
||||
errors.append(f"date mismatch rows: {bad_date}")
|
||||
ratio = (row_n / listed_n) if listed_n else 1.0
|
||||
if dataset == "daily" and listed_n and ratio < float(quality.get("daily_row_ratio") or 0.98):
|
||||
errors.append(f"row ratio {ratio:.4f} < {quality.get('daily_row_ratio')}")
|
||||
null_fields = ("open", "high", "low", "close", "amount") if dataset in {"daily", "index_daily"} else ()
|
||||
if null_fields and rows:
|
||||
nulls = sum(1 for row in rows if any(row.get(field) is None for field in null_fields))
|
||||
null_rate = nulls / row_n
|
||||
if null_rate >= float(quality.get("null_rate_max") or 0.01):
|
||||
errors.append(f"null rate {null_rate:.4f}")
|
||||
if dataset in SOFT_DATASETS and row_n == 0:
|
||||
warnings.append("empty soft dataset")
|
||||
hard_fail = bool(errors) and dataset in HARD_DATASETS.union({"daily", "valuation", "index_daily"})
|
||||
if dataset in SOFT_DATASETS:
|
||||
hard_fail = bool(dup or bad_date)
|
||||
return {
|
||||
"rows": row_n,
|
||||
"listed": listed_n,
|
||||
"ratio": round(ratio, 4),
|
||||
"errors": errors,
|
||||
"warnings": warnings,
|
||||
"hard_fail": hard_fail,
|
||||
"soft_fail": bool(warnings) and not hard_fail,
|
||||
"batch_id": batch_id,
|
||||
}
|
||||
|
||||
def publish(self, dataset: str, trade_date: str, batch_id: str, state: str = "published") -> None:
|
||||
copy_sql = EOD_COPY[dataset]
|
||||
published_at = isoformat(self.clock())
|
||||
with self.db.write() as connection:
|
||||
current = connection.execute(
|
||||
"SELECT active_batch FROM publications WHERE dataset = ? AND trade_date = ?",
|
||||
(dataset, trade_date),
|
||||
).fetchone()
|
||||
prev = str(current["active_batch"]) if current else None
|
||||
connection.execute(copy_sql, (batch_id,))
|
||||
if self.before_commit:
|
||||
self.before_commit()
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO publications(dataset, trade_date, active_batch, prev_batch, state, published_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(dataset, trade_date) DO UPDATE SET
|
||||
prev_batch=excluded.prev_batch,
|
||||
active_batch=excluded.active_batch,
|
||||
state=excluded.state,
|
||||
published_at=excluded.published_at
|
||||
""",
|
||||
(dataset, trade_date, batch_id, prev, state, published_at),
|
||||
)
|
||||
max_gen = connection.execute(
|
||||
"SELECT COALESCE(MAX(generation), 0) AS g FROM publication_history WHERE dataset = ? AND trade_date = ?",
|
||||
(dataset, trade_date),
|
||||
).fetchone()
|
||||
generation = int(max_gen["g"]) + 1
|
||||
connection.execute(
|
||||
"INSERT OR REPLACE INTO publication_history(dataset, trade_date, batch_id, published_at, generation) VALUES (?,?,?,?,?)",
|
||||
(dataset, trade_date, batch_id, published_at, generation),
|
||||
)
|
||||
keep = int(self.settings.quality.get("publication_generations") or 3)
|
||||
stale = connection.execute(
|
||||
"""
|
||||
SELECT batch_id FROM publication_history
|
||||
WHERE dataset = ? AND trade_date = ?
|
||||
ORDER BY generation DESC
|
||||
""",
|
||||
(dataset, trade_date),
|
||||
).fetchall()
|
||||
for row in stale[keep:]:
|
||||
connection.execute(
|
||||
"DELETE FROM publication_history WHERE dataset = ? AND trade_date = ? AND batch_id = ?",
|
||||
(dataset, trade_date, row["batch_id"]),
|
||||
)
|
||||
|
||||
def rollback(self, dataset: str, trade_date: str, actor: str = "admin") -> dict[str, Any]:
|
||||
trade_date = yyyymmdd(trade_date)
|
||||
pub = self.db.fetchone(
|
||||
"SELECT * FROM publications WHERE dataset = ? AND trade_date = ?",
|
||||
(dataset, trade_date),
|
||||
)
|
||||
if not pub or not pub.get("prev_batch"):
|
||||
raise ValueError("没有可回滚的上一批次")
|
||||
target = pub["prev_batch"]
|
||||
published_at = isoformat(self.clock())
|
||||
with self.db.write() as connection:
|
||||
connection.execute(
|
||||
"""
|
||||
UPDATE publications
|
||||
SET prev_batch = active_batch, active_batch = ?, published_at = ?, state = 'published'
|
||||
WHERE dataset = ? AND trade_date = ?
|
||||
""",
|
||||
(target, published_at, dataset, trade_date),
|
||||
)
|
||||
self.audit(actor, "rollback", f"{dataset}:{trade_date}", json.dumps({"to": target, "from": pub["active_batch"]}))
|
||||
return {"dataset": dataset, "trade_date": trade_date, "active_batch": target, "prev_batch": pub["active_batch"]}
|
||||
|
||||
def active_batch(self, dataset: str, trade_date: str) -> str | None:
|
||||
row = self.db.fetchone(
|
||||
"SELECT active_batch FROM publications WHERE dataset = ? AND trade_date = ?",
|
||||
(dataset, trade_date),
|
||||
)
|
||||
return str(row["active_batch"]) if row else None
|
||||
|
||||
def cleanup(self) -> dict[str, int]:
|
||||
staging_days = int(self.settings.quality.get("staging_retain_days") or 14)
|
||||
job_days = int(self.settings.quality.get("job_run_retain_days") or 90)
|
||||
cutoff_staging = add_days(yyyymmdd(self.clock()), -staging_days)
|
||||
cutoff_jobs = add_days(yyyymmdd(self.clock()), -job_days)
|
||||
deleted = 0
|
||||
with self.db.write() as connection:
|
||||
for dataset, (_eod, staging) in DATASET_TABLES.items():
|
||||
cur = connection.execute(
|
||||
f"DELETE FROM {staging} WHERE trade_date < ?",
|
||||
(cutoff_staging,),
|
||||
)
|
||||
deleted += cur.rowcount
|
||||
connection.execute("DELETE FROM job_runs WHERE started_at < ?", (cutoff_jobs,))
|
||||
connection.execute("DELETE FROM src_calls WHERE created_at < ?", (cutoff_jobs,))
|
||||
return {"staging_deleted": deleted}
|
||||
|
||||
def audit(self, actor: str, action: str, target: str = "", detail: str = "") -> None:
|
||||
self.db.execute(
|
||||
"INSERT INTO audit_log(actor, action, target, detail, created_at) VALUES (?,?,?,?,?)",
|
||||
(actor, action, target, detail, isoformat(self.clock())),
|
||||
)
|
||||
|
||||
def _fetch_dataset(self, dataset: str, trade_date: str) -> list[dict[str, Any]]:
|
||||
if dataset == "daily":
|
||||
raw = self._guarded_fetch("daily", {"trade_date": trade_date})
|
||||
factors = {
|
||||
(row["ts_code"], row["trade_date"]): finite_number(row.get("adj_factor"))
|
||||
for row in self._guarded_fetch("adj_factor", {"trade_date": trade_date})
|
||||
}
|
||||
return [
|
||||
normalize_daily(row, adj_factor=factors.get((str(row.get("ts_code") or "").upper(), str(row.get("trade_date") or ""))))
|
||||
for row in raw
|
||||
]
|
||||
if dataset == "index_daily":
|
||||
rows: list[dict[str, Any]] = []
|
||||
for ts_code in DEFAULT_INDEX_CODES:
|
||||
raw = self._guarded_fetch("index_daily", {"ts_code": ts_code, "trade_date": trade_date})
|
||||
rows.extend(self.adapter.normalize("index_daily", raw))
|
||||
return rows
|
||||
api_dataset = dataset
|
||||
raw = self._guarded_fetch(api_dataset, {"trade_date": trade_date})
|
||||
return self.adapter.normalize(api_dataset, raw)
|
||||
|
||||
def _guarded_fetch(self, dataset: str, params: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
if not self.breaker.allow():
|
||||
raise AdapterError("Tushare circuit open")
|
||||
self.bucket.acquire()
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
# For daily we want RAW tushare rows so adj_factor can be merged later.
|
||||
rows = self.adapter.fetch(dataset, params)
|
||||
latency = round((time.perf_counter() - started) * 1000)
|
||||
self.breaker.record_success()
|
||||
self._log_call(dataset, True, latency, "")
|
||||
self._persist_health("ok")
|
||||
return rows
|
||||
except Exception as exc:
|
||||
latency = round((time.perf_counter() - started) * 1000)
|
||||
self.breaker.record_failure(str(exc))
|
||||
self._log_call(dataset, False, latency, str(exc))
|
||||
self._persist_health("error", str(exc))
|
||||
raise
|
||||
|
||||
def _stage(self, dataset: str, batch_id: str, rows: list[dict[str, Any]]) -> None:
|
||||
sql, mapper = STAGING_INSERT[dataset]
|
||||
with self.db.write() as connection:
|
||||
connection.execute(
|
||||
f"DELETE FROM {DATASET_TABLES[dataset][1]} WHERE batch_id = ?",
|
||||
(batch_id,),
|
||||
)
|
||||
connection.executemany(sql, [mapper(row, batch_id) for row in rows])
|
||||
|
||||
def _set_batch(
|
||||
self,
|
||||
batch_id: str,
|
||||
dataset: str,
|
||||
trade_date: str,
|
||||
state: str,
|
||||
attempt: int,
|
||||
rows_in: int | None = None,
|
||||
rows_out: int | None = None,
|
||||
quality: dict[str, Any] | None = None,
|
||||
error: str | None = None,
|
||||
finished: bool = False,
|
||||
) -> None:
|
||||
now = isoformat(self.clock())
|
||||
existing = self.db.fetchone("SELECT batch_id FROM batches WHERE batch_id = ?", (batch_id,))
|
||||
payload = json.dumps(quality, ensure_ascii=False) if quality else None
|
||||
with self.db.write() as connection:
|
||||
if existing is None:
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO batches(batch_id, dataset, trade_date, state, attempt, rows_in, rows_out, quality_json, started_at, finished_at, error)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(batch_id, dataset, trade_date, state, attempt, rows_in, rows_out, payload, now, now if finished else None, error),
|
||||
)
|
||||
else:
|
||||
connection.execute(
|
||||
"""
|
||||
UPDATE batches SET state=?, attempt=?,
|
||||
rows_in=COALESCE(?, rows_in), rows_out=COALESCE(?, rows_out),
|
||||
quality_json=COALESCE(?, quality_json),
|
||||
finished_at=CASE WHEN ? THEN ? ELSE finished_at END,
|
||||
error=COALESCE(?, error)
|
||||
WHERE batch_id = ?
|
||||
""",
|
||||
(state, attempt, rows_in, rows_out, payload, 1 if finished else 0, now, error, batch_id),
|
||||
)
|
||||
|
||||
def _log_call(self, endpoint: str, ok: bool, latency_ms: int, error: str) -> None:
|
||||
self.db.execute(
|
||||
"INSERT INTO src_calls(provider, endpoint, ok, latency_ms, error, created_at) VALUES (?,?,?,?,?,?)",
|
||||
("tushare", endpoint, 1 if ok else 0, latency_ms, error, isoformat(self.clock())),
|
||||
)
|
||||
|
||||
def _persist_health(self, state: str, error: str = "") -> None:
|
||||
snap = self.breaker.snapshot()
|
||||
self.db.execute(
|
||||
"""
|
||||
INSERT INTO src_health(provider, endpoint_class, state, last_ok_at, last_error, consec_failures, opened_at, cooldown_until)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(provider, endpoint_class) DO UPDATE SET
|
||||
state=excluded.state, last_ok_at=excluded.last_ok_at, last_error=excluded.last_error,
|
||||
consec_failures=excluded.consec_failures, opened_at=excluded.opened_at, cooldown_until=excluded.cooldown_until
|
||||
""",
|
||||
(
|
||||
"tushare", "pro",
|
||||
snap.state,
|
||||
isoformat(self.clock()) if state == "ok" else None,
|
||||
error or snap.last_error,
|
||||
snap.consec_failures,
|
||||
isoformat(self.clock()) if snap.state == "open" else None,
|
||||
None,
|
||||
),
|
||||
)
|
||||
@@ -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)}
|
||||
@@ -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
|
||||
@@ -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"},
|
||||
)
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user