This commit is contained in:
2026-09-15 20:02:05 +08:00
parent bfe89ba122
commit 04daeff141
37 changed files with 2674 additions and 1426 deletions

View File

@@ -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__':