71 lines
2.7 KiB
Python
71 lines
2.7 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)
|
|
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", "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
|