dev zt
This commit is contained in:
@@ -9,15 +9,10 @@ from unittest.mock import patch
|
||||
|
||||
from libs.state import State, StateItem
|
||||
from sdk import DealItem, PositionItem
|
||||
from strategy.zt.state import DONE, READY, SOLD, TState
|
||||
|
||||
|
||||
class OrderBookTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
clock = patch('strategy.zt.state.datetime')
|
||||
self.clock = clock.start()
|
||||
self.addCleanup(clock.stop)
|
||||
self.clock.now.return_value = datetime(2026, 9, 1)
|
||||
self.tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(self.tmp.cleanup)
|
||||
self.path = Path(self.tmp.name) / 'state.db'
|
||||
@@ -31,46 +26,18 @@ class OrderBookTests(unittest.TestCase):
|
||||
trade_date=date, trade_time='10:00:00', remark=prefix + 'order1|zt',
|
||||
)
|
||||
|
||||
def test_partial_fills_restart_dedup_and_daily_cycle(self):
|
||||
state = TState(self.path)
|
||||
state.reconcile([PositionItem(stock_code='600000.SH', volume=200, open_price=10)], [])
|
||||
self.assertEqual(state.deals, [])
|
||||
first = self.deal('sell', 'd1', 40, 12)
|
||||
second = self.deal('sell', 'd2', 60, 13)
|
||||
state.reconcile([], [first])
|
||||
state = TState(self.path)
|
||||
self.assertEqual(state.items['600000.SH'].phase, SOLD)
|
||||
self.assertEqual(state.items['600000.SH'].sell_qty, 40)
|
||||
state.reconcile([], [first, first, second])
|
||||
self.assertEqual(len(state.deals), 2)
|
||||
self.assertAlmostEqual(state.items['600000.SH'].sell_price, 12.6)
|
||||
self.clock.now.return_value = datetime.fromisoformat('2026-09-02')
|
||||
state.reconcile([], [first, second])
|
||||
self.assertEqual(len(state.deals), 2)
|
||||
self.assertEqual(state.items['600000.SH'].phase, SOLD)
|
||||
b1 = self.deal('buy', 'd3', 40, 11, '2026-09-02')
|
||||
b2 = self.deal('buy', 'd4', 60, 10, '2026-09-02')
|
||||
state.reconcile([], [b1])
|
||||
self.assertEqual(state.items['600000.SH'].phase, SOLD)
|
||||
state.reconcile([], [b1, b2])
|
||||
self.assertEqual(TState(self.path).items['600000.SH'].phase, DONE)
|
||||
self.assertAlmostEqual(state.items['600000.SH'].buy_cost, 10.4)
|
||||
self.clock.now.return_value = datetime.fromisoformat('2026-09-03')
|
||||
state.reconcile([], [])
|
||||
item = TState(self.path).items['600000.SH']
|
||||
self.assertEqual((item.phase, item.base_qty, item.base_cost, item.sell_qty), (READY, 200, 10, 0))
|
||||
|
||||
def test_json_is_never_read(self):
|
||||
legacy = self.path.with_suffix('.json')
|
||||
legacy.write_text('invalid JSON', encoding='utf-8')
|
||||
book = State(self.path)
|
||||
self.assertIsNone(book.load())
|
||||
self.assertEqual((book.items, book.deals, book.deals_sys_ids), ({}, {}, set()))
|
||||
self.assertEqual((book.state, book.deals, book.deals_sys_ids), ({}, {}, set()))
|
||||
self.assertEqual(legacy.read_text(encoding='utf-8'), 'invalid JSON')
|
||||
|
||||
def test_sync_deals_deduplicates_batch_and_restart(self):
|
||||
book = State(self.path)
|
||||
self.assertEqual((book.items, book.deals, book.deals_sys_ids), ({}, {}, set()))
|
||||
self.assertEqual((book.state, book.deals, book.deals_sys_ids), ({}, {}, set()))
|
||||
first = self.deal('base', 'd1', 40, 10, '20260901')
|
||||
second = self.deal('base', 'd2', 60, 12)
|
||||
book.sync_deals([first, first, second])
|
||||
@@ -107,62 +74,12 @@ class OrderBookTests(unittest.TestCase):
|
||||
writer.sync_state([PositionItem(stock_code='600000.SH', volume=100)])
|
||||
writer.sync_deals([self.deal('base', 'd1', 100, 10)])
|
||||
book.load()
|
||||
self.assertEqual(book.items['600000.SH']['base_qty'], 100)
|
||||
self.assertEqual(book.state['600000.SH']['base_qty'], 100)
|
||||
self.assertEqual(book.deals_sys_ids, {'d1'})
|
||||
self.assertEqual(book.deals['d1']['remark'], 'zt-base-order1|zt')
|
||||
|
||||
def test_first_start_after_full_sale_keeps_buyback_quantity(self):
|
||||
state = TState(self.path)
|
||||
state.reconcile([], [self.deal('sell', 'd1', 100, 12)])
|
||||
item = state.items['600000.SH']
|
||||
self.assertEqual((item.base_qty, item.sell_qty, item.phase), (100, 100, SOLD))
|
||||
state.reconcile([], [self.deal('buy', 'd2', 100, 11)])
|
||||
self.assertEqual(TState(self.path).items['600000.SH'].phase, DONE)
|
||||
|
||||
def test_failed_insert_rolls_back_memory_and_database(self):
|
||||
state = TState(self.path)
|
||||
with closing(sqlite3.connect(self.path)) as db:
|
||||
db.execute("""CREATE TRIGGER fail_insert BEFORE INSERT ON deals
|
||||
BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
|
||||
fill = self.deal('base', 'd1', 100, 10)
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
state.reconcile([], [fill])
|
||||
self.assertFalse(state.items)
|
||||
self.assertFalse(state.deals)
|
||||
self.assertFalse(TState(self.path).items)
|
||||
with closing(sqlite3.connect(self.path)) as db:
|
||||
db.execute('DROP TRIGGER fail_insert')
|
||||
state.reconcile([], [fill])
|
||||
self.assertEqual(TState(self.path).items['600000.SH'].base_qty, 100)
|
||||
|
||||
def test_schema_and_unique_execution(self):
|
||||
state = TState(self.path)
|
||||
state.reconcile([], [self.deal('base', 'd1', 100, 10)])
|
||||
with closing(sqlite3.connect(self.path)) as db:
|
||||
tables = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='table'")}
|
||||
self.assertEqual(tables, {'state', 'deals', 'sqlite_sequence'})
|
||||
columns = {row[1] for row in db.execute('PRAGMA table_info(deals)')}
|
||||
self.assertEqual(columns, {'id', 'order_local_id', 'is_arch', *(field.name for field in fields(DealItem))})
|
||||
indexes = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='index'")}
|
||||
self.assertTrue({'idx_state_stock_code', 'idx_deals_order_sys_id',
|
||||
'idx_deals_order_ref', 'idx_deals_stock_code_date', 'idx_deals_date_time'} <= indexes)
|
||||
for index, expected in (
|
||||
('idx_deals_order_sys_id', ['order_sys_id']),
|
||||
('idx_deals_order_ref', ['order_local_id']),
|
||||
('idx_deals_stock_code_date', ['stock_code']),
|
||||
('idx_deals_date_time', ['trade_date']),
|
||||
):
|
||||
self.assertEqual([row[2] for row in db.execute(f'PRAGMA index_info({index})')], expected)
|
||||
self.assertNotIn('kind', state.deals[0])
|
||||
self.assertEqual(TState(self.path).deals, state.deals)
|
||||
state.deals.append(dict(state.deals[0]))
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
state.save()
|
||||
self.assertEqual(len(state.deals), 1)
|
||||
state.items['600000.SH'].base_cost = float('inf')
|
||||
with self.assertRaises(ValueError):
|
||||
state.save()
|
||||
self.assertEqual(state.items['600000.SH'].base_cost, 10)
|
||||
|
||||
def test_position_columns_defaults_indexes_and_stable_id(self):
|
||||
store = State(self.path)
|
||||
@@ -176,7 +93,7 @@ class OrderBookTests(unittest.TestCase):
|
||||
position = PositionItem(stock_code='600000.SH', volume=100, open_price=10,
|
||||
stock_name='stock', can_use_volume=100, float_profit=-2.5)
|
||||
store.sync_state([position])
|
||||
saved = store.items[position.stock_code]
|
||||
saved = store.state[position.stock_code]
|
||||
first_id = saved['id']
|
||||
self.assertEqual(saved['base_qty'], 100)
|
||||
self.assertEqual(saved['base_price'], 10)
|
||||
@@ -186,20 +103,20 @@ class OrderBookTests(unittest.TestCase):
|
||||
position.volume = 200
|
||||
position.open_price = 12
|
||||
store.sync_state([position])
|
||||
self.assertEqual(store.items[position.stock_code]['id'], first_id)
|
||||
self.assertEqual(store.items[position.stock_code], saved)
|
||||
self.assertEqual(State(self.path).items[position.stock_code], saved)
|
||||
self.assertEqual(store.state[position.stock_code]['id'], first_id)
|
||||
self.assertEqual(store.state[position.stock_code], saved)
|
||||
self.assertEqual(State(self.path).state[position.stock_code], saved)
|
||||
store.sync_state([position, PositionItem(stock_code='600001.SH', volume=100)])
|
||||
self.assertEqual(store.items[position.stock_code], saved)
|
||||
self.assertEqual(store.items['600001.SH']['base_qty'], 100)
|
||||
self.assertEqual(store.state[position.stock_code], saved)
|
||||
self.assertEqual(store.state['600001.SH']['base_qty'], 100)
|
||||
position.volume = 0
|
||||
store.sync_state([position, PositionItem(stock_code='600002.SH')])
|
||||
self.assertEqual(store.items, {})
|
||||
self.assertEqual(State(self.path).items, {})
|
||||
self.assertEqual(store.state, {})
|
||||
self.assertEqual(State(self.path).state, {})
|
||||
store.sync_state([PositionItem(stock_code='600001.SH', volume=100)])
|
||||
self.assertGreater(store.items['600001.SH']['id'], first_id)
|
||||
self.assertGreater(store.state['600001.SH']['id'], first_id)
|
||||
store.sync_state([])
|
||||
self.assertEqual(store.items, {})
|
||||
self.assertEqual(store.state, {})
|
||||
self.assertEqual(store.deals, saved_deals)
|
||||
|
||||
def test_state_fields_survive_restart_and_sync(self):
|
||||
@@ -211,37 +128,23 @@ class OrderBookTests(unittest.TestCase):
|
||||
added_order_local_id='added-1', added_qty=50, added_price=9,
|
||||
added_created_at='2026-09-08T10:30:00',
|
||||
))
|
||||
book.save({row['stock_code']: row})
|
||||
saved = book.items[row['stock_code']]
|
||||
with closing(book._connect()) as db, db:
|
||||
db.execute(
|
||||
f"INSERT INTO state ({', '.join(row)}) VALUES ({', '.join(':' + key for key in row)})",
|
||||
row,
|
||||
)
|
||||
book.load()
|
||||
saved = book.state[row['stock_code']]
|
||||
self.assertEqual({k: v for k, v in saved.items() if k != 'id'}, row)
|
||||
book = State(self.path)
|
||||
book.sync_state([PositionItem(stock_code=row['stock_code'], volume=150, open_price=9.5)])
|
||||
self.assertEqual(book.items[row['stock_code']], saved)
|
||||
row['added_qty'] = -1
|
||||
self.assertEqual(book.state[row['stock_code']], saved)
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
book.save({row['stock_code']: row})
|
||||
self.assertEqual(State(self.path).items[row['stock_code']], saved)
|
||||
with closing(book._connect()) as db, db:
|
||||
db.execute('UPDATE state SET added_qty = -1')
|
||||
self.assertEqual(State(self.path).state[row['stock_code']], saved)
|
||||
|
||||
def test_base_split_fills_and_snapshot_do_not_double_count(self):
|
||||
state = TState(self.path)
|
||||
first = self.deal('base', 'd1', 40, 10)
|
||||
state.reconcile([PositionItem(stock_code='600000.SH', volume=40, open_price=10)], [first])
|
||||
second = self.deal('base', 'd2', 60, 12)
|
||||
state.reconcile([PositionItem(stock_code='600000.SH', volume=100, open_price=11.2)], [first, second])
|
||||
self.assertEqual(state.items['600000.SH'].base_qty, 100)
|
||||
self.assertAlmostEqual(state.items['600000.SH'].base_cost, 11.2)
|
||||
self.assertEqual(len(state.deals), 2)
|
||||
|
||||
def test_date_normalization_and_unrelated_strategy(self):
|
||||
state = TState(self.path)
|
||||
first = self.deal('base', 'd1', 100, 10, '20260901')
|
||||
other = self.deal('base', 'd2', 100, 10)
|
||||
other.remark = 'trend-base-order'
|
||||
state.reconcile([], [first, other])
|
||||
first.trade_date = '2026-09-01'
|
||||
state.reconcile([], [first])
|
||||
self.assertEqual(len(state.deals), 1)
|
||||
self.assertEqual(state.deals[0]['trade_date'], '2026-09-01')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
@@ -22,6 +22,11 @@ class ArchivingTests(unittest.TestCase):
|
||||
(code, order, order, flag, amount / qty, qty, amount, '2026-09-08', time),
|
||||
)
|
||||
|
||||
def test_schema_only_has_state_and_deals(self):
|
||||
with closing(self.book._connect()) as db:
|
||||
tables = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='table'")}
|
||||
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)])
|
||||
original = dict(self.book.state['600000.SH'])
|
||||
@@ -45,7 +50,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual(row['added_qty'], 150)
|
||||
self.assertAlmostEqual(row['added_price'], 1620 / 150)
|
||||
self.assertEqual(row['added_order_local_id'], 'second')
|
||||
self.assertEqual(row['added_order_local_id'], 'late')
|
||||
|
||||
def test_sell_added_then_clear_base(self):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
@@ -80,11 +85,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.book.archiving()
|
||||
self.assertEqual(self.book.state, {})
|
||||
|
||||
def test_excess_sell_rolls_back_and_other_directions_are_skipped(self):
|
||||
self.insert_deal('other', 100, 1000, '09:59:00', flag=23)
|
||||
self.book.archiving()
|
||||
self.assertEqual(self.book.state, {})
|
||||
self.assertEqual(self.book.deals['other']['is_arch'], 0)
|
||||
def test_excess_sell_rolls_back(self):
|
||||
self.insert_deal('buy', 50, 500, '10:00:00')
|
||||
self.insert_deal('sell', 100, 1200, '10:01:00', flag=49)
|
||||
errors = self.book.archiving()
|
||||
@@ -93,6 +94,17 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertEqual(restarted.state, {})
|
||||
self.assertTrue(all(deal['is_arch'] == 0 for deal in restarted.deals.values()))
|
||||
|
||||
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.assertEqual(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.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
|
||||
|
||||
def test_new_state_and_failed_mark_roll_back_together(self):
|
||||
self.insert_deal('first', 100, 1000, '10:00:00')
|
||||
self.insert_deal('second', 100, 1200, '10:01:00', code='600001.SH')
|
||||
@@ -115,7 +127,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_snapshot_matching_buy_is_not_added(self):
|
||||
def test_equal_quantity_buy_is_added_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',
|
||||
@@ -124,7 +136,7 @@ class ArchivingTests(unittest.TestCase):
|
||||
)])
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 0))
|
||||
self.assertEqual((row['base_qty'], row['added_qty']), (100, 100))
|
||||
self.assertEqual(self.book.deals['first']['is_arch'], 1)
|
||||
restarted = State(self.book.path)
|
||||
self.assertEqual(restarted.archiving(), {})
|
||||
@@ -134,19 +146,20 @@ class ArchivingTests(unittest.TestCase):
|
||||
db.execute("UPDATE state SET status = 'CUSTOM' WHERE stock_code = '600000.SH'")
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 100, 'CUSTOM'))
|
||||
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 200, 'CUSTOM'))
|
||||
|
||||
def test_old_archived_history_without_baseline_is_not_reapplied(self):
|
||||
def test_archived_history_is_not_reapplied(self):
|
||||
self.insert_deal('old', 100, 1000, '10:00:00')
|
||||
with closing(self.book._connect()) as db, db:
|
||||
db.execute('UPDATE deals SET is_arch = 1')
|
||||
self.insert_deal('new', 50, 500, '10:01:00')
|
||||
errors = self.book.archiving()
|
||||
self.assertIn('baseline', errors['600000.SH'])
|
||||
self.assertEqual(errors, {})
|
||||
self.assertEqual(self.book.state['600000.SH']['added_qty'], 50)
|
||||
self.assertEqual(self.book.deals['old']['is_arch'], 1)
|
||||
self.assertEqual(self.book.deals['new']['is_arch'], 0)
|
||||
self.assertEqual(self.book.deals['new']['is_arch'], 1)
|
||||
|
||||
def test_late_buy_replays_after_liquidation_and_restart(self):
|
||||
def test_late_buy_is_incremental_after_liquidation_and_restart(self):
|
||||
self.insert_deal('buy', 100, 1000, '10:00:00')
|
||||
self.insert_deal('sell', 100, 1500, '10:02:00', flag=49)
|
||||
self.book.archiving()
|
||||
@@ -156,11 +169,12 @@ class ArchivingTests(unittest.TestCase):
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
row = self.book.state['600000.SH']
|
||||
self.assertEqual(row['added_qty'], 100)
|
||||
self.assertEqual(row['added_price'], 15)
|
||||
self.assertEqual(row['added_price'], 20)
|
||||
|
||||
def test_snapshot_deletion_before_sell_archiving(self):
|
||||
def test_archive_sell_before_syncing_empty_positions(self):
|
||||
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
|
||||
self.insert_deal('sell', 100, 1000, '10:01:00', flag=49)
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
self.book.sync_state([])
|
||||
self.book = State(self.book.path)
|
||||
self.assertEqual(self.book.archiving(), {})
|
||||
|
||||
54
py-client/tests/test_zt_state.py
Normal file
54
py-client/tests/test_zt_state.py
Normal file
@@ -0,0 +1,54 @@
|
||||
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()
|
||||
136
py-client/tests/test_zt_trading.py
Normal file
136
py-client/tests/test_zt_trading.py
Normal file
@@ -0,0 +1,136 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from config import AccountConfig
|
||||
from libs.grid_take_profit import GridState
|
||||
from libs.state import 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
|
||||
|
||||
|
||||
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.orders.busy.return_value = False
|
||||
self.run.orders.new_order_id.side_effect = lambda kind: f'zt-{kind}-order'
|
||||
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, 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_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()
|
||||
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()
|
||||
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()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user