fix(paper): create 接入回放(engine.run) + 日线 parquet 文件名 sh/sz 前缀+_daily

This commit is contained in:
2026-07-07 14:47:20 +08:00
parent 88bd001a97
commit eb6aa34b82
2 changed files with 60 additions and 3 deletions
+58 -2
View File
@@ -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 给 PaperEngineengine 需 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()
+2 -1
View File
@@ -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)