fix zt&state.py

This commit is contained in:
2026-09-12 13:42:23 +08:00
parent bcb03e2ed9
commit 7a7049ce44
9 changed files with 677 additions and 363 deletions

View File

@@ -0,0 +1,86 @@
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_stale_cache_duplicate_preserves_original_and_imports_new_trade(self):
writer = State(self.store.path)
first = self.deal()
writer.sync_deals([first])
first.price = 20
first.trade_amount = 2000
self.store.sync_deals([first, self.deal('second')])
self.assertEqual(self.store.deals['first']['price'], 10)
self.assertEqual(self.store.deals_sys_ids, {'first', 'second'})
self.assertEqual(State(self.store.path).deals, self.store.deals)
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()