Files
big-qmt/py-client/libs/state.py
2026-09-12 13:42:23 +08:00

348 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""SQLite 策略状态与成交存储;每个数据库仅使用一个写入者,不做数据迁移。"""
import math
import logging as log
import sqlite3
from contextlib import closing
from dataclasses import asdict, dataclass, fields
from datetime import datetime
from pathlib import Path
from sdk import DealItem, PositionItem
FLAG_BUY = 48
FLAG_SELL = 49
UNATTRIBUTED_PREFIX = '__unattributed__:'
SCHEMA = """
-- 策略状态base_ 表示底仓added_ 表示补仓。
CREATE TABLE IF NOT EXISTS state (
id INTEGER PRIMARY KEY AUTOINCREMENT, -- 状态记录主键
stock_code TEXT NOT NULL, -- 证券代码
status TEXT NOT NULL DEFAULT '', -- 策略状态,由策略定义取值
base_order_local_id TEXT NOT NULL DEFAULT '', -- 底仓本地委托编号
base_qty INTEGER NOT NULL DEFAULT 0 CHECK (base_qty >= 0), -- 底仓数量
base_price REAL NOT NULL DEFAULT 0, -- 底仓价格
base_created_at TEXT NOT NULL DEFAULT '', -- 底仓创建时间
added_order_local_id TEXT NOT NULL DEFAULT '', -- 补仓本地委托编号
added_qty INTEGER NOT NULL DEFAULT 0 CHECK (added_qty >= 0), -- 补仓数量
added_price REAL NOT NULL DEFAULT 0, -- 补仓价格
added_created_at TEXT NOT NULL DEFAULT '' -- 补仓创建时间
);
-- 每个证券仅保留一条策略状态。
CREATE UNIQUE INDEX IF NOT EXISTS idx_state_stock_code ON state (stock_code);
-- 成交记录独立保存,不随状态删除。
CREATE TABLE IF NOT EXISTS deals (
id INTEGER PRIMARY KEY AUTOINCREMENT,
stock_code TEXT NOT NULL,
order_sys_id TEXT NOT NULL CHECK (order_sys_id <> ''),
order_local_id TEXT NOT NULL CHECK (order_local_id <> ''),
ref INTEGER NOT NULL DEFAULT 0,
order_ref TEXT NOT NULL DEFAULT '',
direction INTEGER NOT NULL DEFAULT 0,
offset_flag INTEGER NOT NULL CHECK (offset_flag IN (23, 24, 48, 49)),
price REAL NOT NULL CHECK (price >= 0),
volume INTEGER NOT NULL CHECK (volume > 0),
trade_amount REAL NOT NULL CHECK (trade_amount > 0),
trade_date TEXT NOT NULL,
trade_time TEXT NOT NULL,
remark TEXT NOT NULL DEFAULT '',
close_profit REAL NOT NULL DEFAULT 0,
is_arch INTEGER DEFAULT 0
);
CREATE UNIQUE INDEX IF NOT EXISTS idx_deals_order_sys_id ON deals (order_sys_id);
CREATE INDEX IF NOT EXISTS idx_deals_order_ref ON deals (order_local_id);
CREATE INDEX IF NOT EXISTS idx_deals_stock_code_date ON deals (stock_code);
CREATE INDEX IF NOT EXISTS idx_deals_date_time ON deals (trade_date);
"""
@dataclass(slots=True)
class StateItem:
"""策略状态字段;同步账户底仓时无法获知的委托编号留空。"""
stock_code: str = '' # 证券代码
status: str = '' # 策略状态
base_order_local_id: str = '' # 底仓本地委托编号
base_qty: int = 0 # 底仓数量
base_price: float = 0.0 # 底仓价格
base_created_at: str = '' # 底仓创建时间
added_order_local_id: str = '' # 补仓本地委托编号
added_qty: int = 0 # 补仓数量
added_price: float = 0.0 # 补仓价格
added_created_at: str = '' # 补仓创建时间
_STATE_COLUMNS = tuple(field.name for field in fields(StateItem))
_UPSERT_STATE = (
f"INSERT INTO state ({', '.join(_STATE_COLUMNS)}) "
f"VALUES ({', '.join(':' + key for key in _STATE_COLUMNS)}) "
"ON CONFLICT(stock_code) DO UPDATE SET "
+ ', '.join(f'{key} = excluded.{key}' for key in _STATE_COLUMNS if key != 'stock_code')
)
class State:
"""单写入者使用的 SQLite 存储;公开缓存仅在事务成功后替换。
sync_state 同步持仓基准sync_deals 保存成交archiving 记入增量。
同证券未归档成交全部为买入且数量等于当前总持仓时,视为已计入
快照,只标记归档;其他成交作为增量处理。本类不做表结构迁移。
"""
def __init__(self, path: str | Path) -> None:
self.path = Path(path)
self.state: dict[str, dict] = {}
self.deals: dict[str, dict] = {}
self.deals_sys_ids: set[str] = set()
self.blocked_codes: set[str] = set()
self.path.parent.mkdir(parents=True, exist_ok=True)
with closing(self._connect()) as db:
db.executescript(SCHEMA)
self.load()
def _connect(self) -> sqlite3.Connection:
db = sqlite3.connect(self.path, timeout=30)
db.row_factory = sqlite3.Row
return db
@staticmethod
def _read_state(db: sqlite3.Connection) -> dict[str, dict]:
return {row['stock_code']: dict(row) for row in db.execute('SELECT * FROM state')}
@staticmethod
def _read_deals(db: sqlite3.Connection) -> dict[str, dict]:
return {
row['order_sys_id']: dict(row)
for row in db.execute('SELECT * FROM deals ORDER BY id')
}
def get_by_code(self, code: str) -> dict:
"""返回缓存中的状态;不存在时返回空字典。"""
return self.state.get(code, {})
def load(self) -> None:
"""在同一个读事务内加载两张表,全部成功后再发布缓存。"""
with closing(self._connect()) as db, db:
db.execute('BEGIN')
state = self._read_state(db)
deals = self._read_deals(db)
self.state, self.deals, self.deals_sys_ids = state, deals, set(deals)
def load_state(self) -> None:
"""只刷新状态缓存。"""
with closing(self._connect()) as db, db:
db.execute('BEGIN')
state = self._read_state(db)
self.state = state
def load_deals(self) -> None:
"""只刷新成交缓存及其去重编号集合。"""
with closing(self._connect()) as db, db:
db.execute('BEGIN')
deals = self._read_deals(db)
self.deals, self.deals_sys_ids = deals, set(deals)
@staticmethod
def _insert_deals(db: sqlite3.Connection, deals: list[DealItem]) -> None:
"""共享事务内保存成交;空本地编号使用明确的待核对标记。"""
existing = {row[0] for row in db.execute('SELECT order_sys_id FROM deals')}
new_deals: dict[str, DealItem] = {}
for deal in deals:
if deal.order_sys_id not in existing:
new_deals.setdefault(deal.order_sys_id, deal)
if not new_deals:
return
today = datetime.now().date().isoformat()
values = []
for deal in new_deals.values():
amount = deal.trade_amount if deal.trade_amount > 0 else deal.price * deal.volume
if not math.isfinite(deal.price) or not math.isfinite(amount):
raise ValueError('Trade price and amount must be finite')
date = deal.trade_date or today
if len(date) == 8 and date.isdigit():
date = f'{date[:4]}-{date[4:6]}-{date[6:]}'
values.append((
deal.stock_code, deal.order_sys_id,
deal.get_local_order_id.strip() or f'{UNATTRIBUTED_PREFIX}{deal.order_sys_id}',
deal.ref, deal.order_ref, deal.direction, deal.offset_flag,
deal.price, deal.volume, amount, date, deal.trade_time,
deal.remark, deal.close_profit,
))
# 唯一键负责最终去重,避免旧缓存造成重复插入失败。
db.executemany("""
INSERT INTO deals (
stock_code, order_sys_id, order_local_id, ref, order_ref,
direction, offset_flag, price, volume, trade_amount,
trade_date, trade_time, remark, close_profit, is_arch
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 0)
ON CONFLICT(order_sys_id) DO NOTHING
""", values)
def sync_deals(self, deals: list[DealItem]) -> None:
"""按成交编号只追加;同批保留首笔,任一无效成交回滚整批。"""
if all(deal.order_sys_id in self.deals_sys_ids for deal in deals):
return
with closing(self._connect()) as db, db:
db.execute('BEGIN IMMEDIATE')
self._insert_deals(db, deals)
cached_deals = self._read_deals(db)
self.deals, self.deals_sys_ids = cached_deals, set(cached_deals)
def sync_state(self, positions: list[PositionItem]) -> None:
"""同步完整持仓:新增底仓、保留已有状态、删除已清仓证券。
此接口不会标记成交。已有库存对应的卖出须先归档,再传入
清仓快照,以免删除归档所需的库存。
"""
created_at = datetime.now().isoformat(timespec='seconds')
holdings = {item.stock_code: item for item in positions if item.volume > 0}
with closing(self._connect()) as db, db:
db.execute('BEGIN IMMEDIATE')
existing = {row['stock_code'] for row in db.execute('SELECT stock_code FROM state')}
values = []
for code, item in holdings.items():
if code in existing:
continue
if not math.isfinite(item.open_price):
raise ValueError('Base price must be finite')
values.append((code, item.volume, item.open_price, created_at))
db.executemany(
'DELETE FROM state WHERE stock_code = ?',
[(code,) for code in existing if code not in holdings],
)
db.executemany(
'INSERT INTO state (stock_code, base_qty, base_price, base_created_at) '
'VALUES (?, ?, ?, ?)', values,
)
state = self._read_state(db)
self.state = state
@staticmethod
def _archive_stock(db: sqlite3.Connection, code: str, snapshot_dedup: bool = True) -> None:
"""在调用方的保存点内完成一只证券的记账及成交标记。"""
current = db.execute('SELECT * FROM state WHERE stock_code = ?', (code,)).fetchone()
state = dict(current) if current else asdict(StateItem(stock_code=code))
pending = db.execute("""
SELECT * FROM deals WHERE stock_code = ? AND is_arch = 0
ORDER BY trade_date, REPLACE(trade_time, ':', ''), id
""", (code,)).fetchall()
if any(deal['order_local_id'].startswith(UNATTRIBUTED_PREFIX) for deal in pending):
raise ValueError('Unattributed trade requires reconciliation')
total_qty = state['base_qty'] + state['added_qty']
if (
snapshot_dedup and total_qty > 0
and all(deal['offset_flag'] == FLAG_BUY for deal in pending)
and sum(deal['volume'] for deal in pending) == total_qty
):
# 按快照去重规则,仅标记这批成交,原仓位数量和成本均保留。
# 卖出和买卖混合批次不能用此规则,否则会漏掉真实减仓。
db.executemany(
'UPDATE deals SET is_arch = 1 WHERE id = ? AND is_arch = 0',
[(deal['id'],) for deal in pending],
)
return
for deal in pending:
if deal['offset_flag'] == FLAG_BUY:
bucket = 'base' if deal['order_local_id'].startswith('zt-base-') else 'added'
total = state[f'{bucket}_qty'] + deal['volume']
state[f'{bucket}_price'] = (
state[f'{bucket}_qty'] * state[f'{bucket}_price'] + deal['trade_amount']
) / total
state[f'{bucket}_qty'] = total
state[f'{bucket}_order_local_id'] = deal['order_local_id']
state[f'{bucket}_created_at'] = f"{deal['trade_date']} {deal['trade_time']}".strip()
elif deal['offset_flag'] == FLAG_SELL:
qty = deal['volume']
total = state['base_qty'] + state['added_qty']
if qty > total:
raise ValueError(f'Sell volume {qty} exceeds recorded holdings {total}')
if qty < state['added_qty']:
state['added_qty'] -= qty
else:
state['base_qty'] = total - qty
state['added_qty'] = 0
state['added_price'] = 0.0
state['added_order_local_id'] = state['added_created_at'] = ''
if state['base_qty'] == 0:
state.update(asdict(StateItem(stock_code=code)))
else:
raise ValueError(f"Unsupported trade flag: {deal['offset_flag']}")
if state['base_qty'] + state['added_qty'] == 0:
db.execute('DELETE FROM state WHERE stock_code = ?', (code,))
else:
# 保留原记录主键及策略状态,包括同批清仓后重新建仓。
if current:
state['status'] = current['status']
db.execute(_UPSERT_STATE, state)
db.execute(
'UPDATE deals SET is_arch = 1 WHERE stock_code = ? '
'AND is_arch = 0', (code,),
)
def _archive_pending(self, db: sqlite3.Connection, snapshot_dedup: bool = True) -> None:
"""保留现有逐证券保存点ZT 增量路径禁用数量相等推断。"""
codes = db.execute(
'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0'
).fetchall()
for row in codes:
code = row['stock_code']
db.execute('SAVEPOINT archive_stock')
try:
self._archive_stock(db, code, snapshot_dedup)
except (ValueError, sqlite3.IntegrityError) as exc:
db.execute('ROLLBACK TO archive_stock')
log.warning('[归档] %s 失败,保留未归档成交:%s', code, exc)
finally:
db.execute('RELEASE archive_stock')
def archiving(self) -> None:
"""归档未处理成交;单证券失败回滚并保留重试,不影响其他证券。"""
with closing(self._connect()) as db, db:
db.execute('BEGIN IMMEDIATE')
self._archive_pending(db)
state = self._read_state(db)
deals = self._read_deals(db)
self.state, self.deals, self.deals_sys_ids = state, deals, set(deals)
def sync_account(self, positions: list[PositionItem], deals: list[DealItem], *, initialize: bool = False) -> None:
"""ZT 原子同步。初始化使用已核对快照;恢复只应用增量、不覆盖库存。
空库存且无成交历史时允许初始化。快照缺失或
未归档成交只隔离对应证券,后续成交到齐并核对一致后自动恢复。
"""
holdings = {p.stock_code: p for p in positions if p.volume > 0}
with closing(self._connect()) as db, db:
db.execute('BEGIN IMMEDIATE')
if initialize:
if (db.execute('SELECT 1 FROM state LIMIT 1').fetchone()
or db.execute('SELECT 1 FROM deals LIMIT 1').fetchone()):
raise ValueError('ZT state already initialized; cannot overwrite existing data')
created_at = datetime.now().isoformat(timespec='seconds')
for code, position in holdings.items():
if not code or not math.isfinite(position.open_price) or position.open_price <= 0:
raise ValueError('Initial position code and cost must be valid')
db.execute(_UPSERT_STATE, asdict(StateItem(
stock_code=code, base_qty=position.volume,
base_price=position.open_price, base_created_at=created_at,
)))
self._insert_deals(db, deals)
# 这些成交已包含在已核对的基准中,不再增减快照库存。
db.execute('UPDATE deals SET is_arch = 1')
else:
self._insert_deals(db, deals)
self._archive_pending(db, snapshot_dedup=False)
state = self._read_state(db)
cached_deals = self._read_deals(db)
blocked = {d['stock_code'] for d in cached_deals.values() if d['is_arch'] != 1}
for code in set(state) | set(holdings):
row = state.get(code, {})
recorded = row.get('base_qty', 0) + row.get('added_qty', 0)
actual = holdings[code].volume if code in holdings else 0
if recorded != actual:
blocked.add(code)
self.state, self.deals, self.deals_sys_ids = state, cached_deals, set(cached_deals)
self.blocked_codes = blocked
if blocked:
log.warning('[ZT 同步] 以下证券状态待核对,暂停交易:%s', ', '.join(sorted(blocked)))