Files
big-qmt/py-client/tests/test_zt_audit_fixes.py
2026-09-12 16:25:34 +08:00

133 lines
6.7 KiB
Python

import tempfile
import unittest
from contextlib import closing
from dataclasses import replace
from datetime import datetime, timedelta
from pathlib import Path
from types import SimpleNamespace as NS
from unittest.mock import Mock, patch
from libs.grid_take_profit import GridState
from libs.order import OrderBook
from libs.snapshot import get_collector_snapshot
from libs.state import State
from sdk import Assets, DealItem, OrderItem, PositionItem
from strategy.zt import boot
from strategy.zt.profit import ZTProfitTracker
class ZTAuditFixTests(unittest.TestCase):
def setUp(self):
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
self.store = State(Path(tmp.name) / 'state.db')
self.code = '600000.SH'
def deal(self, code, identity):
return DealItem(stock_code=code, order_sys_id=identity,
remark=f'zt-base-{identity}|zt', offset_flag=48,
volume=100, price=10, trade_amount=1000)
def position(self, code):
return PositionItem(stock_code=code, volume=100, open_price=10)
def test_zt_cancels_only_owned_orders_and_tracks_all(self):
stamp = datetime.now() - timedelta(minutes=2)
orders = [OrderItem(stock_code=f'60000{i}.SH', order_sys_id=str(i), remark=remark,
order_status=50, offset_flag=23,
insert_date=stamp.strftime('%Y%m%d'), insert_time=stamp.strftime('%H%M%S'))
for i, remark in enumerate(['zt-base-own|zt', 'zt-SELL-own|zt',
'zt-added-own|zt', 'IPO-new|ipo', '', 'TREN-BUY-other'])]
client = Mock()
book = OrderBook()
book.refresh(client, orders, cancel_prefix='zt-')
self.assertEqual([c.args[0] for c in client.cancel_by_id.call_args_list], ['0', '1', '2'])
self.assertEqual(book.data, orders)
self.assertTrue(all(book.busy(o.stock_code, 'BUY') for o in orders))
client.reset_mock()
book.refresh(client, orders)
self.assertEqual([c.args[0] for c in client.cancel_by_id.call_args_list], ['0', '1', '2', '4', '5'])
def test_invalid_trade_persists_blocks_only_its_stock_and_recovers(self):
good = '600001.SH'
bad = replace(self.deal(self.code, 'bad'), price=float('nan'))
valid = self.deal(good, 'good')
positions = [self.position(c) for c in (self.code, good)]
self.store.sync_account(positions, [bad, valid])
self.assertEqual(self.store.blocked_codes, {self.code})
self.assertEqual(self.store.state[good]['base_qty'], 100)
with closing(self.store._connect()) as db:
payload = db.execute('SELECT payload FROM zt_rejected_deals').fetchone()[0]
self.assertIn('bad', payload)
self.assertIn('NaN', payload)
self.store = State(self.store.path)
self.store.sync_account(positions, [valid])
self.assertEqual(self.store.blocked_codes, {self.code})
self.store.sync_account(positions, [replace(bad, price=10), valid])
self.assertEqual(self.store.blocked_codes, set())
self.assertEqual(self.store.state[self.code]['base_qty'], 100)
def test_invalid_fields_do_not_block_other_stocks(self):
for field, value in [('volume', 0), ('volume', 1.5), ('offset_flag', 99),
('trade_amount', float('inf')), ('price', -1)]:
with self.subTest(field=field):
bad = replace(self.deal(self.code, f'bad-{field}'), **{field: value})
good = self.deal('600001.SH', 'good')
self.store.sync_account([self.position('600001.SH')], [bad, good])
self.assertEqual(self.store.state['600001.SH']['base_qty'], 100)
self.assertIn(self.code, self.store.blocked_codes)
def test_bad_stock_does_not_archive_other_trades_until_corrected(self):
first = self.deal(self.code, 'first')
bad = replace(self.deal(self.code, 'bad'), price=float('nan'))
position = replace(self.position(self.code), volume=200)
self.store.sync_account([position], [first, bad])
self.assertEqual(self.store.deals['first']['is_arch'], 0)
self.assertNotIn(self.code, self.store.state)
self.store.sync_account([position], [replace(bad, price=10)])
self.assertEqual(self.store.state[self.code]['base_qty'], 200)
self.assertEqual(self.store.blocked_codes, set())
def test_profit_basis_changes_reset_peak_but_partial_sell_does_not(self):
tracker = ZTProfitTracker()
position = self.position(self.code)
row = dict(base_qty=100, base_order_local_id='one', base_created_at='now', added_qty=0)
state = NS(blocked_codes=set(), get_by_code=lambda code: row)
tracker.sync_positions([position], state)
self.assertEqual(tracker.observe(self.code, 20).state, GridState.ARMED)
row['base_qty'] = 50
tracker.sync_positions([replace(position, volume=50)], state)
self.assertEqual(tracker.observe(self.code, 19).state, GridState.RETREAT)
for update, cost in [({'added_qty': 100, 'added_price': 11}, 10),
({'added_qty': 0}, 10), ({}, 12),
({'base_order_local_id': 'reopened'}, 12)]:
row.update(update)
tracker.sync_positions([replace(position, open_price=cost)], state)
self.assertEqual(tracker.observe(self.code, 10).state, GridState.ARMED)
tracker.sync_positions([], state)
tracker.sync_positions([position], state)
self.assertEqual(tracker.observe(self.code, 5).state, GridState.ARMED)
def test_run_once_updates_collector_before_market_fetch(self):
positions = [self.position(self.code)]
self.store.sync_account(positions, [], initialize=True)
client = Mock()
run = NS(client=client, account_cfg=NS(account_id='zt-test', min_cash_ratio=0.1),
orders=Mock(), profit_tracker=ZTProfitTracker())
client.deals.return_value = []
client.portfolio.return_value = NS(assets=Assets(20000, 10000),
positions={self.code: positions[0]}, orders=[])
client.full_tick.side_effect = RuntimeError('no market data')
with patch.object(boot, 'trading_time', return_value=True), \
patch.object(boot, 'market_allow_open', return_value=True):
boot.RunOnce(run, self.store, [])
snapshot = get_collector_snapshot()
self.assertEqual(snapshot[0], 'zt-test')
self.assertEqual(snapshot[1].total, 20000)
self.assertEqual(snapshot[2], positions)
run.orders.refresh.assert_called_once_with(client, [], cancel_prefix='zt-')
if __name__ == '__main__':
unittest.main()