From eb6aa34b82fdfbbebba63cb7cc11a75b824f67bd Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 7 Jul 2026 14:47:20 +0800 Subject: [PATCH] =?UTF-8?q?fix(paper):=20create=20=E6=8E=A5=E5=85=A5?= =?UTF-8?q?=E5=9B=9E=E6=94=BE(engine.run)=20+=20=E6=97=A5=E7=BA=BF=20parqu?= =?UTF-8?q?et=20=E6=96=87=E4=BB=B6=E5=90=8D=20sh/sz=20=E5=89=8D=E7=BC=80+?= =?UTF-8?q?=5Fdaily?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_api/routes_paper.py | 60 ++++++++++++++++++++++++++++++++++++-- sanguo_data/datareader.py | 3 +- 2 files changed, 60 insertions(+), 3 deletions(-) diff --git a/sanguo_api/routes_paper.py b/sanguo_api/routes_paper.py index 9249f54..fee67c9 100644 --- a/sanguo_api/routes_paper.py +++ b/sanguo_api/routes_paper.py @@ -55,12 +55,21 @@ class PaperCreateRequest(BaseModel): @router.post("/paper/create", dependencies=[Depends(verify_token)]) def create_paper(req: PaperCreateRequest): - from sanguo_trader.persistence import init_db, save_account + from sanguo_trader.persistence import init_db, save_account, update_account_status db = _db_path["path"] or ":memory:" init_db(db) aid = save_account(db, req.model_dump()) - return {"account_id": aid, "status": "created"} + status = "created" + if req.mode == "replay": # 回放模式:建 account 后同步跑 engine.run + try: + _run_replay(db, aid, req) + status = "done" + update_account_status(db, aid, "done") + except Exception as e: # noqa: BLE001 — 回放失败不阻塞 create,状态记 failed + update_account_status(db, aid, "failed", str(e)) + status = "failed" + return {"account_id": aid, "status": status} @router.get("/paper/{aid}", dependencies=[Depends(verify_token)]) @@ -96,3 +105,50 @@ def get_strategies(aid: int): from sanguo_trader.persistence import list_strategy_summary return list_strategy_summary(_db_path["path"], aid) + + +class _DataSourceWrapper: + """包装 iter_bars 给 PaperEngine(engine 需 data_source.iter_bars 接口)。""" + + def __init__(self, cfg): + self.cfg = cfg + + def iter_bars(self, symbols, start, end, interval, adjust="qfq", cfg=None): + from sanguo_trader.data_source import iter_bars + + return iter_bars(symbols, start, end, interval, adjust, cfg or self.cfg) + + +def _run_replay(db, aid, req: PaperCreateRequest): + """构造引擎 + 跑回放(容器内有 vnpy_ctastrategy + NAS parquet,本机仅空转)。""" + from sanguo_trader.account import Account + from sanguo_trader.cta_adapter import PaperCtaEngine + from sanguo_trader.engine import PaperEngine + from sanguo_trader.models import AccountConfig + from sanguo_trader.strategy_runner import StrategyRunner + from sanguo_data.config import find_config_path, load_config + from sanguo_data.datareader import guess_exchange + from .strategy_registry import get_strategy_class + + data_cfg = load_config(find_config_path()) + acc_cfg = AccountConfig( + initial_capital=req.initial_capital, rate=req.rate, slippage=req.slippage, + pricetick=req.pricetick, stamp_duty_rate=req.stamp_duty_rate, + transfer_fee_rate=req.transfer_fee_rate, min_commission=req.min_commission, + ) + account = Account(req.initial_capital) + runners: list = [] + for s in req.strategies: + cls = get_strategy_class(s.name) + if cls is None: + continue # 策略不可用(本机无 vnpy_ctastrategy)→ 跳过 + cta = PaperCtaEngine(s.name, match_session=s.match_session, + listing_days=s.listing_days) + vt_symbol = f"{s.symbol}.{guess_exchange(s.symbol).value}" + strat = cls(cta, s.name, vt_symbol, s.params) # CtaTemplate(cta_engine, name, vt_symbol, setting) + cta.set_strategy(strat) + runners.append(StrategyRunner(s.name, strategy=strat, paper_cta_engine=cta, + symbol=s.symbol)) + pe = PaperEngine(account, runners, _DataSourceWrapper(data_cfg), acc_cfg, + db, aid, req.symbols, req.start, req.end, req.interval) + pe.run() diff --git a/sanguo_data/datareader.py b/sanguo_data/datareader.py index ff5d6bf..da76889 100644 --- a/sanguo_data/datareader.py +++ b/sanguo_data/datareader.py @@ -19,8 +19,9 @@ def read_parquet_daily(symbol: str, start: str, end: str, cfg) -> list[BarData]: start_dt = datetime.strptime(start, "%Y-%m-%d") end_dt = datetime.strptime(end, "%Y-%m-%d") bars: list[BarData] = [] + prefix = "sh" if guess_exchange(symbol) == Exchange.SSE else "sz" for year in range(start_dt.year, end_dt.year + 1): - f = daily_dir / str(year) / f"{symbol}.parquet" + f = daily_dir / str(year) / f"{prefix}{symbol}_daily.parquet" if not f.exists(): continue df = pd.read_parquet(f)