import tempfile import unittest from concurrent.futures import ThreadPoolExecutor from contextlib import closing from datetime import datetime 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.state import FLAG_BUY, State from sdk import Assets, DealItem, PositionItem, Tick from strategy.zt import boot from strategy.zt.positions import manage_positions class ZTTradingTests(unittest.TestCase): def setUp(self): self.code = '600000.SH' self.run = NS(account_cfg=NS(account_id='test', strategy='zt', buy_value=1000, excluded_codes=[], enable_loss_add_position=False, min_cash_ratio=0.1), orders=Mock(), client=Mock(), profit_tracker=Mock(), add_watch=Mock()) self.run.orders.busy.return_value = False self.run.orders.place.return_value = True self.run.profit_tracker.observe.return_value.state = GridState.RETREAT self.run.add_watch.triggered.return_value = True def manage(self, added=0, usable=500, road=0, cost=10, added_cost=10, price=11): position = PositionItem(stock_code=self.code, volume=1000, can_use_volume=usable, on_road_volume=road, open_price=cost) state = NS(blocked_codes=set(), get_by_code=lambda code: dict( base_qty=500, added_qty=added, added_price=added_cost)) manage_positions(self.run, {self.code: Tick(last_price=price)}, [position], True, 1500, state) def test_added_position_is_capped_by_sellable_inventory(self): for added, usable, expected in [(500, 100, 100), (100, 500, 100), (0, 500, 500)]: with self.subTest(added=added, usable=usable): self.run.orders.place.reset_mock() self.manage(added=added, usable=usable) self.assertEqual(self.run.orders.place.call_args.args[1].volume, expected) def test_zero_sellable_does_not_divide_by_default_added_cost(self): with patch('strategy.zt.positions.log.exception') as error: self.manage(usable=0, added_cost=0) error.assert_not_called() self.run.orders.place.assert_not_called() def test_unavailable_shares_do_not_disable_loss_management(self): self.run.account_cfg.enable_loss_add_position = True self.manage(usable=0, cost=20, price=10, added_cost=0) self.assertEqual(self.run.orders.place.call_args.args[1].op, 23) def test_on_road_shares_do_not_disable_available_base(self): self.manage(road=100) self.assertEqual(self.run.orders.place.call_args.args[1].volume, 500) def test_added_cost_is_used_even_if_base_cost_is_higher(self): self.manage(added=100, cost=20, added_cost=10, price=11) self.assertEqual(self.run.orders.place.call_args.args[1].volume, 100) def test_invalid_selected_cost_never_trades(self): for cost in [0, -1, float('nan'), float('inf')]: with self.subTest(cost=cost), patch('strategy.zt.positions.log.exception') as error: self.manage(added=100, added_cost=cost) error.assert_not_called() self.run.orders.place.assert_not_called() def test_run_once_quarantines_manual_trade_but_manages_good_stock(self): with tempfile.TemporaryDirectory() as tmp, ThreadPoolExecutor(max_workers=2) as executor: store = State(Path(tmp) / 'state.db') good = '600001.SH' positions = [PositionItem(stock_code=c, volume=100, can_use_volume=100, open_price=10) for c in [self.code, good]] store.sync_account(positions, [], initialize=True) manual = DealItem(stock_code=self.code, order_sys_id='manual', remark='', offset_flag=FLAG_BUY, volume=100, price=10, trade_amount=1000) self.run.executor = executor self.run.client.deals.return_value = [manual] self.run.client.portfolio.return_value = NS( assets=Assets(10000, 10000), positions={p.stock_code: p for p in positions}, orders=[]) self.run.client.full_tick.return_value = {p.stock_code: Tick(last_price=11) for p in positions} with patch.object(boot, 'datetime') as clock, patch.object(boot, 'market_allow_open', return_value=True): clock.now.return_value = datetime(2026, 9, 11, 10) boot.RunOnce(self.run, store, []) self.assertEqual(store.blocked_codes, {self.code}) self.run.orders.refresh.assert_called_once() self.assertEqual(self.run.orders.place.call_count, 1) self.assertEqual(self.run.orders.place.call_args.args[1].code, good) def test_run_once_does_not_reopen_quarantined_sold_out_code(self): with tempfile.TemporaryDirectory() as tmp, ThreadPoolExecutor(max_workers=2) as executor: store = State(Path(tmp) / 'state.db') store.sync_account([], [], initialize=True) self.run.executor = executor self.run.client.deals.return_value = [DealItem( stock_code=self.code, order_sys_id='manual', remark='', offset_flag=FLAG_BUY, volume=100, price=10, trade_amount=1000)] self.run.client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={}, orders=[]) self.run.client.full_tick.return_value = {} with patch.object(boot, 'datetime') as clock, patch.object(boot, 'market_allow_open', return_value=True), \ patch.object(boot, 'open_signal') as opened: clock.now.return_value = datetime(2026, 9, 11, 10) boot.RunOnce(self.run, store, [NS(code=self.code)]) opened.assert_not_called() def start(self, client, directory): self.run.account_cfg.grid_step_pct = 1 global_cfg = NS(qmt_base_url='unused', qmt_token='', qmt_data_dir=directory) self.run.account_cfg.signal_allow = [] with patch.object(boot, 'Client', return_value=client), \ patch.object(boot.config, 'global_config', global_cfg), \ patch.object(boot.config, 'account_config', self.run.account_cfg), \ patch.object(boot, 'init_signals', return_value=[]), \ patch.object(boot, 'cache_portfolio'), patch.object(boot, 'Overview'), \ patch.object(boot.time, 'localtime', return_value=NS(tm_hour=15, tm_min=0, tm_sec=0)): boot.StartZT() def test_start_initializes_once_without_snapshot_retry_loop(self): client = Mock() client.deals.return_value = [] client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={self.code: PositionItem(stock_code=self.code, volume=100, open_price=10)}, orders=[]) with tempfile.TemporaryDirectory() as tmp: self.start(client, tmp) store = State(Path(tmp) / 'zt_test_state.db') self.assertEqual(store.state[self.code]['base_qty'], 100) with closing(store._connect()) as db: self.assertIsNone(db.execute("SELECT 1 FROM sqlite_master WHERE name='state_meta'").fetchone()) self.assertEqual(client.portfolio.call_count, 1) self.assertEqual(client.deals.call_count, 2) client.reset_mock() self.start(client, tmp) self.assertEqual(client.portfolio.call_count, 1) self.assertEqual(client.deals.call_count, 1) def test_start_rejects_changed_deals_without_writing_baseline(self): client = Mock() client.deals.side_effect = [[], [DealItem(order_sys_id='new')]] client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={}, orders=[]) with tempfile.TemporaryDirectory() as tmp: with self.assertRaises(RuntimeError): self.start(client, tmp) store = State(Path(tmp) / 'zt_test_state.db') self.assertEqual((store.state, store.deals), ({}, {})) client.close.assert_called_once() if __name__ == '__main__': unittest.main()