diff --git a/py-client/libs/state.py b/py-client/libs/state.py index d01898d..2cd66ec 100644 --- a/py-client/libs/state.py +++ b/py-client/libs/state.py @@ -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 - def get_by_code(self,code: str) -> dict: - s = self.state.get(code,{}) - return s + @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: + """返回缓存中的状态;不存在时返回空字典。""" + 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 + today = datetime.now().date().isoformat() + values = [] + for deal in new_deals.values(): + amount = deal.trade_amount if deal.trade_amount > 0 else deal.price * deal.volume + 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:]}' + 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: - 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 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() - + 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), - ) - self.load_state() - + 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], + ) + db.executemany( + 'INSERT INTO state (stock_code, base_qty, base_price, base_created_at) ' + 'VALUES (?, ?, ?, ?)', values, + ) + state = self._read_state(db) + self.state = state + + @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)) + 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'] + 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() + 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: + db.execute('DELETE FROM state WHERE stock_code = ?', (code,)) + else: + # 保留原记录主键及策略状态,包括同批清仓后重新建仓。 + 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', (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: + db.execute('RELEASE archive_stock') + 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: - 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 底仓委托号约定识别,无需调用方传入规则。 - bucket = 'base' if deal['order_local_id'].startswith('zt-base-') else 'added' - total = state[f'{bucket}_qty'] + qty - # 新均价 =(原数量 × 原均价 + 本次成交金额)÷ 买入后总数量。 - 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: - 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'] + 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, - ) - # 持仓处理成功后才标记成交,防止下次重复加仓或重复扣减。 - db.execute( - 'UPDATE deals SET is_arch = 1 WHERE stock_code = ? ' - 'AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49)', (code,) - ) - 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() + 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))) diff --git a/py-client/strategy/zt/boot.py b/py-client/strategy/zt/boot.py index 482f757..78bd36f 100644 --- a/py-client/strategy/zt/boot.py +++ b/py-client/strategy/zt/boot.py @@ -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,38 +29,35 @@ 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 ) - # 获取本策略的信号开仓数据 + # 获取本策略的信号开仓数据 signals = init_signals( 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) diff --git a/py-client/strategy/zt/positions.py b/py-client/strategy/zt/positions.py index b1f952e..02c0f11 100644 --- a/py-client/strategy/zt/positions.py +++ b/py-client/strategy/zt/positions.py @@ -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) diff --git a/py-client/tests/test_deal_model.py b/py-client/tests/test_deal_model.py index be5e690..3e3bd23 100644 --- a/py-client/tests/test_deal_model.py +++ b/py-client/tests/test_deal_model.py @@ -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'}) diff --git a/py-client/tests/test_state_archiving.py b/py-client/tests/test_state_archiving.py index 17394e5..91919a8 100644 --- a/py-client/tests/test_state_archiving.py +++ b/py-client/tests/test_state_archiving.py @@ -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') diff --git a/py-client/tests/test_state_snapshot.py b/py-client/tests/test_state_snapshot.py new file mode 100644 index 0000000..ed48a2c --- /dev/null +++ b/py-client/tests/test_state_snapshot.py @@ -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() diff --git a/py-client/tests/test_state_storage.py b/py-client/tests/test_state_storage.py new file mode 100644 index 0000000..0ee9005 --- /dev/null +++ b/py-client/tests/test_state_storage.py @@ -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() diff --git a/py-client/tests/test_zt_state.py b/py-client/tests/test_zt_state.py index 412e377..defa667 100644 --- a/py-client/tests/test_zt_state.py +++ b/py-client/tests/test_zt_state.py @@ -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() diff --git a/py-client/tests/test_zt_trading.py b/py-client/tests/test_zt_trading.py index 5b03176..65b80d3 100644 --- a/py-client/tests/test_zt_trading.py +++ b/py-client/tests/test_zt_trading.py @@ -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__':