"""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, **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): 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, *, 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" params.update({f"{field}>=": date_from, f"{field}<": date_to, "order": field}) if field == "dateCreated": params["dateCreated<"] = min(before, date_to) 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 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