refactor: unify AI capabilities and workspace integration

This commit is contained in:
yuxuanhui
2026-09-08 19:28:44 +08:00
parent 3d26827b49
commit b604e6050e
23 changed files with 1821 additions and 805 deletions
+157
View File
@@ -0,0 +1,157 @@
"""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