"""中控(Fleet)专用 API:X-Fleet-Token 鉴权,不开放资金/下单。""" from __future__ import annotations 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 {} legs: list[dict] = [] if pos.get("has_position") or str(pos.get("status") or "") in ( "open", "half_open", "option_closed_perp_pending", "opening", ): if 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"), } ) if pos.get("option_side") or pos.get("option_inst_id"): legs.append( { "kind": "option", "side": pos.get("option_side"), "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"), } ) return { "ok": True, "mode": settings.mode, "env_name": settings.env_name, "exchange": exchange_name, "sim": settings.is_sim, "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"), "leverage": st.get("leverage"), "perp_margin_mode": st.get("perp_margin_mode"), "perp_qty_eth": st.get("perp_qty_eth"), "option_qty_eth": st.get("option_qty_eth"), "sizing_mode": st.get("sizing_mode"), "risk_last_k": st.get("risk_last_k"), "risk_sizing_locked": st.get("risk_sizing_locked"), }, "position": { "status": pos.get("status") or ("open" if pos.get("has_position") else "flat"), "has_position": bool(pos.get("has_position")), "group_id": pos.get("group_id"), "open_at_ms": pos.get("open_at_ms"), "initial_premium": pos.get("initial_premium"), "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"), "strike": pos.get("strike"), "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.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() @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,请稍后探活", }