refactor: unify AI capabilities and workspace integration
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user