Files
eth_hedge_sim/backend/app/api/fleet.py
T
2026-07-30 14:58:30 +08:00

347 lines
12 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.
"""中控(Fleet)专用 APIX-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,请稍后探活",
}