Files
qlib/backend/app/cli/sync.py
T
Simon bbb5c1ea52 feat(cli): sync export — 日线按年导出 Parquet(data/parquet/stock_daily/<year>.parquet)
- 对应 ROADMAP §1.4「历史时序大数据转 Parquet」;pyarrow 随 qlib 依赖已可用
- 用法:uv run python -m app.cli.sync export [--years 2023,2024]
- 实测:9680 行 → 2023/2024 两个 parquet(data/parquet 已被 gitignore)
- ruff clean / pytest 86 passed
2026-09-06 18:21:55 +08:00

270 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Phase 1 数据同步 CLI(Tushare 首选 → SQLite)。
用法(cd backend):
uv run python -m app.cli.sync basic
uv run python -m app.cli.sync calendar --start 20240101 --end 20241231
uv run python -m app.cli.sync daily --symbols 600519.SH,000001.SZ --start 20240101
uv run python -m app.cli.sync daily --all --start 20240101 # 全市场
uv run python -m app.cli.sync financial --all
uv run python -m app.cli.sync verify --symbol 600519.SH # 新浪交叉验证
本模块是组装层(composition root):在此装配 Provider / Repository / Session,
业务层代码仍只依赖抽象(domain.repositories / domain.providers)。
"""
from __future__ import annotations
import argparse
import sys
from datetime import date, datetime, timedelta
from sqlalchemy import select
from app.core.config import get_settings
from app.infrastructure.data_sources.errors import DataSourceError
from app.infrastructure.data_sources.failover import FailoverProvider
from app.infrastructure.data_sources.sina import SinaProvider
from app.infrastructure.data_sources.tushare import TushareProvider
from app.infrastructure.persistence.sqlalchemy.models.market import StockModel
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyAdjustFactorRepository,
SqlAlchemyDailyBarRepository,
SqlAlchemyFinancialRepository,
SqlAlchemyStockRepository,
SqlAlchemySyncLogRepository,
SqlAlchemyTradingCalendarRepository,
)
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
_DATE_FMT = "%Y%m%d"
def _parse_day(text: str) -> date:
return datetime.strptime(text, _DATE_FMT).date()
def _failover_provider(session):
"""Tushare 单源经 FailoverProvider 包装:每次尝试写 sync_log(AGENT.md §7 审计)。"""
audit_repo = SqlAlchemySyncLogRepository(session)
primary = TushareProvider(token=get_settings().tushare_token)
return FailoverProvider(primary, fallback=None, audit=audit_repo.add)
def _session_ctx():
return SessionLocal()
def cmd_basic(args) -> int:
with _session_ctx() as session:
provider = _failover_provider(session)
stocks = provider.get_stock_basic()
repo = SqlAlchemyStockRepository(session)
touched = repo.upsert_many(stocks)
session.commit()
print(f"[basic] 拉取 {len(stocks)} 只,落库 {touched} 条")
return 0
def cmd_calendar(args) -> int:
start = _parse_day(args.start)
end = _parse_day(args.end)
with _session_ctx() as session:
provider = _failover_provider(session)
days = provider.get_trade_cal(start, end)
repo = SqlAlchemyTradingCalendarRepository(session)
touched = repo.upsert_many(days)
session.commit()
open_days = sum(1 for d in days if d.is_open)
print(f"[calendar] {start}~{end} 共 {len(days)} 条(交易日 {open_days}),落库 {touched} 条")
return 0
def _symbols_of(args) -> list[str]:
if getattr(args, "all", False):
with _session_ctx() as session:
symbols = list(session.scalars(select(StockModel.symbol).order_by(StockModel.symbol)))
if not symbols:
print("[error] stock 表为空,请先运行:python -m app.cli.sync basic")
sys.exit(2)
return symbols
return [s.strip() for s in args.symbols.split(",") if s.strip()]
def cmd_daily(args) -> int:
symbols = _symbols_of(args)
start = _parse_day(args.start) if args.start else date(2005, 1, 1)
end = _parse_day(args.end) if args.end else date.today()
total = 0
with _session_ctx() as session:
provider = _failover_provider(session)
bar_repo = SqlAlchemyDailyBarRepository(session)
factor_repo = SqlAlchemyAdjustFactorRepository(session)
for i, symbol in enumerate(symbols, start=1):
begin = start
if args.resume:
latest = bar_repo.latest_date(symbol)
if latest is not None:
begin = max(begin, latest + timedelta(days=1))
if begin > end:
continue
try:
bars = provider.get_daily(symbol, begin, end)
factors = provider.get_adjust_factor(symbol, begin, end)
bar_repo.upsert_many(bars)
factor_repo.upsert_many(factors)
total += len(bars)
if i % 100 == 0:
session.commit()
print(f" ... {i}/{len(symbols)} {symbol} 累计 {total} 根")
except DataSourceError as exc:
print(f" [warn] {symbol} 拉取失败: {exc}", file=sys.stderr)
session.commit()
print(f"[daily] {len(symbols)} 只股票合计写入 {total} 根日线(含复权因子)")
return 0
def cmd_financial(args) -> int:
symbols = _symbols_of(args)
total = 0
with _session_ctx() as session:
provider = _failover_provider(session)
fin_repo = SqlAlchemyFinancialRepository(session)
for symbol in symbols:
rows = provider.get_financial(symbol)
fin_repo.upsert_many(rows)
total += len(rows)
session.commit()
print(f"[financial] {len(symbols)} 只股票合计写入 {total} 条财务指标快照")
return 0
def cmd_verify(args) -> int:
"""新浪交叉验证:取新浪最新前复权收盘,与本地最新交易日对照。
注意:新浪为前复权口径,数值不直接等于本地不复权收盘,
本命令仅用于确认新浪可用性 / 最新交易日,不把新浪数据并入主库。
"""
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
)
sina = SinaProvider()
end = date.today()
start = end - timedelta(days=20)
try:
bars = sina.get_daily(args.symbol, start, end)
except DataSourceError as exc:
print(f"[verify] 新浪不可用: {exc}", file=sys.stderr)
return 1
if not bars:
print(f"[verify] 新浪最近无数据({args.symbol})")
return 1
latest = max(bars, key=lambda b: b.trade_date)
with _session_ctx() as session:
local = SqlAlchemyDailyBarRepository(session).latest_date(args.symbol)
print(
f"[verify] {args.symbol}: 新浪最新 {latest.trade_date} 收盘(前复权) {latest.close};"
f"本地最新交易日 {local}"
)
return 0
def cmd_export(args) -> int:
"""把 SQLite 日线按年导出为 Parquet(data/parquet/stock_daily/<year>.parquet)。"""
from pathlib import Path
import pandas as pd
from app.infrastructure.persistence.sqlalchemy.models.market import StockDailyModel
settings = get_settings()
out_root = settings.storage.get("parquet_dir") or Path("data/parquet")
out_root.mkdir(parents=True, exist_ok=True)
total = 0
years = [int(y) for y in (args.years or "").split(",") if y.strip()] or None
with _session_ctx() as session:
all_bars = session.execute(
select(StockDailyModel).order_by(StockDailyModel.trade_date)
).scalars()
frame = pd.DataFrame(
[
{
"symbol": b.symbol,
"trade_date": b.trade_date,
"open": float(b.open) if b.open is not None else None,
"high": float(b.high) if b.high is not None else None,
"low": float(b.low) if b.low is not None else None,
"close": float(b.close) if b.close is not None else None,
"volume": float(b.volume) if b.volume is not None else None,
"amount": float(b.amount) if b.amount is not None else None,
}
for b in all_bars
]
)
if frame.empty:
print("[export] 无日线数据,请先运行 sync daily")
return 0
frame["trade_date"] = pd.to_datetime(frame["trade_date"])
out_dir = out_root / "stock_daily"
out_dir.mkdir(parents=True, exist_ok=True)
for year, group in frame.groupby(frame["trade_date"].dt.year):
if years and int(year) not in years:
continue
path = out_dir / f"{year}.parquet"
group.sort_values(["symbol", "trade_date"]).to_parquet(path, index=False)
total += len(group)
print(f"[export] {year} → {path}({len(group)} 行)")
print(f"[export] 合计 {total} 行 → {out_dir}")
return 0
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(prog="app.cli.sync", description="Tushare 数据同步 CLI")
sub = parser.add_subparsers(dest="command", required=True)
p_basic = sub.add_parser("basic", help="同步股票基础信息")
p_basic.set_defaults(func=cmd_basic)
p_cal = sub.add_parser("calendar", help="同步交易日历")
p_cal.add_argument("--start", required=True, help="YYYYMMDD")
p_cal.add_argument("--end", required=True, help="YYYYMMDD")
p_cal.set_defaults(func=cmd_calendar)
p_daily = sub.add_parser("daily", help="同步日线与复权因子")
p_daily.add_argument("--symbols", default="", help="600519.SH,000001.SZ")
p_daily.add_argument("--all", action="store_true", help="遍历 stock 表全部股票")
p_daily.add_argument("--start", default="", help="YYYYMMDD(默认 20050101)")
p_daily.add_argument("--end", default="", help="YYYYMMDD(默认今天)")
p_daily.add_argument("--resume", action="store_true", help="从本地最新交易日续传")
p_daily.set_defaults(func=cmd_daily)
p_fin = sub.add_parser("financial", help="同步财务指标快照")
p_fin.add_argument("--symbols", default="")
p_fin.add_argument("--all", action="store_true")
p_fin.set_defaults(func=cmd_financial)
p_verify = sub.add_parser("verify", help="新浪交叉验证最新行情")
p_verify.add_argument("--symbol", required=True)
p_verify.set_defaults(func=cmd_verify)
p_export = sub.add_parser("export", help="日线按年导出 Parquet(data/parquet)")
p_export.add_argument("--years", default="", help="逗号分隔年份,留空导出全部")
p_export.set_defaults(func=cmd_export)
return parser
def main(argv: list[str] | None = None) -> int:
args = build_parser().parse_args(argv)
try:
return args.func(args)
except DataSourceError as exc:
print(f"[error] {exc}", file=sys.stderr)
return 1
except KeyboardInterrupt:
print("\n[interrupt] 已中止", file=sys.stderr)
return 130
if __name__ == "__main__":
raise SystemExit(main())