301 lines
10 KiB
Python
301 lines
10 KiB
Python
"""实盘模拟 SQLite 持久化(4 表,spec §live-api)。
|
||
|
||
表:live_accounts / live_trades / live_positions / live_balance
|
||
WAL 模式支持 supervisor 进程写 + API 进程读(DB 解耦,spec §live-api)。
|
||
|
||
设计参考 ``sanguo_trader/persistence.py``(paper_* 表),但:
|
||
- account.setting 存 JSON 字符串(策略参数透传给 CtaTemplate.update_setting)
|
||
- status: stopped | running(API 改字段,supervisor 轮询该字段决定起停)
|
||
- positions 为覆盖式快照(supervisor 定时把 OMS PositionData 落库,不做增量)
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import sqlite3
|
||
from datetime import datetime, timezone
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
SCHEMA = """
|
||
CREATE TABLE IF NOT EXISTS live_accounts (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
name TEXT,
|
||
account TEXT,
|
||
vt_symbol TEXT,
|
||
strategy_class TEXT,
|
||
strategy_name TEXT,
|
||
setting TEXT,
|
||
status TEXT,
|
||
interval TEXT,
|
||
initial_capital REAL,
|
||
connect_wait_sec INTEGER,
|
||
init_wait_sec INTEGER,
|
||
mini_path TEXT,
|
||
error_msg TEXT,
|
||
created_at TEXT,
|
||
updated_at TEXT,
|
||
strategy_type TEXT DEFAULT 'cta',
|
||
pool TEXT,
|
||
max_pool INTEGER,
|
||
benchmark TEXT
|
||
);
|
||
CREATE TABLE IF NOT EXISTS live_trades (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
account_id INTEGER,
|
||
strategy_name TEXT,
|
||
symbol TEXT,
|
||
direction TEXT,
|
||
offset TEXT,
|
||
price REAL,
|
||
volume REAL,
|
||
traded_at TEXT,
|
||
vt_tradeid TEXT
|
||
);
|
||
CREATE TABLE IF NOT EXISTS live_positions (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
account_id INTEGER,
|
||
symbol TEXT,
|
||
volume REAL,
|
||
frozen REAL,
|
||
avg_price REAL,
|
||
updated_at TEXT,
|
||
UNIQUE(account_id, symbol)
|
||
);
|
||
CREATE TABLE IF NOT EXISTS live_balance (
|
||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||
account_id INTEGER,
|
||
date TEXT,
|
||
cash REAL,
|
||
market_value REAL,
|
||
total REAL
|
||
);
|
||
CREATE INDEX IF NOT EXISTS idx_live_trades_account ON live_trades(account_id);
|
||
CREATE INDEX IF NOT EXISTS idx_live_balance_account ON live_balance(account_id, date);
|
||
"""
|
||
|
||
|
||
def _now() -> str:
|
||
return datetime.now(timezone.utc).isoformat()
|
||
|
||
|
||
def init_db(db_path: str) -> None:
|
||
Path(db_path).parent.mkdir(parents=True, exist_ok=True)
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.executescript(SCHEMA)
|
||
# 轻量迁移:旧库补组合实盘列(新库 CREATE 已含,ALTER 报错忽略)
|
||
for col, ddl in (
|
||
("strategy_type", "TEXT DEFAULT 'cta'"),
|
||
("pool", "TEXT"),
|
||
("max_pool", "INTEGER"),
|
||
("benchmark", "TEXT"),
|
||
):
|
||
try:
|
||
conn.execute(f"ALTER TABLE live_accounts ADD COLUMN {col} {ddl}")
|
||
except sqlite3.OperationalError:
|
||
pass # 列已存在
|
||
conn.execute("PRAGMA journal_mode=WAL")
|
||
conn.commit()
|
||
|
||
|
||
# ----------------- live_accounts CRUD -----------------
|
||
|
||
def save_account(db_path: str, account: dict[str, Any]) -> int:
|
||
with sqlite3.connect(db_path) as conn:
|
||
cur = conn.execute(
|
||
"""INSERT INTO live_accounts
|
||
(name, account, vt_symbol, strategy_class, strategy_name, setting,
|
||
status, interval, initial_capital, connect_wait_sec, init_wait_sec,
|
||
mini_path, created_at, updated_at,
|
||
strategy_type, pool, max_pool, benchmark)
|
||
VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)""",
|
||
(
|
||
account.get("name", "live"),
|
||
account.get("account", ""),
|
||
account.get("vt_symbol", ""),
|
||
account.get("strategy_class", "AShareDoubleMaStrategy"),
|
||
account.get("strategy_name", ""),
|
||
json.dumps(account.get("setting", {})),
|
||
account.get("status", "stopped"),
|
||
account.get("interval", "15m"),
|
||
account.get("initial_capital", 1_000_000),
|
||
int(account.get("connect_wait_sec", 10)),
|
||
int(account.get("init_wait_sec", 60)),
|
||
account.get("mini_path", ""),
|
||
_now(), _now(),
|
||
account.get("strategy_type", "cta"),
|
||
account.get("pool", ""),
|
||
int(account.get("max_pool", 0) or 0),
|
||
account.get("benchmark", ""),
|
||
),
|
||
)
|
||
conn.commit()
|
||
return cur.lastrowid
|
||
|
||
|
||
def list_accounts(db_path: str) -> list[dict]:
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.row_factory = sqlite3.Row
|
||
cur = conn.execute("SELECT * FROM live_accounts ORDER BY id DESC")
|
||
return [dict(r) for r in cur.fetchall()]
|
||
|
||
|
||
def get_account(db_path: str, account_id: int) -> dict | None:
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.row_factory = sqlite3.Row
|
||
cur = conn.execute(
|
||
"SELECT * FROM live_accounts WHERE id=?", (account_id,)
|
||
)
|
||
row = cur.fetchone()
|
||
return dict(row) if row else None
|
||
|
||
|
||
def update_account_status(
|
||
db_path: str, account_id: int, status: str, error_msg: str = ""
|
||
) -> None:
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.execute(
|
||
"UPDATE live_accounts SET status=?, error_msg=?, updated_at=? WHERE id=?",
|
||
(status, error_msg, _now(), account_id),
|
||
)
|
||
conn.commit()
|
||
|
||
|
||
def list_running_accounts(db_path: str) -> list[dict]:
|
||
"""supervisor 轮询:取所有 status=running 的实例。"""
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.row_factory = sqlite3.Row
|
||
cur = conn.execute(
|
||
"SELECT * FROM live_accounts WHERE status=? ORDER BY id", ("running",)
|
||
)
|
||
return [dict(r) for r in cur.fetchall()]
|
||
|
||
|
||
# ----------------- live_trades -----------------
|
||
|
||
def save_trade(db_path: str, account_id: int, trade: dict[str, Any]) -> int:
|
||
with sqlite3.connect(db_path) as conn:
|
||
cur = conn.execute(
|
||
"""INSERT INTO live_trades
|
||
(account_id, strategy_name, symbol, direction, offset,
|
||
price, volume, traded_at, vt_tradeid)
|
||
VALUES (?,?,?,?,?,?,?,?,?)""",
|
||
(
|
||
account_id, trade.get("strategy_name", ""),
|
||
trade.get("symbol", ""), trade.get("direction", ""),
|
||
trade.get("offset", ""), trade.get("price", 0),
|
||
trade.get("volume", 0), trade.get("traded_at", ""),
|
||
trade.get("vt_tradeid", ""),
|
||
),
|
||
)
|
||
conn.commit()
|
||
return cur.lastrowid
|
||
|
||
|
||
def list_trades(db_path: str, account_id: int) -> list[dict]:
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.row_factory = sqlite3.Row
|
||
cur = conn.execute(
|
||
"SELECT * FROM live_trades WHERE account_id=? ORDER BY id",
|
||
(account_id,),
|
||
)
|
||
return [dict(r) for r in cur.fetchall()]
|
||
|
||
|
||
# ----------------- live_positions (覆盖式快照) -----------------
|
||
|
||
def save_positions(
|
||
db_path: str, account_id: int, positions: dict[str, dict]
|
||
) -> None:
|
||
"""覆盖式落库。positions = {symbol: {volume, frozen, avg_price}}。
|
||
|
||
supervisor 每 snapshot_interval_sec 调一次,把 OMS 最新 PositionData 覆盖落库。
|
||
只保留 volume>0 的持仓。
|
||
"""
|
||
now = _now()
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.execute("DELETE FROM live_positions WHERE account_id=?", (account_id,))
|
||
conn.executemany(
|
||
"""INSERT INTO live_positions
|
||
(account_id, symbol, volume, frozen, avg_price, updated_at)
|
||
VALUES (?,?,?,?,?,?)""",
|
||
[
|
||
(account_id, sym, p["volume"], p.get("frozen", 0),
|
||
p["avg_price"], now)
|
||
for sym, p in positions.items() if p.get("volume", 0) > 0
|
||
],
|
||
)
|
||
conn.commit()
|
||
|
||
|
||
def load_positions(db_path: str, account_id: int) -> list[dict]:
|
||
"""API 读持仓快照 → [{symbol, volume, frozen, avg_price, updated_at}]。"""
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.row_factory = sqlite3.Row
|
||
cur = conn.execute(
|
||
"SELECT symbol, volume, frozen, avg_price, updated_at "
|
||
"FROM live_positions WHERE account_id=?",
|
||
(account_id,),
|
||
)
|
||
return [dict(r) for r in cur.fetchall()]
|
||
|
||
|
||
# ----------------- live_balance -----------------
|
||
|
||
def save_balance(
|
||
db_path: str, account_id: int, date: str, cash: float,
|
||
market_value: float, total: float
|
||
) -> None:
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.execute(
|
||
"""INSERT INTO live_balance
|
||
(account_id, date, cash, market_value, total)
|
||
VALUES (?,?,?,?,?)""",
|
||
(account_id, date, cash, market_value, total),
|
||
)
|
||
conn.commit()
|
||
|
||
|
||
def list_balance(db_path: str, account_id: int) -> list[dict]:
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.row_factory = sqlite3.Row
|
||
cur = conn.execute(
|
||
"SELECT * FROM live_balance WHERE account_id=? ORDER BY date, id",
|
||
(account_id,),
|
||
)
|
||
return [dict(r) for r in cur.fetchall()]
|
||
|
||
|
||
def get_last_balance(db_path: str, account_id: int) -> dict | None:
|
||
"""最新一条账户快照(API /live/{aid}/account)。"""
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.row_factory = sqlite3.Row
|
||
cur = conn.execute(
|
||
"SELECT account_id, date, cash, market_value, total "
|
||
"FROM live_balance WHERE account_id=? ORDER BY id DESC LIMIT 1",
|
||
(account_id,),
|
||
)
|
||
row = cur.fetchone()
|
||
return dict(row) if row else None
|
||
|
||
|
||
def get_first_balance(db_path: str, account_id: int) -> dict | None:
|
||
"""最早一条账户快照(收益率基线;首快照 total = baseline)。"""
|
||
with sqlite3.connect(db_path) as conn:
|
||
conn.row_factory = sqlite3.Row
|
||
cur = conn.execute(
|
||
"SELECT account_id, date, cash, market_value, total "
|
||
"FROM live_balance WHERE account_id=? ORDER BY id ASC LIMIT 1",
|
||
(account_id,),
|
||
)
|
||
row = cur.fetchone()
|
||
return dict(row) if row else None
|
||
|
||
|
||
__all__ = [
|
||
"init_db", "save_account", "list_accounts", "get_account",
|
||
"update_account_status", "list_running_accounts",
|
||
"save_trade", "list_trades",
|
||
"save_positions", "load_positions",
|
||
"save_balance", "list_balance", "get_last_balance", "get_first_balance",
|
||
]
|