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 types import SimpleNamespace
from unittest.mock import Mock
from libs.order import OrderBook as ActiveOrders
from libs.orderbook import OrderBook
from libs.state import State
from sdk.models import Assets, DealItem, OrderItem, PositionItem
from sdk.portfolio import PortfolioMixin
@@ -76,16 +76,21 @@ class ApiModelTests(unittest.TestCase):
deal = self.client.deals()[0]
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / 'state.db'
book = OrderBook(path)
book = State(path)
book.sync_deals([deal])
loaded = OrderBook(path).deals['sys1']
loaded = State(path).deals['sys1']
self.assertEqual(loaded['volume'], deal.volume)
self.assertEqual(loaded['trade_date'], '2026-09-07')
deal.trade_amount = 0
self.assertEqual(OrderBook.deal_record(deal)['trade_amount'], 1000)
deal.price = 0
with self.assertRaises(ValueError):
OrderBook.deal_record(deal)
deal.order_sys_id = 'sys2'
deal.trade_amount = 0
book.sync_deals([deal, deal])
self.assertEqual(book.deals['sys2']['trade_amount'], 1000)
self.assertEqual(deal.trade_amount, 0)
deal.order_sys_id = 'sys3'
deal.price = 0
with self.assertRaises(ValueError):
book.sync_deals([deal])
self.assertEqual(set(State(path).deals), {'sys1', 'sys2'})
if __name__ == '__main__':

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)

View File

@@ -0,0 +1,183 @@
import sqlite3
import tempfile
import unittest
from contextlib import closing
from pathlib import Path
from libs.state import State
from sdk import DealItem, PositionItem
class ArchivingTests(unittest.TestCase):
def setUp(self):
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
self.book = State(Path(tmp.name) / 'state.db')
def insert_deal(self, order, qty, amount, time, code='600000.SH', flag=48):
with closing(self.book._connect()) as db, db:
db.execute(
'INSERT INTO deals (stock_code, order_sys_id, order_local_id, offset_flag, '
'price, volume, trade_amount, trade_date, trade_time) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)',
(code, order, order, flag, amount / qty, qty, amount, '2026-09-08', time),
)
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'])
self.insert_deal('first', 40, 400, '10:00:00')
self.insert_deal('second', 60, 720, '10:01:00')
self.book.archiving()
row = self.book.state['600000.SH']
self.assertEqual(row['added_qty'], 100)
self.assertAlmostEqual(row['added_price'], 11.2)
self.assertEqual(row['added_order_local_id'], 'second')
self.assertEqual(row['added_created_at'], '2026-09-08 10:01:00')
for key in original:
if not key.startswith('added_'):
self.assertEqual(row[key], original[key])
self.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
restarted = State(self.book.path)
restarted.archiving()
self.assertEqual(restarted.state, self.book.state)
self.insert_deal('late', 50, 500, '09:59:00')
self.book.archiving()
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')
def test_sell_added_then_clear_base(self):
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
self.insert_deal('buy1', 40, 400, '09:59:00')
self.insert_deal('buy', 60, 600, '10:00:00')
self.insert_deal('partial', 40, 480, '10:01:00', flag=49)
self.book.archiving()
row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['base_price'], row['added_qty'], row['added_price']),
(100, 8, 60, 10))
self.insert_deal('sell_added', 60, 720, '10:02:00', flag=49)
self.book.archiving()
row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['added_qty'], row['added_price']), (100, 0, 0))
self.assertEqual((row['added_order_local_id'], row['added_created_at']), ('', ''))
self.insert_deal('sell_base', 100, 1200, '10:03:00', flag=49)
self.book.archiving()
self.assertEqual(self.book.state, {})
self.book.archiving()
self.assertEqual(State(self.book.path).state, {})
self.assertTrue(all(deal['is_arch'] == 1 for deal in self.book.deals.values()))
def test_sell_crosses_into_base_then_liquidates(self):
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=100, open_price=8)])
self.insert_deal('buy', 50, 500, '10:00:00')
self.insert_deal('sell', 80, 960, '10:01:00', flag=49)
self.book.archiving()
row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['base_price'], row['added_qty']), (70, 8, 0))
self.insert_deal('buy_again', 30, 300, '10:02:00')
self.insert_deal('sell_all', 100, 1200, '10:03:00', flag=49)
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)
self.insert_deal('buy', 50, 500, '10:00:00')
self.insert_deal('sell', 100, 1200, '10:01:00', flag=49)
errors = self.book.archiving()
self.assertIn('600000.SH', errors)
restarted = State(self.book.path)
self.assertEqual(restarted.state, {})
self.assertTrue(all(deal['is_arch'] == 0 for deal in restarted.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')
self.book.load()
with closing(self.book._connect()) as db:
db.execute("""CREATE TRIGGER fail_archive BEFORE UPDATE OF is_arch ON deals
WHEN OLD.order_sys_id = 'second'
BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
errors = self.book.archiving()
self.assertIn('600001.SH', errors)
restarted = State(self.book.path)
self.assertEqual(set(restarted.state), {'600000.SH'})
self.assertEqual(restarted.state, self.book.state)
self.assertEqual(restarted.deals['first']['is_arch'], 1)
self.assertEqual(restarted.deals['second']['is_arch'], 0)
with closing(self.book._connect()) as db:
db.execute('DROP TRIGGER fail_archive')
self.book.archiving()
self.assertEqual(len(self.book.state), 2)
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):
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',
offset_flag=48, volume=100, price=8, trade_amount=800,
trade_date='20260908', trade_time='100000',
)])
self.assertEqual(self.book.archiving(), {})
row = self.book.state['600000.SH']
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.assertEqual(restarted.archiving(), {})
self.assertEqual(restarted.state, self.book.state)
self.insert_deal('new_buy', 100, 1000, '10:01:00')
with closing(self.book._connect()) as db, db:
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'))
def test_old_archived_history_without_baseline_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(self.book.deals['old']['is_arch'], 1)
self.assertEqual(self.book.deals['new']['is_arch'], 0)
def test_late_buy_replays_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()
self.assertEqual(self.book.state, {})
self.book = State(self.book.path)
self.insert_deal('late', 100, 2000, '100100')
self.assertEqual(self.book.archiving(), {})
row = self.book.state['600000.SH']
self.assertEqual(row['added_qty'], 100)
self.assertEqual(row['added_price'], 15)
def test_snapshot_deletion_before_sell_archiving(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.book.sync_state([])
self.book = State(self.book.path)
self.assertEqual(self.book.archiving(), {})
self.assertEqual(self.book.state, {})
self.assertEqual(self.book.deals['sell']['is_arch'], 1)
def test_bad_stock_does_not_block_good_stock_and_can_retry(self):
self.insert_deal('bad_sell', 100, 1500, '10:02:00', flag=49)
self.insert_deal('good_buy', 100, 1000, '10:00:00', code='600001.SH')
self.assertIn('600000.SH', self.book.archiving())
self.assertEqual(self.book.deals['good_buy']['is_arch'], 1)
self.assertEqual(self.book.deals['bad_sell']['is_arch'], 0)
self.insert_deal('late_buy', 100, 1000, '10:01:00')
self.assertEqual(self.book.archiving(), {})
self.assertEqual(self.book.deals['bad_sell']['is_arch'], 1)
self.assertNotIn('600000.SH', self.book.state)
if __name__ == '__main__':
unittest.main()