212 lines
8.9 KiB
Python
212 lines
8.9 KiB
Python
|
|
"""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, 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.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.client.cookies.clear()
|
||
|
|
|
||
|
|
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
|
||
|
|
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.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.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):
|
||
|
|
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)
|
||
|
|
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 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")
|