Files
worldquant-alpha-system/backend/app/research_access/contracts.py
T
yuxuanhui 7e990b9a69
Deploy production / deploy (push) Successful in 57s
feat: 增加仅检查 MCP 工具并完善已提交 Alpha 指标展示
2026-09-11 13:14:27 +08:00

194 lines
5.5 KiB
Python

"""Bounded direct research inputs; unknown properties are rejected at the interface."""
from datetime import date, datetime
from typing import Annotated, Literal
from pydantic import Field, model_validator
from ..backtests.contracts import Candidate, SimulationSettings
from ..catalog.contracts import CatalogFilters, Scope
from ..schemas import Contract
Identifier = Annotated[str, Field(min_length=1, max_length=100)]
RunId = Annotated[str, Field(min_length=1, max_length=36)]
AlphaId = Annotated[str, Field(min_length=1, max_length=100, pattern=r"^[A-Za-z0-9_-]+$")]
class Empty(Contract):
pass
class Authentication(Contract):
action: Literal["connect", "verify"] = "connect"
class ConnectionReference(Contract):
job_id: RunId | None = None
class Page(Contract):
limit: int = Field(default=25, ge=1, le=100)
offset: int = Field(default=0, ge=0)
class CompleteSettings(SimulationSettings):
@model_validator(mode="before")
@classmethod
def complete(cls, value):
if isinstance(value, dict) and set(cls.model_fields) - value.keys():
raise ValueError("必须提供每项完整设置;先读取 get_research_capabilities")
return value
model_config = {"json_schema_extra": {"required": list(SimulationSettings.model_fields)}}
class DirectCandidate(Candidate):
settings: CompleteSettings
class Provenance(Contract):
reference: str | None = Field(default=None, max_length=200)
batch_id: str | None = Field(default=None, max_length=200)
hypothesis: str | None = Field(default=None, max_length=2000)
parent_run_id: RunId | None = None
class Submit(Contract):
name: str = Field(min_length=1, max_length=200)
candidates: list[DirectCandidate] = Field(min_length=1, max_length=100)
idempotency_key: Identifier
duplicate_policy: Literal["reject", "rerun"] = "reject"
source: Provenance = Field(default_factory=Provenance)
@model_validator(mode="after")
def unique_ids(self):
if len({c.client_item_id for c in self.candidates}) != len(self.candidates):
raise ValueError("client_item_id 必须唯一")
return self
class Control(Contract):
run_id: RunId
action: Literal["pause", "resume", "stop", "recover"]
expected_version: int = Field(ge=1)
idempotency_key: Identifier
class CatalogSearch(Contract):
filters: CatalogFilters
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
class Scopes(Contract):
kind: Literal["scopes"]
class SettingOptions(Page):
kind: Literal["settings"]
class Operators(Page):
kind: Literal["operators"]
q: str = Field(default="", max_length=300)
category: str | None = None
class Availability(Contract):
kind: Literal["field_availability"]
field_id: Identifier
scope: Scope
class Metadata(Contract):
query: Annotated[Scopes | SettingOptions | Operators | Availability, Field(discriminator="kind")]
class CatalogRefresh(Contract):
kind: Literal["catalog"]
scope: Scope
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
class OperatorsRefresh(Contract):
kind: Literal["operators"]
class SettingsRefresh(Contract):
kind: Literal["settings"]
class PnlRefresh(Contract):
kind: Literal["pnl"]
alpha_ids: list[Identifier] = Field(min_length=1, max_length=100)
class Refresh(Contract):
query: Annotated[
CatalogRefresh | OperatorsRefresh | SettingsRefresh | Availability | PnlRefresh,
Field(discriminator="kind"),
]
class JobReference(Page):
job_id: RunId
class SelfCorrelationCheck(Contract):
alpha_ids: list[AlphaId] = Field(min_length=1, max_length=100)
class SelfCorrelationReference(Contract):
alpha_id: AlphaId
class SubmissionCheck(Contract):
alpha_id: AlphaId
snapshot: str = Field(pattern=r"^[a-f0-9]{64}$")
descriptions: dict[str, str] = Field(min_length=1, max_length=2)
class History(Page):
source: str | None = Field(default=None, max_length=100)
reference: str | None = Field(default=None, max_length=200)
status: str | None = Field(default=None, max_length=30)
created_from: datetime | None = None
created_to: datetime | None = None
scope: Scope | None = None
q: str = Field(default="", max_length=300)
candidates: list[DirectCandidate] | None = Field(default=None, min_length=1, max_length=100)
@model_validator(mode="after")
def dates(self):
for value in (self.created_from, self.created_to):
if value and not value.tzinfo:
raise ValueError("时间须包含时区")
if self.created_from and self.created_to and self.created_from > self.created_to:
raise ValueError("起始时间不能晚于结束时间")
return self
class RunReference(Contract):
run_id: RunId
after: int | None = Field(default=None, ge=0)
event_limit: int = Field(default=25, ge=1, le=100)
class Results(Page):
run_id: RunId
item_ids: list[RunId] | None = Field(default=None, min_length=1, max_length=100)
class Artifact(Page):
item_id: RunId
kind: Literal["snapshot", "pnl"]
date_from: date | None = None
date_to: date | None = None
@model_validator(mode="after")
def dates(self):
if self.kind == "snapshot" and (self.date_from or self.date_to):
raise ValueError("日期筛选仅用于 PnL")
if self.date_from and self.date_to and self.date_from > self.date_to:
raise ValueError("起始日期不能晚于结束日期")
return self