393 lines
19 KiB
Python
393 lines
19 KiB
Python
"""SQLite 策略状态与成交存储;每个数据库仅使用一个写入者,不做数据迁移。"""
|
||
|
||
import math
|
||
import json
|
||
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],
|
||
existing_ids: set[str] | None = None) -> None:
|
||
"""共享事务内保存成交;空本地编号使用明确的待核对标记。"""
|
||
existing = existing_ids if existing_ids is not None else {
|
||
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,
|
||
blocked_codes: set[str] | None = None) -> None:
|
||
"""保留现有逐证券保存点;ZT 增量路径禁用数量相等推断。"""
|
||
codes = db.execute(
|
||
'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0'
|
||
).fetchall()
|
||
for row in codes:
|
||
code = row['stock_code']
|
||
if blocked_codes and code in blocked_codes:
|
||
continue
|
||
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')
|
||
db.execute('''CREATE TABLE IF NOT EXISTS zt_rejected_deals (
|
||
stock_code TEXT NOT NULL, order_sys_id TEXT NOT NULL,
|
||
payload TEXT NOT NULL, error TEXT NOT NULL,
|
||
PRIMARY KEY (stock_code, order_sys_id)
|
||
)''')
|
||
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_account_deals(db, deals)
|
||
rejected = {row[0] for row in db.execute('SELECT stock_code FROM zt_rejected_deals')}
|
||
self._archive_pending(db, snapshot_dedup=False, blocked_codes=rejected)
|
||
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}
|
||
blocked.update(row[0] for row in db.execute('SELECT stock_code FROM zt_rejected_deals'))
|
||
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)))
|
||
|
||
def _insert_account_deals(self, db: sqlite3.Connection, deals: list[DealItem]) -> None:
|
||
"""仅 ZT 增量路径隔离坏成交;保留原始字段,缺席回报不会解除隔离。"""
|
||
existing = {row[0] for row in db.execute('SELECT order_sys_id FROM deals')}
|
||
for deal in deals:
|
||
if not deal.stock_code:
|
||
raise ValueError('Trade without stock code cannot be isolated')
|
||
db.execute('SAVEPOINT insert_zt_deal')
|
||
try:
|
||
if not deal.stock_code or not deal.order_sys_id:
|
||
raise ValueError('Trade code and order ID are required')
|
||
if (deal.offset_flag not in (FLAG_BUY, FLAG_SELL)
|
||
or not math.isfinite(deal.volume) or deal.volume <= 0
|
||
or int(deal.volume) != deal.volume
|
||
or not math.isfinite(deal.price) or deal.price < 0
|
||
or not math.isfinite(deal.trade_amount)
|
||
or not math.isfinite(deal.close_profit)):
|
||
raise ValueError('Invalid trade direction, volume or amount')
|
||
self._insert_deals(db, [deal], existing)
|
||
except (ValueError, TypeError, OverflowError, sqlite3.IntegrityError) as exc:
|
||
db.execute('ROLLBACK TO insert_zt_deal')
|
||
db.execute('''INSERT OR REPLACE INTO zt_rejected_deals
|
||
(stock_code, order_sys_id, payload, error) VALUES (?, ?, ?, ?)''',
|
||
(deal.stock_code, deal.order_sys_id, json.dumps(asdict(deal), ensure_ascii=False), str(exc)))
|
||
log.warning('[ZT 成交] 隔离证券=%s,委托=%s,原因=%s',
|
||
deal.stock_code, deal.order_sys_id, exc)
|
||
else:
|
||
existing.add(deal.order_sys_id)
|
||
db.execute('DELETE FROM zt_rejected_deals WHERE stock_code = ? AND order_sys_id = ?',
|
||
(deal.stock_code, deal.order_sys_id))
|
||
finally:
|
||
db.execute('RELEASE insert_zt_deal')
|