fix zt&state.py
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
import ast
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from dataclasses import asdict, fields
|
||||
@@ -88,7 +89,7 @@ class ApiModelTests(unittest.TestCase):
|
||||
self.assertEqual(deal.trade_amount, 0)
|
||||
deal.order_sys_id = 'sys3'
|
||||
deal.price = 0
|
||||
with self.assertRaises(ValueError):
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
book.sync_deals([deal])
|
||||
self.assertEqual(set(State(path).deals), {'sys1', 'sys2'})
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ import unittest
|
||||
from contextlib import closing
|
||||
from pathlib import Path
|
||||
|
||||
from libs.state import State
|
||||
from libs.state import FLAG_BUY, FLAG_SELL, State
|
||||
from sdk import DealItem, PositionItem
|
||||
|
||||
|
||||
@@ -28,7 +28,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertEqual(tables, {'state', 'deals', 'sqlite_sequence'})
|
||||
|
||||
def test_accumulates_once_and_preserves_base(self):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=200, open_price=8)])
|
||||
original = dict(self.book.state['600000.SH'])
|
||||
self.insert_deal('first', 40, 400, '10:00:00')
|
||||
self.insert_deal('second', 60, 720, '10:01:00')
|
||||
@@ -53,9 +53,9 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertEqual(row['added_order_local_id'], 'late')
|
||||
|
||||
def test_no_argument_archiving_recognizes_base_and_added_orders(self):
|
||||
self.insert_deal('zt-base-first', 100, 1000, '10:00:00', flag=23)
|
||||
self.insert_deal('zt-base-first', 100, 1000, '10:00:00', flag=FLAG_BUY)
|
||||
self.insert_deal('zt-base-second', 100, 1200, '10:01:00', flag=48)
|
||||
self.insert_deal('zt-t-buy-first', 100, 900, '10:02:00', flag=23)
|
||||
self.insert_deal('zt-t-buy-first', 100, 900, '10:02:00', flag=FLAG_BUY)
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['base_price']), (200, 11))
|
||||
@@ -111,13 +111,13 @@ class ArchivingTests(unittest.TestCase):
|
||||
|
||||
def test_stock_buy_and_sell_flags(self):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
self.insert_deal('buy', 100, 1000, '10:00:00', flag=23)
|
||||
self.insert_deal('sell', 50, 600, '10:01:00', flag=24)
|
||||
self.insert_deal('buy', 100, 1000, '10:00:00', flag=FLAG_BUY)
|
||||
self.insert_deal('sell', 50, 600, '10:01:00', flag=FLAG_SELL)
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 50))
|
||||
self.assertEqual(self.book.deals['buy']['offset_flag'], 23)
|
||||
self.assertEqual(self.book.deals['sell']['offset_flag'], 24)
|
||||
self.assertEqual(self.book.deals['buy']['offset_flag'], FLAG_BUY)
|
||||
self.assertEqual(self.book.deals['sell']['offset_flag'], FLAG_SELL)
|
||||
self.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
|
||||
|
||||
def test_new_state_and_failed_mark_roll_back_together(self):
|
||||
@@ -143,7 +143,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertEqual(self.book.state['600000.SH']['base_qty'], 0)
|
||||
self.assertEqual(self.book.state['600000.SH']['added_qty'], 100)
|
||||
|
||||
def test_equal_quantity_buy_is_added_and_preserves_status(self):
|
||||
def test_equal_quantity_buy_only_marks_and_preserves_status(self):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
self.book.sync_deals([DealItem(
|
||||
stock_code='600000.SH', order_sys_id='first', remark='base1|test',
|
||||
@@ -152,17 +152,17 @@ class ArchivingTests(unittest.TestCase):
|
||||
)])
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 100))
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 0))
|
||||
self.assertEqual(self.book.deals['first']['is_arch'], 1)
|
||||
restarted = State(self.book.path)
|
||||
self.assertIsNone(restarted.archiving())
|
||||
self.assertEqual(restarted.state, self.book.state)
|
||||
self.insert_deal('new_buy', 100, 1000, '10:01:00')
|
||||
self.insert_deal('new_buy', 50, 500, '10:01:00')
|
||||
with closing(self.book._connect()) as db, db:
|
||||
db.execute("UPDATE state SET status = 'CUSTOM' WHERE stock_code = '600000.SH'")
|
||||
self.assertIsNone(self.book.archiving())
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 200, 'CUSTOM'))
|
||||
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 50, 'CUSTOM'))
|
||||
|
||||
def test_archived_history_is_not_reapplied(self):
|
||||
self.insert_deal('old', 100, 1000, '10:00:00')
|
||||
|
||||
92
py-client/tests/test_state_snapshot.py
Normal file
92
py-client/tests/test_state_snapshot.py
Normal file
@@ -0,0 +1,92 @@
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from contextlib import closing
|
||||
from pathlib import Path
|
||||
|
||||
from libs.state import FLAG_BUY, FLAG_SELL, State
|
||||
from sdk import DealItem, PositionItem
|
||||
|
||||
|
||||
class SnapshotArchiveTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
self.store = State(Path(tmp.name) / 'state.db')
|
||||
self.code = '600000.SH'
|
||||
|
||||
def deal(self, identity, qty, flag=FLAG_BUY, code=None):
|
||||
return DealItem(
|
||||
stock_code=code or self.code, order_sys_id=identity,
|
||||
remark=f'zt-base-{identity}|zt', offset_flag=flag,
|
||||
volume=qty, price=12, trade_amount=qty * 12,
|
||||
trade_date='2026-09-12', trade_time='100000',
|
||||
)
|
||||
|
||||
def snapshot(self, qty):
|
||||
self.store.sync_state([PositionItem(stock_code=self.code, volume=qty, open_price=10)])
|
||||
|
||||
def test_matching_partial_fills_only_mark_and_survive_restart(self):
|
||||
self.snapshot(100)
|
||||
saved = dict(self.store.state[self.code])
|
||||
self.store.sync_deals([self.deal('one', 40), self.deal('two', 60)])
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code], saved)
|
||||
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
|
||||
restarted = State(self.store.path)
|
||||
restarted.archiving()
|
||||
self.assertEqual(restarted.state[self.code], saved)
|
||||
|
||||
def test_total_includes_added_holdings(self):
|
||||
self.snapshot(100)
|
||||
with closing(self.store._connect()) as db, db:
|
||||
db.execute("UPDATE state SET added_qty=50, added_price=9, status='CUSTOM'")
|
||||
self.store.load()
|
||||
saved = dict(self.store.state[self.code])
|
||||
self.store.sync_deals([self.deal('one', 150)])
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code], saved)
|
||||
self.assertEqual(self.store.deals['one']['is_arch'], 1)
|
||||
|
||||
def test_nonmatching_and_other_stock_are_incremental(self):
|
||||
self.snapshot(100)
|
||||
self.store.sync_deals([self.deal('one', 40), self.deal('other', 60, code='600001.SH')])
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code]['base_qty'], 140)
|
||||
self.assertEqual(self.store.state['600001.SH']['base_qty'], 60)
|
||||
|
||||
def test_matching_sell_still_liquidates(self):
|
||||
self.snapshot(100)
|
||||
self.store.sync_deals([self.deal('sell', 100, FLAG_SELL)])
|
||||
self.store.archiving()
|
||||
self.assertNotIn(self.code, self.store.state)
|
||||
self.assertEqual(self.store.deals['sell']['is_arch'], 1)
|
||||
|
||||
def test_mixed_batch_with_matching_gross_volume_is_not_skipped(self):
|
||||
self.snapshot(100)
|
||||
self.store.sync_deals([self.deal('buy', 40), self.deal('sell', 60, FLAG_SELL)])
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code]['base_qty'], 80)
|
||||
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
|
||||
|
||||
def test_failed_mark_rolls_back_entire_stock_and_retries(self):
|
||||
self.snapshot(100)
|
||||
saved = dict(self.store.state[self.code])
|
||||
self.store.sync_deals([self.deal('one', 40), self.deal('two', 60)])
|
||||
with closing(self.store._connect()) as db, db:
|
||||
db.execute("""CREATE TRIGGER fail_mark BEFORE UPDATE OF is_arch ON deals
|
||||
WHEN OLD.order_sys_id='two'
|
||||
BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
|
||||
with self.assertLogs(level='WARNING'):
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code], saved)
|
||||
self.assertTrue(all(d['is_arch'] == 0 for d in self.store.deals.values()))
|
||||
with closing(self.store._connect()) as db, db:
|
||||
db.execute('DROP TRIGGER fail_mark')
|
||||
self.store.archiving()
|
||||
self.assertEqual(self.store.state[self.code], saved)
|
||||
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
86
py-client/tests/test_state_storage.py
Normal file
86
py-client/tests/test_state_storage.py
Normal 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()
|
||||
@@ -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()
|
||||
|
||||
@@ -1,151 +1,149 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import closing
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from types import SimpleNamespace as NS
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from config import AccountConfig
|
||||
from libs.grid_take_profit import GridState
|
||||
from libs.order import OrderBook
|
||||
from libs.state import State
|
||||
from libs.state import FLAG_BUY, State
|
||||
from sdk import Assets, DealItem, PositionItem, Tick
|
||||
from strategy.zt import boot
|
||||
from strategy.zt.open import open_signal
|
||||
from strategy.zt.positions import manage_positions, t_rounds
|
||||
from strategy.zt.positions import manage_positions
|
||||
|
||||
|
||||
class ZTTradingTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(tmp.cleanup)
|
||||
self.store = State(Path(tmp.name) / 'state.db')
|
||||
self.code = '600000.SH'
|
||||
self.cfg = AccountConfig(account_id='test', buy_value=2000, zt_sell_ratio=0.5)
|
||||
self.run = SimpleNamespace(account_cfg=self.cfg, orders=Mock(), client=Mock(),
|
||||
profit_tracker=Mock(), add_watch=Mock(), open_watch=Mock())
|
||||
self.run = NS(account_cfg=NS(account_id='test', strategy='zt', buy_value=1000,
|
||||
excluded_codes=[], enable_loss_add_position=False,
|
||||
min_cash_ratio=0.1),
|
||||
orders=Mock(), client=Mock(), profit_tracker=Mock(), add_watch=Mock())
|
||||
self.run.orders.busy.return_value = False
|
||||
self.run.orders.new_order_id.side_effect = lambda prefix, kind: f'{prefix}-{kind}-order'
|
||||
self.run.orders.place.return_value = True
|
||||
self.run.profit_tracker.observe.return_value.state = GridState.RETREAT
|
||||
self.run.add_watch.triggered.return_value = True
|
||||
self.run.open_watch.triggered.return_value = True
|
||||
self.position = PositionItem(stock_code=self.code, volume=200, can_use_volume=200, open_price=10)
|
||||
boot.sync_account_state(self.store, [self.position], [], initialize=True)
|
||||
|
||||
def fill(self, kind, order, qty, price=10, date='2026-09-09'):
|
||||
return DealItem(stock_code=self.code, order_sys_id=order, remark=f'zt-{kind}-{order}|zt',
|
||||
offset_flag=24 if kind == 't-sell' else 23,
|
||||
volume=qty, price=price, trade_amount=qty * price,
|
||||
trade_date=date, trade_time='100000')
|
||||
def manage(self, added=0, usable=500, road=0, cost=10, added_cost=10, price=11):
|
||||
position = PositionItem(stock_code=self.code, volume=1000, can_use_volume=usable,
|
||||
on_road_volume=road, open_price=cost)
|
||||
state = NS(blocked_codes=set(), get_by_code=lambda code: dict(
|
||||
base_qty=500, added_qty=added, added_price=added_cost))
|
||||
manage_positions(self.run, {self.code: Tick(last_price=price)}, [position], True, 1500, state)
|
||||
|
||||
def manage(self, price=11, available=10000, positions=None, force=False, today='2026-09-09'):
|
||||
return manage_positions(self.run, self.store, {self.code: Tick(last_price=price)},
|
||||
[self.position] if positions is None else positions,
|
||||
t_rounds(self.store), available, today, force)
|
||||
def test_added_position_is_capped_by_sellable_inventory(self):
|
||||
for added, usable, expected in [(500, 100, 100), (100, 500, 100), (0, 500, 500)]:
|
||||
with self.subTest(added=added, usable=usable):
|
||||
self.run.orders.place.reset_mock()
|
||||
self.manage(added=added, usable=usable)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, expected)
|
||||
|
||||
def test_sell_only_available_shares_and_no_loss_sell(self):
|
||||
self.position.can_use_volume = 0
|
||||
self.manage()
|
||||
self.run.orders.place.assert_not_called()
|
||||
self.position.can_use_volume = 100
|
||||
self.manage(price=9)
|
||||
self.run.orders.place.assert_not_called()
|
||||
self.manage()
|
||||
request = self.run.orders.place.call_args.args[1]
|
||||
self.assertEqual((request.op, request.volume), (24, 100))
|
||||
|
||||
def test_full_sale_restart_and_force_buyback_without_price_or_market_gate(self):
|
||||
sell = self.fill('t-sell', 's1', 200, price=11)
|
||||
boot.sync_account_state(self.store, [], [sell])
|
||||
self.store = State(self.store.path)
|
||||
self.cfg.zt_max_price = 10
|
||||
self.run.add_watch.triggered.return_value = False
|
||||
remaining = self.manage(price=12, positions=[], force=True)
|
||||
request = self.run.orders.place.call_args.args[1]
|
||||
self.assertEqual((request.op, request.volume), (23, 200))
|
||||
self.assertAlmostEqual(remaining, 10000 - 12 * 200 * 1.01)
|
||||
|
||||
def test_partial_fills_once_and_completed_round_blocks_same_day_sale(self):
|
||||
deals = [self.fill('t-sell', 's1', 40, 11), self.fill('t-sell', 's2', 60, 12)]
|
||||
self.position.volume = 100
|
||||
boot.sync_account_state(self.store, [self.position], deals + deals)
|
||||
item = t_rounds(self.store)[self.code]
|
||||
self.assertEqual(item['sold'], 100)
|
||||
self.assertEqual(item['amount'], 1160)
|
||||
self.manage(price=10)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 100)
|
||||
deals.append(self.fill('t-buy', 'b1', 100))
|
||||
self.position.volume = 200
|
||||
boot.sync_account_state(self.store, [self.position], deals)
|
||||
self.run.orders.place.reset_mock()
|
||||
self.manage(price=11)
|
||||
self.run.orders.place.assert_not_called()
|
||||
self.manage(price=11, today='2026-09-10')
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].op, 24)
|
||||
|
||||
def test_cross_day_debt_and_insufficient_cash(self):
|
||||
boot.sync_account_state(self.store, [], [self.fill('t-sell', 's1', 200, date='2026-09-08')])
|
||||
self.manage(positions=[], available=100, force=True)
|
||||
self.run.orders.place.assert_not_called()
|
||||
self.manage(positions=[], force=True)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 200)
|
||||
|
||||
def test_delayed_snapshot_does_not_delete_or_recreate_holdings(self):
|
||||
boot.sync_account_state(self.store, [], [])
|
||||
self.assertEqual(self.store.state[self.code]['base_qty'], 200)
|
||||
sell = self.fill('t-sell', 's1', 200)
|
||||
boot.sync_account_state(self.store, [self.position], [sell])
|
||||
self.assertNotIn(self.code, self.store.state)
|
||||
self.manage()
|
||||
def test_zero_sellable_does_not_divide_by_default_added_cost(self):
|
||||
with patch('strategy.zt.positions.log.exception') as error:
|
||||
self.manage(usable=0, added_cost=0)
|
||||
error.assert_not_called()
|
||||
self.run.orders.place.assert_not_called()
|
||||
|
||||
def test_base_fills_stay_in_base_bucket(self):
|
||||
self.store.sync_state([])
|
||||
deals = [self.fill('base', 'b1', 100), self.fill('base', 'b2', 100, 12)]
|
||||
boot.sync_account_state(self.store, [self.position], deals)
|
||||
row = self.store.state[self.code]
|
||||
self.assertEqual((row['base_qty'], row['base_price'], row['added_qty']), (200, 11, 0))
|
||||
|
||||
def test_run_once_queries_sold_out_code_and_never_opens_with_debt(self):
|
||||
sell = self.fill('t-sell', 's1', 200, 11)
|
||||
self.run.client.deals.return_value = [sell]
|
||||
self.run.client.portfolio.return_value = SimpleNamespace(assets=Assets(10000, 10000), positions={}, orders=[])
|
||||
self.run.client.full_tick.return_value = {self.code: Tick(last_price=12)}
|
||||
with patch.object(boot, 'datetime') as clock, patch.object(boot, 'collector_push'), \
|
||||
patch.object(boot, 'open_signal') as opened, patch.object(boot, 'market_allow_open') as market:
|
||||
clock.now.return_value = datetime(2026, 9, 9, 14, 50)
|
||||
boot.RunOnce(self.run, self.store, [])
|
||||
self.run.client.full_tick.assert_called_once_with([self.code])
|
||||
opened.assert_not_called()
|
||||
market.assert_not_called()
|
||||
def test_unavailable_shares_do_not_disable_loss_management(self):
|
||||
self.run.account_cfg.enable_loss_add_position = True
|
||||
self.manage(usable=0, cost=20, price=10, added_cost=0)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].op, 23)
|
||||
|
||||
def test_open_budget_includes_buffer_and_star_minimum(self):
|
||||
with patch('strategy.zt.open.datetime') as clock:
|
||||
clock.now.return_value = datetime(2026, 9, 9, 10)
|
||||
remaining = open_signal(self.run, {self.code: Tick(last_price=10)},
|
||||
[SimpleNamespace(code=self.code)], 2000)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 100)
|
||||
self.assertEqual(remaining, 990)
|
||||
self.run.orders.place.reset_mock()
|
||||
open_signal(self.run, {'688001.SH': Tick(last_price=10)},
|
||||
[SimpleNamespace(code='688001.SH')], 2000)
|
||||
self.run.orders.place.assert_not_called()
|
||||
def test_on_road_shares_do_not_disable_available_base(self):
|
||||
self.manage(road=100)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 500)
|
||||
|
||||
def test_real_order_id_is_recognized_by_state_sync(self):
|
||||
orders = OrderBook('zt')
|
||||
self.run.orders.new_order_id.side_effect = orders.new_order_id
|
||||
with patch('strategy.zt.open.datetime') as clock:
|
||||
clock.now.return_value = datetime(2026, 9, 9, 10)
|
||||
open_signal(self.run, {self.code: Tick(last_price=10)},
|
||||
[SimpleNamespace(code=self.code)], 2000)
|
||||
request = self.run.orders.place.call_args.args[1]
|
||||
self.assertTrue(request.order_id.startswith('zt-base-'))
|
||||
deal = self.fill('base', 'b1', 100)
|
||||
deal.remark = request.order_id + '|zt'
|
||||
self.store.sync_state([])
|
||||
boot.sync_account_state(self.store, [], [deal])
|
||||
self.assertEqual(self.store.state[self.code]['base_qty'], 100)
|
||||
def test_added_cost_is_used_even_if_base_cost_is_higher(self):
|
||||
self.manage(added=100, cost=20, added_cost=10, price=11)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 100)
|
||||
|
||||
def test_invalid_selected_cost_never_trades(self):
|
||||
for cost in [0, -1, float('nan'), float('inf')]:
|
||||
with self.subTest(cost=cost), patch('strategy.zt.positions.log.exception') as error:
|
||||
self.manage(added=100, added_cost=cost)
|
||||
error.assert_not_called()
|
||||
self.run.orders.place.assert_not_called()
|
||||
|
||||
def test_run_once_quarantines_manual_trade_but_manages_good_stock(self):
|
||||
with tempfile.TemporaryDirectory() as tmp, ThreadPoolExecutor(max_workers=2) as executor:
|
||||
store = State(Path(tmp) / 'state.db')
|
||||
good = '600001.SH'
|
||||
positions = [PositionItem(stock_code=c, volume=100, can_use_volume=100, open_price=10)
|
||||
for c in [self.code, good]]
|
||||
store.sync_account(positions, [], initialize=True)
|
||||
manual = DealItem(stock_code=self.code, order_sys_id='manual', remark='',
|
||||
offset_flag=FLAG_BUY, volume=100, price=10, trade_amount=1000)
|
||||
self.run.executor = executor
|
||||
self.run.client.deals.return_value = [manual]
|
||||
self.run.client.portfolio.return_value = NS(
|
||||
assets=Assets(10000, 10000), positions={p.stock_code: p for p in positions}, orders=[])
|
||||
self.run.client.full_tick.return_value = {p.stock_code: Tick(last_price=11) for p in positions}
|
||||
with patch.object(boot, 'datetime') as clock, patch.object(boot, 'market_allow_open', return_value=True):
|
||||
clock.now.return_value = datetime(2026, 9, 11, 10)
|
||||
boot.RunOnce(self.run, store, [])
|
||||
self.assertEqual(store.blocked_codes, {self.code})
|
||||
self.run.orders.refresh.assert_called_once()
|
||||
self.assertEqual(self.run.orders.place.call_count, 1)
|
||||
self.assertEqual(self.run.orders.place.call_args.args[1].code, good)
|
||||
|
||||
def test_run_once_does_not_reopen_quarantined_sold_out_code(self):
|
||||
with tempfile.TemporaryDirectory() as tmp, ThreadPoolExecutor(max_workers=2) as executor:
|
||||
store = State(Path(tmp) / 'state.db')
|
||||
store.sync_account([], [], initialize=True)
|
||||
self.run.executor = executor
|
||||
self.run.client.deals.return_value = [DealItem(
|
||||
stock_code=self.code, order_sys_id='manual', remark='', offset_flag=FLAG_BUY,
|
||||
volume=100, price=10, trade_amount=1000)]
|
||||
self.run.client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={}, orders=[])
|
||||
self.run.client.full_tick.return_value = {}
|
||||
with patch.object(boot, 'datetime') as clock, patch.object(boot, 'market_allow_open', return_value=True), \
|
||||
patch.object(boot, 'open_signal') as opened:
|
||||
clock.now.return_value = datetime(2026, 9, 11, 10)
|
||||
boot.RunOnce(self.run, store, [NS(code=self.code)])
|
||||
opened.assert_not_called()
|
||||
|
||||
def start(self, client, directory):
|
||||
self.run.account_cfg.grid_step_pct = 1
|
||||
global_cfg = NS(qmt_base_url='unused', qmt_token='', qmt_data_dir=directory)
|
||||
self.run.account_cfg.signal_allow = []
|
||||
with patch.object(boot, 'Client', return_value=client), \
|
||||
patch.object(boot.config, 'global_config', global_cfg), \
|
||||
patch.object(boot.config, 'account_config', self.run.account_cfg), \
|
||||
patch.object(boot, 'init_signals', return_value=[]), \
|
||||
patch.object(boot, 'cache_portfolio'), patch.object(boot, 'Overview'), \
|
||||
patch.object(boot.time, 'localtime', return_value=NS(tm_hour=15, tm_min=0, tm_sec=0)):
|
||||
boot.StartZT()
|
||||
|
||||
def test_start_initializes_once_without_snapshot_retry_loop(self):
|
||||
client = Mock()
|
||||
client.deals.return_value = []
|
||||
client.portfolio.return_value = NS(assets=Assets(10000, 10000),
|
||||
positions={self.code: PositionItem(stock_code=self.code, volume=100, open_price=10)}, orders=[])
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
self.start(client, tmp)
|
||||
store = State(Path(tmp) / 'zt_test_state.db')
|
||||
self.assertEqual(store.state[self.code]['base_qty'], 100)
|
||||
with closing(store._connect()) as db:
|
||||
self.assertIsNone(db.execute("SELECT 1 FROM sqlite_master WHERE name='state_meta'").fetchone())
|
||||
self.assertEqual(client.portfolio.call_count, 1)
|
||||
self.assertEqual(client.deals.call_count, 2)
|
||||
client.reset_mock()
|
||||
self.start(client, tmp)
|
||||
self.assertEqual(client.portfolio.call_count, 1)
|
||||
self.assertEqual(client.deals.call_count, 1)
|
||||
|
||||
def test_start_rejects_changed_deals_without_writing_baseline(self):
|
||||
client = Mock()
|
||||
client.deals.side_effect = [[], [DealItem(order_sys_id='new')]]
|
||||
client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={}, orders=[])
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
with self.assertRaises(RuntimeError):
|
||||
self.start(client, tmp)
|
||||
store = State(Path(tmp) / 'zt_test_state.db')
|
||||
self.assertEqual((store.state, store.deals), ({}, {}))
|
||||
client.close.assert_called_once()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
Reference in New Issue
Block a user