fix bug
This commit is contained in:
125
py-client/tests/zt_harness.py
Normal file
125
py-client/tests/zt_harness.py
Normal file
@@ -0,0 +1,125 @@
|
||||
"""ZT 新路径测试脚手架:真实 Runtime/OrderBook/DipWatch + 模拟客户端。"""
|
||||
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace as NS
|
||||
from unittest.mock import Mock
|
||||
|
||||
from libs.grid_take_profit import GridTrailingTracker
|
||||
from libs.order import OrderBook
|
||||
from libs.runtime import Runtime
|
||||
from libs.watch import DipWatch
|
||||
from sdk import Assets, DealItem, PositionItem, Tick
|
||||
from strategy.zt.rounds import Round, RoundStore, start_round
|
||||
|
||||
ACCOUNT = 'zt-test'
|
||||
|
||||
|
||||
def account_cfg(**overrides):
|
||||
cfg = NS(account_id=ACCOUNT, strategy='zt', host_key='test',
|
||||
grid_step_pct=1.0, zt_open_hands=1, zt_sell_ratio=0.5,
|
||||
zt_buy_fall_pct=1.0, zt_max_price=200.0, zt_t_band_pct=1.0,
|
||||
zt_max_hold_days=5, min_cash_ratio=0.1,
|
||||
excluded_codes=[], signal_allow=['dcm'], buy_value=10000.0)
|
||||
for key, value in overrides.items():
|
||||
setattr(cfg, key, value)
|
||||
return cfg
|
||||
|
||||
|
||||
def global_cfg(**overrides):
|
||||
cfg = NS(qmt_base_url='http://unused', qmt_token='', api_host='http://unused',
|
||||
qmt_data_dir='.', signals={})
|
||||
for key, value in overrides.items():
|
||||
setattr(cfg, key, value)
|
||||
return cfg
|
||||
|
||||
|
||||
class Fixture:
|
||||
"""一套隔离的账户快照、轮次存储与运行上下文。"""
|
||||
|
||||
def __init__(self, **cfg_overrides):
|
||||
self.tmp = tempfile.TemporaryDirectory()
|
||||
self.path = Path(self.tmp.name)
|
||||
self.rounds_path = self.path / f'zt_{ACCOUNT}_rounds.json'
|
||||
self.store = RoundStore(self.rounds_path)
|
||||
self.account_cfg = account_cfg(**cfg_overrides)
|
||||
self.assets = Assets(total=100000.0, available=100000.0)
|
||||
self.positions = {}
|
||||
self.orders = []
|
||||
self.deals = []
|
||||
self.ticks = {}
|
||||
self.client = self._client()
|
||||
self.run = Runtime(
|
||||
client=self.client, global_cfg=global_cfg(),
|
||||
account_cfg=self.account_cfg,
|
||||
orders=OrderBook(cancel_timeout_sec=300),
|
||||
open_watch=DipWatch(expire_seconds=600, rebound_threshold=0.0),
|
||||
add_watch=DipWatch(expire_seconds=600, rebound_threshold=0.0),
|
||||
profit_tracker=GridTrailingTracker(self.account_cfg.grid_step_pct),
|
||||
)
|
||||
|
||||
def _client(self):
|
||||
client = Mock()
|
||||
client.deals.side_effect = lambda: self.deals
|
||||
client.portfolio.side_effect = lambda: NS(
|
||||
assets=self.assets, positions=self.positions, orders=self.orders)
|
||||
client.full_tick.side_effect = lambda codes: dict(self.ticks)
|
||||
return client
|
||||
|
||||
def cleanup(self):
|
||||
self.tmp.cleanup()
|
||||
|
||||
# ---- 便捷构造 ----
|
||||
def hold(self, code='600000.SH', volume=1000, price=10.0, can_use=None):
|
||||
position = PositionItem(stock_code=code, volume=volume, open_price=price,
|
||||
can_use_volume=volume if can_use is None else can_use)
|
||||
self.positions[code] = position
|
||||
return position
|
||||
|
||||
def quote(self, code='600000.SH', price=10.0):
|
||||
self.ticks[code] = Tick(last_price=price)
|
||||
return self.ticks[code]
|
||||
|
||||
def deal(self, local_id, volume, price, sys_id=None, code='600000.SH'):
|
||||
return DealItem(stock_code=code, order_sys_id=sys_id or f'{local_id}-{volume}',
|
||||
remark=local_id, offset_flag=48, volume=volume, price=price,
|
||||
trade_amount=price * volume, trade_date='20260915',
|
||||
trade_time='100000')
|
||||
|
||||
def prime(self, watch, code, price):
|
||||
"""让 DipWatch 先建立观察点,下一次同价或更高价即满足反弹确认。"""
|
||||
watch.triggered('prime', code, price)
|
||||
|
||||
def tick(self, signals=()):
|
||||
"""跑一轮 RunOnce,绕过真实时钟的交易时段判断。"""
|
||||
from unittest.mock import patch
|
||||
from strategy.zt import boot
|
||||
with patch.object(boot, 'trading_time', return_value=True):
|
||||
boot.RunOnce(self.run, self.store, list(signals))
|
||||
return self.store
|
||||
|
||||
def own_base(self, code='600000.SH', qty=1000, cost=10.0, today='2026-09-15'):
|
||||
"""把该证券标记为"本策略自己建仓"(模拟建仓腿已成交)。"""
|
||||
item = Round(code=code, base_qty=qty, base_cost=cost, base_date=today,
|
||||
base_source='opened', phase='CLOSED', outcome='base')
|
||||
self.store.put(item)
|
||||
self.store.save()
|
||||
return item
|
||||
|
||||
def open_round(self, code='600000.SH', kind='LONG_T', order_id='zt-entry-1',
|
||||
today='2026-09-15', **fields):
|
||||
"""写入一条已提交开仓腿的轮次记录。"""
|
||||
item = Round(code=code, base_qty=fields.pop('base_qty', 1000),
|
||||
base_cost=fields.pop('base_cost', 10.0),
|
||||
base_source=fields.pop('base_source', 'opened'))
|
||||
start_round(item, kind, today)
|
||||
item.entry_order_id = order_id
|
||||
for key, value in fields.items():
|
||||
setattr(item, key, value)
|
||||
self.store.put(item)
|
||||
self.store.save()
|
||||
return item
|
||||
|
||||
@property
|
||||
def placed(self):
|
||||
return [call.kwargs for call in self.client.passorder.call_args_list]
|
||||
Reference in New Issue
Block a user