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()