"""SQLite 策略状态与成交存储;每个数据库仅使用一个写入者,不做数据迁移。""" import math import logging as log import sqlite3 from contextlib import closing from dataclasses import asdict, dataclass from datetime import datetime from pathlib import Path from sdk import DealItem, PositionItem 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 = '' # 补仓创建时间 class State: """保存策略状态和只追加的成交记录,仅创建新表,不迁移旧数据。""" 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.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') state = {row['stock_code']: dict(row) for row in db.execute('SELECT * FROM state')} deals = {row['order_sys_id']: dict(row) for row in db.execute('SELECT * FROM deals ORDER BY id')} self.state = state self.deals = deals self.deals_sys_ids = set(deals) def sync_deals(self, deals: list[DealItem], *, archived: bool = False) -> None: """按成交编号去重;初始化底仓时,已包含在快照内的成交可直接标记归档。""" new_deals = {} for deal in deals: if deal.order_sys_id not in self.deals_sys_ids and deal.order_sys_id not in new_deals: new_deals[deal.order_sys_id] = deal if not new_deals: return with closing(self._connect()) as db, db: for deal in new_deals.values(): order_id = deal.get_local_order_id if not order_id: raise ValueError('Local order ID is required') amount = deal.trade_amount if deal.trade_amount > 0 else deal.price * deal.volume if not math.isfinite(amount) or amount <= 0: raise ValueError('Trade amount must be positive and finite') date = deal.trade_date or datetime.now().date().isoformat() if len(date) == 8 and date.isdigit(): date = f'{date[:4]}-{date[4:6]}-{date[6:]}' # 直接读取模型字段,金额和日期的补全不修改传入模型。 db.execute( '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 (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)', (deal.stock_code, deal.order_sys_id, order_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, int(archived)), ) self.load() def archiving(self, *, base_order_prefix: str = '') -> dict[str, str]: """将未归档成交累加到当前底仓和加仓;失败证券保留记录供重试。""" errors = {} with closing(self._connect()) as db, db: db.execute('BEGIN IMMEDIATE') codes = db.execute( 'SELECT DISTINCT stock_code FROM deals WHERE is_arch = 0 AND offset_flag IN (23, 24, 48, 49)' ).fetchall() for entry in codes: code = entry['stock_code'] db.execute('SAVEPOINT archive_stock') try: current = db.execute('SELECT * FROM state WHERE stock_code = ?', (code,)).fetchone() state = dict(current) if current else asdict(StateItem(stock_code=code)) deals = db.execute( 'SELECT * FROM deals WHERE stock_code = ? AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49) ' "ORDER BY trade_date, REPLACE(trade_time, ':', ''), id", (code,) ).fetchall() for deal in deals: qty = deal['volume'] if deal['offset_flag'] in (23, 48): bucket = 'base' if base_order_prefix and deal['order_local_id'].startswith(base_order_prefix) else 'added' total = state[f'{bucket}_qty'] + qty 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() else: 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'] + state['added_qty'] == 0: state = asdict(StateItem(stock_code=code)) if state['base_qty'] + state['added_qty'] == 0: db.execute('DELETE FROM state WHERE stock_code = ?', (code,)) else: # 更新持仓,保留策略状态及已有记录主键。 state['status'] = current['status'] if current else state['status'] state.pop('id', None) columns = tuple(state) db.execute( f"INSERT INTO state ({', '.join(columns)}) " f"VALUES ({', '.join(':' + key for key in columns)}) " 'ON CONFLICT(stock_code) DO UPDATE SET ' + ', '.join(f'{key} = excluded.{key}' for key in columns if key != 'stock_code'), state, ) db.execute( 'UPDATE deals SET is_arch = 1 WHERE stock_code = ? ' 'AND is_arch = 0 AND offset_flag IN (23, 24, 48, 49)', (code,) ) except (ValueError, sqlite3.IntegrityError) as exc: db.execute('ROLLBACK TO archive_stock') errors[code] = str(exc) log.warning('[归档] %s 失败,保留未归档成交:%s', code, exc) finally: db.execute('RELEASE archive_stock') self.load() return errors def sync_state(self, positions: list[PositionItem], *, remove_missing: bool = True) -> None: """同步完整持仓:无状态则插入底仓,已有则保留,清仓则删除。 数量为零或未出现在完整持仓列表中的证券视为已清仓;空列表清空状态。 底仓已包含的历史成交不应再次归档;后续成交须先归档,再同步持仓。 remove_missing=False 时仅接纳新底仓,减仓由成交归档处理。 """ # 传入完整账户持仓;同步时间作为新增底仓的创建时间。 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') existing = {row['stock_code'] for row in db.execute('SELECT stock_code FROM state')} if remove_missing: db.executemany( 'DELETE FROM state WHERE stock_code = ?', [(code,) for code in existing if code not in holdings], ) 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') # 只插入底仓字段,补仓字段使用数据库默认值。 db.execute( 'INSERT INTO state ' '(stock_code, status, base_order_local_id, base_qty, base_price, base_created_at) ' "VALUES (?, '', '', ?, ?, ?)", (code, item.volume, item.open_price, created_at), ) self.load()