2026-09-08 12:43:00 +08:00
|
|
|
"""Explicit snapshot and field-binding contracts for research producers."""
|
|
|
|
|
|
|
|
|
|
from typing import Literal
|
|
|
|
|
|
|
|
|
|
from pydantic import Field, model_validator
|
|
|
|
|
|
|
|
|
|
from ..backtests.contracts import SimulationSettings, Source
|
|
|
|
|
from ..catalog.contracts import Scope
|
|
|
|
|
from ..schemas import Contract
|
2026-09-08 21:29:21 +08:00
|
|
|
from .expressions import PLACEHOLDER
|
2026-09-08 12:43:00 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
class ResearchInputSelection(Contract):
|
|
|
|
|
scope: Scope
|
|
|
|
|
dataset_id: str = Field(min_length=1, max_length=200)
|
|
|
|
|
collection_version: str = Field(min_length=1, max_length=36)
|
|
|
|
|
field_ids: list[str] = Field(min_length=1, max_length=100)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class InputPageArgs(Contract):
|
|
|
|
|
input_id: str = Field(min_length=1, max_length=36)
|
|
|
|
|
limit: int = Field(default=25, ge=1, le=100)
|
|
|
|
|
offset: int = Field(default=0, ge=0)
|
|
|
|
|
q: str = Field(default="", max_length=300)
|
|
|
|
|
field_type: str | None = Field(default=None, max_length=100)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class FieldBinding(Contract):
|
|
|
|
|
field_id: str = Field(min_length=1, max_length=200, pattern=r"^[A-Za-z_][A-Za-z0-9_]*$")
|
|
|
|
|
field_type: Literal["MATRIX", "VECTOR", "GROUP"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class ResearchCandidate(Contract):
|
|
|
|
|
client_item_id: str = Field(min_length=1, max_length=100)
|
|
|
|
|
expression_template: str = Field(min_length=1, max_length=20000)
|
|
|
|
|
bindings: dict[str, FieldBinding] = Field(min_length=1, max_length=100)
|
|
|
|
|
settings: SimulationSettings
|
|
|
|
|
|
|
|
|
|
@model_validator(mode="after")
|
|
|
|
|
def complete_bindings(self):
|
|
|
|
|
placeholders = set(PLACEHOLDER.findall(self.expression_template))
|
|
|
|
|
remainder = PLACEHOLDER.sub("", self.expression_template)
|
|
|
|
|
if placeholders != set(self.bindings) or "{" in remainder or "}" in remainder:
|
|
|
|
|
raise ValueError("模板占位符必须与字段绑定逐一对应,例如 rank({price})")
|
|
|
|
|
return self
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class ChatboxResearchInput(Contract):
|
|
|
|
|
"""Chatbox provenance is supplied by the server, never by model arguments."""
|
|
|
|
|
|
|
|
|
|
name: str = Field(min_length=1, max_length=200)
|
|
|
|
|
hypothesis: str = Field(min_length=1, max_length=2000)
|
|
|
|
|
template_input_id: str = Field(min_length=1, max_length=36)
|
|
|
|
|
candidates: list[ResearchCandidate] = Field(min_length=1, max_length=100)
|
|
|
|
|
|
|
|
|
|
@model_validator(mode="after")
|
|
|
|
|
def unique_candidates(self):
|
|
|
|
|
if len({item.client_item_id for item in self.candidates}) != len(self.candidates):
|
|
|
|
|
raise ValueError("client_item_id 在候选集合内必须唯一")
|
|
|
|
|
return self
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class ResearchPreviewInput(ChatboxResearchInput):
|
|
|
|
|
source: Source = Field(default_factory=Source)
|