import sqlite3 import tempfile import unittest from pathlib import Path from unittest.mock import patch from libs.state import FLAG_BUY, FLAG_SELL, State, UNATTRIBUTED_PREFIX from sdk import DealItem, PositionItem class ZTStateTests(unittest.TestCase): def setUp(self): tmp = tempfile.TemporaryDirectory() self.addCleanup(tmp.cleanup) self.state = State(Path(tmp.name) / 'state.db') self.code = '600000.SH' def position(self, qty, code=None): return PositionItem(stock_code=code or self.code, volume=qty, open_price=10) def deal(self, identity, qty=100, flag=FLAG_BUY, code=None, remark=None): return DealItem(stock_code=code or self.code, order_sys_id=identity, remark=f'zt-base-{identity}|zt' if remark is None else remark, offset_flag=flag, volume=qty, price=10, trade_amount=qty*10, trade_date='20260912', trade_time='100000') def test_initial_snapshot_and_equal_size_increment_are_distinct(self): old = self.deal('old') self.state.sync_account([self.position(100)], [old], initialize=True) self.assertEqual(self.state.state[self.code]['base_qty'], 100) self.assertEqual(self.state.deals['old']['is_arch'], 1) self.state = State(self.state.path) new = self.deal('new', remark='zt-added-new|zt') for _ in range(2): self.state.sync_account([self.position(200)], [old, new, new]) row = self.state.state[self.code] self.assertEqual((row['base_qty'], row['added_qty']), (100, 100)) self.assertEqual(self.state.blocked_codes, set()) def test_initial_mixed_trades_are_already_in_snapshot(self): self.state.sync_account([self.position(150)], [self.deal('buy', 200), self.deal('sell', 50, FLAG_SELL)], initialize=True) self.assertEqual(self.state.state[self.code]['base_qty'], 150) self.assertTrue(all(d['is_arch'] == 1 for d in self.state.deals.values())) def test_restart_full_sell_is_archived_before_reconciliation(self): self.state.sync_account([self.position(100)], [], initialize=True) self.state = State(self.state.path) self.state.sync_account([], [self.deal('sell', flag=FLAG_SELL)]) self.assertEqual(self.state.state, {}) self.assertEqual(self.state.deals['sell']['is_arch'], 1) self.assertEqual(self.state.blocked_codes, set()) def test_empty_initialized_account_survives_restart(self): self.state.sync_account([], [], initialize=True) self.state = State(self.state.path) # 空账户重复初始化仍为空,无需额外标记表。 self.state.sync_account([], [], initialize=True) self.state.sync_account([self.position(100)], [self.deal('new')]) self.assertEqual(self.state.state[self.code]['base_qty'], 100) def test_initialization_cannot_overwrite_existing_holdings(self): self.state.sync_account([self.position(100)], [], initialize=True) with self.assertRaises(ValueError): self.state.sync_account([], [], initialize=True) self.assertEqual(self.state.state[self.code]['base_qty'], 100) def test_initialization_failure_is_atomic(self): invalid = self.position(100) invalid.open_price = float('inf') with self.assertRaises(ValueError): self.state.sync_account([invalid], [self.deal('one')], initialize=True) restarted = State(self.state.path) self.assertEqual((restarted.state, restarted.deals), ({}, {})) with patch.object(self.state, '_read_deals', side_effect=sqlite3.OperationalError('read failed')): with self.assertRaises(sqlite3.OperationalError): self.state.sync_account([self.position(100)], [self.deal('one')], initialize=True) self.assertEqual(State(self.state.path).deals, {}) self.assertEqual(self.state.state, {}) def test_lagging_snapshot_never_deletes_or_recreates_inventory(self): self.state.sync_account([self.position(100)], [], initialize=True) self.state.sync_account([], []) self.assertEqual(self.state.state[self.code]['base_qty'], 100) self.assertEqual(self.state.blocked_codes, {self.code}) self.state.sync_account([self.position(100)], []) self.assertEqual(self.state.blocked_codes, set()) sell = self.deal('sell', flag=FLAG_SELL) self.state.sync_account([self.position(100)], [sell]) self.assertNotIn(self.code, self.state.state) self.assertEqual(self.state.blocked_codes, {self.code}) self.state.sync_account([], [sell]) self.assertEqual(self.state.blocked_codes, set()) def test_blank_remark_persists_and_isolates_only_affected_stock(self): self.state.sync_account([self.position(100)], [], initialize=True) manual = self.deal('manual', remark=' |') good = self.deal('good', code='600001.SH') positions = [self.position(200), self.position(100, '600001.SH')] for _ in range(2): self.state.sync_account(positions, [manual, good]) self.assertEqual(self.state.blocked_codes, {self.code}) self.assertEqual(self.state.state[self.code]['base_qty'], 100) self.assertEqual(self.state.state['600001.SH']['base_qty'], 100) self.assertEqual(self.state.deals['manual']['is_arch'], 0) self.assertTrue(self.state.deals['manual']['order_local_id'].startswith(UNATTRIBUTED_PREFIX)) self.assertEqual(self.state.deals['manual']['remark'], ' |') self.state = State(self.state.path) def test_legacy_database_is_restored_without_reinitialization(self): self.state.sync_state([self.position(100)]) self.state.sync_account([], [self.deal('sell', flag=FLAG_SELL)]) self.assertEqual(self.state.state, {}) self.assertEqual(self.state.deals['sell']['is_arch'], 1) def test_failed_archive_blocks_only_stock_and_retries(self): self.state.sync_account([self.position(100)], [], initialize=True) sell = self.deal('sell', qty=200, flag=FLAG_SELL) self.state.sync_account([], [sell]) self.assertEqual(self.state.state[self.code]['base_qty'], 100) self.assertEqual(self.state.blocked_codes, {self.code}) buy = self.deal('buy', remark='zt-added-buy|zt') buy.trade_time = '095900' self.state.sync_account([], [sell, buy]) self.assertEqual(self.state.state, {}) self.assertEqual(self.state.blocked_codes, set()) if __name__ == '__main__': unittest.main()