Files
sanguo_vnpy_v2/sanguo_trader/engine.py
T
claude_dev 3aca14f723 refactor(engine): 抽出step()单根推进(C-S3实走入口,task2基础)
run()循环体抽为step(bar_date,bars,prev_close,pending)→(pending,closes);
run()改为调step。实走scheduler每日喂当日bar调step单步推进。
- 回放行为不变(test_engine原3用例pass)
- 加test_engine_step_single_bar_advances: day1信号缓冲/day2撮合
- trader全量111/111通过
2026-07-07 23:41:18 +08:00

133 lines
5.7 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.
"""PaperEngine 模拟盘主循环(逐根 bar 重放 + 双层记账 + 持久化,spec §4/§9)。
run():逐 bar → T+1 解冻 → 撮合上一根 next_open pending(用当前 bar)→ 喂策略
on_bar 收新单 → current_close 当根撮合 / next_open 缓冲到下根 → 盯市 → 入库。
adjust 默认 raw(真实价):撮合/涨跌停/成交价/信号共用一套真实价 bar。
除权缺口对 MA 信号的影响留分期项(分红除权)处理。
"""
import logging
import pandas as pd
from .account import Account
from .matcher import cross_order
from .models import MatchSession, OrderSide, PaperTrade
from .persistence import save_daily_balance, save_trade, update_checkpoint
from .strategy_runner import StrategyRunner
logger = logging.getLogger(__name__)
def _to_series(bar) -> pd.Series:
return pd.Series({
"open": bar.open_price, "high": bar.high_price,
"low": bar.low_price, "close": bar.close_price,
"date": getattr(bar, "datetime", ""),
})
def _trade_to_dict(t: PaperTrade, bar_date) -> dict:
return {
"strategy_id": t.strategy_id, "symbol": t.symbol,
"direction": t.side.value, "price": t.price, "volume": t.volume,
"commission": t.commission, "stamp_duty": t.stamp_duty,
"transfer_fee": t.transfer_fee,
"bar_date": str(bar_date), "match_session": t.match_session.value,
}
class PaperEngine:
def __init__(self, account: Account, runners: list[StrategyRunner],
data_source, cfg, db_path: str, account_id: int,
symbols: list[str], start: str, end: str,
interval: str = "d", adjust: str = "raw") -> None:
self.account = account
self.runners = runners
self.data_source = data_source
self.cfg = cfg
self.db_path = db_path
self.account_id = account_id
self.symbols = symbols
self.start = start
self.end = end
self.interval = interval
self.adjust = adjust
def step(self, bar_date, bars, prev_close, pending):
"""单根 bar 推进(回放 run 循环调;C-S3 实走 scheduler 每日调)。
返回 (新 pending, 当根 closes)——实走每日喂当日 bar 调一次。
"""
self._bar_count = getattr(self, "_bar_count", 0) + 1
self.account.unfreeze_all()
for r in self.runners:
r.unfreeze_all()
# 1. 撮合上一根 pendingnext_open,用当前 bar
if pending:
for order, runner in pending:
self._match(order, runner, bars, prev_close, bar_date)
pending = []
# 2. 喂策略 on_bar → 收新单
for runner in self.runners:
sym = runner.symbol
if sym and sym in bars:
runner.paper_cta_engine.on_bar(bars[sym])
for order in runner.paper_cta_engine.pop_orders():
if order.match_session == MatchSession.NEXT_OPEN:
pending.append((order, runner))
else: # current_close 当根撮合
self._match(order, runner, bars, prev_close, bar_date)
# 3. 盯市 + 入库
closes = {s: bars[s].close_price for s in bars}
self.account.mark_to_market(closes)
save_daily_balance(
self.db_path, self.account_id, str(bar_date),
self.account.cash, self.account.market_value, self.account.equity,
is_checkpoint=(self._bar_count % 500 == 0),
)
update_checkpoint(self.db_path, self.account_id, str(bar_date))
return pending, closes
def run(self) -> None:
prev_close: dict[str, float] = {}
pending: list = [] # [(order, runner)] next_open 待下根撮合
for bar_date, bars in self.data_source.iter_bars(
self.symbols, self.start, self.end, self.interval, self.adjust, None
):
pending, closes = self.step(bar_date, bars, prev_close, pending)
prev_close = closes
def _match(self, order, runner, bars, prev_close, bar_date) -> None:
if order.symbol not in bars:
return
match_bar = _to_series(bars[order.symbol])
pc = prev_close.get(order.symbol, order.price)
result = cross_order(order, match_bar, pc, self.cfg)
if isinstance(result, PaperTrade):
if result.side == OrderSide.SELL:
# A 股不能做空:SELL 超过可卖持仓 → 拒单(不开空仓)
pos = self.account.positions.get(order.symbol)
if pos is None or pos.available < result.volume:
save_trade(self.db_path, self.account_id,
{"strategy_id": order.strategy_id, "symbol": order.symbol,
"bar_date": str(bar_date)},
rejected=True,
reject_reason="insufficient_position_no_short")
return
if self.account.cash_enough(result):
self.account.apply_trade(result)
runner.apply_trade(result)
save_trade(self.db_path, self.account_id,
_trade_to_dict(result, bar_date))
else:
save_trade(self.db_path, self.account_id,
{"strategy_id": order.strategy_id, "symbol": order.symbol,
"bar_date": str(bar_date)},
rejected=True, reject_reason="insufficient_cash")
else: # PaperReject
save_trade(self.db_path, self.account_id,
{"strategy_id": order.strategy_id, "symbol": order.symbol,
"bar_date": str(bar_date)},
rejected=True, reject_reason=result.reason)