Files
worldquant-alpha-system/backend/app/ai/contracts.py
T

71 lines
2.8 KiB
Python
Raw Normal View History

"""Public AI contracts. Client input never contains provider history or tool results."""
from typing import Literal
from urllib.parse import urlsplit
from pydantic import Field, SecretStr, field_validator
from ..catalog.contracts import Scope
from ..schemas import AlphaFilters, Contract
class ModelSettingsInput(Contract):
base_url: str = Field(max_length=2000)
api_key: SecretStr | None = None
model: str = Field(min_length=1, max_length=200)
protocol: Literal["chat_completions", "responses"] = "chat_completions"
enabled: bool = False
@field_validator("base_url")
@classmethod
def valid_url(cls, value):
value = value.strip().rstrip("/")
url = urlsplit(value)
if (
url.scheme not in ("http", "https")
or not url.hostname
or url.username
or url.password
or url.query
or url.fragment
):
raise ValueError("请输入不含账户、查询参数或片段的 HTTP/HTTPS API 根地址")
if url.path.endswith(("/chat/completions", "/responses")):
raise ValueError("请填写 API 根地址,例如 https://example.com/v1")
return value
class PageContext(Contract):
page: Literal["alphas", "account", "datasets", "backtests", "operators", "templates", "features", "variants"] = "alphas"
research_asset_id: str | None = Field(default=None, max_length=36)
research_experiment_id: str | None = Field(default=None, max_length=36)
catalog_scope: Scope | None = None
dataset_id: str | None = Field(default=None, min_length=1, max_length=200)
field_id: str | None = Field(default=None, min_length=1, max_length=200)
collection_version: str | None = Field(default=None, min_length=1, max_length=36)
template_input_id: str | None = Field(default=None, min_length=1, max_length=36)
unsaved_field_selection: bool = False
backtest_run_id: str | None = Field(default=None, max_length=36)
backtest_preview_id: str | None = Field(default=None, max_length=36)
backtest_draft_id: str | None = Field(default=None, max_length=36)
alpha_id: str | None = Field(default=None, max_length=100, pattern=r"^[A-Za-z0-9_-]+$")
selected_ids: list[str] = Field(default_factory=list, max_length=100)
filters: AlphaFilters = Field(default_factory=AlphaFilters)
@field_validator("selected_ids")
@classmethod
def check_ids(cls, value):
from ..schemas import valid_ids
return valid_ids(value) if value else []
class RunInput(Contract):
request_id: str = Field(min_length=1, max_length=100, pattern=r"^[A-Za-z0-9_-]+$")
message: str = Field(min_length=1, max_length=20000)
context: PageContext = Field(default_factory=PageContext)
class Decision(Contract):
approved: bool