224 lines
8.0 KiB
Python
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}
|