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

@@ -1,81 +1,131 @@
import sqlite3
import tempfile
import unittest
from unittest.mock import patch
from pathlib import Path
from unittest.mock import patch
from libs.state import State
from libs.state import FLAG_BUY, FLAG_SELL, State, UNATTRIBUTED_PREFIX
from sdk import DealItem, PositionItem
from strategy.zt.boot import sync_account_state
class ZTStateTests(unittest.TestCase):
def setUp(self):
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
self.state = State(Path(tmp.name) / 'zt_test_state.db')
self.state = State(Path(tmp.name) / 'state.db')
self.code = '600000.SH'
def position(self, qty):
return PositionItem(stock_code='600000.SH', volume=qty, open_price=10)
def position(self, qty, code=None):
return PositionItem(stock_code=code or self.code, volume=qty, open_price=10)
def deal(self, order, qty, flag=23, strategy='zt'):
return DealItem(
stock_code='600000.SH', order_sys_id=order,
remark=f'{strategy}-buy-{order}|{strategy}', offset_flag=flag,
volume=qty, price=10, trade_amount=qty * 10,
trade_date='20260909', trade_time='100000',
)
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_incremental_deals_after_restart(self):
historical = self.deal('old', 100)
unrelated = self.deal('trend', 100, strategy='trend')
sync_account_state(self.state, [self.position(100)], [historical, unrelated], initialize=True)
self.assertEqual(set(self.state.deals), {'old'})
self.assertEqual(self.state.state['600000.SH']['base_qty'], 100)
self.assertEqual(self.state.state['600000.SH']['added_qty'], 0)
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)
bought = self.deal('new', 100, flag=48)
new = self.deal('new', remark='zt-added-new|zt')
for _ in range(2):
sync_account_state(self.state, [self.position(200)], [historical, bought, unrelated])
row = self.state.state['600000.SH']
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))
sold = self.deal('sell', 200, flag=24)
sync_account_state(self.state, [], [historical, bought, sold])
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_archive_failure_preserves_holdings_for_retry(self):
sync_account_state(self.state, [self.position(100)], [], initialize=True)
with self.assertRaisesRegex(ValueError, 'ZT'):
sync_account_state(self.state, [], [self.deal('sell', 200, flag=49)])
self.assertEqual(self.state.state['600000.SH']['base_qty'], 100)
self.assertEqual(self.state.deals['sell']['is_arch'], 0)
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_failed_initialization_leaves_original_database_empty(self):
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):
sync_account_state(self.state, [invalid], [self.deal('old', 100)], initialize=True)
self.state.sync_account([invalid], [self.deal('one')], initialize=True)
restarted = State(self.state.path)
self.assertEqual((restarted.state, restarted.deals), ({}, {}))
sync_account_state(restarted, [self.position(100)], [self.deal('old', 100)], initialize=True)
self.assertEqual(restarted.deals['old']['is_arch'], 1)
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_initialization_cannot_overwrite_existing_database(self):
sync_account_state(self.state, [self.position(100)], [], initialize=True)
with self.assertRaises(ValueError):
sync_account_state(self.state, [], [], initialize=True)
self.assertEqual(State(self.state.path).state['600000.SH']['base_qty'], 100)
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_no_new_deals_still_retries_failed_archiving(self):
sync_account_state(self.state, [self.position(100)], [], initialize=True)
sell = self.deal('sell', 100, flag=24)
with patch.object(self.state, 'archiving'):
with self.assertRaises(ValueError):
sync_account_state(self.state, [], [sell])
sync_account_state(self.state, [], [sell])
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()