Files
worldquant-alpha-system/backend/app/worldquant.py
T
yuxuanhui e256d6fef1
Deploy production / deploy (push) Successful in 56s
feat: add Super Alpha research, management and MCP workflows
2026-09-13 12:32:16 +08:00

526 lines
24 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. Explicit descriptions allow PATCH; Alpha submission is not exposed.
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, poll_attempts=None, wait_for_retry_header=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
attempts = self.settings.retry_attempts if poll_attempts is None else poll_attempts
for attempt in range(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 (wait_for_retry_header or self.retry_delay(response.headers["Retry-After"], 0) > 0)
):
if attempt + 1 == 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 patch_descriptions(self, alpha_id, descriptions):
"""Write only reviewed descriptions; ambiguous writes require GET reconciliation.
No blind transport/5xx retry: the first PATCH may already have succeeded.
A job retry reads the current Alpha before deciding whether PATCH is needed.
"""
if not re.fullmatch(r"[A-Za-z0-9_-]{1,100}", alpha_id) or not descriptions or not set(descriptions) <= {"regular", "selection", "combo"}:
raise WqError("Description 写入参数无效", "invalid_description")
if any(not isinstance(value, dict) or set(value) != {"description"} or not isinstance(value["description"], str) for value in descriptions.values()):
raise WqError("仅允许写入 Description", "invalid_description")
if not self.credentials:
raise WqError("请先连接 WorldQuant", "disconnected")
if not self.authenticated:
await self.authenticate(*self.credentials)
for attempt in range(2):
generation = self.auth_generation
try:
response = await self.client.patch(f"/alphas/{alpha_id}", json=descriptions)
except httpx.TransportError:
raise WqError("Description 写回结果未知;重试任务将先核对平台内容", "description_unknown") from None
if response.status_code == 401 and attempt == 0:
await self.authenticate(*self.credentials, stale_generation=generation)
continue
if response.status_code in (200, 204):
return
raise WqError(f"Description 写回未确认(HTTP {response.status_code});可重试任务核对平台内容", "description_write_failed")
async def submission_check(self, alpha_id):
"""cnhkmcp /check contract: wait on Retry-After and return only is.checks."""
if not re.fullmatch(r"[A-Za-z0-9_-]{1,100}", alpha_id):
raise WqError("Alpha ID 无效", "invalid_alpha")
raw = await self._read_json(
"GET", f"/alphas/{alpha_id}/check", poll_attempts=self.settings.pnl_poll_attempts,
wait_for_retry_header=True,
)
metrics = raw.get("is")
checks = metrics.get("checks") if isinstance(metrics, dict) else None
if not isinstance(checks, list) or any(not isinstance(item, dict) for item in checks):
raise WqError("平台检查尚无有效 is.checks 结果,请稍后重试", "pending")
return checks
async def pnl(self, alpha_id):
"""Wait for slow PnL generation separately from transport-error retries.
Each pending response uses the platform's Retry-After delay. Exhausting
the bounded polling budget raises WqError with code 'pending'.
"""
return await self._read_json(
"GET", f"/alphas/{alpha_id}/recordsets/pnl", poll_attempts=self.settings.pnl_poll_attempts
)
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 run_super_selection(self, query):
"""Read cnhk super-selection contract with bounded async retries and shared authentication."""
allowed = {"selection", "instrumentType", "region", "delay", "selectionLimit", "selectionHandling"}
if set(query) != allowed:
raise WqError("Selection 参数不完整或包含未知键", "invalid_selection")
return await self._read_json("GET", "/simulations/super-selection", params=query, allow_list=True,
wait_for_retry_header=True)
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