"""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