"""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