This commit is contained in:
2026-09-07 00:27:33 +08:00
parent fdcbdc7869
commit d37f9edefc
11 changed files with 839 additions and 384 deletions

View File

@@ -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 买回不计为补仓。

170
py-client/libs/orderbook.py Normal file
View File

@@ -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()

View File

@@ -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")),
)

View File

@@ -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()

View File

@@ -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),

View File

@@ -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)

View File

@@ -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

View File

@@ -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

View File

@@ -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

View File

@@ -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()

View File

@@ -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()