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

@@ -4,12 +4,16 @@ import math
import logging as log
import sqlite3
from contextlib import closing
from dataclasses import asdict, dataclass
from dataclasses import asdict, dataclass, fields
from datetime import datetime
from pathlib import Path
from sdk import DealItem, PositionItem
FLAG_BUY = 48
FLAG_SELL = 49
UNATTRIBUTED_PREFIX = '__unattributed__:'
SCHEMA = """
-- 策略状态base_ 表示底仓added_ 表示补仓。
CREATE TABLE IF NOT EXISTS state (
@@ -69,14 +73,29 @@ class StateItem:
added_created_at: str = '' # 补仓创建时间
_STATE_COLUMNS = tuple(field.name for field in fields(StateItem))
_UPSERT_STATE = (
f"INSERT INTO state ({', '.join(_STATE_COLUMNS)}) "
f"VALUES ({', '.join(':' + key for key in _STATE_COLUMNS)}) "
"ON CONFLICT(stock_code) DO UPDATE SET "
+ ', '.join(f'{key} = excluded.{key}' for key in _STATE_COLUMNS if key != 'stock_code')
)
class State:
"""保存策略状态和只追加的成交记录,仅创建新表,不迁移旧数据。"""
"""单写入者使用的 SQLite 存储;公开缓存仅在事务成功后替换。
sync_state 同步持仓基准sync_deals 保存成交archiving 记入增量。
同证券未归档成交全部为买入且数量等于当前总持仓时,视为已计入
快照,只标记归档;其他成交作为增量处理。本类不做表结构迁移。
"""
def __init__(self, path: str | Path) -> None:
self.path = Path(path)
self.state: dict[str, dict] = {}
self.deals: dict[str, dict] = {}
self.deals_sys_ids: set[str] = set()
self.blocked_codes: set[str] = set()
self.path.parent.mkdir(parents=True, exist_ok=True)
with closing(self._connect()) as db:
db.executescript(SCHEMA)
@@ -87,176 +106,242 @@ class State:
db.row_factory = sqlite3.Row
return db
@staticmethod
def _read_state(db: sqlite3.Connection) -> dict[str, dict]:
return {row['stock_code']: dict(row) for row in db.execute('SELECT * FROM state')}
@staticmethod
def _read_deals(db: sqlite3.Connection) -> dict[str, dict]:
return {
row['order_sys_id']: dict(row)
for row in db.execute('SELECT * FROM deals ORDER BY id')
}
def get_by_code(self, code: str) -> dict:
s = self.state.get(code,{})
return s
"""返回缓存中的状态;不存在时返回空字典。"""
return self.state.get(code, {})
def load(self) -> None:
"""缓存状态表和成交记录"""
self.load_state()
self.load_deals()
def load_state(self) -> None:
"""缓存状态表。"""
"""在同一个读事务内加载两张表,全部成功后再发布缓存。"""
with closing(self._connect()) as db, db:
db.execute('BEGIN')
state = {row['stock_code']: dict(row) for row in db.execute('SELECT * FROM state')}
state = self._read_state(db)
deals = self._read_deals(db)
self.state, self.deals, self.deals_sys_ids = state, deals, set(deals)
def load_state(self) -> None:
"""只刷新状态缓存。"""
with closing(self._connect()) as db, db:
db.execute('BEGIN')
state = self._read_state(db)
self.state = state
def load_deals(self) -> None:
"""缓存成交记录"""
"""只刷新成交缓存及其去重编号集合"""
with closing(self._connect()) as db, db:
db.execute('BEGIN')
deals = {row['order_sys_id']: dict(row) for row in db.execute('SELECT * FROM deals ORDER BY id')}
self.deals = deals
self.deals_sys_ids = set(deals)
deals = self._read_deals(db)
self.deals, self.deals_sys_ids = deals, set(deals)
def sync_deals(self, deals: list[DealItem]) -> None:
"""按成交编号去重;初始化底仓时,已包含在快照内的成交可直接标记归档。"""
new_deals = {}
@staticmethod
def _insert_deals(db: sqlite3.Connection, deals: list[DealItem]) -> None:
"""共享事务内保存成交;空本地编号使用明确的待核对标记。"""
existing = {row[0] for row in db.execute('SELECT order_sys_id FROM deals')}
new_deals: dict[str, DealItem] = {}
for deal in deals:
if deal.order_sys_id not in self.deals_sys_ids and deal.order_sys_id not in new_deals:
new_deals[deal.order_sys_id] = deal
if deal.order_sys_id not in existing:
new_deals.setdefault(deal.order_sys_id, deal)
if not new_deals:
return
with closing(self._connect()) as db, db:
today = datetime.now().date().isoformat()
values = []
for deal in new_deals.values():
order_id = deal.get_local_order_id
if not order_id:
raise ValueError('Local order ID is required')
amount = deal.trade_amount if deal.trade_amount > 0 else deal.price * deal.volume
if not math.isfinite(amount) or amount <= 0:
raise ValueError('Trade amount must be positive and finite')
date = deal.trade_date or datetime.now().date().isoformat()
if not math.isfinite(deal.price) or not math.isfinite(amount):
raise ValueError('Trade price and amount must be finite')
date = deal.trade_date or today
if len(date) == 8 and date.isdigit():
date = f'{date[:4]}-{date[4:6]}-{date[6:]}'
# 直接读取模型字段,金额和日期的补全不修改传入模型。
db.execute(
'INSERT INTO deals (stock_code, order_sys_id, order_local_id, ref, '
'order_ref, direction, offset_flag, price, volume, trade_amount, '
'trade_date, trade_time, remark, close_profit, is_arch) '
'VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)',
(deal.stock_code, deal.order_sys_id, order_id, deal.ref,
deal.order_ref, deal.direction, deal.offset_flag, deal.price, deal.volume,
amount, date, deal.trade_time, deal.remark, deal.close_profit, 0),
)
self.load_deals()
values.append((
deal.stock_code, deal.order_sys_id,
deal.get_local_order_id.strip() or f'{UNATTRIBUTED_PREFIX}{deal.order_sys_id}',
deal.ref, deal.order_ref, deal.direction, deal.offset_flag,
deal.price, deal.volume, amount, date, deal.trade_time,
deal.remark, deal.close_profit,
))
# 唯一键负责最终去重,避免旧缓存造成重复插入失败。
db.executemany("""
INSERT INTO deals (
stock_code, order_sys_id, order_local_id, ref, order_ref,
direction, offset_flag, price, volume, trade_amount,
trade_date, trade_time, remark, close_profit, is_arch
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 0)
ON CONFLICT(order_sys_id) DO NOTHING
""", values)
def sync_deals(self, deals: list[DealItem]) -> None:
"""按成交编号只追加;同批保留首笔,任一无效成交回滚整批。"""
if all(deal.order_sys_id in self.deals_sys_ids for deal in deals):
return
with closing(self._connect()) as db, db:
db.execute('BEGIN IMMEDIATE')
self._insert_deals(db, deals)
cached_deals = self._read_deals(db)
self.deals, self.deals_sys_ids = cached_deals, set(cached_deals)
def sync_state(self, positions: list[PositionItem]) -> None:
"""同步完整持仓:无状态则插入底仓,已有则保留,清仓则删除
"""同步完整持仓:新增底仓、保留已有状态、删除已清仓证券
此接口不会标记成交。已有库存对应的卖出须先归档,再传入
清仓快照,以免删除归档所需的库存。
"""
# 传入完整账户持仓;同步时间作为新增底仓的创建时间。
created_at = datetime.now().isoformat(timespec='seconds')
holdings = {item.stock_code: item for item in positions if item.volume > 0}
with closing(self._connect()) as db, db:
db.execute('BEGIN')
db.execute('BEGIN IMMEDIATE')
existing = {row['stock_code'] for row in db.execute('SELECT stock_code FROM state')}
db.executemany(
'DELETE FROM state WHERE stock_code = ?',
[(code,) for code in existing if code not in holdings],
)
values = []
for code, item in holdings.items():
if code in existing:
continue
if not math.isfinite(item.open_price):
raise ValueError('Base price must be finite')
# 只插入底仓字段,补仓字段使用数据库默认值。
db.execute(
'INSERT INTO state '
'(stock_code, status, base_order_local_id, base_qty, base_price, base_created_at) '
"VALUES (?, '', '', ?, ?, ?)",
(code, item.volume, item.open_price, created_at),
values.append((code, item.volume, item.open_price, created_at))
db.executemany(
'DELETE FROM state WHERE stock_code = ?',
[(code,) for code in existing if code not in holdings],
)
self.load_state()
db.executemany(
'INSERT INTO state (stock_code, base_qty, base_price, base_created_at) '
'VALUES (?, ?, ?, ?)', values,
)
state = self._read_state(db)
self.state = state
def archiving(self) -> None:
"""将未归档成交累加到当前底仓和加仓;失败证券保留记录供重试。
归档是把成交反映到持仓状态中,再标记为已处理,不会删除成交记录。
本地委托号以 zt-base- 开头的买入计入底仓,其他买入计入补仓。
不返回结果;失败原因记录到日志,对应成交保留未归档标记供重试。
"""
with closing(self._connect()) as db, db:
# 提前取得数据库写入锁,让持仓更新和成交标记在同一事务内完成。
db.execute('BEGIN IMMEDIATE')
# 只找尚未处理的成交23、48 是买入24、49 是卖出。
codes = db.execute(
'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0'
).fetchall()
for entry in codes:
code = entry['stock_code']
# 每只股票设一个回滚点;这只处理失败时,不撤销其他股票的结果。
db.execute('SAVEPOINT archive_stock')
try:
@staticmethod
def _archive_stock(db: sqlite3.Connection, code: str, snapshot_dedup: bool = True) -> None:
"""在调用方的保存点内完成一只证券的记账及成交标记。"""
current = db.execute('SELECT * FROM state WHERE stock_code = ?', (code,)).fetchone()
# 已有持仓就接着计算;没有记录则从底仓、补仓均为零开始。
state = dict(current) if current else asdict(StateItem(stock_code=code))
# 按成交日期、时间、记录编号依次处理,保证先买后卖等顺序正确。
deals = db.execute(
'SELECT * FROM deals WHERE stock_code = ? AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49) '
"ORDER BY trade_date, REPLACE(trade_time, ':', ''), id", (code,)
).fetchall()
for deal in deals:
qty = deal['volume']
if deal['offset_flag'] in (23, 48):
# 按已有的 ZT 底仓委托号约定识别,无需调用方传入规则。
pending = db.execute("""
SELECT * FROM deals WHERE stock_code = ? AND is_arch = 0
ORDER BY trade_date, REPLACE(trade_time, ':', ''), id
""", (code,)).fetchall()
if any(deal['order_local_id'].startswith(UNATTRIBUTED_PREFIX) for deal in pending):
raise ValueError('Unattributed trade requires reconciliation')
total_qty = state['base_qty'] + state['added_qty']
if (
snapshot_dedup and total_qty > 0
and all(deal['offset_flag'] == FLAG_BUY for deal in pending)
and sum(deal['volume'] for deal in pending) == total_qty
):
# 按快照去重规则,仅标记这批成交,原仓位数量和成本均保留。
# 卖出和买卖混合批次不能用此规则,否则会漏掉真实减仓。
db.executemany(
'UPDATE deals SET is_arch = 1 WHERE id = ? AND is_arch = 0',
[(deal['id'],) for deal in pending],
)
return
for deal in pending:
if deal['offset_flag'] == FLAG_BUY:
bucket = 'base' if deal['order_local_id'].startswith('zt-base-') else 'added'
total = state[f'{bucket}_qty'] + qty
# 新均价 =(原数量 × 原均价 + 本次成交金额)÷ 买入后总数量。
total = state[f'{bucket}_qty'] + deal['volume']
state[f'{bucket}_price'] = (
state[f'{bucket}_qty'] * state[f'{bucket}_price'] + deal['trade_amount']
) / total
state[f'{bucket}_qty'] = total
# 保存这一类仓位最近一次买入的委托号和成交时间。
state[f'{bucket}_order_local_id'] = deal['order_local_id']
state[f'{bucket}_created_at'] = f"{deal['trade_date']} {deal['trade_time']}".strip()
else:
elif deal['offset_flag'] == FLAG_SELL:
qty = deal['volume']
total = state['base_qty'] + state['added_qty']
# 卖出不能超过本地记录的总持仓;不一致时留待排查后重试。
if qty > total:
raise ValueError(f'Sell volume {qty} exceeds recorded holdings {total}')
# 卖出优先扣补仓;补仓未卖完时只减数量,保留原均价。
if qty < state['added_qty']:
state['added_qty'] -= qty
else:
# 补仓全部卖完后,剩余卖出量扣底仓,并清空补仓信息。
state['base_qty'] = total - qty
state['added_qty'] = 0
state['added_price'] = 0.0
state['added_order_local_id'] = state['added_created_at'] = ''
if state['base_qty'] == 0:
state.update(asdict(StateItem(stock_code=code)))
else:
raise ValueError(f"Unsupported trade flag: {deal['offset_flag']}")
if state['base_qty'] + state['added_qty'] == 0:
# 中途清仓先重置,后续若又买入,就从零重新累计。
state = asdict(StateItem(stock_code=code))
if state['base_qty'] + state['added_qty'] == 0:
# 全部成交处理完仍无持仓,则删除该股票的策略状态。
db.execute('DELETE FROM state WHERE stock_code = ?', (code,))
else:
# 更新持仓,保留策略状态及已有记录主键
state['status'] = current['status'] if current else state['status']
state.pop('id', None)
columns = tuple(state)
# 没有该股票就新增,已有则更新仓位字段,不替换原记录主键。
db.execute(
f"INSERT INTO state ({', '.join(columns)}) "
f"VALUES ({', '.join(':' + key for key in columns)}) "
'ON CONFLICT(stock_code) DO UPDATE SET '
+ ', '.join(f'{key} = excluded.{key}' for key in columns if key != 'stock_code'),
state,
)
# 持仓处理成功后才标记成交,防止下次重复加仓或重复扣减。
# 保留原记录主键及策略状态,包括同批清仓后重新建仓
if current:
state['status'] = current['status']
db.execute(_UPSERT_STATE, state)
db.execute(
'UPDATE deals SET is_arch = 1 WHERE stock_code = ? '
'AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49)', (code,)
'AND is_arch = 0', (code,),
)
def _archive_pending(self, db: sqlite3.Connection, snapshot_dedup: bool = True) -> None:
"""保留现有逐证券保存点ZT 增量路径禁用数量相等推断。"""
codes = db.execute(
'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0'
).fetchall()
for row in codes:
code = row['stock_code']
db.execute('SAVEPOINT archive_stock')
try:
self._archive_stock(db, code, snapshot_dedup)
except (ValueError, sqlite3.IntegrityError) as exc:
# 撤销这只股票的全部归档修改,成交仍保持未归档,供下次重试。
db.execute('ROLLBACK TO archive_stock')
log.warning('[归档] %s 失败,保留未归档成交:%s', code, exc)
finally:
# 释放当前股票的回滚点;整个事务在退出外层 with 时提交。
db.execute('RELEASE archive_stock')
# 数据库提交完成后刷新内存缓存,让策略读到最新持仓和归档标记。
self.load()
def archiving(self) -> None:
"""归档未处理成交;单证券失败回滚并保留重试,不影响其他证券。"""
with closing(self._connect()) as db, db:
db.execute('BEGIN IMMEDIATE')
self._archive_pending(db)
state = self._read_state(db)
deals = self._read_deals(db)
self.state, self.deals, self.deals_sys_ids = state, deals, set(deals)
def sync_account(self, positions: list[PositionItem], deals: list[DealItem], *, initialize: bool = False) -> None:
"""ZT 原子同步。初始化使用已核对快照;恢复只应用增量、不覆盖库存。
空库存且无成交历史时允许初始化。快照缺失或
未归档成交只隔离对应证券,后续成交到齐并核对一致后自动恢复。
"""
holdings = {p.stock_code: p for p in positions if p.volume > 0}
with closing(self._connect()) as db, db:
db.execute('BEGIN IMMEDIATE')
if initialize:
if (db.execute('SELECT 1 FROM state LIMIT 1').fetchone()
or db.execute('SELECT 1 FROM deals LIMIT 1').fetchone()):
raise ValueError('ZT state already initialized; cannot overwrite existing data')
created_at = datetime.now().isoformat(timespec='seconds')
for code, position in holdings.items():
if not code or not math.isfinite(position.open_price) or position.open_price <= 0:
raise ValueError('Initial position code and cost must be valid')
db.execute(_UPSERT_STATE, asdict(StateItem(
stock_code=code, base_qty=position.volume,
base_price=position.open_price, base_created_at=created_at,
)))
self._insert_deals(db, deals)
# 这些成交已包含在已核对的基准中,不再增减快照库存。
db.execute('UPDATE deals SET is_arch = 1')
else:
self._insert_deals(db, deals)
self._archive_pending(db, snapshot_dedup=False)
state = self._read_state(db)
cached_deals = self._read_deals(db)
blocked = {d['stock_code'] for d in cached_deals.values() if d['is_arch'] != 1}
for code in set(state) | set(holdings):
row = state.get(code, {})
recorded = row.get('base_qty', 0) + row.get('added_qty', 0)
actual = holdings[code].volume if code in holdings else 0
if recorded != actual:
blocked.add(code)
self.state, self.deals, self.deals_sys_ids = state, cached_deals, set(cached_deals)
self.blocked_codes = blocked
if blocked:
log.warning('[ZT 同步] 以下证券状态待核对,暂停交易:%s', ', '.join(sorted(blocked)))

View File

@@ -2,14 +2,11 @@
import logging as log
import time
from contextlib import closing
from datetime import datetime
from pathlib import Path
from tempfile import TemporaryDirectory
from concurrent.futures import Future, ThreadPoolExecutor
import config
from libs.calc import trading_time
from libs.collector import collector_push
from libs.grid_take_profit import GridTrailingTracker
from libs.market import market_allow_open
from libs.order import OrderBook
@@ -18,9 +15,9 @@ from libs.runtime import Runtime
from libs.signal import SignalItem, init_signals
from libs.state import State
from libs.watch import DipWatch
from sdk import Client, DealItem, PositionItem
from sdk import Client
from .open import open_signal
from .positions import manage_positions, t_rounds
from .positions import manage_positions
from libs.snapshot import cache_portfolio
@@ -32,15 +29,13 @@ def StartZT() -> None:
)
executor = None
try:
portfolio = client.portfolio()
assets = portfolio.assets
positions = list(portfolio.positions.values())
state = State(Path(config.global_config.qmt_data_dir) / f'zt_{config.account_config.account_id}_state.db')
executor = ThreadPoolExecutor(max_workers=2, thread_name_prefix="zt")
run = Runtime(
client=client, global_cfg=config.global_config, account_cfg=config.account_config,
orders=OrderBook('zt'), open_watch=DipWatch(), add_watch=DipWatch(),
profit_tracker=GridTrailingTracker(config.account_config.grid_step_pct),
executor=executor
)
# 获取本策略的信号开仓数据
@@ -48,22 +43,21 @@ def StartZT() -> None:
config.global_config,
config.account_config.signal_allow,
)
log.info(
"[启动] Trend策略已启动账户=%s,信号=%d,持仓=%d",
config.account_config.account_id,
len(signals),
len(positions),
)
initialize = not state.state and not state.deals
deals = client.deals()
portfolio = client.portfolio()
if initialize and {d.order_sys_id: d for d in deals} != {
d.order_sys_id: d for d in client.deals()
}:
raise RuntimeError('ZT 初始化期间成交发生变化,请重新启动')
assets = portfolio.assets
positions = list(portfolio.positions.values())
log.info('[启动] ZT策略已启动账户=%s,信号=%d,持仓=%d',
config.account_config.account_id, len(signals), len(positions))
cache_portfolio(config.account_config.account_id, assets, positions, deals)
state.load()
state.sync_deals(deals)
state.sync_state(positions)
state.archiving()
state.sync_account(positions, deals, initialize=initialize)
run.orders.refresh(client, portfolio.orders)
Overview(assets, positions, config.account_config)
executor = ThreadPoolExecutor(max_workers=2, thread_name_prefix="zt")
DEFAULT_TICK_INTERVAL = 30
while True:
@@ -118,8 +112,7 @@ def RunOnce(run: Runtime, state: State, signals: list[SignalItem]) -> None:
positions = list(portfolio.positions.values())
position_codes = list(portfolio.positions)
state.sync_deals(deals)
state.archiving()
state.sync_account(positions, deals)
run.orders.refresh(run.client, portfolio.orders)
except Exception:
log.exception("[Portfolio] 刷新账户快照失败")
@@ -143,7 +136,7 @@ def RunOnce(run: Runtime, state: State, signals: list[SignalItem]) -> None:
allow_open: list[SignalItem] = []
allow_codes: list[str] = []
for signal in signals:
if signal.code not in portfolio.positions:
if signal.code not in portfolio.positions and signal.code not in state.blocked_codes:
allow_open.append(signal)
allow_codes.append(signal.code)

View File

@@ -1,6 +1,7 @@
"""趋势策略持仓止盈与分级补仓。"""
from dataclasses import dataclass
import math
from libs.calc import calc_buy_volume, calculate_min_profit_rate
from libs.grid_take_profit import GridState
@@ -45,9 +46,9 @@ def manage_positions(
continue
if (
not code
or position.open_price <= 0
or position.volume <= 0
or tick is None
or not math.isfinite(tick.last_price)
or tick.last_price <= 0
):
log.warning(
@@ -57,16 +58,24 @@ def manage_positions(
)
continue
if code in state.blocked_codes:
log.warning('[Position] %s 状态待核对,暂停该证券交易', code)
continue
posState = state.get_by_code(position.stock_code)
if not posState:
continue
volume = position.can_use_volume
target_qty = posState.get('base_qty', 0)
cost_price = position.open_price
if posState.get('added_qty',0) >=100:
volume = posState.get('added_qty',0)
if posState.get('added_qty', 0) > 0:
target_qty = posState['added_qty']
cost_price = posState.get('added_price',0)
# 在途股份不影响已有可卖库存;补仓、底仓均受柜台可卖上限约束。
volume = max(0, min(target_qty, position.can_use_volume, position.volume))
if not math.isfinite(cost_price) or cost_price <= 0:
log.warning('[Position] %s 成本无效,暂停该证券交易', code)
continue
pnl_rate = round(
(tick.last_price - cost_price) / cost_price * 100,
@@ -86,7 +95,8 @@ def manage_positions(
if runtime.account_cfg.enable_loss_add_position and market_ok:
loss_decision = handle_loss(
runtime=runtime,
position=position,
stock_code=position.stock_code,
volume=volume,
tick=tick,
pnl_rate=pnl_rate,
available=available,
@@ -99,7 +109,7 @@ def manage_positions(
strTag = "-"
if pnl_rate >= minimum_profit:
strTag = ""
elif pnl_rate< LOSS_TIERS[0]:
elif pnl_rate< LOSS_TIERS:
strTag = ""
if strTag != "-":
@@ -148,9 +158,8 @@ def handle_profit(
if runtime.orders.busy(stock_code, "SELL"):
return TradeDecision(False, "卖出委托处理中")
volume = volume % 100
if volume <= 0:
return TradeDecision(False, "无可用整手持仓")
return TradeDecision(False, "无可用持仓")
order_id = runtime.orders.new_order_id("zt","SELL")
request = PlaceOrderRequest(
op=OP_SELL,
@@ -199,7 +208,7 @@ def handle_loss(
if not runtime.orders.place(runtime.client, request):
return TradeDecision(False, "补仓订单委托失败")
runtime.add_watch.forget(position.stock_code)
runtime.add_watch.forget(stock_code)
return TradeDecision(True, f"[补仓买入] {volume} 股,订单={order_id}", amount)

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_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)
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(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)
self.manage(added=added, usable=usable)
self.assertEqual(self.run.orders.place.call_args.args[1].volume, expected)
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)
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_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)
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)
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_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_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__':