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

158 lines
5.5 KiB
Python
Raw Normal View History

"""Capability contracts shared by domain adapters and the AI executor.
Handlers receive business operations, never model history or client approval data.
The executor owns authorization, savepoints, audit commits and after-commit timing.
"""
from __future__ import annotations
from collections.abc import Awaitable, Callable, Iterable
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Any, Literal, get_args
from fastapi import HTTPException
from fastapi.encoders import jsonable_encoder
from pydantic import Field
from ..schemas import Contract
if TYPE_CHECKING:
from ..business import Business
from ..jobs import Runner
from ..worldquant import WqClient
class EmptyArgs(Contract):
pass
class ResultMetadata(Contract):
source: str = "local_database"
observed_at: datetime
nulls: str = "null 表示来源未提供,不等于零"
units: dict[str, str] = Field(
default_factory=lambda: {
"turnover": "比例,0.15 = 15%",
"returns": "比例",
"drawdown": "比例",
"margin": "比例",
"pnl": "供应商原始累计值,未提供货币/规模单位",
}
)
@dataclass(frozen=True)
class ToolContext:
business: Business
platform_client: WqClient | None = None
Handler = Callable[[ToolContext, Any], Awaitable[dict]]
ConfirmedHandler = Callable[[ToolContext, Any, dict], Awaitable[dict]]
Notification = Callable[["Runner", dict], Awaitable[None]]
Effect = Literal["query", "prepare", "confirm"]
Resource = Literal["alphas", "datasets", "backtests", "jobs", "account"]
@dataclass(frozen=True, kw_only=True)
class Capability:
"""One complete tool definition; invalid policy combinations fail at assembly.
``invoke`` accepts untrusted arguments for query/prepare and returns unabridged
business data. Confirmed handlers are only called by AIRuntime after its gate.
"""
name: str
schema: type[Contract]
description: str
label: str
renderer: str
effect: Effect
handler: Handler | None = None
preview: Handler | None = None
execute: ConfirmedHandler | None = None
after_commit: Notification | None = None
refresh: tuple[Resource, ...] = ()
source: str = "local_database"
def __post_init__(self):
if not self.name or not self.label or not self.renderer:
raise ValueError("Capability needs a name, label and renderer")
if any(resource not in get_args(Resource) for resource in self.refresh):
raise ValueError("Capability refresh target must be a workspace resource")
if self.effect == "confirm":
if self.handler is not None or self.preview is None or self.execute is None:
raise ValueError("Confirmed capability needs preview and execute only")
elif self.effect in ("query", "prepare"):
if self.handler is None or any((self.preview, self.execute, self.after_commit)):
raise ValueError("Query/prepare capability needs a handler and cannot notify execution")
if self.effect == "query" and self.refresh:
raise ValueError("Queries cannot invalidate business resources")
else:
raise ValueError("Capability needs an explicit effect")
@property
def requires_confirmation(self):
return self.effect == "confirm"
def presentation(self):
return {
"label": self.label,
"renderer": self.renderer,
"effect": self.effect,
"refresh": list(self.refresh),
}
async def invoke(self, context: ToolContext, arguments: dict):
"""Validate query/prepare input; raise 409 if used to bypass confirmation."""
if self.requires_confirmation:
raise HTTPException(409, "此能力必须先预览并确认")
data = await self.handler(context, self.schema.model_validate(arguments))
return jsonable_encoder(
{
**data,
"_meta": ResultMetadata(
source=self.source, observed_at=datetime.now(timezone.utc)
).model_dump(mode="json"),
}
)
def assemble(groups: Iterable[Iterable[Capability]]) -> dict[str, Capability]:
"""Assemble explicit domain definitions, rejecting ambiguous tool names."""
result = {}
for group in groups:
for capability in group:
if capability.name in result:
raise ValueError(f"Duplicate capability: {capability.name}")
result[capability.name] = capability
return result
def model_result(value):
"""Bound model context without mutating persisted data; expose any truncation."""
truncated = False
def bound(item):
nonlocal truncated
if isinstance(item, str) and len(item) > 2000:
truncated = True
return item[:2000] + "…(已截断)"
if isinstance(item, list):
truncated = truncated or len(item) > 100
return [bound(v) for v in item[:100]]
if isinstance(item, dict):
truncated = truncated or len(item) > 100
return {k: bound(v) for k, v in list(item.items())[:100]}
return item
result = bound(value)
if isinstance(result, dict) and truncated:
result["_meta"] = {
**result.get("_meta", {}),
"truncated": True,
"detail": "模型摘要已截断;完整内容保留在业务记录,可按引用分页读取",
}
return result