55 lines
2.3 KiB
Python
55 lines
2.3 KiB
Python
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from libs.state import State
|
|
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')
|
|
|
|
def position(self, qty):
|
|
return PositionItem(stock_code='600000.SH', 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 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)
|
|
self.state = State(self.state.path)
|
|
bought = self.deal('new', 100, flag=48)
|
|
for _ in range(2):
|
|
sync_account_state(self.state, [self.position(200)], [historical, bought, unrelated])
|
|
row = self.state.state['600000.SH']
|
|
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.state, {})
|
|
self.assertEqual(self.state.deals['sell']['is_arch'], 1)
|
|
|
|
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)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|