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