"""中控(Fleet)专用 API:X-Fleet-Token 鉴权,不开放资金/下单。""" from __future__ import annotations import asyncio import hashlib import hmac import logging import os import secrets import subprocess import threading import time from pathlib import Path from typing import Annotated from fastapi import APIRouter, Depends, Header, HTTPException, status from pydantic import BaseModel, Field from ..config import get_settings from ..credentials import get_credentials from ..models.db import get_db from ..strategy import get_engine from .auth import require_user logger = logging.getLogger(__name__) router = APIRouter(prefix="/api/fleet", tags=["fleet"]) _SETTING_HASH = "fleet_api_token_hash" _TICKET_TTL_SEC = 60 _tickets: dict[str, dict] = {} _tickets_lock = threading.Lock() _update_lock = threading.Lock() _update_state: dict = {"running": False, "started_at_ms": 0, "last_error": ""} def _hash_token(token: str) -> str: return hashlib.sha256(token.encode("utf-8")).hexdigest() def fleet_token_configured(db=None) -> bool: db = db or get_db() h = (db.get_setting(_SETTING_HASH, "") or "").strip() return bool(h) def set_fleet_token(plain: str, db=None) -> None: db = db or get_db() plain = (plain or "").strip() if not plain: db.set_setting(_SETTING_HASH, "") return if len(plain) < 16: raise ValueError("中控 API Token 至少 16 位") db.set_setting(_SETTING_HASH, _hash_token(plain)) def clear_fleet_token(db=None) -> None: set_fleet_token("", db) def require_fleet_token( x_fleet_token: Annotated[str | None, Header(alias="X-Fleet-Token")] = None, authorization: Annotated[str | None, Header()] = None, ) -> str: db = get_db() stored = (db.get_setting(_SETTING_HASH, "") or "").strip() if not stored: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="策略机未配置中控 API Token", ) provided = (x_fleet_token or "").strip() if not provided and authorization: auth = authorization.strip() if auth.lower().startswith("fleet "): provided = auth[6:].strip() if not provided or not hmac.compare_digest(stored, _hash_token(provided)): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="invalid fleet token", ) return provided def _repo_root() -> Path: return Path(__file__).resolve().parents[3] def _purge_tickets() -> None: now = time.time() dead = [k for k, v in _tickets.items() if float(v.get("exp", 0)) < now] for k in dead: _tickets.pop(k, None) def create_login_ticket(username: str) -> tuple[str, int]: with _tickets_lock: _purge_tickets() ticket = secrets.token_urlsafe(32) _tickets[ticket] = {"exp": time.time() + _TICKET_TTL_SEC, "u": username} return ticket, _TICKET_TTL_SEC def consume_login_ticket(ticket: str) -> str: ticket = (ticket or "").strip() if not ticket: raise HTTPException(status_code=401, detail="invalid ticket") with _tickets_lock: _purge_tickets() meta = _tickets.pop(ticket, None) if not meta: raise HTTPException(status_code=401, detail="ticket invalid or used") if float(meta.get("exp", 0)) < time.time(): raise HTTPException(status_code=401, detail="ticket expired") username = str(meta.get("u") or "").strip() if not username: raise HTTPException(status_code=401, detail="invalid ticket") return username class FleetTokenBody(BaseModel): token: str = Field(default="", max_length=256) @router.get("/meta") async def fleet_meta(_user: Annotated[str, Depends(require_user)]) -> dict: return { "configured": fleet_token_configured(), "hint": "在中控生成 Token 后粘贴到此保存;用于远程启停、更新与免密登录。", } @router.put("/token") async def put_fleet_token( body: FleetTokenBody, _user: Annotated[str, Depends(require_user)], ) -> dict: try: set_fleet_token(body.token) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) from e return {"ok": True, "configured": fleet_token_configured()} @router.delete("/token") async def delete_fleet_token(_user: Annotated[str, Depends(require_user)]) -> dict: clear_fleet_token() return {"ok": True, "configured": False} @router.get("/status") async def fleet_status(_tok: Annotated[str, Depends(require_fleet_token)]) -> dict: settings = get_settings() try: from ..exchange.runtime import load_runtime_settings from ..strategy.session import get_session rt = load_runtime_settings() exchange_name = rt.exchange sess = get_session() snap = sess.snapshot() if sess else None except Exception: exchange_name = settings.exchange snap = None try: st = get_engine().state() except Exception: st = {} pos = st.get("position") if isinstance(st.get("position"), dict) else {} hedge_mode = str( pos.get("hedge_mode") or st.get("hedge_mode") or "" ).strip().lower() is_oo = hedge_mode == "option_option" or bool(pos.get("option2_inst_id")) legs: list[dict] = [] if pos.get("has_position") or str(pos.get("status") or "") in ( "open", "half_open", "option_closed_perp_pending", "opening", ): # 期期无永续腿;有 perp 才展示 if not is_oo and (pos.get("perp_side") or pos.get("perp_inst_id")): legs.append( { "kind": "perp", "side": pos.get("perp_side"), "inst_id": pos.get("perp_inst_id"), "qty": pos.get("perp_qty_eth"), "avg_px": pos.get("perp_entry_px"), "mark_px": pos.get("perp_mark_px"), "upl": pos.get("perp_upl"), "margin": pos.get("perp_margin"), "premium": None, } ) if pos.get("option_side") or pos.get("option_inst_id"): legs.append( { "kind": "option", "side": pos.get("option_side") or ("call" if is_oo else None), "inst_id": pos.get("option_inst_id"), "qty": pos.get("option_qty_eth"), "avg_px": pos.get("option_entry_px"), "mark_px": pos.get("option_mark_px"), "upl": pos.get("option_upl"), "margin": None, "premium": pos.get("initial_premium"), "strike": pos.get("strike"), } ) if pos.get("option2_inst_id") or pos.get("option2_side"): legs.append( { "kind": "option", "side": pos.get("option2_side") or "put", "inst_id": pos.get("option2_inst_id"), "qty": pos.get("option2_qty_eth"), "avg_px": pos.get("option2_entry_px"), "mark_px": pos.get("option2_mark_px"), "upl": pos.get("option2_upl"), "margin": None, "premium": pos.get("initial_premium2"), "strike": pos.get("strike2"), } ) # 风控展示字段以引擎 state 为准;缺省时回落 settings 表(避免旧进程漏字段) db = get_db() def _sf(key: str, default: float) -> float: try: return float(db.get_setting(key, str(default)) or default) except Exception: return float(default) def _pick(key: str, default: float | None = None): if key in st and st.get(key) is not None: return st.get(key) if default is None: return None return _sf(key, default) latest_funds = 0.0 try: from ..sim.funds_wallets import FundsWallets latest_funds = float(FundsWallets(db).total_usdt_equiv()) except Exception: try: from ..sim.ledger import Ledger latest_funds = float(Ledger(db).snapshot().get("equity") or 0) except Exception: latest_funds = 0.0 residuals: list[dict] = [] try: residuals = get_engine().matcher.list_residual_options_enriched() except Exception: residuals = [] return { "ok": True, "mode": settings.mode, "env_name": settings.env_name, "exchange": exchange_name, "sim": settings.is_sim, "latest_funds": latest_funds, "residuals": residuals, "market_connected": bool(snap.connected) if snap else False, "pair": snap.pair.to_dict() if snap and snap.pair else None, "index_px": ( pos.get("index_px") if pos.get("index_px") is not None else (getattr(snap, "index_px", None) if snap else None) ), "updated_at_ms": snap.updated_at_ms if snap else None, "strategy": { "running": st.get("running"), "phase": st.get("phase"), "rounds_done": st.get("rounds_done"), "last_error": st.get("last_error"), "group_id": pos.get("group_id"), "rest_left_sec": st.get("rest_left_sec"), "exit_mode": st.get("exit_mode"), "exit_target_usdt": st.get("exit_target_usdt"), "net_profit_target": st.get("net_profit_target"), "premium_exit_multiple": _pick("premium_exit_multiple"), "semi_auto_enabled": st.get("semi_auto_enabled"), "semi_armed": st.get("semi_armed"), "semi_view_side": st.get("semi_view_side"), "semi_option_move_points": st.get("semi_option_move_points"), "semi_perp_exit_unit": st.get("semi_perp_exit_unit"), "semi_net_exit_target": st.get("semi_net_exit_target"), "semi_min_option_hours": st.get("semi_min_option_hours"), "semi_min_option_leverage": st.get("semi_min_option_leverage"), "semi_moneyness": st.get("semi_moneyness"), "semi_otm_max_offset": st.get("semi_otm_max_offset"), "semi_perp_unit": st.get("semi_perp_unit"), "semi_option_unit": st.get("semi_option_unit"), "leverage": _pick("leverage", float(settings.leverage)), "min_option_leverage": _pick( "min_option_leverage", float(settings.min_option_leverage) ), "min_option_hours": _pick( "min_option_hours", float(settings.min_option_hours) ), "perp_margin_mode": st.get("perp_margin_mode"), "perp_qty_eth": st.get("perp_qty_eth"), "option_qty_eth": st.get("option_qty_eth"), "oo_put_qty_eth": ( st.get("oo_put_qty_eth") if st.get("oo_put_qty_eth") is not None else _sf("oo_put_qty_eth", 0.0) or None ), "sizing_mode": st.get("sizing_mode"), "risk_last_k": st.get("risk_last_k"), "risk_sizing_locked": st.get("risk_sizing_locked"), "risk_sizing_preview": st.get("risk_sizing_preview"), "risk_loss_pct": _pick("risk_loss_pct", 1.0), "risk_perp_unit": _pick("risk_perp_unit", 1.0), "risk_option_unit": _pick("risk_option_unit", 2.0), "risk_exit_unit": _pick("risk_exit_unit", 15.0), "martingale_enabled": st.get("martingale_enabled"), "martingale_doubles": st.get("martingale_doubles"), "risk_effective_loss_pct": st.get("risk_effective_loss_pct"), "oo_amplitude_pct": _pick( "oo_amplitude_pct", float(settings.oo_amplitude_pct) ), "oo_amplitude_hours": _pick( "oo_amplitude_hours", float(settings.oo_amplitude_hours) ), "oo_amplitude_filter_enabled": ( str( st.get("oo_amplitude_filter_enabled") if "oo_amplitude_filter_enabled" in st else db.get_setting( "oo_amplitude_filter_enabled", str(settings.oo_amplitude_filter_enabled), ) ) .strip() .lower() in ("1", "true", "yes", "on") ), "oo_min_option_hours": _pick( "oo_min_option_hours", float(settings.oo_min_option_hours) ), "oo_min_leverage": _pick( "oo_min_leverage", float(settings.oo_min_leverage) ), "oo_reward_ratio": _pick( "oo_reward_ratio", float(settings.oo_reward_ratio) ), "hedge_mode": ( hm if ( hm := str( st.get("hedge_mode") or ("option_option" if is_oo else None) or db.get_setting("hedge_mode", settings.hedge_mode) or settings.hedge_mode or "perp_option" ) .strip() .lower() ) in ("perp_option", "option_option") else "perp_option" ), }, "position": { "status": pos.get("status") or ("open" if pos.get("has_position") else "flat"), "has_position": bool(pos.get("has_position")), "hedge_mode": "option_option" if is_oo else "perp_option", "group_id": pos.get("group_id"), "open_at_ms": pos.get("open_at_ms"), "initial_premium": pos.get("initial_premium"), "initial_premium2": pos.get("initial_premium2"), "perp_margin": pos.get("perp_margin"), "exit_target_usdt": pos.get("exit_target_usdt"), "net_pnl": pos.get("net_pnl"), "perp_upl": pos.get("perp_upl"), "option_upl": pos.get("option_upl"), "option2_upl": pos.get("option2_upl"), "strike": pos.get("strike"), "strike2": pos.get("strike2"), "expiry_ymd": pos.get("expiry_ymd"), "legs": legs, }, "update": { "running": bool(_update_state.get("running")), "started_at_ms": _update_state.get("started_at_ms") or 0, "last_error": _update_state.get("last_error") or "", }, } @router.get("/stats") async def fleet_stats(_tok: Annotated[str, Depends(require_fleet_token)]) -> dict: """中控拉取策略机整体统计(同 /api/stats/summary,Fleet Token 鉴权)。""" from .stats import build_stats_summary return build_stats_summary() @router.post("/start") async def fleet_start(_tok: Annotated[str, Depends(require_fleet_token)]) -> dict: return await get_engine().start() @router.post("/pause") async def fleet_pause(_tok: Annotated[str, Depends(require_fleet_token)]) -> dict: return await get_engine().pause() class ResidualCloseBody(BaseModel): group_id: str = Field(min_length=1, max_length=128) @router.post("/residual/close") async def fleet_residual_close( body: ResidualCloseBody, _tok: Annotated[str, Depends(require_fleet_token)], ) -> dict: """中控手动平单条残留:只验流动性,不验权利金回收比例。""" matcher = get_engine().matcher result = await asyncio.to_thread(matcher.close_residual_manual, body.group_id) if not result.ok: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=result.detail or "平残留失败", ) return { "ok": True, "detail": result.detail, "data": result.data, "liquidity_wait": result.liquidity_wait, } @router.post("/issue-login") async def fleet_issue_login(_tok: Annotated[str, Depends(require_fleet_token)]) -> dict: username, _ = get_credentials() ticket, ttl = create_login_ticket(username) return { "ok": True, "ticket": ticket, "expires_in": ttl, "login_path": f"/fleet-login?ticket={ticket}", } def _run_update_job() -> None: root = _repo_root() script = root / "deploy" / "lib" / "update.sh" if not script.is_file(): script = root / "deploy" / "pull_and_restart.sh" try: if os.name == "nt": _update_state["last_error"] = "update script requires bash (Linux deploy host)" logger.error("fleet update skipped: not a Linux deploy host") return if not script.is_file(): _update_state["last_error"] = f"update script missing: {script}" logger.error("fleet update: %s", _update_state["last_error"]) return logger.info("fleet update starting: %s", script) proc = subprocess.run( ["bash", str(script)], cwd=str(root), capture_output=True, text=True, timeout=600, env={**os.environ, "DEBIAN_FRONTEND": "noninteractive"}, ) if proc.returncode != 0: err = (proc.stderr or proc.stdout or "")[-2000:] _update_state["last_error"] = f"exit={proc.returncode} {err}" logger.error("fleet update failed: %s", _update_state["last_error"]) else: _update_state["last_error"] = "" logger.info("fleet update finished ok") except Exception as e: _update_state["last_error"] = str(e) logger.exception("fleet update exception") finally: _update_state["running"] = False @router.post("/update") async def fleet_update(_tok: Annotated[str, Depends(require_fleet_token)]) -> dict: """接受更新请求:后台跑 deploy update(会 reload 本进程)。""" with _update_lock: if _update_state.get("running"): return { "ok": True, "accepted": False, "running": True, "msg": "更新已在进行中", } _update_state["running"] = True _update_state["started_at_ms"] = int(time.time() * 1000) _update_state["last_error"] = "" def _deferred() -> None: time.sleep(0.8) _run_update_job() threading.Thread(target=_deferred, name="fleet-update", daemon=True).start() return { "ok": True, "accepted": True, "running": True, "msg": "已接受更新,进程即将 reload,请稍后探活", }