diff --git a/docs/zt.md b/docs/zt.md index 6e82dd9..d331b01 100644 --- a/docs/zt.md +++ b/docs/zt.md @@ -12,6 +12,46 @@ zt_max_price: 200 - 仅 `dcm` 信号可建立底仓;建仓使用反弹确认,跳过价格高于 `zt_max_price` 的股票。 - 对 dcm 底仓,盈利网格出现回撤时卖出 `zt_sell_ratio` 对应的可用整手;不卖出超过记录底仓的数量。 -- 卖单全部成交后,价格较卖出价回落 `zt_buy_fall_pct`,并经反弹确认,买回同等数量。 +- 每笔卖出成交直接累计做 T 数量;活动委托结束后,价格较卖出均价回落 `zt_buy_fall_pct`,并经反弹确认,买回实际卖出数量。 - 每只股票每日只做一轮;14:50 后不再开新卖单,已卖未买的仓位强制按市价买回,避免隔夜净减仓。 - 仅支持标准 A 股的先卖后买,不把当日新买入股票作为可卖库存。 + +## SQLite 状态存储 + +`libs/orderbook.py` 使用标准库 `sqlite3`,数据库路径为 +`{qmt_data_dir}/zt_{account_id}_state.db`,每个账户/策略独立存储,由一个策略实例串行更新。 +启动时仅创建当前表结构和索引,不执行数据迁移或旧 JSON 导入。 + +`OrderBook` 初始化后调用 `load()`,填充 `positions`(按代码)、`deals`(按系统订单号) +和 `deals_sys_ids`(系统订单号集合)。`sync_deals(list[DealItem])` 根据该集合过滤已保存及 +同批重复成交,再以单个事务批量插入,提交成功后刷新缓存;不修改持仓。 +`sys_order_id` 对应 API 的 `m_strOrderSysID`,`local_order_id` 从 `m_strRemark` 的首段提取。 + +| 表 | 字段与用途 | 索引 | +| --- | --- | --- | +| `positions` | 自增 `id`、股票代码 `code`、底仓订单/数量/成本 `base_order_id/base_qty/base_cost`、补仓订单/次数/数量/成本 `added_order_id/added_num/added_qty/added_cost`、状态 `status` | `code` 唯一索引;`base_order_id`;`added_order_id` | +| `deals` | `id`、系统/本地订单号、证券信息、方向、API 状态、剩余/成交数量、委托日期时间、备注、委托价格及成交均价/金额,字段映射见下表 | `id` 主键;`sys_order_id` 唯一;`local_order_id`;`(code, insert_date)`;`(insert_date, insert_time)` | + +SDK 的 `DealItem` 与成交表业务字段一致: + +| API 字段 | SDK / SQLite 字段 | +| --- | --- | +| `m_strOrderSysID` | `sys_order_id` | +| `m_strInstrumentID`、`m_strExchangeID` | `instrument_id`、`exchange_id`,组合生成 `code` | +| `m_strInstrumentName` | `name` | +| `m_nOffsetFlag` | `offset_flag`,解析生成买卖方向 `side` | +| `m_nOrderStatus` | `status` | +| `m_nVolumeTotal`、`m_nVolumeTraded` | `remaining_volume`、`traded_volume` | +| `m_nOrderTime` | `order_time` | +| `m_strInsertDate`、`m_strInsertTime` | `insert_date`、`insert_time` | +| `m_strRemark` | `remark`,提取 `local_order_id` | +| `m_dPrice`、`m_dTradePrice`、`m_dTradeAmount` | `price`、`trade_price`、`trade_amount` | + +API 返回已成交数据,每个系统订单号仅入库一次;`status` 保存 API 原值,不增加确认流程。 +成交金额缺失时仅使用成交均价乘成交数量补足,不使用委托价格。 +买卖方向使用 `side`,ZT 通过本地委托号前缀区分底仓买入与做 T 买回。 +ZT 根据实际成交数量及金额更新持仓,日期取 API 提供的 `insert_date`,入库时规范为 `YYYY-MM-DD`。 +首次接管持仓仅写持仓表。活动委托和重复下单检查由 `libs/order.py` 的委托簿负责。 +持仓快照与新增成交在同一事务内保存,历史成交只追加;写入失败回滚数据库并恢复内存状态。 +持仓更新按 `code` 保留原有自增 ID。ZT 将轮次状态写入 `status`,卖出、买回数量及均价、 +轮次日期在重启时从 `deals` 重建;补仓字段预留,ZT 买回不计为补仓。 diff --git a/py-client/libs/orderbook.py b/py-client/libs/orderbook.py new file mode 100644 index 0000000..8a4219c --- /dev/null +++ b/py-client/libs/orderbook.py @@ -0,0 +1,170 @@ +"""SQLite 状态簿:每个账户/策略使用独立数据库,单个策略串行读写。""" + +from __future__ import annotations + +import math +import sqlite3 +from contextlib import closing +from dataclasses import asdict +from datetime import datetime +from pathlib import Path + +from sdk import DealItem + +SCHEMA = """ +CREATE TABLE IF NOT EXISTS positions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + code TEXT NOT NULL, + base_order_id TEXT NOT NULL DEFAULT '', + base_qty INTEGER NOT NULL DEFAULT 0 CHECK (base_qty >= 0), + base_cost REAL NOT NULL DEFAULT 0.0 CHECK (base_cost >= 0), + added_order_id TEXT NOT NULL DEFAULT '', + added_num INTEGER NOT NULL DEFAULT 0 CHECK (added_num >= 0), + added_qty INTEGER NOT NULL DEFAULT 0 CHECK (added_qty >= 0), + added_cost REAL NOT NULL DEFAULT 0.0 CHECK (added_cost >= 0), + status TEXT NOT NULL +); +CREATE UNIQUE INDEX IF NOT EXISTS idx_positions_code ON positions (code); +CREATE INDEX IF NOT EXISTS idx_positions_base_order_id ON positions (base_order_id); +CREATE INDEX IF NOT EXISTS idx_positions_added_order_id ON positions (added_order_id); + +CREATE TABLE IF NOT EXISTS deals ( + id INTEGER PRIMARY KEY, + sys_order_id TEXT NOT NULL UNIQUE CHECK (sys_order_id <> ''), + local_order_id TEXT NOT NULL, + code TEXT NOT NULL, + instrument_id TEXT NOT NULL, + exchange_id TEXT NOT NULL, + name TEXT NOT NULL, + offset_flag TEXT NOT NULL, + side TEXT NOT NULL CHECK (side IN ('BUY', 'SELL')), + status TEXT NOT NULL, + remaining_volume INTEGER NOT NULL CHECK (remaining_volume >= 0), + traded_volume INTEGER NOT NULL CHECK (traded_volume > 0), + order_time INTEGER NOT NULL, + insert_date TEXT NOT NULL, + insert_time TEXT NOT NULL, + remark TEXT NOT NULL, + price REAL NOT NULL, + trade_price REAL NOT NULL CHECK (trade_price >= 0), + trade_amount REAL NOT NULL CHECK (trade_amount > 0) +); +CREATE INDEX IF NOT EXISTS idx_deals_local_order_id ON deals (local_order_id); +CREATE INDEX IF NOT EXISTS idx_deals_code_date ON deals (code, insert_date); +CREATE INDEX IF NOT EXISTS idx_deals_date_time ON deals (insert_date, insert_time); +""" + + +class OrderBook: + """Persist position snapshots and append-only executions; one writer per database.""" + + def __init__(self, path: str | Path) -> None: + self.path = Path(path) + self.positions: dict[str, dict] = {} + self.deals: dict[str, dict] = {} + self.deals_sys_ids: set[str] = set() + self.path.parent.mkdir(parents=True, exist_ok=True) + with closing(self._connect()) as db: + db.executescript(SCHEMA) + self.load() + + def _connect(self) -> sqlite3.Connection: + db = sqlite3.connect(self.path, timeout=30) + db.row_factory = sqlite3.Row + return db + + def load(self) -> None: + with closing(self._connect()) as db, db: + db.execute('BEGIN') + positions = {row['code']: dict(row) for row in db.execute('SELECT * FROM positions')} + deals = {row['sys_order_id']: dict(row) for row in db.execute('SELECT * FROM deals ORDER BY id')} + self.positions = positions + self.deals = deals + self.deals_sys_ids = set(deals) + self._deal_count = len(deals) + + @staticmethod + def _insert_deals(db: sqlite3.Connection, deals: list[dict]) -> None: + db.executemany( + """INSERT INTO deals + (sys_order_id, local_order_id, code, instrument_id, exchange_id, name, + offset_flag, side, status, remaining_volume, traded_volume, order_time, + insert_date, insert_time, remark, price, trade_price, trade_amount) + VALUES (:sys_order_id, :local_order_id, :code, :instrument_id, :exchange_id, :name, + :offset_flag, :side, :status, :remaining_volume, :traded_volume, :order_time, + :insert_date, :insert_time, :remark, :price, :trade_price, :trade_amount)""", + deals, + ) + + @staticmethod + def deal_record(deal: DealItem) -> dict: + """API 字段转入库记录;委托价格不用于推算成交金额。""" + if not deal.sys_order_id or deal.traded_volume <= 0: + raise ValueError('System order ID and positive traded volume are required') + row = asdict(deal) + if any(isinstance(value, float) and not math.isfinite(value) for value in row.values()): + raise ValueError('Prices and amounts must be finite') + amount = deal.trade_amount if deal.trade_amount > 0 else deal.trade_price * deal.traded_volume + if not math.isfinite(amount) or amount <= 0: + raise ValueError('Execution amount must be positive and finite') + row['trade_amount'] = amount + date = deal.insert_date or datetime.now().date().isoformat() + if len(date) == 8 and date.isdigit(): + date = f'{date[:4]}-{date[4:6]}-{date[6:]}' + row['insert_date'] = date + return row + + def sync_deals(self, deals: list[DealItem]) -> None: + """按系统订单号去重,批量写入已成交数据;失败时不更新缓存。""" + new_deals = [] + seen = self.deals_sys_ids.copy() + for deal in deals: + if deal.sys_order_id in seen: + continue + new_deals.append(self.deal_record(deal)) + seen.add(deal.sys_order_id) + if not new_deals: + return + with closing(self._connect()) as db, db: + self._insert_deals(db, new_deals) + self.load() + + def save(self, items: dict, deals: list[dict]) -> None: + if len(deals) < self._deal_count: + raise ValueError('Execution history is append-only') + new_deals = deals[self._deal_count:] + positions = [ + { + 'base_order_id': '', 'base_qty': 0, 'base_cost': 0.0, + 'added_order_id': '', 'added_num': 0, 'added_qty': 0, 'added_cost': 0.0, + **item, + } + for item in items.values() + ] + for row in [*items.values(), *new_deals]: + if any(isinstance(value, float) and not math.isfinite(value) for value in row.values()): + raise ValueError('Quantities, prices and amounts must be finite') + with closing(self._connect()) as db, db: + # 更新已有证券时保留其自增 ID;仅删除快照中已移除的证券。 + for row in db.execute('SELECT code FROM positions').fetchall(): + if row['code'] not in items: + db.execute('DELETE FROM positions WHERE code = ?', (row['code'],)) + db.executemany( + """INSERT INTO positions + (code, base_order_id, base_qty, base_cost, + added_order_id, added_num, added_qty, added_cost, status) + VALUES (:code, :base_order_id, :base_qty, :base_cost, + :added_order_id, :added_num, :added_qty, :added_cost, :status) + ON CONFLICT(code) DO UPDATE SET + base_order_id = excluded.base_order_id, + base_qty = excluded.base_qty, + base_cost = excluded.base_cost, + added_order_id = excluded.added_order_id, + added_num = excluded.added_num, + added_qty = excluded.added_qty, + added_cost = excluded.added_cost, + status = excluded.status""", + positions, + ) + self._insert_deals(db, new_deals) + self.load() diff --git a/py-client/sdk/models.py b/py-client/sdk/models.py index c957650..822baad 100644 --- a/py-client/sdk/models.py +++ b/py-client/sdk/models.py @@ -67,64 +67,52 @@ class OrderItem: @dataclass(slots=True) class DealItem: - """由 QMT Deal 成交对象解析得到的标准成交记录。""" + """Execution data from the API's fixed order-detail fields.""" - id: str - order_id: str - code: str - side: str - remark: str - traded_at: datetime | None - volume: int - price: float - amount: float + sys_order_id: str = "" local_order_id: str = "" - order_ref: str = "" + code: str = "" + instrument_id: str = "" exchange_id: str = "" name: str = "" - account_id: str = "" - commission: float = 0.0 - trade_date: str = "" - trade_time: str = "" + offset_flag: str = "" + side: str = "" + status: str = "" + remaining_volume: int = 0 + traded_volume: int = 0 + order_time: int = 0 + insert_date: str = "" + insert_time: str = "" + remark: str = "" + price: float = 0.0 + trade_price: float = 0.0 + trade_amount: float = 0.0 @classmethod def from_trade_detail(cls, data: dict[str, Any]) -> "DealItem": - """从 TradeDetailData 的 QMT Deal 原始字段创建成交记录。""" instrument_id = str(data.get("m_strInstrumentID") or "") exchange_id = str(data.get("m_strExchangeID") or "") - code = ( - f"{instrument_id}.{exchange_id}" - if instrument_id and exchange_id - else instrument_id - ) remark = str(data.get("m_strRemark") or "") - trade_date = str(data.get("m_strTradeDate") or "") - trade_time = str(data.get("m_strTradeTime") or "") + offset_flag = str(data.get("m_nOffsetFlag") or "") return cls( - id=str(data.get("m_strTradeID") or ""), - order_id=str(data.get("m_strOrderSysID") or ""), - code=code, - side={"23": "BUY", "24": "SELL", "48": "BUY", "49": "SELL"}.get( - str(data.get("m_nOffsetFlag")), "" - ), - remark=remark, - traded_at=_parse_datetime(trade_date, trade_time), - volume=_number(data.get("m_nVolume"), int), - price=_number(data.get("m_dPrice")), - amount=_number(data.get("m_dTradeAmount")), + sys_order_id=str(data.get("m_strOrderSysID") or ""), local_order_id=remark.split("|", 1)[0] if remark else "", - order_ref=str(data.get("m_strOrderRef") or ""), + code=f"{instrument_id}.{exchange_id}" if instrument_id and exchange_id else instrument_id, + instrument_id=instrument_id, exchange_id=exchange_id, name=str(data.get("m_strInstrumentName") or ""), - account_id=str(data.get("m_strAccountID") or ""), - commission=_number( - data.get( - "m_dCommission", - data.get("m_dComission", data.get("m_dComssion")), - ) - ), - trade_date=trade_date, - trade_time=trade_time, + offset_flag=offset_flag, + side={"23": "BUY", "24": "SELL", "48": "BUY", "49": "SELL"}.get(offset_flag, ""), + status=str(data.get("m_nOrderStatus") or ""), + remaining_volume=_number(data.get("m_nVolumeTotal"), int), + traded_volume=_number(data.get("m_nVolumeTraded"), int), + order_time=_number(data.get("m_nOrderTime"), int), + insert_date=str(data.get("m_strInsertDate") or ""), + insert_time=str(data.get("m_strInsertTime") or ""), + remark=remark, + price=_number(data.get("m_dPrice")), + trade_price=_number(data.get("m_dTradePrice")), + trade_amount=_number(data.get("m_dTradeAmount")), ) diff --git a/py-client/sdk/portfolio.py b/py-client/sdk/portfolio.py index 6a91168..53e117d 100644 --- a/py-client/sdk/portfolio.py +++ b/py-client/sdk/portfolio.py @@ -42,8 +42,8 @@ class PortfolioMixin: return [OrderItem.from_trade_detail(row) for row in data] def deals(self) -> list[DealItem]: - """查询原始 Deal 成交对象并转换为标准成交记录。""" - return [DealItem.from_trade_detail(row) for row in self.org("deal")] + data = self._get_json("/api/portfolio/deal") or [] + return [DealItem.from_trade_detail(row) for row in data] def trade_detail_data(self, datatype: str) -> Any: datatype = str(datatype).strip().lower() diff --git a/py-client/strategy/trend/boot.py b/py-client/strategy/trend/boot.py index dec5ab4..13018e8 100644 --- a/py-client/strategy/trend/boot.py +++ b/py-client/strategy/trend/boot.py @@ -46,7 +46,7 @@ def StartTrend() -> None: config.account_config.signal_allow, ) log.info( - "[启动] 趋势策略已启动,账户=%s,信号=%d,持仓=%d", + "[启动] Trend策略已启动,账户=%s,信号=%d,持仓=%d", config.account_config.account_id, len(signals), len(positions), diff --git a/py-client/strategy/zt/boot.py b/py-client/strategy/zt/boot.py index abdf46e..92670e0 100644 --- a/py-client/strategy/zt/boot.py +++ b/py-client/strategy/zt/boot.py @@ -5,18 +5,21 @@ from __future__ import annotations +from concurrent.futures import Future, ThreadPoolExecutor import logging as log import time from datetime import datetime, time as clock_time +from pathlib import Path import config from libs.calc import trading_time from libs.market import market_allow_open -from libs.signal import init_signals +from libs.signal import SignalItem, init_signals from libs.collector import collector_push from libs.grid_take_profit import GridTrailingTracker from sdk import Client -from libs.order import OrderBook +from libs.overview import Overview +from libs.order import BUSY_STATUSES, OrderBook from libs.watch import DipWatch from libs.runtime import Runtime from .state import TState, SOLD @@ -31,9 +34,11 @@ def StartZT() -> None: config.global_config.qmt_token, config.HTTP_TIMEOUT, ) as client: - state = TState.for_strategy( - config.global_config.qmt_data_dir, "zt", config.account_config.account_id + state = TState( + Path(config.global_config.qmt_data_dir) + / f"zt_{config.account_config.account_id}_state.db" ) + executor = ThreadPoolExecutor(max_workers=3, thread_name_prefix="zt") run = Runtime( client=client, global_cfg=config.global_config, @@ -42,110 +47,153 @@ def StartZT() -> None: open_watch=DipWatch(), add_watch=DipWatch(), profit_tracker=GridTrailingTracker(config.account_config.grid_step_pct), + executor=executor ) - log.info( - "[ZT 启动] 账户=%s,底仓信号=dcm,状态文件=%s", - run.account_cfg.account_id, - state.path, + + portfolio = client.portfolio() + assets = portfolio.assets + positions = list(portfolio.positions.values()) + run.orders.refresh(client, portfolio.orders) + + # 获取本策略的信号开仓数据 + signals = init_signals(config.global_config,["dcm"]) + log.info("[启动] ZT 策略已启动,账户=%s,信号=%d,持仓=%d", + config.account_config.account_id, + len(signals), + len(positions), ) + + Overview(assets, positions, config.account_config) + + DEFAULT_TICK_INTERVAL = 30 while True: - now = datetime.now() - if now.time() >= clock_time(15): - # 收盘前最后一次只读对账,不发新单;未完成买回继续持久保存。 - try: - portfolio = client.portfolio() - deals = client.deals() - state.reconcile( - list(portfolio.positions.values()), - deals, - now.date().isoformat(), - ) - except Exception: - log.exception("[ZT] 收盘对账失败,保留本地待确认记录") - for item in state.items.values(): - if item.phase == SOLD or state.busy(item.code): - log.warning("[ZT] 收盘仍有待完成轮次:%s", item.code) + lt = time.localtime() + if (lt.tm_hour, lt.tm_min, lt.tm_sec) >= (15, 0, 0): + log.info("[Trend] 已到 15:00,结束趋势策略") return + current_sec = lt.tm_sec + + # 计算距离下一个目标时间点(0秒或30秒)的等待时间 + if current_sec < DEFAULT_TICK_INTERVAL: + wait_seconds = DEFAULT_TICK_INTERVAL - current_sec + elif current_sec < 60: + wait_seconds = 60 - current_sec + else: + wait_seconds = DEFAULT_TICK_INTERVAL + + # 等待到目标时间点 + time.sleep(wait_seconds) + # 单轮失败不能杀死唯一的交易定时线程。 try: - RunOnce(run, state) - except Exception: - log.exception("[ZT] 本 tick 执行失败,下一个 tick 继续") - # 计算距离下一个目标时间点(0秒或30秒)的等待时间。 - time.sleep(30 - datetime.now().second % 30) + RunOnce(run, state, signals) + except Exception as e: + log.error( + f"[Trend] 本 tick 执行失败,下一 tick 继续: {e}", exc_info=True + ) - -def RunOnce(run: Runtime, state: TState) -> None: +def RunOnce(run: Runtime, state: TState, signals: list[SignalItem]) -> None: """账户快照 → 成交对账 → 做 T 管理 → dcm 建仓,共用一份资金预算。""" now = datetime.now() - if not trading_time(now) or now.time() >= clock_time(15): + if not trading_time(now): return - today = now.date().isoformat() + started_at = time.monotonic() # 1. 一次获取资产、持仓和订单,并清理过期订单。 - portfolio = run.client.portfolio() - deals = run.client.deals() - positions = list(portfolio.positions.values()) - run.orders.refresh(run.client, portfolio.orders) - # 状态只按真实成交记账,不使用委托状态推算数量和成本。 - state.reconcile(positions, deals, today) - - # 2. 获取本策略的信号开仓数据;信号失败不阻断已有做 T 买回。 try: - signals = init_signals(run.global_cfg, ["dcm"]) + portfolio = run.client.portfolio() + assets = portfolio.assets + deals = run.client.deals() + position_codes = list(portfolio.positions) + positions = list(portfolio.positions.values()) + run.orders.refresh(run.client, portfolio.orders) + state.reconcile(positions,deals) except Exception: - log.exception("[ZT] 获取 dcm 信号失败,本轮只管理已有底仓") - signals = [] - position_codes = { - position.stock_code for position in positions if position.volume > 0 - } - candidates = [ - signal - for signal in signals - if signal.signal_key == "dcm" and signal.code not in position_codes - ] - - # 3. 获取持仓和待开仓证券的实时行情 tick,零持仓的待买回证券也包含在内。 - codes = list( - dict.fromkeys( - list(position_codes) - + list(state.items) - + [signal.code for signal in candidates] - ) - ) - ticks = run.client.full_tick(codes) if codes else {} - now = datetime.now() # 网络请求可能跨过尾盘边界,提交前重新判断。 - if not trading_time(now) or now.time() >= clock_time(15): + log.exception("[Portfolio] 刷新账户快照失败") return - # 4. 先完成买回,避免开底仓抢占资金;交易逻辑串行,状态无需多线程写入。 - available = max(0.0, portfolio.assets.available) - # 未确认买单可能尚未反映在资金快照中,保守预留,宁可少买也不重复使用。 - for pending in state.pending.values(): - if pending.kind != "sell": - tick = ticks.get(pending.code) - if tick is None or tick.last_price <= 0: - available = 0.0 - break - available = max(0.0, available - pending.qty * tick.last_price * 1.01) - force = now.time() >= clock_time(14, 50) - available = manage_positions(run, state, ticks, positions, available, today, force) + futures: list[tuple[str, Future]] = [ + ( + "数据提交", + run.executor.submit( + collector_push, + run.account_cfg.account_id, + assets, + positions, + ), + ) + ] - # 5. 验证可用资金;低于资金安全线时禁止开新仓,尾盘只完成做 T 买回。 - reserve = max(0.0, portfolio.assets.total * run.account_cfg.min_cash_ratio) - # 未完成的卖出/买回可能继续占用资金,不再额外开底仓。 - outstanding = bool(state.pending) or any( - item.phase == SOLD for item in state.items.values() + # 2. 验证可用资金;低于资金安全线时禁止开新仓。 + allow_open_by_cash = ( + assets.available >= assets.total * run.account_cfg.min_cash_ratio ) - if not force and not outstanding and market_allow_open() and available > reserve: - open_signal(run, state, ticks, candidates, available - reserve) + if not allow_open_by_cash: + log.info( + "[Status] 禁止开仓:可用资金不足,可用=%.2f,总资产=%.2f", + assets.available, + assets.total, + ) + + # 3. 获取大盘状态,只有大盘信号允许时才执行开仓。 + market_ok = market_allow_open() + + # 4. 验证有效开仓信号:排除已有持仓和未决订单。 + allow_open: list[SignalItem] = [] + allow_codes: list[str] = [] + for signal in signals: + if signal.code not in position_codes: + allow_open.append(signal) + allow_codes.append(signal.code) + + if allow_open and not market_ok: + log.info("[开仓] 禁止开仓:大盘信号不允许,候选=%d", len(allow_open)) + + # 5. 获取持仓和待开仓证券的实时行情 tick。 + all_codes = list(dict.fromkeys(position_codes + allow_codes)) + try: + ticks = run.client.full_tick(all_codes) + except Exception: + log.exception("[行情] 获取行情失败,代码数量=%d", len(all_codes)) + return - # 6. 数据采集不与交易逻辑争用状态;采集函数自身隔离传输异常。 - collector_push(run.account_cfg.account_id, portfolio.assets, positions) log.info( - "[ZT] 本轮完成,底仓=%d,待确认=%d,耗时=%d毫秒", - len(state.items), - len(state.pending), - int((time.monotonic() - started_at) * 1000), + "[RunOnce] 本轮就绪,持仓=%d,候选=%d,大盘允许=%s,资金允许=%s", + len(positions), + len(allow_open), + market_ok, + allow_open_by_cash, ) + + # 启动线程,开始计算 + # 7. 持仓计算。 + futures.append( + ( + "持仓计算", + run.executor.submit( + manage_positions, run, ticks, positions, market_ok, assets.available + ), + ) + ) + + # 8. 开仓计算:必须同时存在有效信号且大盘允许开仓。 + if allow_open and market_ok and allow_open_by_cash: + futures.append( + ("开仓计算", run.executor.submit(open_signal, run, ticks, allow_open)) + ) + + # 9. 开始执行 + for name, future in futures: + _wait_worker(name, future) + log.info( + "[RunOnce] 本轮完成,耗时=%d毫秒", int((time.monotonic() - started_at) * 1000) + ) + + +def _wait_worker(name: str, future: Future) -> None: + """保留单轮继续运行的语义,分别记录工作线程异常。""" + try: + future.result() + except Exception: + log.exception("[运行] %s线程失败", name) \ No newline at end of file diff --git a/py-client/strategy/zt/open.py b/py-client/strategy/zt/open.py index df685b7..3479004 100644 --- a/py-client/strategy/zt/open.py +++ b/py-client/strategy/zt/open.py @@ -10,7 +10,7 @@ from libs.calc import calc_buy_volume from sdk import OP_BUY from libs.runtime import Runtime from libs.order import PlaceOrderRequest -from .state import PendingOrder, TState +from .state import TState def open_signal(run: Runtime, state: TState, ticks, signals, available: float) -> float: @@ -20,23 +20,18 @@ def open_signal(run: Runtime, state: TState, ticks, signals, available: float) - now = datetime.now() if (now.hour, now.minute) >= (14, 50): break - if item.signal_key != "dcm" or item.code in run.account_cfg.excluded_codes: + if item.code in run.account_cfg.excluded_codes: continue item_state = state.items.get(item.code) if item_state is not None and item_state.base_qty > 0: continue - # 1. 验证信号配置允许开仓的时间区间。 - signal_config = run.global_cfg.signals.get("dcm") - if signal_config is None or not check_timezone(signal_config.timezone): - continue - # 2. 检查该证券是否已有买入委托锁,防止重复下单。 + # 由委托簿检查活动委托,防止重复下单。 if ( - state.busy(item.code) - or run.orders.busy(item.code, "BUY") + run.orders.busy(item.code, "BUY") or run.orders.busy(item.code, "SELL") ): continue - # 3. 验证行情和最新价格是否有效。 + # 行情无效或超过策略价格上限时跳过。 tick = ticks.get(item.code) price = tick.last_price if tick else 0.0 if ( @@ -45,29 +40,20 @@ def open_signal(run: Runtime, state: TState, ticks, signals, available: float) - or price > run.account_cfg.zt_max_price ): continue - # 4. 根据单笔买入金额计算整手开仓数量,预留少量价差和费用。 + # 根据单笔买入金额计算整手数量,并预留少量价差和费用。 budget = min(run.account_cfg.buy_value, available) volume = calc_buy_volume(price, budget) amount = price * volume * 1.01 if volume <= 0 or price * volume > budget or amount > available: continue - # 5. 等待价格从观察低点反弹,防止直接接下跌中的“飞刀”。 + # 等待价格从观察低点反弹,防止直接接下跌中的“飞刀”。 if not run.open_watch.triggered("ZT 建仓", item.code, price): continue order_id = run.orders.new_order_id("base") request = PlaceOrderRequest( OP_BUY, item.code, volume, order_id, "zt", kind="base" ) - state.new_order( - PendingOrder( - order_id, - item.code, - "base", - volume, - datetime.now().date().isoformat(), - ) - ) - # 即使响应丢失,也保留资金预算和 pending,不能继续使用这笔钱。 + # 即使响应丢失,本轮也预留资金;状态簿只在取得实际成交后入账。 available -= amount if run.orders.place(run.client, request): run.open_watch.forget(item.code) @@ -75,45 +61,3 @@ def open_signal(run: Runtime, state: TState, ticks, signals, available: float) - except Exception: log.exception("[ZT 建仓] %s 处理异常,继续后续信号", item.code) return available - - -def check_timezone(timezone: str, now: datetime | None = None) -> bool: - """验证当前时间是否处于配置区间。 - - ``*`` 表示全天允许;多个区间用逗号分隔,例如 - ``9:30-10:30,13:30-14:30``。同时支持跨午夜区间。 - """ - timezone = str(timezone or "").strip() - if timezone == "*": - return True - - current = now or datetime.now() - current_minutes = current.hour * 60 + current.minute - - for section in timezone.split(","): - bounds = section.strip().split("-") - if len(bounds) != 2: - continue - start = _parse_minutes(bounds[0]) - end = _parse_minutes(bounds[1]) - if start is None or end is None: - continue - - if start <= end and start <= current_minutes <= end: - return True - if start > end and (current_minutes >= start or current_minutes <= end): - return True - - return False - - -def _parse_minutes(value: str) -> int | None: - """把 ``时:分`` 转换为当天分钟数,无效值返回 None。""" - try: - hour_text, minute_text = value.strip().split(":") - hour, minute = int(hour_text), int(minute_text) - except (TypeError, ValueError): - return None - if not 0 <= hour <= 23 or not 0 <= minute <= 59: - return None - return hour * 60 + minute diff --git a/py-client/strategy/zt/positions.py b/py-client/strategy/zt/positions.py index fc7de99..f1677cc 100644 --- a/py-client/strategy/zt/positions.py +++ b/py-client/strategy/zt/positions.py @@ -2,7 +2,6 @@ from __future__ import annotations -from datetime import datetime import logging as log import math @@ -10,7 +9,7 @@ from libs.grid_take_profit import GridState from sdk import OP_BUY, OP_SELL, PositionItem from libs.order import PlaceOrderRequest from libs.runtime import Runtime -from .state import PendingOrder, READY, SOLD, TState +from .state import READY, SOLD, TState def manage_positions( @@ -26,11 +25,7 @@ def manage_positions( by_code = {position.stock_code: position for position in positions} for code, state in list(state_store.items.items()): try: - now = datetime.now() - if now.hour >= 15: - break - force_buy_back = force_buy_back or (now.hour, now.minute) >= (14, 50) - if code in run.account_cfg.excluded_codes or state_store.busy(code): + if code in run.account_cfg.excluded_codes: continue if run.orders.busy(code, "BUY") or run.orders.busy(code, "SELL"): continue @@ -52,11 +47,11 @@ def manage_positions( continue if state.phase == SOLD: available = _try_buy_back( - run, state_store, state, price, available, today, force_buy_back + run, state, price, available, force_buy_back ) elif state.phase == READY and position and not force_buy_back: if price <= run.account_cfg.zt_max_price: - _try_sell(run, state_store, state, position, price, today) + _try_sell(run, state, position, price, today) except Exception: log.exception("[ZT 持仓] %s 处理异常,继续后续证券", code) return available @@ -64,7 +59,6 @@ def manage_positions( def _try_sell( run: Runtime, - state_store: TState, state, position: PositionItem, price: float, @@ -88,18 +82,15 @@ def _try_sell( request = PlaceOrderRequest( OP_SELL, state.code, volume, order_id, "zt", kind="sell" ) - state_store.new_order(PendingOrder(order_id, state.code, "sell", volume, today)) if run.orders.place(run.client, request): log.info("[ZT 卖出] %s %d 股,等待成交后确定买回数量和价格", state.code, volume) def _try_buy_back( run: Runtime, - state_store: TState, state, price: float, available: float, - today: str, force: bool, ) -> float: """按实际卖出均价下跌后反弹买回;尾盘不再受下跌幅度、反弹及价格上限限制。""" @@ -115,18 +106,14 @@ def _try_buy_back( return available order_id = run.orders.new_order_id("t-buy") request = PlaceOrderRequest(OP_BUY, state.code, volume, order_id, "zt", kind="buy") - state_store.new_order(PendingOrder(order_id, state.code, "buy", volume, today)) - # pending 已落盘,任何请求结果都预留资金;下一轮再从柜台快照确认。 + # 本轮预留资金;状态簿只在取得实际成交后入账。 available -= amount - try: - if run.orders.place(run.client, request): - run.add_watch.forget(state.code) - log.info( - "[ZT 买回] %s %d 股,%s", - state.code, - volume, - "尾盘强制买回" if force else "下跌后反弹", - ) - except Exception: - log.exception("[ZT 买回] %s 请求结果未知,保留 pending 和预算", state.code) + if run.orders.place(run.client, request): + run.add_watch.forget(state.code) + log.info( + "[ZT 买回] %s %d 股,%s", + state.code, + volume, + "尾盘强制买回" if force else "下跌后反弹", + ) return available diff --git a/py-client/strategy/zt/state.py b/py-client/strategy/zt/state.py index b0d1d4c..fbf6cf8 100644 --- a/py-client/strategy/zt/state.py +++ b/py-client/strategy/zt/state.py @@ -1,15 +1,14 @@ -"""做 T 策略的底仓、待确认委托和实际成交记录。""" +"""做 T 策略的持仓状态和逐笔实际成交记录。""" from __future__ import annotations -import json -import logging as log import math -from dataclasses import asdict, dataclass, field +from dataclasses import dataclass +from datetime import datetime from pathlib import Path -from time import time from sdk import DealItem, PositionItem +from libs.orderbook import OrderBook READY, SOLD, DONE = "READY", "SOLD", "DONE" @@ -21,180 +20,162 @@ class TStateItem: base_cost: float = 0.0 trade_date: str = "" phase: str = READY - sell_order_id: str = "" sell_qty: int = 0 sell_price: float = 0.0 - buy_order_id: str = "" - base_order_id: str = "" buy_qty: int = 0 buy_cost: float = 0.0 - - -@dataclass(slots=True) -class PendingOrder: - order_id: str - code: str - kind: str # base:底仓;sell:做 T 卖出;buy:做 T 买回 - qty: int - trade_date: str - submit_at: float = field(default_factory=time) + id: int = 0 + base_order_id: str = '' + added_order_id: str = '' + added_num: int = 0 + added_qty: int = 0 + added_cost: float = 0.0 class TState: - """交易逻辑串行更新;JSON 保存底仓、待确认委托及成交历史。""" + """Apply actual executions immediately, atomically with their position changes.""" def __init__(self, path: str | Path) -> None: - self.path = Path(path) - self.items: dict[str, TStateItem] = {} - self.pending: dict[str, PendingOrder] = {} - self.records: list[dict] = [] + self._store = OrderBook(path) + self.path = self._store.path self._load() + @staticmethod + def _is_zt_deal(deal: DealItem) -> bool: + return ( + deal.side == 'BUY' and deal.local_order_id.startswith(('zt-base-', 'zt-t-buy-')) + ) or ( + deal.side == 'SELL' and deal.local_order_id.startswith('zt-t-sell-') + ) + + @staticmethod + def _reset(item: TStateItem, date: str) -> bool: + if item.phase == DONE and item.trade_date != date: + item.phase, item.trade_date = READY, '' + item.sell_qty = item.buy_qty = 0 + item.sell_price = item.buy_cost = 0.0 + return True + return False + @classmethod - def for_strategy( - cls, data_dir: str | Path, strategy: str, account_id: str - ) -> TState: - return cls(Path(data_dir) / f"{strategy}_{account_id}_state.json") - - def busy(self, code: str) -> bool: - return any(order.code == code for order in self.pending.values()) - - def new_order(self, order: PendingOrder) -> None: - """下单前落盘;请求超时不能当作失败删除,等待后续委托确认。""" - if self.busy(order.code): - raise ValueError(f"{order.code} 已有待确认委托") - self.pending[order.order_id] = order - try: - self.save() - except Exception: - del self.pending[order.order_id] - raise + def _apply_t_deal(cls, item: TStateItem, deal: dict) -> None: + """实时入账与重启恢复共用同一套做 T 轮次计算。""" + cls._reset(item, deal['insert_date']) + qty, amount = deal['traded_volume'], deal['trade_amount'] + if deal['side'] == 'SELL': + total = item.sell_qty + qty + item.sell_price = (item.sell_qty * item.sell_price + amount) / total + item.sell_qty = total + item.phase = SOLD + else: + total = item.buy_qty + qty + item.buy_cost = (item.buy_qty * item.buy_cost + amount) / total + item.buy_qty = total + item.phase = DONE if total >= item.sell_qty else SOLD + item.trade_date = deal['insert_date'] def reconcile( - self, positions: list[PositionItem], deals: list[DealItem], today: str + self, positions: list[PositionItem], deals: list[DealItem] ) -> None: - """先按实际成交记账,再接管未知持仓;不覆盖已记录的底仓成本。""" - # 同一本地委托可能有多笔成交;按成交编号去重后合并数量和金额。 - by_id: dict[str, dict[str, DealItem]] = {} + """Deduplicate each fill; partial fills do not wait for order completion.""" + today = datetime.now().date().isoformat() + seen = {row['sys_order_id'] for row in self.deals} + rows = [] for deal in deals: - if deal.local_order_id and deal.id: - by_id.setdefault(deal.local_order_id, {})[deal.id] = deal - for order_id, pending in list(self.pending.items()): - side = "SELL" if pending.kind == "sell" else "BUY" - rows = [ - row - for row in by_id.get(order_id, {}).values() - if row.code == pending.code and row.side == side - ] - if not rows: - log.warning("[ZT 状态] 成交暂未查到,保留待确认:%s", order_id) + if not self._is_zt_deal(deal) or deal.sys_order_id in seen: continue - qty = sum(row.volume for row in rows) - # 成交未达到计划数量时继续等待,防止后续成交到达后重复记账。 - if qty != pending.qty: + try: + row = self._store.deal_record(deal) + except ValueError: continue - amounts = [ - row.amount if row.amount > 0 else row.price * row.volume - for row in rows - if row.volume > 0 - ] - if any(not math.isfinite(amount) or amount <= 0 for amount in amounts): - continue - amount = sum(amounts) - cost = amount / qty if qty else 0.0 - item = self.items.setdefault(pending.code, TStateItem(pending.code)) - if pending.kind == "base": - item.base_order_id = order_id - item.base_qty, item.base_cost = qty, cost - elif pending.kind == "sell": - item.trade_date = today # 跨日成交也占用确认当天的一轮。 - item.sell_order_id = order_id - item.sell_qty, item.sell_price = qty, cost - item.buy_qty, item.buy_cost = 0, 0.0 - item.phase = SOLD if qty else READY - else: - total = item.buy_qty + qty - item.buy_cost = ( - (item.buy_qty * item.buy_cost + amount) / total if total else 0.0 + rows.append(row) + seen.add(deal.sys_order_id) + rows.sort(key=lambda r: (r['insert_date'], r['insert_time'])) + modified = False + try: + # Snapshot includes these fills: subtract their net quantity before replay. + net = {} + for row in rows: + net[row['code']] = net.get(row['code'], 0) + ( + row['traded_volume'] if row['side'] == 'BUY' else -row['traded_volume'] ) - item.buy_qty = total - item.buy_order_id = order_id - item.phase = DONE if total >= item.sell_qty else SOLD - if item.phase == DONE: - item.trade_date = today - # 记录真实成交编号,重启后仍可核对本次状态变更的来源。 - self.records.append( - { - **asdict(pending), - "confirmed_date": today, - "filled_qty": qty, - "filled_cost": cost, - "amount": amount, - "deal_ids": [row.id for row in rows], - } - ) - del self.pending[order_id] + for position in positions: + code = position.stock_code + if code in self.items or position.volume <= 0: + continue + if not math.isfinite(position.open_price) or position.open_price <= 0: + continue + qty = max(0, position.volume - net.get(code, 0)) + self.items[code] = TStateItem(code, qty, position.open_price if qty else 0.0) + modified = True - for position in positions: - code = position.stock_code - if position.volume <= 0 or self.busy(code): - continue - if ( - code not in self.items - and math.isfinite(position.open_price) - and position.open_price > 0 - ): - self.items[code] = TStateItem( - code, position.volume, position.open_price - ) - self.records.append( - { - "kind": "import", - "code": code, - "date": today, - "filled_qty": position.volume, - "filled_cost": position.open_price, - } - ) - log.warning( - "[ZT 底仓] 首次接管 %s,使用当前均价,无法还原历史成本", code - ) + # 全部卖出时快照可能已无该证券,按净卖出数量恢复待买回的底仓数量。 + for code, delta in net.items(): + if code not in self.items and delta < 0: + self.items[code] = TStateItem(code, -delta) - for item in self.items.values(): - # 未买回的轮次跨日继续,不删除零持仓的做 T 债务。 - if ( - item.trade_date != today - and item.phase == DONE - and not self.busy(item.code) - ): - item.phase, item.trade_date = READY, "" - item.sell_qty = item.buy_qty = 0 - item.sell_price = item.buy_cost = 0.0 - item.sell_order_id = item.buy_order_id = "" - self.save() + for row in rows: + item = self.items.setdefault(row['code'], TStateItem(row['code'])) + self._reset(item, row['insert_date']) + qty, amount = row['traded_volume'], row['trade_amount'] + if row['local_order_id'].startswith('zt-base-'): + total = item.base_qty + qty + item.base_cost = (item.base_qty * item.base_cost + amount) / total + item.base_qty = total + item.base_order_id = row['local_order_id'] + else: + self._apply_t_deal(item, row) + self.deals.append(row) + modified = True + + for item in self.items.values(): + modified = self._reset(item, today) or modified + if modified: + self.save() + except Exception: + self._load() + raise def save(self) -> None: - self.path.parent.mkdir(parents=True, exist_ok=True) - temporary = self.path.with_suffix(self.path.suffix + ".tmp") - payload = { - "items": {key: asdict(item) for key, item in self.items.items()}, - "pending": {key: asdict(item) for key, item in self.pending.items()}, - "records": self.records, - } - temporary.write_text( - json.dumps(payload, ensure_ascii=False, indent=2, allow_nan=False) + "\n", - encoding="utf-8", - ) - temporary.replace(self.path) + try: + self._store.save( + { + code: { + 'code': item.code, + 'base_order_id': item.base_order_id, + 'base_qty': item.base_qty, + 'base_cost': item.base_cost, + 'added_order_id': item.added_order_id, + 'added_num': item.added_num, + 'added_qty': item.added_qty, + 'added_cost': item.added_cost, + 'status': item.phase, + } + for code, item in self.items.items() + }, + self.deals, + ) + except Exception: + self._load() + raise def _load(self) -> None: - if not self.path.is_file(): - return - raw = json.loads(self.path.read_text(encoding="utf-8")) - self.items = { - code: TStateItem(**item) for code, item in raw["items"].items() - } - self.pending = { - key: PendingOrder(**item) for key, item in raw["pending"].items() - } - self.records = raw["records"] + self._store.load() + self.items = {} + for code, position in self._store.positions.items(): + position = dict(position) + position['phase'] = position.pop('status') + self.items[code] = TStateItem(**position) + self.deals = [ + {key: value for key, value in deal.items() if key != 'id'} + for deal in self._store.deals.values() + ] + # 轮次明细不占用持仓表字段,从已保存的逐笔成交重建。 + for deal in self.deals: + if not deal['local_order_id'].startswith('zt-base-') and deal['code'] in self.items: + self._apply_t_deal(self.items[deal['code']], deal) + for code, item in self.items.items(): + if self._store.positions[code]['status'] == READY: + item.phase, item.trade_date = READY, '' + item.sell_qty = item.buy_qty = 0 + item.sell_price = item.buy_cost = 0.0 diff --git a/py-client/tests/test_deal_model.py b/py-client/tests/test_deal_model.py new file mode 100644 index 0000000..81b0fbe --- /dev/null +++ b/py-client/tests/test_deal_model.py @@ -0,0 +1,78 @@ +import sqlite3 +import tempfile +import unittest +from contextlib import closing +from dataclasses import asdict +from pathlib import Path + +from libs.orderbook import OrderBook +from sdk.models import DealItem +from sdk.portfolio import PortfolioMixin + + +class DealModelTests(unittest.TestCase): + def setUp(self): + self.raw = { + 'm_strOrderSysID': 'sys-123', + 'm_strInstrumentID': '600000', + 'm_strExchangeID': 'SH', + 'm_strInstrumentName': 'Test stock', + 'm_nOffsetFlag': '24', + 'm_nOrderStatus': '56', + 'm_nVolumeTotal': '0', + 'm_nVolumeTraded': '200', + 'm_nOrderTime': '101530', + 'm_strInsertDate': '20260907', + 'm_strInsertTime': '10:15:30', + 'm_strRemark': 'zt-t-sell-local1|zt', + 'm_dPrice': '15.8', + 'm_dTradePrice': '15.6', + 'm_dTradeAmount': '3120', + } + + def test_all_api_fields_are_parsed(self): + self.assertEqual(asdict(DealItem.from_trade_detail(self.raw)), { + 'sys_order_id': 'sys-123', 'local_order_id': 'zt-t-sell-local1', + 'code': '600000.SH', 'instrument_id': '600000', 'exchange_id': 'SH', + 'name': 'Test stock', 'offset_flag': '24', 'side': 'SELL', 'status': '56', + 'remaining_volume': 0, 'traded_volume': 200, 'order_time': 101530, + 'insert_date': '20260907', 'insert_time': '10:15:30', + 'remark': 'zt-t-sell-local1|zt', 'price': 15.8, + 'trade_price': 15.6, 'trade_amount': 3120.0, + }) + + def test_api_deals_response_returns_deal_items(self): + client = PortfolioMixin() + for response in ({'deals': [self.raw]}, [self.raw]): + with self.subTest(response_type=type(response).__name__): + client._get_json = lambda path: response + deals = client.deals() + self.assertIsInstance(deals[0], DealItem) + self.assertEqual(deals[0].sys_order_id, 'sys-123') + self.assertEqual(deals[0].traded_volume, 200) + + def test_model_matches_sql_columns_and_restart(self): + deal = DealItem.from_trade_detail(self.raw) + with tempfile.TemporaryDirectory() as tmp: + path = Path(tmp) / 'state.db' + book = OrderBook(path) + book.sync_deals([deal]) + with closing(sqlite3.connect(path)) as db: + columns = {row[1] for row in db.execute('PRAGMA table_info(deals)')} + self.assertEqual(columns, {'id', *asdict(deal)}) + loaded = OrderBook(path).deals['sys-123'] + expected = asdict(deal) + expected['insert_date'] = '2026-09-07' + self.assertEqual({key: value for key, value in loaded.items() if key != 'id'}, expected) + + def test_missing_amount_uses_execution_price_not_order_price(self): + deal = DealItem.from_trade_detail({**self.raw, 'm_dTradeAmount': '0'}) + row = OrderBook.deal_record(deal) + self.assertEqual(row['trade_amount'], 3120) + deal.trade_price = 0 + with self.assertRaises(ValueError): + OrderBook.deal_record(deal) + + +if __name__ == '__main__': + unittest.main() diff --git a/py-client/tests/test_orderbook.py b/py-client/tests/test_orderbook.py new file mode 100644 index 0000000..7b7a027 --- /dev/null +++ b/py-client/tests/test_orderbook.py @@ -0,0 +1,219 @@ +import sqlite3 +import tempfile +import unittest +from contextlib import closing +from pathlib import Path +from datetime import datetime +from unittest.mock import patch + +from libs.orderbook import OrderBook +from sdk import DealItem, PositionItem +from strategy.zt.state import DONE, READY, SOLD, TState + + +class OrderBookTests(unittest.TestCase): + def setUp(self): + clock = patch('strategy.zt.state.datetime') + self.clock = clock.start() + self.addCleanup(clock.stop) + self.clock.now.return_value = datetime(2026, 9, 1) + self.tmp = tempfile.TemporaryDirectory() + self.addCleanup(self.tmp.cleanup) + self.path = Path(self.tmp.name) / 'state.db' + + def deal(self, kind, sys_order_id, qty, price, date='2026-09-01'): + prefix = {'base': 'zt-base-', 'sell': 'zt-t-sell-', 'buy': 'zt-t-buy-'}[kind] + return DealItem.from_trade_detail({ + 'm_strOrderSysID': sys_order_id, + 'm_strInstrumentID': '600000', 'm_strExchangeID': 'SH', + 'm_strInstrumentName': 'Test stock', + 'm_nOffsetFlag': '24' if kind == 'sell' else '23', + 'm_nOrderStatus': '56', 'm_nVolumeTotal': '0', + 'm_nVolumeTraded': str(qty), 'm_nOrderTime': '100000', + 'm_strInsertDate': date, 'm_strInsertTime': '10:00:00', + 'm_strRemark': prefix + 'order1|zt', + 'm_dPrice': str(price + 1), 'm_dTradePrice': str(price), + 'm_dTradeAmount': str(qty * price), + }) + + def test_partial_fills_restart_dedup_and_daily_cycle(self): + state = TState(self.path) + state.reconcile([PositionItem(stock_code='600000.SH', volume=200, open_price=10)], []) + self.assertEqual(state.deals, []) + first = self.deal('sell', 'd1', 40, 12) + second = self.deal('sell', 'd2', 60, 13) + state.reconcile([], [first]) + state = TState(self.path) + self.assertEqual(state.items['600000.SH'].phase, SOLD) + self.assertEqual(state.items['600000.SH'].sell_qty, 40) + state.reconcile([], [first, first, second]) + self.assertEqual(len(state.deals), 2) + self.assertAlmostEqual(state.items['600000.SH'].sell_price, 12.6) + self.clock.now.return_value = datetime.fromisoformat('2026-09-02') + state.reconcile([], [first, second]) + self.assertEqual(len(state.deals), 2) + self.assertEqual(state.items['600000.SH'].phase, SOLD) + b1 = self.deal('buy', 'd3', 40, 11, '2026-09-02') + b2 = self.deal('buy', 'd4', 60, 10, '2026-09-02') + state.reconcile([], [b1]) + self.assertEqual(state.items['600000.SH'].phase, SOLD) + state.reconcile([], [b1, b2]) + self.assertEqual(TState(self.path).items['600000.SH'].phase, DONE) + self.assertAlmostEqual(state.items['600000.SH'].buy_cost, 10.4) + self.clock.now.return_value = datetime.fromisoformat('2026-09-03') + state.reconcile([], []) + item = TState(self.path).items['600000.SH'] + self.assertEqual((item.phase, item.base_qty, item.base_cost, item.sell_qty), (READY, 200, 10, 0)) + + def test_json_is_never_read(self): + legacy = self.path.with_suffix('.json') + legacy.write_text('invalid JSON', encoding='utf-8') + book = OrderBook(self.path) + self.assertIsNone(book.load()) + self.assertEqual((book.positions, book.deals, book.deals_sys_ids), ({}, {}, set())) + self.assertEqual(legacy.read_text(encoding='utf-8'), 'invalid JSON') + + def test_sync_deals_deduplicates_batch_and_restart(self): + book = OrderBook(self.path) + self.assertEqual((book.positions, book.deals, book.deals_sys_ids), ({}, {}, set())) + first = self.deal('base', 'd1', 40, 10, '20260901') + second = self.deal('base', 'd2', 60, 12) + with patch.object(book, '_insert_deals', wraps=book._insert_deals) as insert: + book.sync_deals([first, first, second]) + insert.assert_called_once() + self.assertEqual(len(insert.call_args.args[1]), 2) + self.assertEqual(book.deals_sys_ids, {'d1', 'd2'}) + self.assertEqual(book.deals['d1']['insert_date'], '2026-09-01') + self.assertEqual(book.deals['d2']['traded_volume'], 60) + book = OrderBook(self.path) + self.assertEqual(book.deals_sys_ids, {'d1', 'd2'}) + with patch.object(book, '_insert_deals') as insert: + book.sync_deals([first, second]) + book.sync_deals([]) + insert.assert_not_called() + self.assertEqual(len(book.deals), 2) + + def test_sync_deals_failure_rolls_back_entire_batch_and_cache(self): + book = OrderBook(self.path) + first = self.deal('base', 'd1', 100, 10) + invalid = self.deal('base', 'd2', 100, 10) + invalid.side = 'INVALID' + with self.assertRaises(sqlite3.IntegrityError): + book.sync_deals([first, invalid]) + self.assertEqual(book.deals, {}) + self.assertEqual(book.deals_sys_ids, set()) + self.assertEqual(OrderBook(self.path).deals, {}) + invalid.side = 'BUY' + book.sync_deals([first, invalid]) + self.assertEqual(book.deals_sys_ids, {'d1', 'd2'}) + + def test_load_refreshes_all_caches(self): + book = OrderBook(self.path) + writer = OrderBook(self.path) + writer.save({'600000.SH': {'code': '600000.SH', 'status': READY}}, []) + writer.sync_deals([self.deal('base', 'd1', 100, 10)]) + book.load() + self.assertEqual(book.positions['600000.SH']['status'], READY) + self.assertEqual(book.deals_sys_ids, {'d1'}) + self.assertEqual(book.deals['d1']['local_order_id'], 'zt-base-order1') + + def test_first_start_after_full_sale_keeps_buyback_quantity(self): + state = TState(self.path) + state.reconcile([], [self.deal('sell', 'd1', 100, 12)]) + item = state.items['600000.SH'] + self.assertEqual((item.base_qty, item.sell_qty, item.phase), (100, 100, SOLD)) + state.reconcile([], [self.deal('buy', 'd2', 100, 11)]) + self.assertEqual(TState(self.path).items['600000.SH'].phase, DONE) + + def test_failed_insert_rolls_back_memory_and_database(self): + state = TState(self.path) + with closing(sqlite3.connect(self.path)) as db: + db.execute("""CREATE TRIGGER fail_insert BEFORE INSERT ON deals + BEGIN SELECT RAISE(ABORT, 'test failure'); END""") + fill = self.deal('base', 'd1', 100, 10) + with self.assertRaises(sqlite3.IntegrityError): + state.reconcile([], [fill]) + self.assertFalse(state.items) + self.assertFalse(state.deals) + self.assertFalse(TState(self.path).items) + with closing(sqlite3.connect(self.path)) as db: + db.execute('DROP TRIGGER fail_insert') + state.reconcile([], [fill]) + self.assertEqual(TState(self.path).items['600000.SH'].base_qty, 100) + + def test_schema_and_unique_execution(self): + state = TState(self.path) + state.reconcile([], [self.deal('base', 'd1', 100, 10)]) + with closing(sqlite3.connect(self.path)) as db: + tables = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='table'")} + self.assertEqual(tables, {'positions', 'deals', 'sqlite_sequence'}) + columns = {row[1] for row in db.execute('PRAGMA table_info(deals)')} + self.assertFalse({'kind', 'confirmed_date', 'deal_ids', 'deal_id', 'order_id', 'qty', 'filled_qty', 'filled_cost', 'amount', 'trade_date', 'trade_time'} & columns) + self.assertTrue({'sys_order_id', 'local_order_id'} <= columns) + indexes = {row[0] for row in db.execute("SELECT name FROM sqlite_master WHERE type='index'")} + self.assertTrue({'idx_positions_code', 'idx_positions_base_order_id', + 'idx_positions_added_order_id', 'idx_deals_local_order_id', + 'idx_deals_code_date', 'idx_deals_date_time'} <= indexes) + self.assertNotIn('kind', state.deals[0]) + self.assertEqual(TState(self.path).deals, state.deals) + state.deals.append(dict(state.deals[0])) + with self.assertRaises(sqlite3.IntegrityError): + state.save() + self.assertEqual(len(state.deals), 1) + state.items['600000.SH'].base_cost = float('inf') + with self.assertRaises(ValueError): + state.save() + self.assertEqual(state.items['600000.SH'].base_cost, 10) + + def test_position_columns_defaults_indexes_and_stable_id(self): + store = OrderBook(self.path) + with closing(sqlite3.connect(self.path)) as db: + columns = [row[1] for row in db.execute('PRAGMA table_info(positions)')] + self.assertEqual(columns, [ + 'id', 'code', 'base_order_id', 'base_qty', 'base_cost', + 'added_order_id', 'added_num', 'added_qty', 'added_cost', 'status', + ]) + indexes = {row[1] for row in db.execute('PRAGMA index_list(positions)')} + self.assertEqual(indexes, { + 'idx_positions_code', 'idx_positions_base_order_id', 'idx_positions_added_order_id', + }) + store.save({'600000.SH': {'code': '600000.SH', 'status': READY}}, []) + position = store.positions['600000.SH'] + first_id = position['id'] + self.assertGreater(first_id, 0) + self.assertEqual(position['base_order_id'], '') + self.assertEqual(position['base_qty'], 0) + self.assertEqual(position['added_cost'], 0.0) + position.update(base_order_id='base1', base_qty=200, base_cost=10.5, + added_order_id='add1', added_num=1, added_qty=100, added_cost=9.0, + status='ACTIVE') + store.save({'600000.SH': position}, []) + self.assertEqual(store.positions['600000.SH'], position) + store.save({}, []) + store.save({'600001.SH': {'code': '600001.SH', 'status': READY}}, []) + self.assertGreater(store.positions['600001.SH']['id'], first_id) + + def test_base_split_fills_and_snapshot_do_not_double_count(self): + state = TState(self.path) + first = self.deal('base', 'd1', 40, 10) + state.reconcile([PositionItem(stock_code='600000.SH', volume=40, open_price=10)], [first]) + second = self.deal('base', 'd2', 60, 12) + state.reconcile([PositionItem(stock_code='600000.SH', volume=100, open_price=11.2)], [first, second]) + self.assertEqual(state.items['600000.SH'].base_qty, 100) + self.assertAlmostEqual(state.items['600000.SH'].base_cost, 11.2) + self.assertEqual(len(state.deals), 2) + + def test_date_normalization_and_unrelated_strategy(self): + state = TState(self.path) + first = self.deal('base', 'd1', 100, 10, '20260901') + other = self.deal('base', 'd2', 100, 10) + other.local_order_id = 'trend-base-order' + state.reconcile([], [first, other]) + first.insert_date = '2026-09-01' + state.reconcile([], [first]) + self.assertEqual(len(state.deals), 1) + self.assertEqual(state.deals[0]['insert_date'], '2026-09-01') + + +if __name__ == '__main__': + unittest.main()