import tempfile import unittest from unittest.mock import patch from pathlib import Path from libs.state import State from sdk import DealItem, PositionItem from strategy.zt.boot import sync_account_state class ZTStateTests(unittest.TestCase): def setUp(self): tmp = tempfile.TemporaryDirectory() self.addCleanup(tmp.cleanup) self.state = State(Path(tmp.name) / 'zt_test_state.db') def position(self, qty): return PositionItem(stock_code='600000.SH', volume=qty, open_price=10) def deal(self, order, qty, flag=23, strategy='zt'): return DealItem( stock_code='600000.SH', order_sys_id=order, remark=f'{strategy}-buy-{order}|{strategy}', offset_flag=flag, volume=qty, price=10, trade_amount=qty * 10, trade_date='20260909', trade_time='100000', ) def test_initial_snapshot_and_incremental_deals_after_restart(self): historical = self.deal('old', 100) unrelated = self.deal('trend', 100, strategy='trend') sync_account_state(self.state, [self.position(100)], [historical, unrelated], initialize=True) self.assertEqual(set(self.state.deals), {'old'}) self.assertEqual(self.state.state['600000.SH']['base_qty'], 100) self.assertEqual(self.state.state['600000.SH']['added_qty'], 0) self.state = State(self.state.path) bought = self.deal('new', 100, flag=48) for _ in range(2): sync_account_state(self.state, [self.position(200)], [historical, bought, unrelated]) row = self.state.state['600000.SH'] self.assertEqual((row['base_qty'], row['added_qty']), (100, 100)) sold = self.deal('sell', 200, flag=24) sync_account_state(self.state, [], [historical, bought, sold]) self.assertEqual(self.state.state, {}) self.assertEqual(self.state.deals['sell']['is_arch'], 1) def test_archive_failure_preserves_holdings_for_retry(self): sync_account_state(self.state, [self.position(100)], [], initialize=True) with self.assertRaisesRegex(ValueError, 'ZT'): sync_account_state(self.state, [], [self.deal('sell', 200, flag=49)]) self.assertEqual(self.state.state['600000.SH']['base_qty'], 100) self.assertEqual(self.state.deals['sell']['is_arch'], 0) def test_failed_initialization_leaves_original_database_empty(self): invalid = self.position(100) invalid.open_price = float('inf') with self.assertRaises(ValueError): sync_account_state(self.state, [invalid], [self.deal('old', 100)], initialize=True) restarted = State(self.state.path) self.assertEqual((restarted.state, restarted.deals), ({}, {})) sync_account_state(restarted, [self.position(100)], [self.deal('old', 100)], initialize=True) self.assertEqual(restarted.deals['old']['is_arch'], 1) def test_initialization_cannot_overwrite_existing_database(self): sync_account_state(self.state, [self.position(100)], [], initialize=True) with self.assertRaises(ValueError): sync_account_state(self.state, [], [], initialize=True) self.assertEqual(State(self.state.path).state['600000.SH']['base_qty'], 100) def test_no_new_deals_still_retries_failed_archiving(self): sync_account_state(self.state, [self.position(100)], [], initialize=True) sell = self.deal('sell', 100, flag=24) with patch.object(self.state, 'archiving'): with self.assertRaises(ValueError): sync_account_state(self.state, [], [sell]) sync_account_state(self.state, [], [sell]) self.assertEqual(self.state.state, {}) self.assertEqual(self.state.deals['sell']['is_arch'], 1) if __name__ == '__main__': unittest.main()