fix bug
This commit is contained in:
@@ -1,149 +1,85 @@
|
||||
import sqlite3
|
||||
import tempfile
|
||||
"""委托簿:在途状态、方向锁、以及撤单范围。"""
|
||||
|
||||
import unittest
|
||||
from contextlib import closing
|
||||
from dataclasses import asdict, fields
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
from datetime import datetime, timedelta
|
||||
from unittest.mock import Mock
|
||||
|
||||
from libs.state import FLAG_BUY, FLAG_SELL, State, StateItem
|
||||
from sdk import DealItem, PositionItem
|
||||
from libs.order import BUSY_STATUSES, OrderBook, TRACKED_STATUSES
|
||||
from sdk import OrderItem
|
||||
|
||||
|
||||
class OrderBookTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.tmp = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(self.tmp.cleanup)
|
||||
self.path = Path(self.tmp.name) / 'state.db'
|
||||
|
||||
def deal(self, kind, sys_order_id, qty, price, date='2026-09-01'):
|
||||
prefix = {'base': 'zt-base-', 'sell': 'zt-t-sell-', 'buy': 'zt-t-buy-'}[kind]
|
||||
return DealItem(
|
||||
order_sys_id=sys_order_id, stock_code='600000.SH',
|
||||
offset_flag=FLAG_SELL if kind == 'sell' else FLAG_BUY,
|
||||
volume=qty, price=price, trade_amount=qty * price,
|
||||
trade_date=date, trade_time='10:00:00', remark=prefix + 'order1|zt',
|
||||
)
|
||||
def order(index, remark, status=50, side=23, age_minutes=30):
|
||||
stamp = datetime.now() - timedelta(minutes=age_minutes)
|
||||
return OrderItem(stock_code=f'60000{index}.SH', order_sys_id=f'sys{index}',
|
||||
remark=remark, order_status=status, offset_flag=side,
|
||||
insert_date=stamp.strftime('%Y%m%d'),
|
||||
insert_time=stamp.strftime('%H%M%S'))
|
||||
|
||||
|
||||
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.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.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])
|
||||
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
|
||||
self.assertEqual(book.deals['d1']['trade_date'], '2026-09-01')
|
||||
self.assertEqual(book.deals['d2']['volume'], 60)
|
||||
self.assertEqual(book.deals['d2']['order_local_id'], 'zt-base-order1')
|
||||
book = State(self.path)
|
||||
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
|
||||
self.assertEqual(book.deals['d1']['order_local_id'], 'zt-base-order1')
|
||||
with patch.object(book, '_connect') as connect:
|
||||
book.sync_deals([first, second])
|
||||
book.sync_deals([])
|
||||
connect.assert_not_called()
|
||||
self.assertEqual(len(book.deals), 2)
|
||||
|
||||
def test_sync_deals_failure_rolls_back_entire_batch_and_cache(self):
|
||||
book = State(self.path)
|
||||
first = self.deal('base', 'd1', 100, 10)
|
||||
invalid = self.deal('base', 'd2', 100, 10)
|
||||
invalid.offset_flag = -1
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
book.sync_deals([first, invalid])
|
||||
self.assertEqual(book.deals, {})
|
||||
self.assertEqual(book.deals_sys_ids, set())
|
||||
self.assertEqual(State(self.path).deals, {})
|
||||
invalid.offset_flag = FLAG_BUY
|
||||
book.sync_deals([first, invalid])
|
||||
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
|
||||
|
||||
def test_load_refreshes_all_caches(self):
|
||||
book = State(self.path)
|
||||
writer = State(self.path)
|
||||
writer.sync_state([PositionItem(stock_code='600000.SH', volume=100)])
|
||||
writer.sync_deals([self.deal('base', 'd1', 100, 10)])
|
||||
book.load()
|
||||
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 cancelled(client):
|
||||
return [call.args[0] for call in client.cancel_by_id.call_args_list]
|
||||
|
||||
|
||||
class CancelScopeTests(unittest.TestCase):
|
||||
"""撤单必须限本策略前缀;防重则继续看全账户在途。"""
|
||||
|
||||
def test_zt_prefix_cancels_only_its_own_orders(self):
|
||||
orders = [order(0, 'zt-base-own'), order(1, 'zt-SELL-own'),
|
||||
order(2, 'zt-entry-own'), order(3, 'IPO-new'),
|
||||
order(4, ''), order(5, 'TREN-BUY-other')]
|
||||
client = Mock()
|
||||
book = OrderBook()
|
||||
book.refresh(client, orders, cancel_prefix='zt-')
|
||||
self.assertEqual(cancelled(client), ['sys0', 'sys1', 'sys2'])
|
||||
self.assertEqual(book.data, orders) # 撤单后仍保留在途锁
|
||||
self.assertTrue(all(book.busy(o.stock_code, 'BUY') for o in orders))
|
||||
|
||||
def test_default_prefix_still_cancels_every_non_ipo_order(self):
|
||||
orders = [order(0, 'zt-base-own'), order(1, 'TREN-BUY-other'),
|
||||
order(2, 'IPO-new'), order(3, '')]
|
||||
client = Mock()
|
||||
OrderBook().refresh(client, orders)
|
||||
self.assertEqual(cancelled(client), ['sys0', 'sys1', 'sys3'])
|
||||
|
||||
def test_fresh_and_unreportable_orders_are_never_cancelled(self):
|
||||
# 48(未报)在跟踪集合内但不可撤;刚提交的委托也不撤。
|
||||
orders = [order(0, 'zt-entry-fresh', age_minutes=0),
|
||||
order(1, 'zt-entry-filled', status=56),
|
||||
order(2, 'zt-entry-unreported', status=48)]
|
||||
client = Mock()
|
||||
book = OrderBook()
|
||||
book.refresh(client, orders, cancel_prefix='zt-')
|
||||
client.cancel_by_id.assert_not_called()
|
||||
self.assertEqual(book.data, orders)
|
||||
|
||||
|
||||
def test_position_columns_defaults_indexes_and_stable_id(self):
|
||||
store = State(self.path)
|
||||
store.sync_deals([self.deal('base', 'd1', 100, 10)])
|
||||
saved_deals = dict(store.deals)
|
||||
with closing(sqlite3.connect(self.path)) as db:
|
||||
columns = {row[1] for row in db.execute('PRAGMA table_info(state)')}
|
||||
self.assertEqual(columns, {'id', *(field.name for field in fields(StateItem))})
|
||||
indexes = {row[1] for row in db.execute('PRAGMA index_list(state)')}
|
||||
self.assertEqual(indexes, {'idx_state_stock_code'})
|
||||
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.state[position.stock_code]
|
||||
first_id = saved['id']
|
||||
self.assertEqual(saved['base_qty'], 100)
|
||||
self.assertEqual(saved['base_price'], 10)
|
||||
self.assertEqual(saved['added_qty'], 0)
|
||||
self.assertEqual(saved['base_order_local_id'], '')
|
||||
self.assertTrue(saved['base_created_at'])
|
||||
position.volume = 200
|
||||
position.open_price = 12
|
||||
store.sync_state([position])
|
||||
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.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.state, {})
|
||||
self.assertEqual(State(self.path).state, {})
|
||||
store.sync_state([PositionItem(stock_code='600001.SH', volume=100)])
|
||||
self.assertGreater(store.state['600001.SH']['id'], first_id)
|
||||
store.sync_state([])
|
||||
self.assertEqual(store.state, {})
|
||||
self.assertEqual(store.deals, saved_deals)
|
||||
class BusyLockTests(unittest.TestCase):
|
||||
def test_only_busy_statuses_lock_a_direction(self):
|
||||
book = OrderBook()
|
||||
client = Mock()
|
||||
book.refresh(client, [order(0, 'zt-entry-a', status=50)], cancel_prefix='zt-')
|
||||
self.assertTrue(book.busy('600000.SH', 'BUY'))
|
||||
book.refresh(client, [order(1, 'zt-entry-b', status=56)], cancel_prefix='zt-')
|
||||
self.assertFalse(book.busy('600001.SH', 'BUY'))
|
||||
|
||||
def test_state_fields_survive_restart_and_sync(self):
|
||||
book = State(self.path)
|
||||
row = asdict(StateItem(
|
||||
stock_code='600000.SH', status='READY',
|
||||
base_order_local_id='base-1', base_qty=100, base_price=10,
|
||||
base_created_at='2026-09-08T09:30:00',
|
||||
added_order_local_id='added-1', added_qty=50, added_price=9,
|
||||
added_created_at='2026-09-08T10:30:00',
|
||||
))
|
||||
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.state[row['stock_code']], saved)
|
||||
with self.assertRaises(sqlite3.IntegrityError):
|
||||
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_busy_and_tracked_status_sets_are_disjoint_as_designed(self):
|
||||
self.assertNotIn('56', BUSY_STATUSES)
|
||||
self.assertTrue(BUSY_STATUSES <= TRACKED_STATUSES)
|
||||
self.assertIn('56', TRACKED_STATUSES)
|
||||
|
||||
def test_place_marks_the_direction_busy_before_submitting(self):
|
||||
book = OrderBook()
|
||||
book.busy_cache.set('BUY-600000.SH', True, timeout=180)
|
||||
self.assertTrue(book.busy('600000.SH', 'BUY'))
|
||||
self.assertFalse(book.busy('600000.SH', 'SELL'))
|
||||
|
||||
def test_unknown_offset_flag_never_places_an_order(self):
|
||||
from libs.order import PlaceOrderRequest
|
||||
book = OrderBook()
|
||||
client = Mock()
|
||||
self.assertFalse(book.place(client, PlaceOrderRequest(99, '600000.SH', 100,
|
||||
'zt-x', 'zt')))
|
||||
client.passorder.assert_not_called()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
Reference in New Issue
Block a user