Files

301 lines
10 KiB
Python
Raw Permalink 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 持久化(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 | runningAPI 改字段,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",
]