76 lines
3.3 KiB
Python
76 lines
3.3 KiB
Python
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()
|