fix: fetch dataset scope options from WorldQuant platform

This commit is contained in:
yuxuanhui
2026-09-08 18:57:03 +08:00
parent 80cb7d2b80
commit 3d26827b49
13 changed files with 367 additions and 53 deletions
+1 -1
View File
@@ -304,7 +304,7 @@ class AIRuntime:
call.preview = await preview_tool(business, name, args)
call.status = "pending"
else:
call.result = jsonable_encoder(await read_tool(business, name, args))
call.result = jsonable_encoder(await read_tool(business, name, args, self.runner.client))
call.status = "completed"
except HTTPException as exc:
call.result, call.status = {"error": exc.detail}, "failed"
+8 -4
View File
@@ -6,7 +6,7 @@ from typing import Literal
from pydantic import Field
from ..backtests.contracts import ControlInput, PreviewInput, RerunInput, StartInput
from ..catalog.contracts import UNIVERSES, CatalogFilters, Scope
from ..catalog.contracts import CatalogFilters, Scope
from ..research.contracts import (
ChatboxResearchInput,
InputPageArgs,
@@ -105,7 +105,7 @@ class BacktestRerunArgs(BacktestRunArgs):
CATALOG = {
"get_catalog_scopes": (EmptyArgs, "读取本版支持的研究范围组合,不表示账户已获平台权限。"),
"get_catalog_scopes": (EmptyArgs, "从平台读取当前账户可用的研究范围组合。"),
"search_catalog": (
CatalogSearchArgs,
"分页查询本地目录。省略 dataset_id 查询数据集;提供 dataset_id 查询其字段、类型和完整集合版本。无缓存时说明需在数据集页同步,不编造字段。",
@@ -200,11 +200,13 @@ def bounded(value):
return value
async def read_tool(business, name, args):
async def read_tool(business, name, args, platform_client=None):
from datetime import timezone
if name == "get_catalog_scopes":
data = {"universes": UNIVERSES, "instrument_type": "EQUITY", "delays": [0, 1]}
from ..catalog.platform import platform_options
data = await platform_options(platform_client)
elif name == "search_catalog":
data = await business.catalog.search(args.filters, args.dataset_id)
data.update(
@@ -271,6 +273,8 @@ async def read_tool(business, name, args):
if isinstance(data, list):
data = {"items": data[:20]}
data["_meta"] = ResultMetadata(observed_at=datetime.now(timezone.utc)).model_dump(mode="json")
if name == "get_catalog_scopes":
data["_meta"]["source"] = "worldquant_platform"
return bounded(data)
+6 -20
View File
@@ -16,29 +16,15 @@ def utc_timestamp(value: datetime) -> datetime:
UTCTimestamp = Annotated[datetime, AfterValidator(utc_timestamp)]
# Supported research scopes, not an assertion about a connected account's permissions.
UNIVERSES = {
"USA": ["TOP3000", "TOP1000", "TOP500", "TOP200"],
"CHN": ["TOP2000"],
"EUR": ["TOP2500", "TOP1200"],
"ASI": ["TOP1000"],
"GLB": ["TOP3000"],
"JPN": ["TOP1600"],
"HKG": ["TOP800"],
}
# Platform membership is checked before synchronization, not while reading historical caches.
ScopeName = Annotated[str, Field(min_length=1, max_length=100, pattern=r"^[^|\s]+$")]
class Scope(Contract):
instrument_type: Literal["EQUITY"] = "EQUITY"
region: str
universe: str
delay: int = Field(ge=0, le=1)
@model_validator(mode="after")
def valid_scope(self):
if self.universe not in UNIVERSES.get(self.region, []):
raise ValueError("不支持的 Region / Universe 组合")
return self
instrument_type: ScopeName = "EQUITY"
region: ScopeName
universe: ScopeName
delay: int = Field(ge=0)
def key(self):
return f"{self.instrument_type}|{self.region}|{self.universe}|{self.delay}"
+30
View File
@@ -0,0 +1,30 @@
"""Shared platform scope discovery for HTTP and AI consumers."""
from fastapi import HTTPException
from ..worldquant import WqError
async def platform_options(client):
"""Return live account choices, surfacing sanitized upstream failures to callers."""
if client is None:
raise HTTPException(409, "请先连接 WorldQuant")
try:
return await client.get_platform_setting_options()
except WqError as exc:
raise HTTPException(
409 if exc.code in ("disconnected", "verification_required") else 502, str(exc)
) from None
async def validate_platform_scope(client, scope):
"""Reject unsupported combinations before creating a synchronization job."""
data = await platform_options(client)
if not any(
row["instrument_type"] == scope.instrument_type
and row["region"] == scope.region
and row["delay"] == scope.delay
and scope.universe in row["universes"]
for row in data["instrument_options"]
):
raise HTTPException(422, "平台不支持此 Instrument Type / Region / Delay / Universe 组合")
+4 -3
View File
@@ -7,7 +7,6 @@ from fastapi import APIRouter, Depends, Query, Request
from ..schemas import JobOutput
from ..security import require_auth
from .contracts import (
UNIVERSES,
CatalogFilters,
CatalogJobInput,
CatalogPage,
@@ -19,14 +18,15 @@ from .contracts import (
NoteOutput,
Scope,
)
from .platform import platform_options, validate_platform_scope
from .service import Catalog
router = APIRouter(prefix="/api/v1/catalog", tags=["catalog"], dependencies=[Depends(require_auth)])
@router.get("/scopes")
async def scopes() -> dict[str, list[str]]:
return UNIVERSES
async def scopes(request: Request):
return await platform_options(request.app.state.runner.client)
@router.get("/datasets", response_model=CatalogPage)
@@ -69,6 +69,7 @@ async def field_note(
@router.post("/sync-jobs", status_code=202, response_model=JobOutput)
async def sync(request: Request, body: CatalogJobInput):
await validate_platform_scope(request.app.state.runner.client, body.scope)
async with request.app.state.sessions.begin() as db:
result = await Catalog(db).create_job(body)
request.app.state.runner.wake.set()
+57 -4
View File
@@ -266,6 +266,10 @@ class WqClient:
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:
@@ -273,7 +277,7 @@ class WqClient:
refreshed = False
for attempt in range(self.settings.retry_attempts):
generation = self.auth_generation
response = await self._request("GET", path, params=params, headers=headers)
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
@@ -359,9 +363,58 @@ class WqClient:
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}
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
+2
View File
@@ -109,6 +109,8 @@ def create_test_app():
},
headers={"Set-Cookie": "mock=only; Path=/"},
)
if request.method == "OPTIONS":
return catalog_response(request) or httpx.Response(404)
if path.startswith("/simulations") or (
path.startswith("/alphas/") and path.rsplit("/", 1)[-1] in simulations.alphas
):
+17
View File
@@ -21,6 +21,8 @@ def field_records(dataset="TEST_FIN", count=123):
def catalog_response(request, fields=None):
path, params = request.url.path, request.url.params
if request.method == "OPTIONS" and path == "/simulations":
return httpx.Response(200, json=platform_response())
if path not in ("/data-sets", "/data-fields"):
return None
assert request.method == "GET"
@@ -53,3 +55,18 @@ def catalog_response(request, fields=None):
rows = rows[:50] + [rows[49]] + rows[50:]
offset, limit = int(params.get("offset", 0)), int(params.get("limit", 50))
return httpx.Response(200, json={"results": rows[offset : offset + limit]})
def platform_response():
"""Synthetic upstream options include a previously unsupported market and delay."""
def values(items):
return [{"value": item} for item in items]
regions = {"USA": ["TOP3000", "TOP1000"], "CHN": ["TOP2000U"], "IND": ["TOP500"]}
children = {"instrumentType": {"choices": values(["EQUITY"])},
"region": {"choices": {"instrumentType": {"EQUITY": values(regions)}}}}
for key in ("universe", "delay"):
children[key] = {"choices": {"instrumentType": {"EQUITY": {"region": {
region: values(universes if key == "universe" else [0, 1] if region == "USA" else [1])
for region, universes in regions.items()
}}}}}
return {"actions": {"POST": {"settings": {"children": children}}}}
+22 -2
View File
@@ -26,6 +26,8 @@ async def catalog(logged_in, app):
401, headers={"WWW-Authenticate": "persona", "Location": "/authentication/persona/test"}
)
return httpx.Response(201, json={"token": {"expiry": 14400}})
if request.method == "OPTIONS":
return catalog_response(request)
assert request.method == "GET"
if request.url.path == "/users/self":
return httpx.Response(200, json={"id": "TEST_USER"})
@@ -209,7 +211,7 @@ async def test_scope_ownership_empty_and_unknown_fields_are_rejected(catalog):
).status_code == 422
assert (await prepare(client, version, dataset_id="TEST_NEWS")).status_code == 409
assert (await prepare(client, version, scope={**SCOPE, "delay": 0})).status_code == 404
assert (await prepare(client, version, scope={**SCOPE, "region": "CHN"})).status_code == 422
assert (await prepare(client, version, scope={**SCOPE, "region": "CHN"})).status_code == 404
subset = await prepare(client, version, selection="explicit", excluded_ids=["TEST_FIN_110"])
assert subset.status_code == 201 and len(subset.json()["field_ids"]) == 122
assert "TEST_FIN_110" not in subset.json()["field_ids"]
@@ -241,9 +243,11 @@ async def test_retry_after_auth_wait_disconnect_and_collection_manifest(catalog)
assert job["status"] == "completed" and delays == [2]
manifest = await search(client, "/datasets/TEST_FIN/collection")
assert manifest["collection_version"] == job["id"] and len(manifest["field_ids"]) == 123
queued = await client.post(BASE + "/sync-jobs", json={"scope": SCOPE, "dataset_id": "TEST_FIN"})
state["persona"] = True
runner.client.authenticated = False
waiting = await sync(catalog, "TEST_FIN")
await runner.execute(queued.json()["id"])
waiting = (await client.get("/api/v1/sync-jobs/" + queued.json()["id"])).json()
assert waiting["status"] == "waiting_auth"
assert (await search(client, "/datasets/TEST_FIN/collection")) == manifest
await runner.disconnect()
@@ -255,3 +259,19 @@ async def test_retry_after_auth_wait_disconnect_and_collection_manifest(catalog)
await runner.execute(connect["id"])
await runner.execute(waiting["id"])
assert (await client.get("/api/v1/sync-jobs/" + waiting["id"])).json()["status"] == "completed"
async def test_dynamic_platform_scopes_and_validation(catalog):
client, runner, _ = catalog
options = (await client.get(BASE + "/scopes")).json()
assert "IND" in options["regions_by_type"]["EQUITY"]
assert (await sync(catalog, scope={**SCOPE, "region": "IND", "universe": "TOP500"}))["status"] == "completed"
invalid = await client.post(BASE + "/sync-jobs", json={"scope": {**SCOPE, "region": "IND"}})
assert invalid.status_code == 422
from app.ai.tools import EmptyArgs, read_tool
ai = await read_tool(None, "get_catalog_scopes", EmptyArgs(), runner.client)
assert ai["instrument_options"] == options["instrument_options"]
assert ai["_meta"]["source"] == "worldquant_platform"
runner.client.disconnect()
assert (await client.get(BASE + "/scopes")).status_code == 409
assert (await search(client))["total"] == 0
+49
View File
@@ -0,0 +1,49 @@
"""Exercise OPTIONS through the real authenticated adapter, without live account traffic."""
import httpx
import pytest
from app.worldquant import WqClient, WqError
from tests.catalog_fake import platform_response
async def test_options_refreshes_auth_and_keeps_new_choices(settings):
calls = []
auth = 0
def handler(request):
nonlocal auth
calls.append((request.method, request.url.path))
if request.url.path == "/authentication":
auth += 1
return httpx.Response(201, json={})
assert request.method == "OPTIONS" and request.url.path == "/simulations"
if auth == 1:
return httpx.Response(401)
data = platform_response()
children = data["actions"]["POST"]["settings"]["children"]
children["delay"]["choices"]["instrumentType"]["EQUITY"]["region"]["IND"] = [{"value": 2}]
return httpx.Response(200, json=data)
client = WqClient(settings, transport=httpx.MockTransport(handler))
try:
await client.authenticate("test@example.com", "test-only")
data = await client.get_platform_setting_options()
assert auth == 2
assert {"instrument_type": "EQUITY", "region": "IND", "delay": 2,
"universes": ["TOP500"]} in data["instrument_options"]
assert all(method == "OPTIONS" or path == "/authentication" for method, path in calls)
finally:
await client.close()
@pytest.mark.parametrize("body", [{}, {"actions": None}, {"actions": {"POST": {"settings": {"children": []}}}}])
async def test_malformed_options_fail_without_static_fallback(settings, body):
client = WqClient(settings, transport=httpx.MockTransport(lambda _: httpx.Response(200, json=body)))
client.credentials, client.authenticated = ("test", "test"), True
try:
with pytest.raises(WqError) as error:
await client.get_platform_setting_options()
assert error.value.code == "invalid_response"
finally:
await client.close()