fix zt&state.py

This commit is contained in:
2026-09-12 13:42:23 +08:00
parent bcb03e2ed9
commit 7a7049ce44
9 changed files with 677 additions and 363 deletions

View File

@@ -1,4 +1,5 @@
import ast
import sqlite3
import tempfile
import unittest
from dataclasses import asdict, fields
@@ -88,7 +89,7 @@ class ApiModelTests(unittest.TestCase):
self.assertEqual(deal.trade_amount, 0)
deal.order_sys_id = 'sys3'
deal.price = 0
with self.assertRaises(ValueError):
with self.assertRaises(sqlite3.IntegrityError):
book.sync_deals([deal])
self.assertEqual(set(State(path).deals), {'sys1', 'sys2'})

View File

@@ -4,7 +4,7 @@ import unittest
from contextlib import closing
from pathlib import Path
from libs.state import State
from libs.state import FLAG_BUY, FLAG_SELL, State
from sdk import DealItem, PositionItem
@@ -28,7 +28,7 @@ class ArchivingTests(unittest.TestCase):
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)])
self.book.sync_state([PositionItem(stock_code='600000.SH', volume=200, 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')
@@ -53,9 +53,9 @@ class ArchivingTests(unittest.TestCase):
self.assertEqual(row['added_order_local_id'], 'late')
def test_no_argument_archiving_recognizes_base_and_added_orders(self):
self.insert_deal('zt-base-first', 100, 1000, '10:00:00', flag=23)
self.insert_deal('zt-base-first', 100, 1000, '10:00:00', flag=FLAG_BUY)
self.insert_deal('zt-base-second', 100, 1200, '10:01:00', flag=48)
self.insert_deal('zt-t-buy-first', 100, 900, '10:02:00', flag=23)
self.insert_deal('zt-t-buy-first', 100, 900, '10:02:00', flag=FLAG_BUY)
self.assertIsNone(self.book.archiving())
row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['base_price']), (200, 11))
@@ -111,13 +111,13 @@ class ArchivingTests(unittest.TestCase):
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.insert_deal('buy', 100, 1000, '10:00:00', flag=FLAG_BUY)
self.insert_deal('sell', 50, 600, '10:01:00', flag=FLAG_SELL)
self.assertIsNone(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.assertEqual(self.book.deals['buy']['offset_flag'], FLAG_BUY)
self.assertEqual(self.book.deals['sell']['offset_flag'], FLAG_SELL)
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):
@@ -143,7 +143,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_equal_quantity_buy_is_added_and_preserves_status(self):
def test_equal_quantity_buy_only_marks_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',
@@ -152,17 +152,17 @@ class ArchivingTests(unittest.TestCase):
)])
self.assertIsNone(self.book.archiving())
row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['added_qty']), (100, 100))
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.assertIsNone(restarted.archiving())
self.assertEqual(restarted.state, self.book.state)
self.insert_deal('new_buy', 100, 1000, '10:01:00')
self.insert_deal('new_buy', 50, 500, '10:01:00')
with closing(self.book._connect()) as db, db:
db.execute("UPDATE state SET status = 'CUSTOM' WHERE stock_code = '600000.SH'")
self.assertIsNone(self.book.archiving())
row = self.book.state['600000.SH']
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 200, 'CUSTOM'))
self.assertEqual((row['base_qty'], row['added_qty'], row['status']), (100, 50, 'CUSTOM'))
def test_archived_history_is_not_reapplied(self):
self.insert_deal('old', 100, 1000, '10:00:00')

View File

@@ -0,0 +1,92 @@
import sqlite3
import tempfile
import unittest
from contextlib import closing
from pathlib import Path
from libs.state import FLAG_BUY, FLAG_SELL, State
from sdk import DealItem, PositionItem
class SnapshotArchiveTests(unittest.TestCase):
def setUp(self):
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
self.store = State(Path(tmp.name) / 'state.db')
self.code = '600000.SH'
def deal(self, identity, qty, flag=FLAG_BUY, code=None):
return DealItem(
stock_code=code or self.code, order_sys_id=identity,
remark=f'zt-base-{identity}|zt', offset_flag=flag,
volume=qty, price=12, trade_amount=qty * 12,
trade_date='2026-09-12', trade_time='100000',
)
def snapshot(self, qty):
self.store.sync_state([PositionItem(stock_code=self.code, volume=qty, open_price=10)])
def test_matching_partial_fills_only_mark_and_survive_restart(self):
self.snapshot(100)
saved = dict(self.store.state[self.code])
self.store.sync_deals([self.deal('one', 40), self.deal('two', 60)])
self.store.archiving()
self.assertEqual(self.store.state[self.code], saved)
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
restarted = State(self.store.path)
restarted.archiving()
self.assertEqual(restarted.state[self.code], saved)
def test_total_includes_added_holdings(self):
self.snapshot(100)
with closing(self.store._connect()) as db, db:
db.execute("UPDATE state SET added_qty=50, added_price=9, status='CUSTOM'")
self.store.load()
saved = dict(self.store.state[self.code])
self.store.sync_deals([self.deal('one', 150)])
self.store.archiving()
self.assertEqual(self.store.state[self.code], saved)
self.assertEqual(self.store.deals['one']['is_arch'], 1)
def test_nonmatching_and_other_stock_are_incremental(self):
self.snapshot(100)
self.store.sync_deals([self.deal('one', 40), self.deal('other', 60, code='600001.SH')])
self.store.archiving()
self.assertEqual(self.store.state[self.code]['base_qty'], 140)
self.assertEqual(self.store.state['600001.SH']['base_qty'], 60)
def test_matching_sell_still_liquidates(self):
self.snapshot(100)
self.store.sync_deals([self.deal('sell', 100, FLAG_SELL)])
self.store.archiving()
self.assertNotIn(self.code, self.store.state)
self.assertEqual(self.store.deals['sell']['is_arch'], 1)
def test_mixed_batch_with_matching_gross_volume_is_not_skipped(self):
self.snapshot(100)
self.store.sync_deals([self.deal('buy', 40), self.deal('sell', 60, FLAG_SELL)])
self.store.archiving()
self.assertEqual(self.store.state[self.code]['base_qty'], 80)
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
def test_failed_mark_rolls_back_entire_stock_and_retries(self):
self.snapshot(100)
saved = dict(self.store.state[self.code])
self.store.sync_deals([self.deal('one', 40), self.deal('two', 60)])
with closing(self.store._connect()) as db, db:
db.execute("""CREATE TRIGGER fail_mark BEFORE UPDATE OF is_arch ON deals
WHEN OLD.order_sys_id='two'
BEGIN SELECT RAISE(ABORT, 'test failure'); END""")
with self.assertLogs(level='WARNING'):
self.store.archiving()
self.assertEqual(self.store.state[self.code], saved)
self.assertTrue(all(d['is_arch'] == 0 for d in self.store.deals.values()))
with closing(self.store._connect()) as db, db:
db.execute('DROP TRIGGER fail_mark')
self.store.archiving()
self.assertEqual(self.store.state[self.code], saved)
self.assertTrue(all(d['is_arch'] == 1 for d in self.store.deals.values()))
if __name__ == '__main__':
unittest.main()

View File

@@ -0,0 +1,86 @@
import sqlite3
import tempfile
import unittest
from dataclasses import asdict
from pathlib import Path
from unittest.mock import patch
from libs.state import FLAG_BUY, State
from sdk import DealItem, PositionItem
class StateStorageTests(unittest.TestCase):
def setUp(self):
tmp = tempfile.TemporaryDirectory()
self.addCleanup(tmp.cleanup)
self.store = State(Path(tmp.name) / 'state.db')
def deal(self, identity='first'):
return DealItem(
stock_code='600000.SH', order_sys_id=identity,
remark=f'zt-base-{identity}|zt', offset_flag=FLAG_BUY,
volume=100, price=10, trade_amount=1000,
trade_date='20260912', trade_time='100000',
)
def test_stale_cache_duplicate_preserves_original_and_imports_new_trade(self):
writer = State(self.store.path)
first = self.deal()
writer.sync_deals([first])
first.price = 20
first.trade_amount = 2000
self.store.sync_deals([first, self.deal('second')])
self.assertEqual(self.store.deals['first']['price'], 10)
self.assertEqual(self.store.deals_sys_ids, {'first', 'second'})
self.assertEqual(State(self.store.path).deals, self.store.deals)
def test_load_failure_does_not_publish_partial_cache(self):
writer = State(self.store.path)
writer.sync_state([PositionItem(stock_code='600000.SH', volume=100)])
with patch.object(self.store, '_read_deals', side_effect=sqlite3.OperationalError('read failed')):
with self.assertRaises(sqlite3.OperationalError):
self.store.load()
self.assertEqual((self.store.state, self.store.deals, self.store.deals_sys_ids), ({}, {}, set()))
self.store.load()
self.assertEqual(self.store.state['600000.SH']['base_qty'], 100)
def test_cache_read_failure_rolls_back_archive_and_can_retry(self):
self.store.sync_deals([self.deal()])
with patch.object(self.store, '_read_deals', side_effect=sqlite3.OperationalError('read failed')):
with self.assertRaises(sqlite3.OperationalError):
self.store.archiving()
restarted = State(self.store.path)
self.assertEqual(restarted.state, {})
self.assertEqual(restarted.deals['first']['is_arch'], 0)
self.assertEqual(self.store.deals, restarted.deals)
self.store.archiving()
self.assertEqual(self.store.state['600000.SH']['base_qty'], 100)
self.assertEqual(self.store.deals['first']['is_arch'], 1)
def test_invalid_snapshot_preserves_existing_holdings(self):
self.store.sync_state([PositionItem(stock_code='600000.SH', volume=100)])
saved = self.store.state
with self.assertRaises(ValueError):
self.store.sync_state([PositionItem(stock_code='600001.SH', volume=100, open_price=float('inf'))])
self.assertEqual(self.store.state, saved)
self.assertEqual(State(self.store.path).state, saved)
def test_normalization_preserves_input_and_rejects_nonfinite_price(self):
deal = self.deal()
deal.trade_amount = 0
original = asdict(deal)
self.store.sync_deals([deal])
self.assertEqual(asdict(deal), original)
self.assertEqual(self.store.deals['first']['trade_amount'], 1000)
self.assertEqual(self.store.deals['first']['trade_date'], '2026-09-12')
for price in (float('inf'), float('-inf'), float('nan')):
with self.subTest(price=price):
invalid = self.deal('invalid')
invalid.price = price
with self.assertRaises(ValueError):
self.store.sync_deals([self.deal('second'), invalid])
self.assertEqual(State(self.store.path).deals_sys_ids, {'first'})
if __name__ == '__main__':
unittest.main()

View File

@@ -1,81 +1,131 @@
import sqlite3
import tempfile
import unittest
from unittest.mock import patch
from pathlib import Path
from unittest.mock import patch
from libs.state import State
from libs.state import FLAG_BUY, FLAG_SELL, State, UNATTRIBUTED_PREFIX
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')
self.state = State(Path(tmp.name) / 'state.db')
self.code = '600000.SH'
def position(self, qty):
return PositionItem(stock_code='600000.SH', volume=qty, open_price=10)
def position(self, qty, code=None):
return PositionItem(stock_code=code or self.code, 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 deal(self, identity, qty=100, flag=FLAG_BUY, code=None, remark=None):
return DealItem(stock_code=code or self.code, order_sys_id=identity,
remark=f'zt-base-{identity}|zt' if remark is None else remark,
offset_flag=flag, volume=qty, price=10, trade_amount=qty*10,
trade_date='20260912', 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)
def test_initial_snapshot_and_equal_size_increment_are_distinct(self):
old = self.deal('old')
self.state.sync_account([self.position(100)], [old], initialize=True)
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
self.assertEqual(self.state.deals['old']['is_arch'], 1)
self.state = State(self.state.path)
bought = self.deal('new', 100, flag=48)
new = self.deal('new', remark='zt-added-new|zt')
for _ in range(2):
sync_account_state(self.state, [self.position(200)], [historical, bought, unrelated])
row = self.state.state['600000.SH']
self.state.sync_account([self.position(200)], [old, new, new])
row = self.state.state[self.code]
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.blocked_codes, set())
def test_initial_mixed_trades_are_already_in_snapshot(self):
self.state.sync_account([self.position(150)],
[self.deal('buy', 200), self.deal('sell', 50, FLAG_SELL)],
initialize=True)
self.assertEqual(self.state.state[self.code]['base_qty'], 150)
self.assertTrue(all(d['is_arch'] == 1 for d in self.state.deals.values()))
def test_restart_full_sell_is_archived_before_reconciliation(self):
self.state.sync_account([self.position(100)], [], initialize=True)
self.state = State(self.state.path)
self.state.sync_account([], [self.deal('sell', flag=FLAG_SELL)])
self.assertEqual(self.state.state, {})
self.assertEqual(self.state.deals['sell']['is_arch'], 1)
self.assertEqual(self.state.blocked_codes, set())
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)
def test_empty_initialized_account_survives_restart(self):
self.state.sync_account([], [], initialize=True)
self.state = State(self.state.path)
# 空账户重复初始化仍为空,无需额外标记表。
self.state.sync_account([], [], initialize=True)
self.state.sync_account([self.position(100)], [self.deal('new')])
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
def test_failed_initialization_leaves_original_database_empty(self):
def test_initialization_cannot_overwrite_existing_holdings(self):
self.state.sync_account([self.position(100)], [], initialize=True)
with self.assertRaises(ValueError):
self.state.sync_account([], [], initialize=True)
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
def test_initialization_failure_is_atomic(self):
invalid = self.position(100)
invalid.open_price = float('inf')
with self.assertRaises(ValueError):
sync_account_state(self.state, [invalid], [self.deal('old', 100)], initialize=True)
self.state.sync_account([invalid], [self.deal('one')], initialize=True)
restarted = State(self.state.path)
self.assertEqual((restarted.state, restarted.deals), ({}, {}))
sync_account_state(restarted, [self.position(100)], [self.deal('old', 100)], initialize=True)
self.assertEqual(restarted.deals['old']['is_arch'], 1)
with patch.object(self.state, '_read_deals', side_effect=sqlite3.OperationalError('read failed')):
with self.assertRaises(sqlite3.OperationalError):
self.state.sync_account([self.position(100)], [self.deal('one')], initialize=True)
self.assertEqual(State(self.state.path).deals, {})
self.assertEqual(self.state.state, {})
def test_initialization_cannot_overwrite_existing_database(self):
sync_account_state(self.state, [self.position(100)], [], initialize=True)
with self.assertRaises(ValueError):
sync_account_state(self.state, [], [], initialize=True)
self.assertEqual(State(self.state.path).state['600000.SH']['base_qty'], 100)
def test_lagging_snapshot_never_deletes_or_recreates_inventory(self):
self.state.sync_account([self.position(100)], [], initialize=True)
self.state.sync_account([], [])
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
self.assertEqual(self.state.blocked_codes, {self.code})
self.state.sync_account([self.position(100)], [])
self.assertEqual(self.state.blocked_codes, set())
sell = self.deal('sell', flag=FLAG_SELL)
self.state.sync_account([self.position(100)], [sell])
self.assertNotIn(self.code, self.state.state)
self.assertEqual(self.state.blocked_codes, {self.code})
self.state.sync_account([], [sell])
self.assertEqual(self.state.blocked_codes, set())
def test_no_new_deals_still_retries_failed_archiving(self):
sync_account_state(self.state, [self.position(100)], [], initialize=True)
sell = self.deal('sell', 100, flag=24)
with patch.object(self.state, 'archiving'):
with self.assertRaises(ValueError):
sync_account_state(self.state, [], [sell])
sync_account_state(self.state, [], [sell])
def test_blank_remark_persists_and_isolates_only_affected_stock(self):
self.state.sync_account([self.position(100)], [], initialize=True)
manual = self.deal('manual', remark=' |')
good = self.deal('good', code='600001.SH')
positions = [self.position(200), self.position(100, '600001.SH')]
for _ in range(2):
self.state.sync_account(positions, [manual, good])
self.assertEqual(self.state.blocked_codes, {self.code})
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
self.assertEqual(self.state.state['600001.SH']['base_qty'], 100)
self.assertEqual(self.state.deals['manual']['is_arch'], 0)
self.assertTrue(self.state.deals['manual']['order_local_id'].startswith(UNATTRIBUTED_PREFIX))
self.assertEqual(self.state.deals['manual']['remark'], ' |')
self.state = State(self.state.path)
def test_legacy_database_is_restored_without_reinitialization(self):
self.state.sync_state([self.position(100)])
self.state.sync_account([], [self.deal('sell', flag=FLAG_SELL)])
self.assertEqual(self.state.state, {})
self.assertEqual(self.state.deals['sell']['is_arch'], 1)
def test_failed_archive_blocks_only_stock_and_retries(self):
self.state.sync_account([self.position(100)], [], initialize=True)
sell = self.deal('sell', qty=200, flag=FLAG_SELL)
self.state.sync_account([], [sell])
self.assertEqual(self.state.state[self.code]['base_qty'], 100)
self.assertEqual(self.state.blocked_codes, {self.code})
buy = self.deal('buy', remark='zt-added-buy|zt')
buy.trade_time = '095900'
self.state.sync_account([], [sell, buy])
self.assertEqual(self.state.state, {})
self.assertEqual(self.state.blocked_codes, set())
if __name__ == '__main__':
unittest.main()

View File

@@ -1,151 +1,149 @@
import tempfile
import unittest
from concurrent.futures import ThreadPoolExecutor
from contextlib import closing
from datetime import datetime
from pathlib import Path
from types import SimpleNamespace
from types import SimpleNamespace as NS
from unittest.mock import Mock, patch
from config import AccountConfig
from libs.grid_take_profit import GridState
from libs.order import OrderBook
from libs.state import State
from libs.state import FLAG_BUY, 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
from strategy.zt.positions import manage_positions
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 = NS(account_cfg=NS(account_id='test', strategy='zt', buy_value=1000,
excluded_codes=[], enable_loss_add_position=False,
min_cash_ratio=0.1),
orders=Mock(), client=Mock(), profit_tracker=Mock(), add_watch=Mock())
self.run.orders.busy.return_value = False
self.run.orders.new_order_id.side_effect = lambda prefix, kind: f'{prefix}-{kind}-order'
self.run.orders.place.return_value = True
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, added=0, usable=500, road=0, cost=10, added_cost=10, price=11):
position = PositionItem(stock_code=self.code, volume=1000, can_use_volume=usable,
on_road_volume=road, open_price=cost)
state = NS(blocked_codes=set(), get_by_code=lambda code: dict(
base_qty=500, added_qty=added, added_price=added_cost))
manage_positions(self.run, {self.code: Tick(last_price=price)}, [position], True, 1500, state)
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_added_position_is_capped_by_sellable_inventory(self):
for added, usable, expected in [(500, 100, 100), (100, 500, 100), (0, 500, 500)]:
with self.subTest(added=added, usable=usable):
self.run.orders.place.reset_mock()
self.manage(added=added, usable=usable)
self.assertEqual(self.run.orders.place.call_args.args[1].volume, expected)
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()
def test_zero_sellable_does_not_divide_by_default_added_cost(self):
with patch('strategy.zt.positions.log.exception') as error:
self.manage(usable=0, added_cost=0)
error.assert_not_called()
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()
def test_unavailable_shares_do_not_disable_loss_management(self):
self.run.account_cfg.enable_loss_add_position = True
self.manage(usable=0, cost=20, price=10, added_cost=0)
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()
def test_on_road_shares_do_not_disable_available_base(self):
self.manage(road=100)
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 500)
def test_real_order_id_is_recognized_by_state_sync(self):
orders = OrderBook('zt')
self.run.orders.new_order_id.side_effect = orders.new_order_id
with patch('strategy.zt.open.datetime') as clock:
clock.now.return_value = datetime(2026, 9, 9, 10)
open_signal(self.run, {self.code: Tick(last_price=10)},
[SimpleNamespace(code=self.code)], 2000)
request = self.run.orders.place.call_args.args[1]
self.assertTrue(request.order_id.startswith('zt-base-'))
deal = self.fill('base', 'b1', 100)
deal.remark = request.order_id + '|zt'
self.store.sync_state([])
boot.sync_account_state(self.store, [], [deal])
self.assertEqual(self.store.state[self.code]['base_qty'], 100)
def test_added_cost_is_used_even_if_base_cost_is_higher(self):
self.manage(added=100, cost=20, added_cost=10, price=11)
self.assertEqual(self.run.orders.place.call_args.args[1].volume, 100)
def test_invalid_selected_cost_never_trades(self):
for cost in [0, -1, float('nan'), float('inf')]:
with self.subTest(cost=cost), patch('strategy.zt.positions.log.exception') as error:
self.manage(added=100, added_cost=cost)
error.assert_not_called()
self.run.orders.place.assert_not_called()
def test_run_once_quarantines_manual_trade_but_manages_good_stock(self):
with tempfile.TemporaryDirectory() as tmp, ThreadPoolExecutor(max_workers=2) as executor:
store = State(Path(tmp) / 'state.db')
good = '600001.SH'
positions = [PositionItem(stock_code=c, volume=100, can_use_volume=100, open_price=10)
for c in [self.code, good]]
store.sync_account(positions, [], initialize=True)
manual = DealItem(stock_code=self.code, order_sys_id='manual', remark='',
offset_flag=FLAG_BUY, volume=100, price=10, trade_amount=1000)
self.run.executor = executor
self.run.client.deals.return_value = [manual]
self.run.client.portfolio.return_value = NS(
assets=Assets(10000, 10000), positions={p.stock_code: p for p in positions}, orders=[])
self.run.client.full_tick.return_value = {p.stock_code: Tick(last_price=11) for p in positions}
with patch.object(boot, 'datetime') as clock, patch.object(boot, 'market_allow_open', return_value=True):
clock.now.return_value = datetime(2026, 9, 11, 10)
boot.RunOnce(self.run, store, [])
self.assertEqual(store.blocked_codes, {self.code})
self.run.orders.refresh.assert_called_once()
self.assertEqual(self.run.orders.place.call_count, 1)
self.assertEqual(self.run.orders.place.call_args.args[1].code, good)
def test_run_once_does_not_reopen_quarantined_sold_out_code(self):
with tempfile.TemporaryDirectory() as tmp, ThreadPoolExecutor(max_workers=2) as executor:
store = State(Path(tmp) / 'state.db')
store.sync_account([], [], initialize=True)
self.run.executor = executor
self.run.client.deals.return_value = [DealItem(
stock_code=self.code, order_sys_id='manual', remark='', offset_flag=FLAG_BUY,
volume=100, price=10, trade_amount=1000)]
self.run.client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={}, orders=[])
self.run.client.full_tick.return_value = {}
with patch.object(boot, 'datetime') as clock, patch.object(boot, 'market_allow_open', return_value=True), \
patch.object(boot, 'open_signal') as opened:
clock.now.return_value = datetime(2026, 9, 11, 10)
boot.RunOnce(self.run, store, [NS(code=self.code)])
opened.assert_not_called()
def start(self, client, directory):
self.run.account_cfg.grid_step_pct = 1
global_cfg = NS(qmt_base_url='unused', qmt_token='', qmt_data_dir=directory)
self.run.account_cfg.signal_allow = []
with patch.object(boot, 'Client', return_value=client), \
patch.object(boot.config, 'global_config', global_cfg), \
patch.object(boot.config, 'account_config', self.run.account_cfg), \
patch.object(boot, 'init_signals', return_value=[]), \
patch.object(boot, 'cache_portfolio'), patch.object(boot, 'Overview'), \
patch.object(boot.time, 'localtime', return_value=NS(tm_hour=15, tm_min=0, tm_sec=0)):
boot.StartZT()
def test_start_initializes_once_without_snapshot_retry_loop(self):
client = Mock()
client.deals.return_value = []
client.portfolio.return_value = NS(assets=Assets(10000, 10000),
positions={self.code: PositionItem(stock_code=self.code, volume=100, open_price=10)}, orders=[])
with tempfile.TemporaryDirectory() as tmp:
self.start(client, tmp)
store = State(Path(tmp) / 'zt_test_state.db')
self.assertEqual(store.state[self.code]['base_qty'], 100)
with closing(store._connect()) as db:
self.assertIsNone(db.execute("SELECT 1 FROM sqlite_master WHERE name='state_meta'").fetchone())
self.assertEqual(client.portfolio.call_count, 1)
self.assertEqual(client.deals.call_count, 2)
client.reset_mock()
self.start(client, tmp)
self.assertEqual(client.portfolio.call_count, 1)
self.assertEqual(client.deals.call_count, 1)
def test_start_rejects_changed_deals_without_writing_baseline(self):
client = Mock()
client.deals.side_effect = [[], [DealItem(order_sys_id='new')]]
client.portfolio.return_value = NS(assets=Assets(10000, 10000), positions={}, orders=[])
with tempfile.TemporaryDirectory() as tmp:
with self.assertRaises(RuntimeError):
self.start(client, tmp)
store = State(Path(tmp) / 'zt_test_state.db')
self.assertEqual((store.state, store.deals), ({}, {}))
client.close.assert_called_once()
if __name__ == '__main__':