Files
stock/serve/backend/app/api/data_sync.py
T
2026-07-06 10:27:00 +08:00

135 lines
5.0 KiB
Python

"""数据同步 — 接收 local 端推送的 Parquet 数据包。"""
from __future__ import annotations
import logging
import tarfile
import tempfile
from io import BytesIO
from pathlib import Path
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",
}
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
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
target.write_bytes(src.read())
extracted += 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", parts, extracted)
return {"ok": True, "parts": sorted(requested), "file_count": extracted}