import tempfile import unittest 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, 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 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_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_state(positions) 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()