67f22626e9
Co-authored-by: Cursor <cursoragent@cursor.com>
347 lines
12 KiB
Python
347 lines
12 KiB
Python
"""中控(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"),
|
||
"premium_exit_multiple": st.get("premium_exit_multiple"),
|
||
"leverage": st.get("leverage"),
|
||
"min_option_leverage": st.get("min_option_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"),
|
||
"risk_loss_pct": st.get("risk_loss_pct"),
|
||
"risk_perp_unit": st.get("risk_perp_unit"),
|
||
"risk_option_unit": st.get("risk_option_unit"),
|
||
"risk_exit_unit": st.get("risk_exit_unit"),
|
||
},
|
||
"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,请稍后探活",
|
||
}
|