fix: fetch dataset scope options from WorldQuant platform
This commit is contained in:
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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}"
|
||||
|
||||
@@ -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 组合")
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user