import sqlite3 import tempfile import unittest from dataclasses import asdict from pathlib import Path from unittest.mock import patch from libs.state import FLAG_BUY, State from sdk import DealItem, PositionItem class StateStorageTests(unittest.TestCase): def setUp(self): tmp = tempfile.TemporaryDirectory() self.addCleanup(tmp.cleanup) self.store = State(Path(tmp.name) / 'state.db') def deal(self, identity='first'): return DealItem( stock_code='600000.SH', order_sys_id=identity, remark=f'zt-base-{identity}|zt', offset_flag=FLAG_BUY, volume=100, price=10, trade_amount=1000, trade_date='20260912', trade_time='100000', ) def test_load_failure_does_not_publish_partial_cache(self): writer = State(self.store.path) writer.sync_state([PositionItem(stock_code='600000.SH', volume=100)]) with patch.object(self.store, '_read_deals', side_effect=sqlite3.OperationalError('read failed')): with self.assertRaises(sqlite3.OperationalError): self.store.load() self.assertEqual((self.store.state, self.store.deals, self.store.deals_sys_ids), ({}, {}, set())) self.store.load() self.assertEqual(self.store.state['600000.SH']['base_qty'], 100) def test_cache_read_failure_rolls_back_archive_and_can_retry(self): self.store.sync_deals([self.deal()]) with patch.object(self.store, '_read_deals', side_effect=sqlite3.OperationalError('read failed')): with self.assertRaises(sqlite3.OperationalError): self.store.archiving() restarted = State(self.store.path) self.assertEqual(restarted.state, {}) self.assertEqual(restarted.deals['first']['is_arch'], 0) self.assertEqual(self.store.deals, restarted.deals) self.store.archiving() self.assertEqual(self.store.state['600000.SH']['base_qty'], 100) self.assertEqual(self.store.deals['first']['is_arch'], 1) def test_invalid_snapshot_preserves_existing_holdings(self): self.store.sync_state([PositionItem(stock_code='600000.SH', volume=100)]) saved = self.store.state with self.assertRaises(ValueError): self.store.sync_state([PositionItem(stock_code='600001.SH', volume=100, open_price=float('inf'))]) self.assertEqual(self.store.state, saved) self.assertEqual(State(self.store.path).state, saved) def test_normalization_preserves_input_and_rejects_nonfinite_price(self): deal = self.deal() deal.trade_amount = 0 original = asdict(deal) self.store.sync_deals([deal]) self.assertEqual(asdict(deal), original) self.assertEqual(self.store.deals['first']['trade_amount'], 1000) self.assertEqual(self.store.deals['first']['trade_date'], '2026-09-12') for price in (float('inf'), float('-inf'), float('nan')): with self.subTest(price=price): invalid = self.deal('invalid') invalid.price = price with self.assertRaises(ValueError): self.store.sync_deals([self.deal('second'), invalid]) self.assertEqual(State(self.store.path).deals_sys_ids, {'first'}) if __name__ == '__main__': unittest.main()