469 lines
21 KiB
Python
469 lines
21 KiB
Python
"""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):
|
||
return await self._read_json("GET", path, params=params, headers=headers)
|
||
|
||
async def _read_json(self, method: str, path: str, *, allow_list=False, **kwargs):
|
||
"""Authenticated read with shared refresh/retry handling; callers use GET or OPTIONS."""
|
||
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(method, path, **kwargs)
|
||
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) and not (allow_list and isinstance(result, list)):
|
||
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 get_pyramid_multipliers(self):
|
||
"""Return current Pyramid multipliers using the shared authenticated session."""
|
||
return await self.get("/users/self/activities/pyramid-multipliers")
|
||
|
||
async def get_pyramid_alphas(self, start_date=None, end_date=None):
|
||
"""Return distribution for optional ISO dates; only missing routes trigger fallback.
|
||
|
||
Paths mirror the cnhk adapter. Access denial, authentication and network
|
||
failures must retain their meaning instead of trying unrelated routes.
|
||
"""
|
||
params = {}
|
||
if start_date is not None:
|
||
params["startDate"] = start_date
|
||
if end_date is not None:
|
||
params["endDate"] = end_date
|
||
for path in (
|
||
"/users/self/activities/pyramid-alphas",
|
||
"/users/self/pyramid/alphas",
|
||
"/activities/pyramid-alphas",
|
||
):
|
||
try:
|
||
return await self.get(path, params=params)
|
||
except WqError as exc:
|
||
if exc.code != "not_found":
|
||
raise
|
||
raise WqError("当前账户的 Pyramid 分布接口不可用", "not_found")
|
||
|
||
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, *, date_from=None, date_to=None):
|
||
# Cover every platform stage; submitted records are not assumed to be OS only.
|
||
status_key = "status" if submission == "UNSUBMITTED" else "status!"
|
||
params = {
|
||
status_key: "UNSUBMITTED",
|
||
"hidden": str(hidden).lower(),
|
||
"limit": 100,
|
||
"offset": offset,
|
||
"order": "dateCreated",
|
||
"dateCreated<": before,
|
||
}
|
||
if date_from is not None and date_to is not None:
|
||
# Daily intervals are [UTC midnight, next midnight). Submitted dates
|
||
# use the actual submission time, independently of creation or stage.
|
||
field = "dateCreated" if submission == "UNSUBMITTED" else "dateSubmitted"
|
||
# BRAIN treats the query separator '=' as part of its comparison
|
||
# syntax: a '>=' key adds an extra '=' and makes the date invalid.
|
||
# Both bounds are inclusive upstream; exclude the next UTC midnight.
|
||
upper = (datetime.fromisoformat(date_to) - timedelta(microseconds=1)).isoformat()
|
||
params.update({f"{field}>": date_from, f"{field}<": upper, "order": field})
|
||
if field == "dateCreated":
|
||
params["dateCreated<"] = min(before, upper)
|
||
return await self.get("/users/self/alphas", params)
|
||
|
||
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)
|
||
|
||
async def operators(self, offset=0):
|
||
"""The operator endpoint has both list and paginated response forms."""
|
||
return await self._read_json("GET", "/operators", allow_list=True, params={"limit": 100, "offset": offset})
|
||
|
||
async def field_availability(self, field_id, scope):
|
||
"""Use a validated identifier, never an arbitrary upstream path or URL."""
|
||
if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", field_id):
|
||
raise WqError("字段标识格式无效", "invalid_field")
|
||
return await self.get(f"/data-fields/{field_id}", {
|
||
"instrumentType": scope.instrument_type, "region": scope.region,
|
||
"universe": scope.universe, "delay": scope.delay,
|
||
})
|
||
|
||
async def research_setting_options(self):
|
||
"""Snapshot full setting choices for constrained research, including neutralization."""
|
||
return await self._read_json("OPTIONS", "/simulations")
|
||
|
||
async def get_platform_setting_options(self):
|
||
"""Read platform choices for the connected account; malformed responses raise WqError."""
|
||
data = await self._read_json("OPTIONS", "/simulations")
|
||
try:
|
||
children = data["actions"]["POST"]["settings"]["children"]
|
||
|
||
def choices(key, instrument=None, region=None):
|
||
value = children[key]["choices"]
|
||
if instrument is not None:
|
||
value = value["instrumentType"][instrument]
|
||
if region is not None:
|
||
value = value["region"][region]
|
||
values = [item["value"] for item in value]
|
||
if not values:
|
||
raise ValueError()
|
||
return values
|
||
|
||
instruments = choices("instrumentType")
|
||
regions = {}
|
||
rows = []
|
||
for instrument in instruments:
|
||
regions[instrument] = choices("region", instrument)
|
||
for region in regions[instrument]:
|
||
universes = choices("universe", instrument, region)
|
||
for delay in choices("delay", instrument, region):
|
||
if type(delay) is not int or delay < 0:
|
||
raise ValueError()
|
||
if not all(
|
||
isinstance(v, str) and v and "|" not in v
|
||
for v in [instrument, region, *universes]
|
||
):
|
||
raise ValueError()
|
||
rows.append(
|
||
dict(instrument_type=instrument, region=region, delay=delay, universes=universes)
|
||
)
|
||
return dict(
|
||
instrument_options=rows,
|
||
instrument_types=instruments,
|
||
regions_by_type=regions,
|
||
total_combinations=len(rows),
|
||
)
|
||
except (KeyError, TypeError, ValueError):
|
||
raise WqError("平台配置选项格式无法识别,请稍后重试", "invalid_response") from None
|