Files
eth_hedge_sim/backend/app/live/okx_trade.py
T
dekun ca66c494e7 Use exchange fees, UPL, funding for LIVE PnL; USDC 1:1 to USDT.
Option open P&L stays local (option net); exit target ignores estimated close fees.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-07-26 21:49:20 +08:00

302 lines
11 KiB
Python

"""OKX V5 私有交易 REST(下单)。"""
from __future__ import annotations
import base64
import hashlib
import hmac
import json
import logging
import time
from dataclasses import dataclass
from typing import Any
import httpx
from ..config import Settings, get_settings
from ..exchange.okx.parse import safe_float
from .rate_limit import RateLimitError, get_throttle, parse_retry_after_header
logger = logging.getLogger(__name__)
@dataclass(slots=True)
class LiveFill:
inst_id: str
side: str
avg_px: float
sz: float # 张或币,取决于合约
fee: float
ord_id: str
raw: dict[str, Any]
class OkxTradeClient:
def __init__(self, settings: Settings | None = None) -> None:
self.settings = settings or get_settings()
proxy = (self.settings.okx_http_proxy or "").strip() or None
self._client = httpx.Client(
base_url=self.settings.okx_rest_base.rstrip("/"),
timeout=20.0,
proxy=proxy,
headers={"Accept": "application/json", "User-Agent": "eth-hedge-live/0.1"},
)
self._ct_val_cache: dict[str, float] = {}
self._throttle = get_throttle("okx_trade", min_interval_sec=1.0)
def close(self) -> None:
self._client.close()
def _ts(self) -> str:
# OKX: ISO8601 with milliseconds
return (
time.strftime("%Y-%m-%dT%H:%M:%S", time.gmtime())
+ f".{int(time.time() * 1000) % 1000:03d}Z"
)
def _sign(self, ts: str, method: str, path: str, body: str) -> str:
secret = (self.settings.okx_api_secret or "").encode("utf-8")
msg = f"{ts}{method.upper()}{path}{body}".encode("utf-8")
dig = hmac.new(secret, msg, hashlib.sha256).digest()
return base64.b64encode(dig).decode("utf-8")
def _headers(self, ts: str, sign: str) -> dict[str, str]:
return {
"OK-ACCESS-KEY": self.settings.okx_api_key or "",
"OK-ACCESS-SIGN": sign,
"OK-ACCESS-TIMESTAMP": ts,
"OK-ACCESS-PASSPHRASE": self.settings.okx_api_passphrase or "",
"Content-Type": "application/json",
}
def _request(
self, method: str, path: str, body: dict[str, Any] | None = None
) -> list[dict[str, Any]]:
self._throttle.before_request()
payload = "" if body is None else json.dumps(body, separators=(",", ":"))
ts = self._ts()
sign = self._sign(ts, method, path, payload)
headers = self._headers(ts, sign)
if method.upper() == "GET":
r = self._client.get(path, headers=headers)
else:
r = self._client.request(method.upper(), path, content=payload, headers=headers)
if r.status_code in (418, 429):
ra = parse_retry_after_header(r.headers)
self._throttle.mark_http(r.status_code, ra)
raise RateLimitError(
f"OKX HTTP {r.status_code}: {r.text[:200]}",
retry_after=self._throttle.remaining_cooldown(),
)
try:
r.raise_for_status()
except httpx.HTTPStatusError as e:
raise RuntimeError(f"OKX HTTP {r.status_code}: {r.text[:300]}") from e
data = r.json()
code = str(data.get("code") or "")
msg = str(data.get("msg") or "")
# OKX 业务层频率类错误
if code != "0":
low = f"{code} {msg}".lower()
if code in ("50011", "50061") or "too many" in low or "频率" in msg:
self._throttle.mark_seconds(20.0)
raise RateLimitError(
f"OKX trade rate-limited code={code} msg={msg}",
retry_after=self._throttle.remaining_cooldown(),
)
raise RuntimeError(
f"OKX trade error code={code} msg={msg} data={data.get('data')}"
)
rows = data.get("data") or []
return [x for x in rows if isinstance(x, dict)]
def get_ct_val(self, inst_id: str, *, inst_type: str) -> float:
if inst_id in self._ct_val_cache:
return self._ct_val_cache[inst_id]
r = self._client.get(
"/api/v5/public/instruments",
params={"instType": inst_type, "instId": inst_id},
)
r.raise_for_status()
body = r.json()
rows = body.get("data") or []
for row in rows:
if str(row.get("instId")) == inst_id:
v = safe_float(row.get("ctVal")) or safe_float(row.get("ctMult"))
if v and v > 0:
self._ct_val_cache[inst_id] = float(v)
return float(v)
default = 0.01
self._ct_val_cache[inst_id] = default
return default
def place_market(
self,
*,
inst_id: str,
side: str, # buy|sell
sz: str,
td_mode: str,
pos_side: str | None = None,
reduce_only: bool = False,
) -> LiveFill:
body: dict[str, Any] = {
"instId": inst_id,
"tdMode": td_mode,
"side": side,
"ordType": "market",
"sz": str(sz),
}
if pos_side:
body["posSide"] = pos_side
if reduce_only:
body["reduceOnly"] = True
rows = self._request("POST", "/api/v5/trade/order", body)
if not rows:
raise RuntimeError("OKX 下单无返回")
ord_id = str(rows[0].get("ordId") or "")
# 查单取均价
fill = self._wait_fill(inst_id, ord_id)
return fill
def _wait_fill(self, inst_id: str, ord_id: str, *, tries: int = 8) -> LiveFill:
path = f"/api/v5/trade/order?instId={inst_id}&ordId={ord_id}"
last: dict[str, Any] = {}
for _ in range(tries):
rows = self._request("GET", path)
if rows:
last = rows[0]
state = str(last.get("state") or "")
avg = safe_float(last.get("avgPx"))
if state in ("filled", "partially_filled") and avg and avg > 0:
sz = safe_float(last.get("accFillSz")) or safe_float(last.get("sz")) or 0.0
fee = abs(safe_float(last.get("fee")) or 0.0)
fee_ccy = str(last.get("feeCcy") or "USDT")
# 订单上 fee 常为空,再查成交明细
if fee <= 0 and ord_id:
fee, fee_ccy = self.sum_fill_fees(inst_id, ord_id)
from .money import abs_fee_usdt
return LiveFill(
inst_id=inst_id,
side=str(last.get("side") or ""),
avg_px=float(avg),
sz=float(sz),
fee=abs_fee_usdt(fee, fee_ccy),
ord_id=ord_id,
raw=last,
)
if state in ("canceled", "failed"):
raise RuntimeError(f"OKX 订单失败 state={state} {last}")
time.sleep(0.25)
raise RuntimeError(f"OKX 订单未成交 ordId={ord_id} last={last}")
def sum_fill_fees(self, inst_id: str, ord_id: str) -> tuple[float, str]:
"""成交明细手续费合计(原币种金额, 币种)。"""
path = f"/api/v5/trade/fills?instId={inst_id}&ordId={ord_id}"
try:
rows = self._request("GET", path)
except Exception as e:
logger.warning("okx fills fee query failed: %s", e)
return 0.0, "USDT"
total = 0.0
ccy = "USDT"
for row in rows:
f = abs(safe_float(row.get("fee")) or 0.0)
total += f
if row.get("feeCcy"):
ccy = str(row.get("feeCcy"))
return total, ccy
def get_perp_upl_usdt(self, inst_id: str, *, pos_side: str | None = None) -> float | None:
"""当前永续未实现盈亏(USDT,1:1)。"""
from .money import to_usdt
try:
rows = self._request(
"GET", f"/api/v5/account/positions?instId={inst_id}"
)
except Exception as e:
logger.warning("okx positions failed: %s", e)
return None
want = (pos_side or "").strip().lower()
for row in rows:
ps = str(row.get("posSide") or "").lower()
pos = safe_float(row.get("pos")) or 0.0
if abs(pos) < 1e-12:
continue
if want and want not in ("net", "") and ps and ps != want and ps != "net":
continue
upl = safe_float(row.get("upl"))
if upl is None:
continue
ccy = str(row.get("ccy") or row.get("settleCcy") or "USDT")
return to_usdt(float(upl), ccy)
return 0.0
def get_funding_usdt(
self, inst_id: str, *, begin_ms: int, end_ms: int | None = None
) -> float:
"""资金费合计(已计入账户的 signed 金额,USDT 1:1)。付费为负。"""
from .money import to_usdt
end = int(end_ms or int(time.time() * 1000))
# type=8 funding fee
path = (
f"/api/v5/account/bills?instType=SWAP&instId={inst_id}"
f"&type=8&begin={int(begin_ms)}&end={end}"
)
total = 0.0
try:
rows = self._request("GET", path)
except Exception as e:
logger.warning("okx funding bills failed: %s", e)
return 0.0
for row in rows:
# balChg / pnl 视接口;资金费常用 pnl 或 balChg
raw = safe_float(row.get("pnl"))
if raw is None:
raw = safe_float(row.get("balChg"))
if raw is None:
continue
ccy = str(row.get("ccy") or "USDT")
total += to_usdt(float(raw), ccy)
return total
def get_closed_perp_pnl_usdt(
self, inst_id: str, *, begin_ms: int, end_ms: int | None = None
) -> float | None:
"""平仓后从历史仓位取已实现盈亏(不含手续费;含部分仓位盈亏)。"""
from .money import to_usdt
end = int(end_ms or int(time.time() * 1000))
path = (
f"/api/v5/account/positions-history?instType=SWAP&instId={inst_id}"
f"&before={end}&after={int(begin_ms)}"
)
try:
# positions-history 用 GET query;部分环境用 before/after 语义相反,失败则返回 None
rows = self._request(
"GET",
f"/api/v5/account/positions-history?instType=SWAP&instId={inst_id}",
)
except Exception as e:
logger.warning("okx positions-history failed: %s", e)
return None
total = 0.0
hit = False
for row in rows:
u_time = int(safe_float(row.get("uTime")) or safe_float(row.get("cTime")) or 0)
if u_time and (u_time < int(begin_ms) - 60_000 or u_time > end + 60_000):
continue
rpnl = safe_float(row.get("realizedPnl"))
if rpnl is None:
rpnl = safe_float(row.get("pnl"))
if rpnl is None:
continue
hit = True
ccy = str(row.get("ccy") or "USDT")
total += to_usdt(float(rpnl), ccy)
return total if hit else None