import sqlite3 import tempfile import unittest from contextlib import closing from dataclasses import asdict, fields from pathlib import Path from unittest.mock import patch from libs.state import FLAG_BUY, FLAG_SELL, State, StateItem from sdk import DealItem, PositionItem class OrderBookTests(unittest.TestCase): def setUp(self): self.tmp = tempfile.TemporaryDirectory() self.addCleanup(self.tmp.cleanup) self.path = Path(self.tmp.name) / 'state.db' def deal(self, kind, sys_order_id, qty, price, date='2026-09-01'): prefix = {'base': 'zt-base-', 'sell': 'zt-t-sell-', 'buy': 'zt-t-buy-'}[kind] return DealItem( order_sys_id=sys_order_id, stock_code='600000.SH', offset_flag=FLAG_SELL if kind == 'sell' else FLAG_BUY, volume=qty, price=price, trade_amount=qty * price, trade_date=date, trade_time='10:00:00', remark=prefix + 'order1|zt', ) def test_json_is_never_read(self): legacy = self.path.with_suffix('.json') legacy.write_text('invalid JSON', encoding='utf-8') book = State(self.path) self.assertIsNone(book.load()) self.assertEqual((book.state, book.deals, book.deals_sys_ids), ({}, {}, set())) self.assertEqual(legacy.read_text(encoding='utf-8'), 'invalid JSON') def test_sync_deals_deduplicates_batch_and_restart(self): book = State(self.path) self.assertEqual((book.state, book.deals, book.deals_sys_ids), ({}, {}, set())) first = self.deal('base', 'd1', 40, 10, '20260901') second = self.deal('base', 'd2', 60, 12) book.sync_deals([first, first, second]) self.assertEqual(book.deals_sys_ids, {'d1', 'd2'}) self.assertEqual(book.deals['d1']['trade_date'], '2026-09-01') self.assertEqual(book.deals['d2']['volume'], 60) self.assertEqual(book.deals['d2']['order_local_id'], 'zt-base-order1') book = State(self.path) self.assertEqual(book.deals_sys_ids, {'d1', 'd2'}) self.assertEqual(book.deals['d1']['order_local_id'], 'zt-base-order1') with patch.object(book, '_connect') as connect: book.sync_deals([first, second]) book.sync_deals([]) connect.assert_not_called() self.assertEqual(len(book.deals), 2) def test_sync_deals_failure_rolls_back_entire_batch_and_cache(self): book = State(self.path) first = self.deal('base', 'd1', 100, 10) invalid = self.deal('base', 'd2', 100, 10) invalid.offset_flag = -1 with self.assertRaises(sqlite3.IntegrityError): book.sync_deals([first, invalid]) self.assertEqual(book.deals, {}) self.assertEqual(book.deals_sys_ids, set()) self.assertEqual(State(self.path).deals, {}) invalid.offset_flag = FLAG_BUY book.sync_deals([first, invalid]) self.assertEqual(book.deals_sys_ids, {'d1', 'd2'}) def test_load_refreshes_all_caches(self): book = State(self.path) writer = State(self.path) writer.sync_state([PositionItem(stock_code='600000.SH', volume=100)]) writer.sync_deals([self.deal('base', 'd1', 100, 10)]) book.load() self.assertEqual(book.state['600000.SH']['base_qty'], 100) self.assertEqual(book.deals_sys_ids, {'d1'}) self.assertEqual(book.deals['d1']['remark'], 'zt-base-order1|zt') def test_position_columns_defaults_indexes_and_stable_id(self): store = State(self.path) store.sync_deals([self.deal('base', 'd1', 100, 10)]) saved_deals = dict(store.deals) with closing(sqlite3.connect(self.path)) as db: columns = {row[1] for row in db.execute('PRAGMA table_info(state)')} self.assertEqual(columns, {'id', *(field.name for field in fields(StateItem))}) indexes = {row[1] for row in db.execute('PRAGMA index_list(state)')} self.assertEqual(indexes, {'idx_state_stock_code'}) position = PositionItem(stock_code='600000.SH', volume=100, open_price=10, stock_name='stock', can_use_volume=100, float_profit=-2.5) store.sync_state([position]) saved = store.state[position.stock_code] first_id = saved['id'] self.assertEqual(saved['base_qty'], 100) self.assertEqual(saved['base_price'], 10) self.assertEqual(saved['added_qty'], 0) self.assertEqual(saved['base_order_local_id'], '') self.assertTrue(saved['base_created_at']) position.volume = 200 position.open_price = 12 store.sync_state([position]) self.assertEqual(store.state[position.stock_code]['id'], first_id) self.assertEqual(store.state[position.stock_code], saved) self.assertEqual(State(self.path).state[position.stock_code], saved) store.sync_state([position, PositionItem(stock_code='600001.SH', volume=100)]) self.assertEqual(store.state[position.stock_code], saved) self.assertEqual(store.state['600001.SH']['base_qty'], 100) position.volume = 0 store.sync_state([position, PositionItem(stock_code='600002.SH')]) self.assertEqual(store.state, {}) self.assertEqual(State(self.path).state, {}) store.sync_state([PositionItem(stock_code='600001.SH', volume=100)]) self.assertGreater(store.state['600001.SH']['id'], first_id) store.sync_state([]) self.assertEqual(store.state, {}) self.assertEqual(store.deals, saved_deals) def test_state_fields_survive_restart_and_sync(self): book = State(self.path) row = asdict(StateItem( stock_code='600000.SH', status='READY', base_order_local_id='base-1', base_qty=100, base_price=10, base_created_at='2026-09-08T09:30:00', added_order_local_id='added-1', added_qty=50, added_price=9, added_created_at='2026-09-08T10:30:00', )) with closing(book._connect()) as db, db: db.execute( f"INSERT INTO state ({', '.join(row)}) VALUES ({', '.join(':' + key for key in row)})", row, ) book.load() saved = book.state[row['stock_code']] self.assertEqual({k: v for k, v in saved.items() if k != 'id'}, row) book = State(self.path) book.sync_state([PositionItem(stock_code=row['stock_code'], volume=150, open_price=9.5)]) self.assertEqual(book.state[row['stock_code']], saved) with self.assertRaises(sqlite3.IntegrityError): with closing(book._connect()) as db, db: db.execute('UPDATE state SET added_qty = -1') self.assertEqual(State(self.path).state[row['stock_code']], saved) if __name__ == '__main__': unittest.main()