optz print.

This commit is contained in:
2026-09-08 15:18:09 +08:00
parent 2a1458e91e
commit 8a3a29268e
12 changed files with 538 additions and 404 deletions

View File

@@ -7,7 +7,7 @@ from pathlib import Path
from datetime import datetime
from unittest.mock import patch
from libs.orderbook import OrderBook
from libs.state import State, StateItem
from sdk import DealItem, PositionItem
from strategy.zt.state import DONE, READY, SOLD, TState
@@ -63,14 +63,14 @@ class OrderBookTests(unittest.TestCase):
def test_json_is_never_read(self):
legacy = self.path.with_suffix('.json')
legacy.write_text('invalid JSON', encoding='utf-8')
book = OrderBook(self.path)
book = State(self.path)
self.assertIsNone(book.load())
self.assertEqual((book.positions, book.deals, book.deals_sys_ids), ({}, {}, set()))
self.assertEqual((book.items, 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 = OrderBook(self.path)
self.assertEqual((book.positions, book.deals, book.deals_sys_ids), ({}, {}, set()))
book = State(self.path)
self.assertEqual((book.items, 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])
@@ -78,7 +78,7 @@ class OrderBookTests(unittest.TestCase):
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 = OrderBook(self.path)
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:
@@ -88,7 +88,7 @@ class OrderBookTests(unittest.TestCase):
self.assertEqual(len(book.deals), 2)
def test_sync_deals_failure_rolls_back_entire_batch_and_cache(self):
book = OrderBook(self.path)
book = State(self.path)
first = self.deal('base', 'd1', 100, 10)
invalid = self.deal('base', 'd2', 100, 10)
invalid.offset_flag = -1
@@ -96,18 +96,18 @@ class OrderBookTests(unittest.TestCase):
book.sync_deals([first, invalid])
self.assertEqual(book.deals, {})
self.assertEqual(book.deals_sys_ids, set())
self.assertEqual(OrderBook(self.path).deals, {})
self.assertEqual(State(self.path).deals, {})
invalid.offset_flag = 23
book.sync_deals([first, invalid])
self.assertEqual(book.deals_sys_ids, {'d1', 'd2'})
def test_load_refreshes_all_caches(self):
book = OrderBook(self.path)
writer = OrderBook(self.path)
writer.sync_positions([PositionItem(stock_code='600000.SH', volume=100)])
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.positions['600000.SH']['volume'], 100)
self.assertEqual(book.items['600000.SH']['base_qty'], 100)
self.assertEqual(book.deals_sys_ids, {'d1'})
self.assertEqual(book.deals['d1']['remark'], 'zt-base-order1|zt')
@@ -140,11 +140,11 @@ class OrderBookTests(unittest.TestCase):
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, {'positions', 'deals', 'sqlite_sequence'})
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', *(field.name for field in fields(DealItem))})
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_positions_stock_code', 'idx_deals_order_sys_id',
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']),
@@ -165,29 +165,62 @@ class OrderBookTests(unittest.TestCase):
self.assertEqual(state.items['600000.SH'].base_cost, 10)
def test_position_columns_defaults_indexes_and_stable_id(self):
store = OrderBook(self.path)
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(positions)')}
self.assertEqual(columns, {'id', *(field.name for field in fields(PositionItem))})
indexes = {row[1] for row in db.execute('PRAGMA index_list(positions)')}
self.assertEqual(indexes, {'idx_positions_stock_code'})
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_positions([position])
saved = store.positions[position.stock_code]
store.sync_state([position])
saved = store.items[position.stock_code]
first_id = saved['id']
self.assertEqual({k: v for k, v in saved.items() if k != 'id'}, asdict(position))
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
store.sync_positions([position])
self.assertEqual(store.positions[position.stock_code]['id'], first_id)
self.assertEqual(store.positions[position.stock_code]['volume'], 200)
store.sync_positions([])
self.assertEqual(store.positions, {})
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)
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)
position.volume = 0
store.sync_state([position, PositionItem(stock_code='600002.SH')])
self.assertEqual(store.items, {})
self.assertEqual(State(self.path).items, {})
store.sync_state([PositionItem(stock_code='600001.SH', volume=100)])
self.assertGreater(store.items['600001.SH']['id'], first_id)
store.sync_state([])
self.assertEqual(store.items, {})
self.assertEqual(store.deals, saved_deals)
store.sync_positions([PositionItem(stock_code='600001.SH')])
self.assertGreater(store.positions['600001.SH']['id'], first_id)
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',
))
book.save({row['stock_code']: row})
saved = book.items[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
with self.assertRaises(sqlite3.IntegrityError):
book.save({row['stock_code']: row})
self.assertEqual(State(self.path).items[row['stock_code']], saved)
def test_base_split_fills_and_snapshot_do_not_double_count(self):
state = TState(self.path)