fix: fetch dataset scope options from WorldQuant platform
This commit is contained in:
@@ -0,0 +1,12 @@
|
|||||||
|
# 动态获取数据集研究范围
|
||||||
|
Type: task
|
||||||
|
Status: resolved
|
||||||
|
|
||||||
|
通过 OPTIONS /simulations 替换硬编码研究范围;页面与 AI 共享来源,发起同步时校验当前组合,历史缓存保留离线读取能力。
|
||||||
|
|
||||||
|
验证:平台响应解析、认证重试、异常响应、接口组合校验、前端类型检查和数据集浏览器回归。
|
||||||
|
|
||||||
|
## Answer
|
||||||
|
已接入账户认证会话的 OPTIONS /simulations,页面及 AI 动态读取范围,同步前按当前平台组合校验。移除静态白名单,历史目录仍可离线读取。前端支持类型/地区/延迟/股票池联动、独立错误与重试。
|
||||||
|
|
||||||
|
验证:后端全量 134 passed;Ruff、前端生产构建及 diff 检查通过。浏览器全量首次 11/13 通过;新增用例时序问题修正后,数据集及工作空间两组 6/6 通过,任务面板遮挡未复现。未调用真实平台、未部署。
|
||||||
@@ -304,7 +304,7 @@ class AIRuntime:
|
|||||||
call.preview = await preview_tool(business, name, args)
|
call.preview = await preview_tool(business, name, args)
|
||||||
call.status = "pending"
|
call.status = "pending"
|
||||||
else:
|
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"
|
call.status = "completed"
|
||||||
except HTTPException as exc:
|
except HTTPException as exc:
|
||||||
call.result, call.status = {"error": exc.detail}, "failed"
|
call.result, call.status = {"error": exc.detail}, "failed"
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ from typing import Literal
|
|||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
|
|
||||||
from ..backtests.contracts import ControlInput, PreviewInput, RerunInput, StartInput
|
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 (
|
from ..research.contracts import (
|
||||||
ChatboxResearchInput,
|
ChatboxResearchInput,
|
||||||
InputPageArgs,
|
InputPageArgs,
|
||||||
@@ -105,7 +105,7 @@ class BacktestRerunArgs(BacktestRunArgs):
|
|||||||
|
|
||||||
|
|
||||||
CATALOG = {
|
CATALOG = {
|
||||||
"get_catalog_scopes": (EmptyArgs, "读取本版支持的研究范围组合,不表示账户已获平台权限。"),
|
"get_catalog_scopes": (EmptyArgs, "从平台读取当前账户可用的研究范围组合。"),
|
||||||
"search_catalog": (
|
"search_catalog": (
|
||||||
CatalogSearchArgs,
|
CatalogSearchArgs,
|
||||||
"分页查询本地目录。省略 dataset_id 查询数据集;提供 dataset_id 查询其字段、类型和完整集合版本。无缓存时说明需在数据集页同步,不编造字段。",
|
"分页查询本地目录。省略 dataset_id 查询数据集;提供 dataset_id 查询其字段、类型和完整集合版本。无缓存时说明需在数据集页同步,不编造字段。",
|
||||||
@@ -200,11 +200,13 @@ def bounded(value):
|
|||||||
return value
|
return value
|
||||||
|
|
||||||
|
|
||||||
async def read_tool(business, name, args):
|
async def read_tool(business, name, args, platform_client=None):
|
||||||
from datetime import timezone
|
from datetime import timezone
|
||||||
|
|
||||||
if name == "get_catalog_scopes":
|
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":
|
elif name == "search_catalog":
|
||||||
data = await business.catalog.search(args.filters, args.dataset_id)
|
data = await business.catalog.search(args.filters, args.dataset_id)
|
||||||
data.update(
|
data.update(
|
||||||
@@ -271,6 +273,8 @@ async def read_tool(business, name, args):
|
|||||||
if isinstance(data, list):
|
if isinstance(data, list):
|
||||||
data = {"items": data[:20]}
|
data = {"items": data[:20]}
|
||||||
data["_meta"] = ResultMetadata(observed_at=datetime.now(timezone.utc)).model_dump(mode="json")
|
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)
|
return bounded(data)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -16,29 +16,15 @@ def utc_timestamp(value: datetime) -> datetime:
|
|||||||
UTCTimestamp = Annotated[datetime, AfterValidator(utc_timestamp)]
|
UTCTimestamp = Annotated[datetime, AfterValidator(utc_timestamp)]
|
||||||
|
|
||||||
|
|
||||||
# Supported research scopes, not an assertion about a connected account's permissions.
|
# Platform membership is checked before synchronization, not while reading historical caches.
|
||||||
UNIVERSES = {
|
ScopeName = Annotated[str, Field(min_length=1, max_length=100, pattern=r"^[^|\s]+$")]
|
||||||
"USA": ["TOP3000", "TOP1000", "TOP500", "TOP200"],
|
|
||||||
"CHN": ["TOP2000"],
|
|
||||||
"EUR": ["TOP2500", "TOP1200"],
|
|
||||||
"ASI": ["TOP1000"],
|
|
||||||
"GLB": ["TOP3000"],
|
|
||||||
"JPN": ["TOP1600"],
|
|
||||||
"HKG": ["TOP800"],
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class Scope(Contract):
|
class Scope(Contract):
|
||||||
instrument_type: Literal["EQUITY"] = "EQUITY"
|
instrument_type: ScopeName = "EQUITY"
|
||||||
region: str
|
region: ScopeName
|
||||||
universe: str
|
universe: ScopeName
|
||||||
delay: int = Field(ge=0, le=1)
|
delay: int = Field(ge=0)
|
||||||
|
|
||||||
@model_validator(mode="after")
|
|
||||||
def valid_scope(self):
|
|
||||||
if self.universe not in UNIVERSES.get(self.region, []):
|
|
||||||
raise ValueError("不支持的 Region / Universe 组合")
|
|
||||||
return self
|
|
||||||
|
|
||||||
def key(self):
|
def key(self):
|
||||||
return f"{self.instrument_type}|{self.region}|{self.universe}|{self.delay}"
|
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 ..schemas import JobOutput
|
||||||
from ..security import require_auth
|
from ..security import require_auth
|
||||||
from .contracts import (
|
from .contracts import (
|
||||||
UNIVERSES,
|
|
||||||
CatalogFilters,
|
CatalogFilters,
|
||||||
CatalogJobInput,
|
CatalogJobInput,
|
||||||
CatalogPage,
|
CatalogPage,
|
||||||
@@ -19,14 +18,15 @@ from .contracts import (
|
|||||||
NoteOutput,
|
NoteOutput,
|
||||||
Scope,
|
Scope,
|
||||||
)
|
)
|
||||||
|
from .platform import platform_options, validate_platform_scope
|
||||||
from .service import Catalog
|
from .service import Catalog
|
||||||
|
|
||||||
router = APIRouter(prefix="/api/v1/catalog", tags=["catalog"], dependencies=[Depends(require_auth)])
|
router = APIRouter(prefix="/api/v1/catalog", tags=["catalog"], dependencies=[Depends(require_auth)])
|
||||||
|
|
||||||
|
|
||||||
@router.get("/scopes")
|
@router.get("/scopes")
|
||||||
async def scopes() -> dict[str, list[str]]:
|
async def scopes(request: Request):
|
||||||
return UNIVERSES
|
return await platform_options(request.app.state.runner.client)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/datasets", response_model=CatalogPage)
|
@router.get("/datasets", response_model=CatalogPage)
|
||||||
@@ -69,6 +69,7 @@ async def field_note(
|
|||||||
|
|
||||||
@router.post("/sync-jobs", status_code=202, response_model=JobOutput)
|
@router.post("/sync-jobs", status_code=202, response_model=JobOutput)
|
||||||
async def sync(request: Request, body: CatalogJobInput):
|
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:
|
async with request.app.state.sessions.begin() as db:
|
||||||
result = await Catalog(db).create_job(body)
|
result = await Catalog(db).create_job(body)
|
||||||
request.app.state.runner.wake.set()
|
request.app.state.runner.wake.set()
|
||||||
|
|||||||
@@ -266,6 +266,10 @@ class WqClient:
|
|||||||
raise WqError("验证会话已失效,请重新连接", "authentication_failed")
|
raise WqError("验证会话已失效,请重新连接", "authentication_failed")
|
||||||
|
|
||||||
async def get(self, path: str, params=None, headers=None):
|
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:
|
if not self.credentials:
|
||||||
raise WqError("请先连接 WorldQuant", "disconnected")
|
raise WqError("请先连接 WorldQuant", "disconnected")
|
||||||
if not self.authenticated:
|
if not self.authenticated:
|
||||||
@@ -273,7 +277,7 @@ class WqClient:
|
|||||||
refreshed = False
|
refreshed = False
|
||||||
for attempt in range(self.settings.retry_attempts):
|
for attempt in range(self.settings.retry_attempts):
|
||||||
generation = self.auth_generation
|
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:
|
if response.status_code == 401 and not refreshed:
|
||||||
await self.authenticate(*self.credentials, stale_generation=generation)
|
await self.authenticate(*self.credentials, stale_generation=generation)
|
||||||
refreshed = True
|
refreshed = True
|
||||||
@@ -359,9 +363,58 @@ class WqClient:
|
|||||||
|
|
||||||
async def catalog_page(self, scope, dataset_id, offset):
|
async def catalog_page(self, scope, dataset_id, offset):
|
||||||
"""Read a single scoped page. IDs are query parameters, never upstream paths."""
|
"""Read a single scoped page. IDs are query parameters, never upstream paths."""
|
||||||
params = {"instrumentType": scope["instrument_type"], "region": scope["region"],
|
params = {
|
||||||
"universe": scope["universe"], "delay": scope["delay"],
|
"instrumentType": scope["instrument_type"],
|
||||||
"limit": 50, "offset": offset}
|
"region": scope["region"],
|
||||||
|
"universe": scope["universe"],
|
||||||
|
"delay": scope["delay"],
|
||||||
|
"limit": 50,
|
||||||
|
"offset": offset,
|
||||||
|
}
|
||||||
if dataset_id is not None:
|
if dataset_id is not None:
|
||||||
params["dataset.id"] = dataset_id
|
params["dataset.id"] = dataset_id
|
||||||
return await self.get("/data-fields" if dataset_id else "/data-sets", params)
|
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
|
||||||
|
|||||||
@@ -109,6 +109,8 @@ def create_test_app():
|
|||||||
},
|
},
|
||||||
headers={"Set-Cookie": "mock=only; Path=/"},
|
headers={"Set-Cookie": "mock=only; Path=/"},
|
||||||
)
|
)
|
||||||
|
if request.method == "OPTIONS":
|
||||||
|
return catalog_response(request) or httpx.Response(404)
|
||||||
if path.startswith("/simulations") or (
|
if path.startswith("/simulations") or (
|
||||||
path.startswith("/alphas/") and path.rsplit("/", 1)[-1] in simulations.alphas
|
path.startswith("/alphas/") and path.rsplit("/", 1)[-1] in simulations.alphas
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -21,6 +21,8 @@ def field_records(dataset="TEST_FIN", count=123):
|
|||||||
|
|
||||||
def catalog_response(request, fields=None):
|
def catalog_response(request, fields=None):
|
||||||
path, params = request.url.path, request.url.params
|
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"):
|
if path not in ("/data-sets", "/data-fields"):
|
||||||
return None
|
return None
|
||||||
assert request.method == "GET"
|
assert request.method == "GET"
|
||||||
@@ -53,3 +55,18 @@ def catalog_response(request, fields=None):
|
|||||||
rows = rows[:50] + [rows[49]] + rows[50:]
|
rows = rows[:50] + [rows[49]] + rows[50:]
|
||||||
offset, limit = int(params.get("offset", 0)), int(params.get("limit", 50))
|
offset, limit = int(params.get("offset", 0)), int(params.get("limit", 50))
|
||||||
return httpx.Response(200, json={"results": rows[offset : offset + limit]})
|
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}}}}
|
||||||
|
|||||||
@@ -26,6 +26,8 @@ async def catalog(logged_in, app):
|
|||||||
401, headers={"WWW-Authenticate": "persona", "Location": "/authentication/persona/test"}
|
401, headers={"WWW-Authenticate": "persona", "Location": "/authentication/persona/test"}
|
||||||
)
|
)
|
||||||
return httpx.Response(201, json={"token": {"expiry": 14400}})
|
return httpx.Response(201, json={"token": {"expiry": 14400}})
|
||||||
|
if request.method == "OPTIONS":
|
||||||
|
return catalog_response(request)
|
||||||
assert request.method == "GET"
|
assert request.method == "GET"
|
||||||
if request.url.path == "/users/self":
|
if request.url.path == "/users/self":
|
||||||
return httpx.Response(200, json={"id": "TEST_USER"})
|
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
|
).status_code == 422
|
||||||
assert (await prepare(client, version, dataset_id="TEST_NEWS")).status_code == 409
|
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, "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"])
|
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 subset.status_code == 201 and len(subset.json()["field_ids"]) == 122
|
||||||
assert "TEST_FIN_110" not in subset.json()["field_ids"]
|
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]
|
assert job["status"] == "completed" and delays == [2]
|
||||||
manifest = await search(client, "/datasets/TEST_FIN/collection")
|
manifest = await search(client, "/datasets/TEST_FIN/collection")
|
||||||
assert manifest["collection_version"] == job["id"] and len(manifest["field_ids"]) == 123
|
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
|
state["persona"] = True
|
||||||
runner.client.authenticated = False
|
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 waiting["status"] == "waiting_auth"
|
||||||
assert (await search(client, "/datasets/TEST_FIN/collection")) == manifest
|
assert (await search(client, "/datasets/TEST_FIN/collection")) == manifest
|
||||||
await runner.disconnect()
|
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(connect["id"])
|
||||||
await runner.execute(waiting["id"])
|
await runner.execute(waiting["id"])
|
||||||
assert (await client.get("/api/v1/sync-jobs/" + waiting["id"])).json()["status"] == "completed"
|
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
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -35,6 +35,28 @@ type Scope = {
|
|||||||
universe: string;
|
universe: string;
|
||||||
delay: number;
|
delay: number;
|
||||||
};
|
};
|
||||||
|
type PlatformOption = Omit<Scope, "universe"> & { universes: string[] };
|
||||||
|
type PlatformOptions = { instrument_options: PlatformOption[] };
|
||||||
|
function normalizeScope(rows: PlatformOption[], preferred: Scope): Scope {
|
||||||
|
const typed = rows.filter(
|
||||||
|
(r) => r.instrument_type === preferred.instrument_type,
|
||||||
|
);
|
||||||
|
const regional = typed.filter((r) => r.region === preferred.region);
|
||||||
|
const row =
|
||||||
|
regional.find((r) => r.delay === preferred.delay) ??
|
||||||
|
regional[0] ??
|
||||||
|
typed[0] ??
|
||||||
|
rows[0];
|
||||||
|
if (!row) return preferred;
|
||||||
|
return {
|
||||||
|
instrument_type: row.instrument_type,
|
||||||
|
region: row.region,
|
||||||
|
delay: row.delay,
|
||||||
|
universe: row.universes.includes(preferred.universe)
|
||||||
|
? preferred.universe
|
||||||
|
: row.universes[0],
|
||||||
|
};
|
||||||
|
}
|
||||||
type Note = { note: string; version: number; updated_at: string };
|
type Note = { note: string; version: number; updated_at: string };
|
||||||
type Entry = {
|
type Entry = {
|
||||||
id: string;
|
id: string;
|
||||||
@@ -136,7 +158,18 @@ export function DatasetPage({
|
|||||||
universe: "TOP3000",
|
universe: "TOP3000",
|
||||||
delay: 1,
|
delay: 1,
|
||||||
});
|
});
|
||||||
const [scopes, setScopes] = useState<Record<string, string[]>>({});
|
const scopeRef = useRef(scope);
|
||||||
|
scopeRef.current = scope;
|
||||||
|
const [scopes, setScopes] = useState<PlatformOption[]>([]);
|
||||||
|
const [scopeError, setScopeError] = useState("");
|
||||||
|
const [scopeRevision, setScopeRevision] = useState(0);
|
||||||
|
const selectedOption = scopes.find(
|
||||||
|
(r) =>
|
||||||
|
r.instrument_type === scope.instrument_type &&
|
||||||
|
r.region === scope.region &&
|
||||||
|
r.delay === scope.delay,
|
||||||
|
);
|
||||||
|
const scopeValid = !!selectedOption?.universes.includes(scope.universe);
|
||||||
const [browse, setBrowse] = useState<Browse>(initialBrowse);
|
const [browse, setBrowse] = useState<Browse>(initialBrowse);
|
||||||
const [fieldBrowse, setFieldBrowse] = useState<Browse>(initialBrowse);
|
const [fieldBrowse, setFieldBrowse] = useState<Browse>(initialBrowse);
|
||||||
const [size, setSize] = useState(25);
|
const [size, setSize] = useState(25);
|
||||||
@@ -370,11 +403,34 @@ export function DatasetPage({
|
|||||||
if (account) setSize(account.page_size);
|
if (account) setSize(account.page_size);
|
||||||
}, [account?.page_size]);
|
}, [account?.page_size]);
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (active)
|
if (!active) return;
|
||||||
api<Record<string, string[]>>("/catalog/scopes")
|
let live = true;
|
||||||
.then(setScopes)
|
setScopes([]);
|
||||||
.catch((e) => setError(e.message));
|
setScopeError("");
|
||||||
}, [active]);
|
api<PlatformOptions>("/catalog/scopes")
|
||||||
|
.then((result) => {
|
||||||
|
if (!live) return;
|
||||||
|
if (!result.instrument_options.length)
|
||||||
|
throw new Error("平台未返回可用研究范围");
|
||||||
|
setScopes(result.instrument_options);
|
||||||
|
const next = normalizeScope(
|
||||||
|
result.instrument_options,
|
||||||
|
scopeRef.current,
|
||||||
|
);
|
||||||
|
if (
|
||||||
|
(Object.keys(next) as (keyof Scope)[]).some(
|
||||||
|
(key) => next[key] !== scopeRef.current[key],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
changeScope(next);
|
||||||
|
})
|
||||||
|
.catch((e) => {
|
||||||
|
if (live) setScopeError(e.message);
|
||||||
|
});
|
||||||
|
return () => {
|
||||||
|
live = false;
|
||||||
|
};
|
||||||
|
}, [active, version, scopeRevision]);
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (!active) return;
|
if (!active) return;
|
||||||
let live = true;
|
let live = true;
|
||||||
@@ -535,6 +591,7 @@ export function DatasetPage({
|
|||||||
if (!field) setSelected(row);
|
if (!field) setSelected(row);
|
||||||
};
|
};
|
||||||
async function sync(datasetId?: string) {
|
async function sync(datasetId?: string) {
|
||||||
|
if (!scopeValid) return;
|
||||||
setBusy(true);
|
setBusy(true);
|
||||||
try {
|
try {
|
||||||
await post("/catalog/sync-jobs", {
|
await post("/catalog/sync-jobs", {
|
||||||
@@ -848,6 +905,7 @@ export function DatasetPage({
|
|||||||
<>
|
<>
|
||||||
<section className="catalog-page">
|
<section className="catalog-page">
|
||||||
{[
|
{[
|
||||||
|
"InstrumentType",
|
||||||
"Region",
|
"Region",
|
||||||
"Universe",
|
"Universe",
|
||||||
"Delay",
|
"Delay",
|
||||||
@@ -865,34 +923,68 @@ export function DatasetPage({
|
|||||||
</span>
|
</span>
|
||||||
))}
|
))}
|
||||||
<div className="catalog-tools">
|
<div className="catalog-tools">
|
||||||
|
<Select
|
||||||
|
aria-labelledby="catalog-label-InstrumentType"
|
||||||
|
value={scope.instrument_type}
|
||||||
|
disabled={!scopes.length}
|
||||||
|
optionList={options([
|
||||||
|
...new Set(scopes.map((r) => r.instrument_type)),
|
||||||
|
])}
|
||||||
|
onChange={(v) =>
|
||||||
|
changeScope(
|
||||||
|
normalizeScope(scopes, {
|
||||||
|
...scope,
|
||||||
|
instrument_type: String(v),
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
/>
|
||||||
<Select
|
<Select
|
||||||
aria-labelledby="catalog-label-Region"
|
aria-labelledby="catalog-label-Region"
|
||||||
value={scope.region}
|
value={scope.region}
|
||||||
optionList={options(Object.keys(scopes))}
|
disabled={!scopes.length}
|
||||||
|
optionList={options([
|
||||||
|
...new Set(
|
||||||
|
scopes
|
||||||
|
.filter((r) => r.instrument_type === scope.instrument_type)
|
||||||
|
.map((r) => r.region),
|
||||||
|
),
|
||||||
|
])}
|
||||||
onChange={(v) =>
|
onChange={(v) =>
|
||||||
changeScope({
|
changeScope(
|
||||||
...scope,
|
normalizeScope(scopes, { ...scope, region: String(v) }),
|
||||||
region: String(v),
|
)
|
||||||
universe: scopes[String(v)][0],
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
/>
|
/>
|
||||||
<Select
|
<Select
|
||||||
aria-labelledby="catalog-label-Universe"
|
aria-labelledby="catalog-label-Universe"
|
||||||
value={scope.universe}
|
value={scope.universe}
|
||||||
optionList={options(scopes[scope.region] ?? [])}
|
disabled={!selectedOption}
|
||||||
|
optionList={options(selectedOption?.universes ?? [])}
|
||||||
onChange={(v) => changeScope({ ...scope, universe: String(v) })}
|
onChange={(v) => changeScope({ ...scope, universe: String(v) })}
|
||||||
/>
|
/>
|
||||||
<Select
|
<Select
|
||||||
aria-labelledby="catalog-label-Delay"
|
aria-labelledby="catalog-label-Delay"
|
||||||
value={scope.delay}
|
value={scope.delay}
|
||||||
optionList={[
|
disabled={!scopes.length}
|
||||||
{ label: "Delay 0", value: 0 },
|
optionList={scopes
|
||||||
{ label: "Delay 1", value: 1 },
|
.filter(
|
||||||
]}
|
(r) =>
|
||||||
onChange={(v) => changeScope({ ...scope, delay: Number(v) })}
|
r.instrument_type === scope.instrument_type &&
|
||||||
|
r.region === scope.region,
|
||||||
|
)
|
||||||
|
.map((r) => ({ label: `Delay ${r.delay}`, value: r.delay }))}
|
||||||
|
onChange={(v) =>
|
||||||
|
changeScope(
|
||||||
|
normalizeScope(scopes, { ...scope, delay: Number(v) }),
|
||||||
|
)
|
||||||
|
}
|
||||||
/>
|
/>
|
||||||
<Button loading={busy} onClick={() => void sync()}>
|
<Button
|
||||||
|
loading={busy}
|
||||||
|
disabled={!scopeValid}
|
||||||
|
onClick={() => void sync()}
|
||||||
|
>
|
||||||
同步目录
|
同步目录
|
||||||
</Button>
|
</Button>
|
||||||
<Button
|
<Button
|
||||||
@@ -907,6 +999,19 @@ export function DatasetPage({
|
|||||||
{formatTime(data.synced_at, account?.timezone)}
|
{formatTime(data.synced_at, account?.timezone)}
|
||||||
</span>
|
</span>
|
||||||
</div>
|
</div>
|
||||||
|
{scopeError && (
|
||||||
|
<Banner
|
||||||
|
type="danger"
|
||||||
|
description={
|
||||||
|
<span>
|
||||||
|
{scopeError}{" "}
|
||||||
|
<Button onClick={() => setScopeRevision((v) => v + 1)}>
|
||||||
|
重试获取选项
|
||||||
|
</Button>
|
||||||
|
</span>
|
||||||
|
}
|
||||||
|
/>
|
||||||
|
)}
|
||||||
{error && <Banner type="danger" description={error} />}
|
{error && <Banner type="danger" description={error} />}
|
||||||
<section className="library-panel">
|
<section className="library-panel">
|
||||||
<div className="catalog-tools">
|
<div className="catalog-tools">
|
||||||
|
|||||||
@@ -255,3 +255,38 @@ test("非首页排除、聊天暂存详情、遮罩逐层关闭和多尺寸", as
|
|||||||
await expect(fields).not.toBeVisible();
|
await expect(fields).not.toBeVisible();
|
||||||
expect(errors).toEqual([]);
|
expect(errors).toEqual([]);
|
||||||
});
|
});
|
||||||
|
|
||||||
|
test("平台选项失败可重试,地区联动限制 Delay 和 Universe", async ({ page }) => {
|
||||||
|
await setup(page);
|
||||||
|
await page.getByLabel("Delay", { exact: true }).click();
|
||||||
|
await page.getByRole("option", { name: /Delay 0/ }).click();
|
||||||
|
await page.getByLabel("Region", { exact: true }).click();
|
||||||
|
await page.getByRole("option", { name: /IND/ }).click();
|
||||||
|
await expect(page.getByLabel("Universe", { exact: true })).toContainText(
|
||||||
|
"TOP500",
|
||||||
|
);
|
||||||
|
await expect(page.getByLabel("Delay", { exact: true })).toContainText(
|
||||||
|
"Delay 1",
|
||||||
|
);
|
||||||
|
await page.getByLabel("Delay", { exact: true }).click();
|
||||||
|
await expect(page.getByRole("option", { name: /Delay 0/ })).toHaveCount(0);
|
||||||
|
await page.keyboard.press("Escape");
|
||||||
|
await page.route("**/api/v1/catalog/scopes", (route) =>
|
||||||
|
route.fulfill({
|
||||||
|
status: 502,
|
||||||
|
contentType: "application/json",
|
||||||
|
body: JSON.stringify({ detail: "测试平台选项不可用" }),
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
await page.getByRole("button", { name: "Alpha 管理", exact: true }).click();
|
||||||
|
await page.getByRole("button", { name: "数据集", exact: true }).click();
|
||||||
|
await expect(page.getByText("测试平台选项不可用")).toBeVisible();
|
||||||
|
await expect(
|
||||||
|
page.getByRole("button", { name: "同步目录", exact: true }),
|
||||||
|
).toBeDisabled();
|
||||||
|
await page.unroute("**/api/v1/catalog/scopes");
|
||||||
|
await page.getByRole("button", { name: "重试获取选项" }).click();
|
||||||
|
await expect(
|
||||||
|
page.getByRole("button", { name: "同步目录", exact: true }),
|
||||||
|
).toBeEnabled();
|
||||||
|
});
|
||||||
|
|||||||
Reference in New Issue
Block a user