Files
big-qmt/py-client/libs/state.py
2026-09-14 16:28:56 +08:00

304 lines
14 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 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 (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);
"""
DELAS_INSERT_SQL = """
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)
"""
@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)
def sync_deals(self, deals: list[DealItem]) -> None:
"""按 order_sys_id 只追加成交时间超过 30 秒的新记录,保留全部历史记录。"""
now = datetime.now()
new_deals: dict[str, DealItem] = {}
for deal in deals:
if deal.order_sys_id in self.deals_sys_ids:
continue
date = (deal.trade_date or now.date().isoformat()).replace('-', '')
time = deal.trade_time.replace(':', '')
traded_at = datetime.strptime(f'{date} {time}', '%Y%m%d %H%M%S')
if (now - traded_at).total_seconds() <= 30:
continue
new_deals.setdefault(deal.order_sys_id, deal)
if not new_deals:
return
# 插入记录。
with closing(self._connect()) as db, db:
db.execute('BEGIN')
today = 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(),
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(DELAS_INSERT_SQL, values)
cache_deals = self._read_deals(db)
self.deals, self.deals_sys_ids = cache_deals, set(cache_deals)
def sync_state(self, positions: list[PositionItem]) -> None:
"""同步完整持仓:新增底仓、保留已有状态、删除已清仓证券。
此接口不会标记成交。已有库存对应的卖出须先归档,再传入
清仓快照,以免删除归档所需的库存。
"""
created_at = datetime.now().isoformat(sep=' ', 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
def merge_deals(self) -> dict[str, dict]:
"""按本地委托编号汇总未归档成交,返回新字典,不修改原始记录。
数量、金额和平仓盈亏累加,价格为总金额除以总数量;
其余字段(包括成交编号和时间)保留同组首笔记录的值。
"""
merged: dict[str, dict] = {}
for deal in self.deals.values():
if deal['is_arch'] != 0:
continue
order_local_id = deal['order_local_id']
if order_local_id not in merged:
merged[order_local_id] = deal.copy()
else:
item = merged[order_local_id]
item['stock_code']=deal['stock_code']
item['volume'] += deal['volume']
item['trade_amount'] += deal['trade_amount']
item['close_profit'] += deal['close_profit']
for item in merged.values():
item['price'] = item['trade_amount'] / item['volume']
return merged
def archiving(self) -> None:
"""先合并缓存中的未归档成交再计算;单证券失败回滚并保留重试。"""
merged = self.merge_deals()
if not merged:
return
with closing(self._connect()) as db, db:
db.execute('BEGIN IMMEDIATE')
for order_local_id, deal in merged.items():
db.execute('SAVEPOINT archive_stock')
code = deal['stock_code']
try:
result = db.execute('SELECT * FROM state WHERE stock_code = ?', (code,)).fetchone()
state = dict(result) if result else asdict(StateItem(stock_code=code))
if deal['offset_flag'] == FLAG_BUY:
bucket = 'base' if deal['order_local_id'].startswith('zt-base-') else 'added'
state[f'{bucket}_price'] = deal['price']
state[f'{bucket}_qty'] = deal['volume']
state[f'{bucket}_order_local_id'] = deal['order_local_id']
state[f'{bucket}_created_at'] = f"{deal['trade_date']} {deal['trade_time']}".strip()
db.execute(_UPSERT_STATE, state)
elif deal['offset_flag'] == FLAG_SELL:
newState = StateItem(stock_code=code)
qty = deal['volume']
total = state['base_qty'] + state['added_qty']
if qty > total:
raise ValueError(f'Sell volume {qty} exceeds recorded holdings {total}')
elif qty == total:
# 清仓底仓与补仓
newState.status='CLEAR'
elif qty == state['added_qty']:
# 清仓补仓
newState.base_qty = state['base_qty']
newState.base_price = state['base_price']
newState.base_order_local_id = state['base_order_local_id']
newState.base_created_at = state['base_created_at']
elif qty == state['base_qty']:
# 清仓底仓
newState.status='CLEAR'
else:
raise ValueError(f'Sell volume {qty} holdings {total}')
if newState.status == 'CLEAR':
db.execute('DELETE FROM state WHERE stock_code = ?', (code,))
else:
# 保留原记录主键及策略状态,包括同批清仓后重新建仓。
db.execute(_UPSERT_STATE, asdict(newState))
db.execute('UPDATE deals SET is_arch = 1 WHERE order_local_id = ? ''AND is_arch = 0', (order_local_id,))
except (ValueError, sqlite3.IntegrityError) as exc:
db.execute('ROLLBACK TO archive_stock')
log.warning('[归档] %s 失败,保留未归档成交:%s', code, exc)
finally:
db.execute('RELEASE archive_stock')
state = self._read_state(db)
deals = self._read_deals(db)
self.state, self.deals, self.deals_sys_ids = state, deals, set(deals)