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

90 lines
3.1 KiB
Python

"""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)
description_model: str = Field(default="", max_length=200)
protocol: Literal["chat_completions", "responses"] = "chat_completions"
enabled: bool = False
@field_validator("description_model")
@classmethod
def clean_description_model(cls, value):
return value.strip()
@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[
"home",
"alphas",
"account",
"datasets",
"backtests",
"operators",
"templates",
"features",
"variants",
"pipeline",
"quantflow",
] = "alphas"
research_run_id: str | None = Field(default=None, max_length=36)
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