Files
eth_hedge_sim/backend/app/api/fleet.py
T

509 lines
18 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 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"),
"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/summaryFleet 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,请稍后探活",
}