42877213ae
- run(): T+1解冻→撮合上根pending(用当前bar)→喂策略收单→current_close当根/next_open缓冲→盯市入库 - 双层记账一致性(总账=分户之和), checkpoint续跑字段 - StrategyRunner +symbol 字段 3 tests passed.
94 lines
3.6 KiB
Python
94 lines
3.6 KiB
Python
"""PaperEngine 主循环测试(mock data_source + mock 策略,spec §4)。"""
|
||
from types import SimpleNamespace
|
||
|
||
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, MatchSession
|
||
from sanguo_trader.persistence import (
|
||
init_db, list_daily_balance, list_trades, save_account,
|
||
)
|
||
from sanguo_trader.strategy_runner import StrategyRunner
|
||
|
||
|
||
def _bar(date, o, h, l, c):
|
||
return SimpleNamespace(open_price=o, high_price=h, low_price=l,
|
||
close_price=c, datetime=date)
|
||
|
||
|
||
class _FakeDataSource:
|
||
def __init__(self, sections):
|
||
self.sections = sections
|
||
|
||
def iter_bars(self, symbols, start, end, interval, adjust="qfq", cfg=None):
|
||
del symbols, start, end, interval, adjust, cfg
|
||
for date, bars in self.sections:
|
||
yield date, bars
|
||
|
||
|
||
class _AlwaysBuyStrategy:
|
||
def __init__(self, engine, vt_symbol):
|
||
self.cta_engine = engine
|
||
self.vt_symbol = vt_symbol
|
||
|
||
def on_bar(self, bar):
|
||
self.cta_engine.send_order(self, "LONG", "OPEN", bar.close_price, 100)
|
||
|
||
|
||
def _build(tmp_path, sections, match_session=MatchSession.NEXT_OPEN):
|
||
db = str(tmp_path / "e.db")
|
||
init_db(db)
|
||
aid = save_account(db, {"name": "t", "initial_capital": 1_000_000})
|
||
cfg = AccountConfig(initial_capital=1_000_000)
|
||
account = Account(1_000_000)
|
||
cta = PaperCtaEngine("s1", match_session=match_session)
|
||
strat = _AlwaysBuyStrategy(cta, "600000.SSE")
|
||
cta.set_strategy(strat)
|
||
runner = StrategyRunner("s1", strategy=strat, paper_cta_engine=cta, symbol="600000")
|
||
pe = PaperEngine(account, [runner], _FakeDataSource(sections), cfg, db, aid,
|
||
symbols=["600000"], start="2024-01-01", end="2024-12-31")
|
||
return pe, db, aid, account, runner
|
||
|
||
|
||
def test_engine_next_open_fills_at_next_bar_open(tmp_path):
|
||
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)}),
|
||
("2024-01-03", {"600000": _bar("2024-01-03", 10.8, 11.5, 10.5, 11.0)}),
|
||
]
|
||
pe, db, aid, *_ = _build(tmp_path, sections)
|
||
pe.run()
|
||
fills = [t for t in list_trades(db, aid) if not t["rejected"]]
|
||
# day1 信号→day2 撮合@10.5(open);day2 信号→day3 撮合@10.8(open);day3 信号无day4
|
||
assert len(fills) == 2
|
||
assert fills[0]["price"] == 10.5
|
||
assert fills[1]["price"] == 10.8
|
||
|
||
|
||
def test_engine_current_close_fills_same_bar(tmp_path):
|
||
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, *_ = _build(tmp_path, sections, match_session=MatchSession.CURRENT_CLOSE)
|
||
pe.run()
|
||
fills = [t for t in list_trades(db, aid) if not t["rejected"]]
|
||
assert len(fills) == 2
|
||
assert fills[0]["price"] == 10.0 # 当根 close
|
||
assert fills[1]["price"] == 10.8
|
||
|
||
|
||
def test_engine_daily_balance_and_consistency(tmp_path):
|
||
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)}),
|
||
("2024-01-03", {"600000": _bar("2024-01-03", 10.8, 11.5, 10.5, 11.0)}),
|
||
]
|
||
pe, db, aid, account, runner = _build(tmp_path, sections)
|
||
pe.run()
|
||
balances = list_daily_balance(db, aid)
|
||
assert len(balances) == 3
|
||
# 总账持仓 = 分户持仓(day2+day3 各买100 = 200)
|
||
assert account.positions["600000"].volume == 200
|
||
assert runner.positions["600000"].volume == 200
|