Files
sanguo_vnpy_v2/sanguo_api/routes_paper.py
T

165 lines
6.0 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.
"""模拟盘 API 路由(spec §10)。
create 建 paper_account(持久化配置);GET 查询净值/成交/状态。
回放执行(engine.run)由 orchestrator 异步触发或容器内同步跑,端到端冒烟在容器
(本机无 NAS parquet + vnpy 完整依赖),本模块只做 account 管理 + 查询。
"""
import sqlite3
from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import BaseModel
from .auth import verify_token as verify_token_impl
router = APIRouter()
_db_path = {"path": None}
def set_db_path(p):
_db_path["path"] = p
if p:
from sanguo_trader.persistence import init_db
init_db(p) # app 启动建表(幂等),保证 GET 查询不报 no such table
async def verify_token(authorization: str | None = Header(None)):
if authorization is None or not authorization.startswith("Bearer "):
raise HTTPException(401, "Missing/invalid authorization")
return verify_token_impl(authorization.split(" ", 1)[1])
class StrategyCfg(BaseModel):
name: str
params: dict = {}
match_session: str = "next_open"
symbol: str
listing_days: int = 0
class PaperCreateRequest(BaseModel):
name: str = "paper"
mode: str = "replay"
interval: str = "d"
symbols: list[str]
strategies: list[StrategyCfg]
initial_capital: float = 1_000_000
rate: float = 0.0003
slippage: float = 0.0
pricetick: float = 0.01
stamp_duty_rate: float = 0.0005
transfer_fee_rate: float = 0.00001
min_commission: float = 5.0
start: str
end: str
@router.post("/paper/create", dependencies=[Depends(verify_token)])
def create_paper(req: PaperCreateRequest):
from sanguo_trader.persistence import init_db, save_account
import threading
db = _db_path["path"] or ":memory:"
init_db(db)
aid = save_account(db, req.model_dump())
status = "created"
if req.mode == "replay": # 回放后台线程跑,create 立即返回(避免阻塞 worker 502)
def _bg():
from sanguo_trader.persistence import update_account_status
try:
_run_replay(db, aid, req)
update_account_status(db, aid, "done")
except Exception as e: # noqa: BLE001
update_account_status(db, aid, "failed", str(e))
threading.Thread(target=_bg, daemon=True).start()
status = "running"
return {"account_id": aid, "status": status}
@router.get("/paper/{aid}", dependencies=[Depends(verify_token)])
def get_paper(aid: int):
db = _db_path["path"]
with sqlite3.connect(db) as conn:
conn.row_factory = sqlite3.Row
row = conn.execute(
"SELECT * FROM paper_accounts WHERE id=?", (aid,)
).fetchone()
if not row:
raise HTTPException(404, "account not found")
return dict(row)
@router.get("/paper/{aid}/equity", dependencies=[Depends(verify_token)])
def get_equity(aid: int):
from sanguo_trader.persistence import list_daily_balance
return list_daily_balance(_db_path["path"], aid)
@router.get("/paper/{aid}/trades", dependencies=[Depends(verify_token)])
def get_trades(aid: int):
from sanguo_trader.persistence import list_trades
return list_trades(_db_path["path"], aid)
@router.get("/paper/{aid}/strategies", dependencies=[Depends(verify_token)])
def get_strategies(aid: int):
"""分策略归因:成交/拒单/费用聚合(spec §7)。"""
from sanguo_trader.persistence import list_strategy_summary
return list_strategy_summary(_db_path["path"], aid)
class _DataSourceWrapper:
"""包装 iter_bars 给 PaperEngineengine 需 data_source.iter_bars 接口)。"""
def __init__(self, cfg):
self.cfg = cfg
def iter_bars(self, symbols, start, end, interval, adjust="qfq", cfg=None):
from sanguo_trader.data_source import iter_bars
return iter_bars(symbols, start, end, interval, adjust, cfg or self.cfg)
def _run_replay(db, aid, req: PaperCreateRequest):
"""构造引擎 + 跑回放(容器内有 vnpy_ctastrategy + NAS parquet,本机仅空转)。"""
from sanguo_trader.account import Account
from sanguo_trader.cta_adapter import PaperCtaEngine
from sanguo_trader.engine import PaperEngine
from sanguo_trader.models import AccountConfig
from sanguo_trader.strategy_runner import StrategyRunner
from sanguo_data.config import find_config_path, load_config
from sanguo_data.datareader import guess_exchange
from .strategy_registry import get_strategy_class
data_cfg = load_config(find_config_path())
acc_cfg = AccountConfig(
initial_capital=req.initial_capital, rate=req.rate, slippage=req.slippage,
pricetick=req.pricetick, stamp_duty_rate=req.stamp_duty_rate,
transfer_fee_rate=req.transfer_fee_rate, min_commission=req.min_commission,
)
account = Account(req.initial_capital)
runners: list = []
for s in req.strategies:
cls = get_strategy_class(s.name)
if cls is None:
continue # 策略不可用(本机无 vnpy_ctastrategy)→ 跳过
cta = PaperCtaEngine(s.name, match_session=s.match_session,
listing_days=s.listing_days, size=100) # A 股 1 手=100 股
vt_symbol = f"{s.symbol}.{guess_exchange(s.symbol).value}"
strat = cls(cta, s.name, vt_symbol, s.params) # CtaTemplate(cta_engine, name, vt_symbol, setting)
strat.trading = True # 允许 send_order(等价 on_start
try:
from vnpy.trader.utility import ArrayManager
if not hasattr(strat, "am"):
strat.am = ArrayManager(20) # 默认 100 根才 inited,短区间不够;用 20 兼容
except Exception:
pass
cta.set_strategy(strat)
runners.append(StrategyRunner(s.name, strategy=strat, paper_cta_engine=cta,
symbol=s.symbol))
pe = PaperEngine(account, runners, _DataSourceWrapper(data_cfg), acc_cfg,
db, aid, req.symbols, req.start, req.end, req.interval)
pe.run()