diff --git a/sanguo_trader/engine.py b/sanguo_trader/engine.py index 1c80247..2535c78 100644 --- a/sanguo_trader/engine.py +++ b/sanguo_trader/engine.py @@ -54,41 +54,48 @@ class PaperEngine: 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. 撮合上一根 pending(next_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 待下根撮合 - bar_count = 0 for bar_date, bars in self.data_source.iter_bars( self.symbols, self.start, self.end, self.interval, self.adjust, None ): - bar_count += 1 - self.account.unfreeze_all() - for r in self.runners: - r.unfreeze_all() - # 1. 撮合上一根 pending(next_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=(bar_count % 500 == 0), - ) - update_checkpoint(self.db_path, self.account_id, str(bar_date)) + pending, closes = self.step(bar_date, bars, prev_close, pending) prev_close = closes def _match(self, order, runner, bars, prev_close, bar_date) -> None: diff --git a/tests/trader/test_engine.py b/tests/trader/test_engine.py index b6e168a..e48ef3c 100644 --- a/tests/trader/test_engine.py +++ b/tests/trader/test_engine.py @@ -91,3 +91,23 @@ def test_engine_daily_balance_and_consistency(tmp_path): # 总账持仓 = 分户持仓(day2+day3 各买100 = 200) assert account.positions["600000"].volume == 200 assert runner.positions["600000"].volume == 200 + + +def test_engine_step_single_bar_advances(tmp_path): + """engine.step 单根推进(C-S3 实走每日入口):day1 信号缓冲,day2 撮合。""" + sections = [ + ("2024-01-01", {"600000": _bar("2024-01-01", 10.0, 10.5, 9.5, 10.0)}), + ("2024-01-02", {"600000": _bar("2024-01-02", 10.5, 11.0, 10.0, 10.8)}), + ] + pe, db, aid, account, runner = _build(tmp_path, sections) + # day1 step:_AlwaysBuyStrategy 买单(NEXT_OPEN)→ 进 pending,当根不撮合 + pending, closes = pe.step("2024-01-01", sections[0][1], {}, []) + assert len(pending) == 1 + assert account.positions.get("600000") is None + assert closes["600000"] == 10.0 + # day2 step:撮合 day1 pending @ open 10.5;策略 on_bar(day2) 又发单进 pending 等 day3 + pending2, closes2 = pe.step("2024-01-02", sections[1][1], closes, pending) + assert len(pending2) == 1 # day2 新信号(无 day3 不撮合) + assert account.positions["600000"].volume == 100 # day1 单 day2 open 10.5 撮合 100 股 + # step 入库(day1+day2 各一条余额) + assert len(list_daily_balance(db, aid)) == 2