304 lines
14 KiB
Python
304 lines
14 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 (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(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)
|