Files
worldquant-alpha-system/backend/app/worldquant.py
T

354 lines
15 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.
"""WorldQuant adapter. Only authentication and explicit backtests allow upstream POST.
No upstream response body or request headers are included in exceptions: they may
contain credentials, cookies, or temporary authentication links.
"""
import asyncio
import math
import random
import re
from contextvars import ContextVar
from datetime import datetime, timedelta, timezone
from email.utils import parsedate_to_datetime
from typing import Awaitable, Callable
from urllib.parse import urljoin, urlparse
import httpx
class WqError(Exception):
def __init__(self, message: str, code: str = "upstream_error"):
super().__init__(message)
self.code = code
class VerificationRequired(WqError):
def __init__(self, url: str):
super().__init__("请在 WorldQuant 完成人工验证,再点击继续验证", "verification_required")
self.url = url
class SimulationDeferred(WqError):
def __init__(self, message, delay=5, code="rate_limited"):
super().__init__(message, code)
self.delay = delay
class WqClient:
def __init__(self, settings, transport=None, sleep=asyncio.sleep):
self.settings = settings
self.client = httpx.AsyncClient(
base_url=settings.wq_base_url,
timeout=settings.request_timeout,
transport=transport,
follow_redirects=False,
)
self.lock = asyncio.Lock()
self.authenticated = False
self.auth_generation = 0
self.credentials: tuple[str, str] | None = None
self.verification_url: str | None = None
self.permissions: list[str] | None = None
self.session_expires_at: datetime | None = None
self.session_duration: float | None = None
self.sleep = sleep
self._retry_hook = ContextVar("wq_retry_hook", default=None)
@property
def on_retry(self) -> Callable[[float], Awaitable[None]] | None:
return self._retry_hook.get()
@on_retry.setter
def on_retry(self, value):
# Sync and simulation tasks share a session, never each other's retry callback.
self._retry_hook.set(value)
def simulation_url(self, value):
"""Accept only same-origin simulation resources; never forward cookies elsewhere."""
base = urlparse(self.settings.wq_base_url)
url = urlparse(urljoin(self.settings.wq_base_url, value))
if (
url.scheme != base.scheme
or url.netloc != base.netloc
or url.query
or url.fragment
or not re.fullmatch(r"/simulations/[A-Za-z0-9_-]+", url.path)
):
raise WqError("模拟引用地址无法确认", "invalid_simulation_url")
return url.geturl()
async def submit_simulations(self, payload):
"""One POST only. Transport/5xx/invalid acknowledgement may already be accepted."""
try:
response = await self.client.post(
"/simulations", json=payload[0] if len(payload) == 1 else payload
)
except httpx.TransportError:
raise WqError("提交结果未知,禁止自动重提,请核对平台任务", "submission_unknown") from None
if response.status_code == 429:
raise SimulationDeferred(
"平台限流,暂停后续提交", self.retry_delay(response.headers.get("Retry-After"), 0)
)
if response.status_code == 401:
self.authenticated = False
raise SimulationDeferred("平台会话过期,重新认证后继续", 2, "session_expired")
if response.status_code in (400, 403, 404, 422):
raise WqError(f"平台拒绝回测提交(HTTP {response.status_code})", "submission_rejected")
if response.status_code != 201 or not response.headers.get("Location"):
raise WqError("平台未返回可靠提交凭证,请核对后再处理", "submission_unknown")
try:
return self.simulation_url(response.headers["Location"])
except WqError:
raise WqError("平台已响应但模拟引用无法确认,禁止自动重提", "submission_unknown") from None
async def poll_simulation(self, url):
response = await self._request("GET", self.simulation_url(url))
if response.status_code == 401:
self.authenticated = False
raise SimulationDeferred("平台会话过期,重新认证后继续", 2, "session_expired")
if response.status_code not in (200, 202):
raise WqError(
f"模拟查询失败(HTTP {response.status_code}),保留原任务", "simulation_unavailable"
)
try:
data = response.json()
if not isinstance(data, dict):
raise ValueError()
except ValueError:
raise WqError("模拟响应格式无法识别,保留原任务", "invalid_response") from None
return data, self.retry_delay(response.headers["Retry-After"], 0) if response.headers.get(
"Retry-After"
) else 0
async def close(self):
await self.client.aclose()
def disconnect(self):
self.authenticated = False
self.credentials = None
self.verification_url = None
self.permissions = None
self.session_expires_at = None
self.session_duration = None
self.client.cookies.clear()
def capture_session(self, response: httpx.Response):
"""Keep only permissions and expiry metadata; never retain the authentication body."""
self.permissions, self.session_expires_at, self.session_duration = None, None, None
try:
data = response.json()
except ValueError:
return
if not isinstance(data, dict):
return
permissions = data.get("permissions")
if isinstance(permissions, list) and all(isinstance(p, str) for p in permissions):
self.permissions = list(dict.fromkeys(permissions))
token = data.get("token")
expiry = token.get("expiry") if isinstance(token, dict) else None
if (
isinstance(expiry, (int, float))
and not isinstance(expiry, bool)
and math.isfinite(expiry)
and 0 <= expiry <= 31536000
):
self.session_duration = expiry
self.session_expires_at = datetime.now(timezone.utc) + timedelta(seconds=expiry)
def session_info(self):
remaining = (
max(0, int((self.session_expires_at - datetime.now(timezone.utc)).total_seconds()))
if self.session_expires_at
else None
)
return {
"authenticated": self.authenticated and remaining != 0,
"expires_at": self.session_expires_at,
"remaining_seconds": remaining,
"total_seconds": self.session_duration,
}
def safe_verification_url(self, response: httpx.Response) -> str:
url = urljoin(str(response.url), response.headers.get("Location", ""))
base = urlparse(self.settings.wq_base_url)
parsed = urlparse(url)
if (
not response.headers.get("Location")
or parsed.scheme != base.scheme
or parsed.netloc != base.netloc
):
raise WqError("平台返回了无法确认的验证地址,请到 BRAIN 官网检查账户", "invalid_verification")
return url
@staticmethod
def retry_delay(value: str | None, attempt: int) -> float:
if value:
try:
seconds = float(value)
except ValueError:
try:
seconds = (parsedate_to_datetime(value) - datetime.now(timezone.utc)).total_seconds()
except (TypeError, ValueError):
seconds = -1
if math.isfinite(seconds) and seconds >= 0:
return seconds
return min(2**attempt + random.uniform(0, 0.5), 30)
async def _request(self, method: str, path: str, **kwargs) -> httpx.Response:
for attempt in range(self.settings.retry_attempts):
try:
response = await self.client.request(method, path, **kwargs)
except (httpx.TimeoutException, httpx.NetworkError):
if attempt + 1 == self.settings.retry_attempts:
raise WqError("WorldQuant 网络请求失败,请稍后重试", "network_error") from None
response = None
if response is not None and response.status_code not in (429, 500, 502, 503, 504):
return response
if attempt + 1 == self.settings.retry_attempts:
raise WqError("WorldQuant 暂时限流或不可用,重试预算已用完", "retry_exhausted")
delay = self.retry_delay(
response.headers.get("Retry-After") if response is not None else None, attempt
)
if self.on_retry:
await self.on_retry(delay)
await self.sleep(delay)
raise WqError("重试失败")
async def authenticate(self, email: str, password: str, force=False, stale_generation=None):
async with self.lock:
if (
self.authenticated
and not force
and (stale_generation is None or stale_generation != self.auth_generation)
):
return
if self.verification_url and not force:
raise VerificationRequired(self.verification_url)
self.credentials = (email, password)
self.authenticated = False
self.permissions, self.session_expires_at, self.session_duration = None, None, None
if force:
self.client.cookies.clear()
self.verification_url = None
response = await self._request("POST", "/authentication", auth=httpx.BasicAuth(email, password))
if (
response.status_code == 401
and response.headers.get("WWW-Authenticate", "").lower() == "persona"
):
self.verification_url = self.safe_verification_url(response)
raise VerificationRequired(self.verification_url)
if response.status_code in (401, 403):
raise WqError("平台认证未通过,请检查凭据或账户访问权限", "authentication_failed")
if response.status_code not in (200, 201):
raise WqError(f"平台认证返回异常状态 {response.status_code}", "authentication_failed")
self.authenticated = True
self.capture_session(response)
self.auth_generation += 1
self.verification_url = None
async def verify(self):
"""Continue the same server-side challenge session after the human visits the URL."""
if not self.verification_url:
if not self.credentials:
raise WqError("验证会话已结束,请重新连接", "authentication_failed")
await self.authenticate(*self.credentials, force=True)
return
response = await self._request("POST", self.verification_url)
if response.status_code in (200, 201):
self.authenticated = True
self.capture_session(response)
self.auth_generation += 1
self.verification_url = None
elif response.status_code in (401, 403, 202):
raise VerificationRequired(self.verification_url)
else:
raise WqError("验证会话已失效,请重新连接", "authentication_failed")
async def get(self, path: str, params=None, headers=None):
if not self.credentials:
raise WqError("请先连接 WorldQuant", "disconnected")
if not self.authenticated:
await self.authenticate(*self.credentials)
refreshed = False
for attempt in range(self.settings.retry_attempts):
generation = self.auth_generation
response = await self._request("GET", path, params=params, headers=headers)
if response.status_code == 401 and not refreshed:
await self.authenticate(*self.credentials, stale_generation=generation)
refreshed = True
continue
if response.status_code in (401, 403):
raise WqError("无权访问该平台资源", "access_denied")
if response.status_code == 404:
raise WqError("平台资源不存在或不可访问", "not_found")
if response.status_code not in (200, 201, 202):
raise WqError(f"WorldQuant 请求返回状态 {response.status_code}")
# Recordsets may return 200/202 with Retry-After before results exist.
if (
response.headers.get("Retry-After")
and self.retry_delay(response.headers["Retry-After"], 0) > 0
):
if attempt + 1 == self.settings.retry_attempts:
raise WqError("平台数据仍在准备,请稍后重试", "pending")
delay = self.retry_delay(response.headers["Retry-After"], attempt)
if self.on_retry:
await self.on_retry(delay)
await self.sleep(delay)
continue
try:
result = response.json()
if not isinstance(result, dict):
raise ValueError()
return result
except ValueError:
raise WqError("平台返回的数据格式无法识别", "invalid_response") from None
raise WqError("平台数据尚未就绪", "pending")
async def profile(self):
return await self.get("/users/self")
async def account_usage(self):
"""Read independent account resources; unavailable sections do not hide the profile."""
from .account_data import usage_snapshot
resources = {
"simulations": "/users/self/activities/simulations",
"submissions": "/users/self/activities/submissions",
"alphas": "/users/self/alphas/summary",
}
data, errors = {}, {}
for key, path in resources.items():
try:
data[key] = await self.get(
path, headers={"Accept": "application/json;version=4.0"} if key == "alphas" else None
)
except VerificationRequired:
raise
except WqError as exc:
if exc.code in ("authentication_failed", "disconnected"):
raise
errors[key] = str(exc)
return usage_snapshot(data, errors)
async def alphas(self, submission, hidden, offset, before):
# Cover every platform stage; submitted records are not assumed to be OS only.
status_key = "status" if submission == "UNSUBMITTED" else "status!"
return await self.get(
"/users/self/alphas",
{
status_key: "UNSUBMITTED",
"hidden": str(hidden).lower(),
"limit": 100,
"offset": offset,
"order": "dateCreated",
"dateCreated<": before,
},
)
async def alpha(self, alpha_id):
return await self.get(f"/alphas/{alpha_id}")
async def pnl(self, alpha_id):
return await self.get(f"/alphas/{alpha_id}/recordsets/pnl")