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

289 lines
12 KiB
Python
Raw Normal View History

"""Read-only WorldQuant adapter. Authentication is the only allowed 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
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 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.on_retry: Callable[[float], Awaitable[None]] | None = None
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")
async def catalog_page(self, scope, dataset_id, offset):
"""Read a single scoped page. IDs are query parameters, never upstream paths."""
params = {"instrumentType": scope["instrument_type"], "region": scope["region"],
"universe": scope["universe"], "delay": scope["delay"],
"limit": 50, "offset": offset}
if dataset_id is not None:
params["dataset.id"] = dataset_id
return await self.get("/data-fields" if dataset_id else "/data-sets", params)