fix zt&state.py
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user