Files

224 lines
8.0 KiB
Python

"""数据同步 — 接收 local 端推送的 Parquet 数据包。"""
from __future__ import annotations
import logging
import os
import tarfile
import tempfile
from io import BytesIO
from pathlib import Path
import polars as pl
from fastapi import APIRouter, Depends, HTTPException, Request, UploadFile, File, Form
from app.config import settings
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/data/sync", tags=["data-sync"])
# 允许同步的子目录列表
SYNCABLE_PARTS = {
"kline_daily", "kline_daily_enriched",
"kline_index_daily", "kline_index_enriched",
"kline_etf_daily", "kline_etf_enriched", "kline_etf_minute",
"kline_minute",
"adj_factor", "adj_factor_etf",
"instruments", "instruments_index", "instruments_etf",
"financials",
"pools",
"ext_data",
}
def _resolve_unique_key(columns: list[str]) -> list[str] | None:
"""根据列名推断去重键。
策略:
- 含 symbol + datetime: 按 [symbol, datetime] 去重(分钟 K)
- 含 symbol + date: 按 [symbol, date] 去重(日 K / 复权因子)
- 含 symbol: 按 [symbol] 去重(标的维表)
- 其他: 无法推断,回退到覆盖
"""
cols = set(columns)
if "symbol" not in cols:
return None
if "datetime" in cols:
return ["symbol", "datetime"]
if "date" in cols:
return ["symbol", "date"]
return ["symbol"]
def _merge_parquet(target: Path, new_bytes: bytes) -> tuple[bool, str]:
"""把 new_bytes 代表的 Parquet 与 target 已有文件合并去重。
返回 (是否成功落盘, 操作描述):
- created: 目标不存在,直接新建
- merged: 与已有文件按主键合并去重(保留后写入的记录)
- empty_new: 上传文件为空,保留旧文件
- overwritten_no_key: 无法推断去重键,回退覆盖
- overwritten_corrupt_existing: 已有文件损坏,回退覆盖
- invalid_new: 上传文件不是合法 Parquet,未落盘
"""
if not target.exists():
target.parent.mkdir(parents=True, exist_ok=True)
target.write_bytes(new_bytes)
return True, "created"
try:
new_df = pl.read_parquet(BytesIO(new_bytes))
except Exception as exc:
logger.warning("failed to parse uploaded parquet %s: %s", target.name, exc)
return False, "invalid_new"
if new_df.is_empty():
return True, "empty_new"
key = _resolve_unique_key(new_df.columns)
if key is None:
logger.warning(
"no unique key for %s (columns=%s), overwriting",
target.name,
new_df.columns,
)
target.write_bytes(new_bytes)
return True, "overwritten_no_key"
try:
existing_df = pl.read_parquet(target)
except Exception as exc:
logger.warning("existing parquet %s corrupt, overwriting: %s", target, exc)
target.write_bytes(new_bytes)
return True, "overwritten_corrupt_existing"
if existing_df.is_empty():
merged = new_df
action = "created"
else:
merged = pl.concat([existing_df, new_df], how="diagonal_relaxed").unique(
subset=key, keep="last"
)
action = "merged"
# 原子写入,防止写入过程中崩溃导致文件损坏
tmp = target.with_suffix(f".tmp-{os.getpid()}")
try:
merged.write_parquet(tmp)
os.replace(tmp, target)
except Exception:
if tmp.exists():
tmp.unlink(missing_ok=True)
raise
return True, action
def _verify_sync_key(request: Request) -> None:
"""简单的鉴权 — 校验 X-Sync-Key 头。"""
key = request.headers.get("X-Sync-Key", "")
expected = getattr(settings, "sync_key", "")
if expected and key != expected:
raise HTTPException(status_code=403, detail="sync key mismatch")
def _refresh_views(request: Request) -> None:
"""重新注册 DuckDB 视图 + 刷新 Polars 缓存。"""
repo = getattr(request.app.state, "repo", None)
if repo is None:
return
try:
d = repo.store.data_dir.as_posix()
views = {
"kline_daily": f"{d}/kline_daily/**/*.parquet",
"kline_enriched": f"{d}/kline_daily_enriched/**/*.parquet",
"kline_index_daily": f"{d}/kline_index_daily/**/*.parquet",
"kline_index_enriched": f"{d}/kline_index_enriched/**/*.parquet",
"kline_etf_daily": f"{d}/kline_etf_daily/**/*.parquet",
"kline_etf_enriched": f"{d}/kline_etf_enriched/**/*.parquet",
"kline_etf_minute": f"{d}/kline_etf_minute/**/*.parquet",
"kline_minute": f"{d}/kline_minute/**/*.parquet",
"adj_factor": f"{d}/adj_factor/**/*.parquet",
"adj_factor_etf": f"{d}/adj_factor_etf/**/*.parquet",
"instruments": f"{d}/instruments/**/*.parquet",
"instruments_index": f"{d}/instruments_index/**/*.parquet",
"instruments_etf": f"{d}/instruments_etf/**/*.parquet",
}
for name, path in views.items():
try:
repo.store.db.execute(
f"CREATE OR REPLACE VIEW {name} AS "
f"SELECT * FROM read_parquet('{path}', union_by_name=true)"
)
except Exception as exc:
logger.warning("view %s refresh failed: %s", name, exc)
repo.store._register_unified_views()
repo.clear_cache()
repo.refresh_cache()
# 清除 API 数据缓存
from app.api.data import invalidate_data_cache
invalidate_data_cache()
except Exception as exc:
logger.warning("cache refresh failed: %s", exc)
@router.post("/upload")
async def upload_sync(
request: Request,
file: UploadFile = File(...),
parts: str = Form(""),
_auth: None = Depends(_verify_sync_key),
):
"""接收 local 端打包的 tar.gz 数据,解压写入 data/ 目录。"""
data_dir: Path = request.app.state.datastore.data_dir
# 解析要同步的 parts
requested = set(p.strip() for p in parts.split(",") if p.strip())
if not requested:
raise HTTPException(status_code=400, detail="parts 不能为空,逗号分隔子目录名")
invalid = requested - SYNCABLE_PARTS
if invalid:
raise HTTPException(
status_code=400,
detail=f"不支持的同步目录: {', '.join(sorted(invalid))}",
)
# 读取上传的 tar.gz
raw = await file.read()
if not raw:
raise HTTPException(status_code=400, detail="上传文件为空")
extracted = 0
action_counter: dict[str, int] = {}
try:
with tarfile.open(fileobj=BytesIO(raw), mode="r:gz") as tar:
for member in tar.getmembers():
# 安全校验: 防止路径穿越
member_path = Path(member.name)
if member_path.is_absolute() or ".." in member_path.parts:
logger.warning("skipping unsafe path: %s", member.name)
continue
# 确定属于哪个 part
top_dir = member_path.parts[0] if member_path.parts else ""
if top_dir not in requested:
continue
target = data_dir / member.name
if member.isfile():
target.parent.mkdir(parents=True, exist_ok=True)
with tar.extractfile(member) as src:
if src is None:
continue
ok, action = _merge_parquet(target, src.read())
if ok:
extracted += 1
action_counter[action] = action_counter.get(action, 0) + 1
except tarfile.TarError as exc:
raise HTTPException(status_code=400, detail=f"压缩包解析失败: {exc}") from exc
# 刷新缓存
_refresh_views(request)
logger.info("sync uploaded: parts=%s files=%d actions=%s", parts, extracted, action_counter)
return {"ok": True, "parts": sorted(requested), "file_count": extracted, "actions": action_counter}