fix zt&state.py
This commit is contained in:
@@ -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__':
|
||||
|
||||
Reference in New Issue
Block a user