diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 46d110b7..edd077c5 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -604,8 +604,14 @@ system (`agent_engine/runtime/hooks/`, start/stop, run start/end/error, tool error, and `transform_tool_result` (lets trusted code truncate/redact/normalize a tool result before it reaches the model) — for auth, policy, audit, and context enrichment, distinct from -LLM-invoked tools. Per-tool `input_policy` (trusted parameter injection) is -still not implemented — see §11. +LLM-invoked tools. Tool results cross the runtime through a provider-independent +normalized value whose model text, structured data, and safe artifact metadata +remain separate and survive versioned, concurrency-safe idempotent replay. Only +the text projection enters the LangChain conversation; structured values remain +available to trusted hooks and the execution ledger (see +[`ADR 0004`](adr/0004-normalized-structured-tool-results.md)). Per-tool +`input_policy` (trusted parameter injection) is still not implemented — see +§11. **Shared tool usage (✅ done, not in the original task list):** `agent_engine/tool_usage/` owns execution metadata: a `ToolUsageRepository` port diff --git a/docs/MCP_AND_TOOLS.md b/docs/MCP_AND_TOOLS.md index 00c7a8ca..5b616f75 100644 --- a/docs/MCP_AND_TOOLS.md +++ b/docs/MCP_AND_TOOLS.md @@ -155,6 +155,66 @@ stops requesting tools. Each call is recorded in the run's tool-usage repository with its `provider` (`"local"` or `"mcp"`), so the origin is tracked for tracing even though it is hidden from the model. +### Normalized tool results + +A successful tool call is represented inside Extra as a provider-independent +`NormalizedToolResult`, not as only a string: + +```python +NormalizedToolResult( + text="Found 2 invoices", + structured={"count": 2}, + artifact={"source": "billing"}, +) +``` + +The fields have separate responsibilities: + +- `text` is the existing model-facing result. MCP text blocks are joined with + deterministic normalization rules. A structured-only result uses canonical + JSON text; a completely empty result uses a stable placeholder. +- `structured` is machine-readable output, including MCP + `structuredContent`. It must be JSON-like and is not copied into a model + message as metadata. For a structured-only result, its canonical JSON is + deliberately used as the model-facing `text` fallback. The shared JSON-safe + value policy limits nesting to 64 levels, total values to 10,000, cumulative + string data to 1,000,000 UTF-8 bytes, individual keys to 1,024 bytes, + cumulative key data to 256,000 bytes, and final canonical JSON to 1 MiB. +- `artifact` carries bounded, relevant artifact metadata. Raw in-memory binary + bodies and oversized values are omitted and replaced by type/size metadata. + The adapter limits depth to 8, each collection to 128 entries, total visited + values to 1,024, individual strings to 8,192 characters, cumulative string + content to 32,768 characters, keys to 256 characters, and final canonical + metadata to 65,536 characters. + +`langchain-mcp-adapters` currently exposes MCP `structuredContent` through the +`ToolMessage.artifact.structured_content` path. Extra reads that provider +contract once at the tool adapter boundary and requires version 0.2 or newer, +where that artifact contract is available. It passes only its own normalized +result deeper into the runtime. Local dictionary/list results are normalized by +the same abstraction; plain string tools remain unchanged. + +Only `text` is appended to the LangChain conversation. The complete normalized +result is available to trusted result hooks and is serialized into the +idempotency ledger as a versioned, JSON-primitive payload. A replay therefore +restores the same text, structured value, and artifact metadata without calling +the provider again. None of the non-text values are added to tool-usage records, +logs, model messages, callbacks, traces, or errors automatically. +The structured-only text fallback is the explicit exception: because that JSON +becomes model text, it is visible wherever ordinary model messages are visible. + +The ledger's atomic claim has one execution owner. Concurrent duplicate callers +wait for that owner and replay its immutable terminal result; terminal rows +cannot be overwritten. Legacy string ledger values restore as text-only +results. Custom repository adapters implement the same claim/wait/complete +contract and must durably serialize the versioned primitive payload. + +Structured values are preserved without application-specific schema validation, +but the runtime enforces its generic JSON-safe shape. A malformed provider +result becomes a controlled failed tool result and is never converted with +`repr()` or an arbitrary `str()` fallback. See +[`ADR 0004`](adr/0004-normalized-structured-tool-results.md). + The engine is driven as an async context manager: `build()` connects MCP servers and discovers tools; `close()` (on context exit) releases them. `run()` does not connect MCP servers on its own — `build()` must run first. @@ -223,6 +283,10 @@ resume updates the existing record instead of adding a second one. Arguments and results are never stored: they may carry sensitive or oversized data, and no consumer of tool usage needs them. +This tool-usage repository is distinct from the private tool-execution +idempotency ledger. Tool usage stores no result values; the execution ledger +stores `NormalizedToolResult` so a replay can reproduce the completed call. + `ToolUsageRepository` is an abstract base class with three operations (`record`, `list_for_run`, `list_for_conversation`). The engine ships a process-local adapter (`InMemoryToolUsageRepository`); a distributed deployment supplies a Redis- or diff --git a/docs/RUNTIME_HOOKS.md b/docs/RUNTIME_HOOKS.md index fb693938..d13dcb4b 100644 --- a/docs/RUNTIME_HOOKS.md +++ b/docs/RUNTIME_HOOKS.md @@ -40,7 +40,7 @@ explicit-ref mode) and whether a returned value is used. | `on_run_error` | when a run fails | the `BaseException` | ignored (never masks the error) | | `before_tool_call` | before every local or MCP tool call | `ToolRequestContext` | ignored (observe-only) | | `after_tool_call` | after a local or MCP tool call **succeeds** | `ToolCallContext` (status `succeeded`) | ignored | -| `transform_tool_result` | after a tool **succeeds**, before its result is appended to the conversation | `ToolResultContext` (carries the `result`) | updated `ToolResultContext` (or `None`) | +| `transform_tool_result` | after a tool **succeeds**, before its result is appended to the conversation | `ToolResultContext` (carries text, structured result, and artifact metadata) | updated `ToolResultContext` (or `None`) | | `on_tool_error` | when a local or MCP tool call **fails** | `ToolCallContext` (status `failed`) | ignored | | `before_mcp_request` | before every outgoing MCP HTTP request | `McpRequestContext` | updated `McpRequestContext` (or `None`) | | `after_mcp_response` | after every MCP HTTP response | `McpResponseContext` | ignored (observe-only) | @@ -343,16 +343,26 @@ EngineContext(system_name, metadata) RunEndContext(run_id, system_name, status, visited, used_tool_count, metadata) ToolRequestContext(agent_id, tool_name, provider, server_id, metadata) ToolCallContext(agent_id, tool_name, provider, server_id, status, latency_ms, error, metadata) -ToolResultContext(agent_id, tool_name, provider, result, server_id, latency_ms, metadata) +ToolResultContext(agent_id, tool_name, provider, result, server_id, latency_ms, metadata, structured_result, artifact) McpRequestContext(server_id, url, operation, tool_name, headers, metadata) McpResponseContext(server_id, url, status_code, operation, tool_name, latency_ms, metadata) ``` -All are frozen dataclasses. `RunContext.replace(**changes)` and -`McpRequestContext.with_headers({...})` and `ToolResultContext.with_result(...)` -return updated copies; the `headers` and `metadata` dicts may also be mutated in -place. Hooks never receive raw graph -state. +All are frozen dataclasses. `RunContext.replace(**changes)`, +`McpRequestContext.with_headers({...})`, and the +`ToolResultContext.with_result(...)`, `with_structured_result(...)`, and +`with_artifact(...)` helpers return updated copies; the `headers` and `metadata` +dicts may also be mutated in place. `with_result(...)` changes only model-facing +text and preserves structured output and artifact metadata. Hooks never receive +raw graph state. + +`ToolResultContext.result` remains a string for existing hooks. +`structured_result` carries provider-independent machine-readable output, and +`artifact` contains structurally bounded metadata after binary and oversized +bodies have been omitted. Structural bounds do not identify application +secrets: hook implementations may inspect, redact, or replace these trusted +runtime values but must not log their contents. See +[`ADR 0004`](adr/0004-normalized-structured-tool-results.md). --- diff --git a/docs/adr/0004-normalized-structured-tool-results.md b/docs/adr/0004-normalized-structured-tool-results.md new file mode 100644 index 00000000..053fe4b4 --- /dev/null +++ b/docs/adr/0004-normalized-structured-tool-results.md @@ -0,0 +1,94 @@ +# ADR 0004 — Preserve structured tool results separately from model text + +- **Status:** Accepted +- **Date:** 2026-09-01 + +## Context + +The tool execution path reduced every successful local or MCP result to a +string. Text extraction kept MCP content blocks readable for the model, but it +discarded MCP `structuredContent`, LangChain artifact metadata, and structured +local-tool returns. The idempotency ledger consequently replayed only text, and +`transform_tool_result` hooks could not inspect machine-readable output. + +## Decision + +Extra owns a provider-agnostic frozen `NormalizedToolResult` with independent +`text`, `structured`, and `artifact` fields. + +- The LangChain tool adapter is the normalization boundary. It extracts text + from content blocks, reads MCP structured output from the current + `ToolMessage.artifact.structured_content` contract (and the compatible + `structuredContent` spelling), and removes LangChain-specific objects before + the result enters the runtime. +- Local JSON object/array results that LangChain serializes into + `ToolMessage.content` are parsed back into `structured` while the serialized + model-facing text remains unchanged. +- The model loop consumes only `NormalizedToolResult.text`. It creates a plain + `ToolMessage` without an artifact, so structured values do not enter model + context, LangChain callbacks, or tracing as hidden message metadata. +- A structured-only result receives deterministic canonical JSON as its model + text. A completely empty result receives a stable placeholder. Unsupported + values fail with controlled model text; provider objects are never rendered + through `repr()` or arbitrary `str()` fallback. The structured-only JSON is + ordinary model text and therefore has the same model/callback visibility as + any other tool text. +- `ToolResultContext.result` remains the text field for hook compatibility and + gains additive `structured_result` and `artifact` fields. `with_result()` + changes text only, so truncating text cannot discard structured output. +- The value object validates JSON-like values, stores canonical JSON privately, + and returns copies from its accessors. Provider-owned nested objects therefore + cannot mutate the authoritative result after normalization. Generic depth, + value-count, string/key byte, integer-size, and final encoded-size budgets + prevent structured payloads from multiplying memory use across hooks, + persistence, and replay. +- The tool-execution repository persists a versioned JSON-primitive payload. + An atomic claim identifies one execution owner; concurrent duplicates wait + for a terminal result and replay it without re-invoking the provider. Terminal + rows cannot be overwritten, and legacy string records remain readable as + text-only results. +- Artifact mappings retain structurally safe, bounded metadata under explicit + per-value, aggregate-size, collection-count, key-length, and nesting limits. + Binary and oversized values are replaced by type/size/omission metadata and + are never copied into model context. This is structural safety, not semantic + secret classification; trusted hooks remain responsible for application-level + redaction before persistence when needed. +- Structured values are not validated against application output schemas in + v1, but they must satisfy the generic JSON-safe runtime contract. Malformed + provider results become controlled failed tool results. + +## Contract changes + +- `ToolExecutionRepository.complete()` accepts the versioned persisted payload, + and the port adds `wait_for_completion()`. Custom repository adapters must + serialize all three fields, atomically create claims, reject terminal + overwrites, and wake duplicate waiters after every owner outcome. +- Tool execution states use `ToolExecutionStatus`, a `StrEnum` that remains + wire-compatible with the existing string values while removing duplicated + state literals from the ledger implementation. +- `ToolResultContext` adds optional `structured_result` and `artifact` fields + plus immutable update helpers. Existing text-only hooks remain source + compatible. +- `ToolInvoker.invoke()` returns `NormalizedToolResult`; this is an internal + engine contract consumed by the model loop. + +There is no YAML, HTTP API, widget, or model-facing text contract change. +The runtime dependency floor for `langchain-mcp-adapters` is 0.2 because the +normalizer relies on that release's structured-content artifact contract. + +## Consequences + +- Text-only local tools behave exactly as before. +- MCP text, structured output, and safe artifact metadata can coexist and + survive idempotent replay. +- Hooks can inspect or deliberately replace structured values without parsing + text, while ordinary text transforms preserve them by default. +- Structured values are runtime data and are not automatically written to tool + usage, hidden model-message metadata, logs, traces, or error messages. The + documented structured-only text fallback is the deliberate model-visible + exception. +- Existing persisted string results remain readable. The shipped execution + repository is process-local, so no database migration is required; external + repository adapters must adopt the expanded port and versioned payload. +- Full multimodal artifact rendering, binary persistence, and output-schema + enforcement remain out of scope. diff --git a/pyproject.toml b/pyproject.toml index 54df3eea..7cd5ea94 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -49,7 +49,7 @@ dependencies = [ "langchain>=0.3", "langchain-core>=0.3", "mcp>=1.27,<2", - "langchain-mcp-adapters>=0.1", + "langchain-mcp-adapters>=0.2", "pyjwt>=2.9", "python-dotenv>=1.0", "pydantic-settings>=2.0", diff --git a/src/agent_engine/approvals/__init__.py b/src/agent_engine/approvals/__init__.py index 5cc366ec..a55bd623 100644 --- a/src/agent_engine/approvals/__init__.py +++ b/src/agent_engine/approvals/__init__.py @@ -68,12 +68,15 @@ RunRecord, RunStatus, ToolExecutionRecord, + ToolExecutionStatus, ) from agent_engine.approvals.sanitization import mask_arguments, mask_sensitive from agent_engine.approvals.session_approval_repository import SessionApprovalRepository from agent_engine.approvals.session_approval_store import SessionApprovalStore from agent_engine.approvals.tool_execution_manager import ( + ToolExecutionClaim, ToolExecutionManager, + ToolExecutionStateError, execution_id_for, ) from agent_engine.approvals.tool_execution_repository import ToolExecutionRepository @@ -110,9 +113,12 @@ "SessionApprovalRepository", "SessionApprovalScope", "SessionApprovalStore", + "ToolExecutionClaim", "ToolExecutionManager", "ToolExecutionRecord", "ToolExecutionRepository", + "ToolExecutionStateError", + "ToolExecutionStatus", "ToolInvocation", "ToolNoLongerExists", "UnauthorizedApprover", diff --git a/src/agent_engine/approvals/in_memory_tool_execution_repository.py b/src/agent_engine/approvals/in_memory_tool_execution_repository.py index 99c9cd64..8e898c0b 100644 --- a/src/agent_engine/approvals/in_memory_tool_execution_repository.py +++ b/src/agent_engine/approvals/in_memory_tool_execution_repository.py @@ -3,31 +3,65 @@ from __future__ import annotations import asyncio +import copy +import dataclasses -from agent_engine.approvals.models import ToolExecutionRecord +from agent_engine.approvals.models import ToolExecutionRecord, ToolExecutionStatus from agent_engine.approvals.tool_execution_repository import ToolExecutionRepository +from agent_engine.runtime.tool_results import PersistedToolResult class InMemoryToolExecutionRepository(ToolExecutionRepository): def __init__(self) -> None: self._records: dict[str, ToolExecutionRecord] = {} + self._completed: dict[str, asyncio.Event] = {} self._lock = asyncio.Lock() async def get(self, execution_id: str) -> ToolExecutionRecord | None: async with self._lock: - return self._records.get(execution_id) + record = self._records.get(execution_id) + return _copy_record(record) if record is not None else None async def start(self, record: ToolExecutionRecord) -> tuple[ToolExecutionRecord, bool]: async with self._lock: existing = self._records.get(record.execution_id) if existing is not None: - return existing, False - self._records[record.execution_id] = record - return record, True + return _copy_record(existing), False + stored = _copy_record(record) + self._records[record.execution_id] = stored + self._completed[record.execution_id] = asyncio.Event() + return _copy_record(stored), True - async def complete(self, execution_id: str, status: str, result: str) -> None: + async def wait_for_completion(self, execution_id: str) -> ToolExecutionRecord: async with self._lock: record = self._records.get(execution_id) - if record is not None: - record.status = status - record.result = result + if record is None: + raise KeyError(f"tool execution not found: {execution_id}") + if record.status != ToolExecutionStatus.STARTED: + return _copy_record(record) + completed = self._completed[execution_id] + await completed.wait() + async with self._lock: + return _copy_record(self._records[execution_id]) + + async def complete( + self, + execution_id: str, + status: ToolExecutionStatus, + result: PersistedToolResult, + ) -> None: + if status == ToolExecutionStatus.STARTED: + raise ValueError("completed tool execution must be terminal") + async with self._lock: + record = self._records.get(execution_id) + if record is None: + raise KeyError(f"tool execution not found: {execution_id}") + if record.status != ToolExecutionStatus.STARTED: + raise ValueError(f"tool execution is already terminal: {execution_id}") + record.status = status + record.result = copy.deepcopy(result) + self._completed[execution_id].set() + + +def _copy_record(record: ToolExecutionRecord) -> ToolExecutionRecord: + return dataclasses.replace(record, result=copy.deepcopy(record.result)) diff --git a/src/agent_engine/approvals/models.py b/src/agent_engine/approvals/models.py index d6fca2c7..55dbba5a 100644 --- a/src/agent_engine/approvals/models.py +++ b/src/agent_engine/approvals/models.py @@ -21,6 +21,7 @@ from agent_engine.approvals.errors import InvalidStateTransition from agent_engine.runtime.tool_models import ToolProviderName +from agent_engine.runtime.tool_results import PersistedToolResult class RunStatus(StrEnum): @@ -39,6 +40,12 @@ class ApprovalStatus(StrEnum): REJECTED = "rejected" +class ToolExecutionStatus(StrEnum): + STARTED = "started" + SUCCEEDED = "succeeded" + FAILED = "failed" + + # Allowed forward transitions. Anything not listed is rejected. Terminal run # states have no outgoing path, and approvals cannot move REJECTED -> APPROVED. _RUN_TRANSITIONS: dict[RunStatus, frozenset[RunStatus]] = { @@ -154,6 +161,6 @@ class ToolExecutionRecord: tool_call_id: str run_id: str tool_name: str - status: str = "started" # started | succeeded | failed - result: str | None = None + status: ToolExecutionStatus = ToolExecutionStatus.STARTED + result: PersistedToolResult | None = None created_at: float = field(default_factory=time.time) diff --git a/src/agent_engine/approvals/tool_execution_manager.py b/src/agent_engine/approvals/tool_execution_manager.py index d1e13a11..e4b4e16a 100644 --- a/src/agent_engine/approvals/tool_execution_manager.py +++ b/src/agent_engine/approvals/tool_execution_manager.py @@ -2,10 +2,13 @@ from __future__ import annotations +import asyncio import hashlib +from dataclasses import dataclass -from agent_engine.approvals.models import ToolExecutionRecord +from agent_engine.approvals.models import ToolExecutionRecord, ToolExecutionStatus from agent_engine.approvals.tool_execution_repository import ToolExecutionRepository +from agent_engine.runtime.tool_results import NormalizedToolResult, ToolResultValidationError def execution_id_for(tool_call_id: str, *, salt: str = "") -> str: @@ -19,6 +22,34 @@ def execution_id_for(tool_call_id: str, *, salt: str = "") -> str: return f"exec_{digest[:24]}" +class ToolExecutionStateError(RuntimeError): + """The execution ledger contains an invalid or incomplete terminal row.""" + + +@dataclass(frozen=True) +class ToolExecutionClaim: + """Whether this caller owns execution or must replay a terminal result.""" + + should_execute: bool + status: ToolExecutionStatus | None = None + result: NormalizedToolResult | None = None + + def __post_init__(self) -> None: + if self.should_execute: + if self.status is not None or self.result is not None: + raise ValueError("an execution owner cannot already have a terminal result") + return + if ( + self.status + not in ( + ToolExecutionStatus.SUCCEEDED, + ToolExecutionStatus.FAILED, + ) + or self.result is None + ): + raise ValueError("a replay claim must contain one terminal result") + + class ToolExecutionManager: """Idempotency ledger for tool executions. @@ -31,17 +62,23 @@ def __init__(self, *, execution_repository: ToolExecutionRepository | None = Non self._executions = execution_repository async def already_executed(self, execution_id: str) -> ToolExecutionRecord | None: + """Return the successful ledger row (backwards-compatible API).""" if self._executions is None: return None record = await self._executions.get(execution_id) - if record is not None and record.status == "succeeded": + if record is not None and record.status == ToolExecutionStatus.SUCCEEDED: return record return None + async def restored_result(self, execution_id: str) -> NormalizedToolResult | None: + """Return a successful execution through Extra's current result model.""" + record = await self.already_executed(execution_id) + return _restore_result(record) if record is not None else None + async def begin_execution( self, execution_id: str, *, tool_call_id: str, run_id: str, tool_name: str ) -> bool: - """Return whether this attempt owns the key and may execute.""" + """Create an execution row without waiting (legacy coordination API).""" if self._executions is None: return True _, created = await self._executions.start( @@ -54,6 +91,67 @@ async def begin_execution( ) return created - async def finish_execution(self, execution_id: str, *, status: str, result: str) -> None: + async def claim_execution( + self, execution_id: str, *, tool_call_id: str, run_id: str, tool_name: str + ) -> ToolExecutionClaim: + """Atomically claim a call or wait for its current owner to finish.""" + if self._executions is None: + return ToolExecutionClaim(should_execute=True) + record, created = await self._executions.start( + ToolExecutionRecord( + execution_id=execution_id, + tool_call_id=tool_call_id, + run_id=run_id, + tool_name=tool_name, + ) + ) + if created: + return ToolExecutionClaim(should_execute=True) + if record.status == ToolExecutionStatus.STARTED: + record = await self._executions.wait_for_completion(execution_id) + return ToolExecutionClaim( + should_execute=False, + status=record.status, + result=_restore_result(record), + ) + + async def finish_execution( + self, + execution_id: str, + *, + status: ToolExecutionStatus, + result: NormalizedToolResult | str, + ) -> None: + if status == ToolExecutionStatus.STARTED: + raise ValueError("finished tool execution must be terminal") if self._executions is not None: - await self._executions.complete(execution_id, status=status, result=result) + normalized = ( + result + if isinstance(result, NormalizedToolResult) + else NormalizedToolResult.text_only(result) + ) + completion = asyncio.create_task( + self._executions.complete( + execution_id, + status=status, + result=normalized.to_persisted(), + ) + ) + try: + await asyncio.shield(completion) + except asyncio.CancelledError: + await asyncio.shield(completion) + raise + + +def _restore_result(record: ToolExecutionRecord) -> NormalizedToolResult: + if record.status == ToolExecutionStatus.STARTED or record.result is None: + raise ToolExecutionStateError( + f"tool execution {record.execution_id} has no terminal result" + ) + try: + return NormalizedToolResult.from_persisted(record.result) + except ToolResultValidationError as exc: + raise ToolExecutionStateError( + f"tool execution {record.execution_id} contains an invalid result" + ) from exc diff --git a/src/agent_engine/approvals/tool_execution_repository.py b/src/agent_engine/approvals/tool_execution_repository.py index df5bf913..d7eefb59 100644 --- a/src/agent_engine/approvals/tool_execution_repository.py +++ b/src/agent_engine/approvals/tool_execution_repository.py @@ -4,7 +4,8 @@ from abc import ABC, abstractmethod -from agent_engine.approvals.models import ToolExecutionRecord +from agent_engine.approvals.models import ToolExecutionRecord, ToolExecutionStatus +from agent_engine.runtime.tool_results import PersistedToolResult class ToolExecutionRepository(ABC): @@ -17,5 +18,15 @@ async def start(self, record: ToolExecutionRecord) -> tuple[ToolExecutionRecord, raise NotImplementedError @abstractmethod - async def complete(self, execution_id: str, status: str, result: str) -> None: + async def wait_for_completion(self, execution_id: str) -> ToolExecutionRecord: + """Wait for the owner of an existing execution to publish its result.""" + raise NotImplementedError + + @abstractmethod + async def complete( + self, + execution_id: str, + status: ToolExecutionStatus, + result: PersistedToolResult, + ) -> None: raise NotImplementedError diff --git a/src/agent_engine/engine/langgraph/execution/model_loop.py b/src/agent_engine/engine/langgraph/execution/model_loop.py index 8abe9de1..2a6dfa00 100644 --- a/src/agent_engine/engine/langgraph/execution/model_loop.py +++ b/src/agent_engine/engine/langgraph/execution/model_loop.py @@ -18,6 +18,7 @@ log_limit, ) from agent_engine.runtime.streaming import current_streams +from agent_engine.runtime.tool_results import NormalizedToolResult logger = logging.getLogger(__name__) @@ -50,7 +51,7 @@ async def run_tool_loop( model: Any, context: ModelContext, node_path: str, - invoke_tool: Callable[[dict[str, Any]], Awaitable[str]], + invoke_tool: Callable[[dict[str, Any]], Awaitable[str | NormalizedToolResult]], *, refresh_execution_context: ExecutionContextRefresher | None = None, ) -> Any: @@ -68,21 +69,28 @@ async def run_tool_loop( context.append(response) for tool_call in response.tool_calls: logger.debug( - "[%s] ← tool_call: %s(%s)", + "[%s] ← tool_call: %s(arguments=%d)", node_path, tool_call["name"], - tool_call["args"], + len(tool_call.get("args") or {}), + ) + raw_result = await invoke_tool(tool_call) + result = ( + raw_result + if isinstance(raw_result, NormalizedToolResult) + else NormalizedToolResult.text_only(raw_result) ) - content = await invoke_tool(tool_call) logger.debug( - "[%s] → tool_result[%s]: %s", + "[%s] → tool_result[%s] chars=%d structured=%s artifact=%s", node_path, tool_call["name"], - content[:300], + len(result.text), + result.has_structured, + result.has_artifact, ) context.append( ToolMessage( - content=content, + content=result.text, tool_call_id=tool_call["id"], name=tool_call["name"], ) diff --git a/src/agent_engine/engine/langgraph/tools/tool_invoker.py b/src/agent_engine/engine/langgraph/tools/tool_invoker.py index aeeb3f9b..ea9fc84a 100644 --- a/src/agent_engine/engine/langgraph/tools/tool_invoker.py +++ b/src/agent_engine/engine/langgraph/tools/tool_invoker.py @@ -12,23 +12,30 @@ from __future__ import annotations +import asyncio import logging import time from dataclasses import dataclass from typing import Any -from langchain_core.messages import ToolMessage from langchain_core.tools import BaseTool from agent_engine.approvals.coordinator import ApprovalCoordinator from agent_engine.approvals.invocation import ToolInvocation +from agent_engine.approvals.models import ToolExecutionStatus from agent_engine.approvals.tool_execution_manager import ( + ToolExecutionClaim, ToolExecutionManager, execution_id_for, ) from agent_engine.core.spec import AgentSpec from agent_engine.engine.langgraph.tools.agent_tool_binding import AgentToolBinding from agent_engine.engine.langgraph.tools.tool_gate import DenyTool, ExecuteTool, ToolGate +from agent_engine.engine.langgraph.tools.tool_result_normalizer import ( + ProviderToolResultError, + ToolResultNormalizationError, + normalize_tool_result, +) from agent_engine.logging_config import log from agent_engine.runtime.execution_limiter import ( ExecutionLimitExceeded, @@ -45,6 +52,7 @@ ) from agent_engine.runtime.hooks.models import ToolStatus from agent_engine.runtime.tool_models import ToolProviderName +from agent_engine.runtime.tool_results import NormalizedToolResult from agent_engine.tool_usage.models import ToolCallIdentity, stable_tool_call_id from agent_engine.tool_usage.tracker import ToolUsageTracker @@ -95,35 +103,6 @@ def _elapsed_ms(start: float) -> int: return int((time.perf_counter() - start) * 1000) -def _extract_result_text(result: object) -> str: - """Turn a tool's return value into the text the model reads. - - A ``response_format="content_and_artifact"`` tool (every MCP tool) returns - ``ToolMessage.content`` as either a plain string or a list of MCP content - blocks (``{"type": "text", "text": ...}``) rather than a single string. - ``str()`` on that list produces its Python repr — visible punctuation and - key names — instead of the text itself, so each shape needs its own - handling rather than one blind stringification. - """ - content = result.content if isinstance(result, ToolMessage) else result - if isinstance(content, str): - return content - if isinstance(content, list): - return "\n".join(_block_text(block) for block in content) - return str(content) - - -def _block_text(block: object) -> str: - """Read the text out of one MCP content block (each block is a flat dict — - none of the standard block types nest another list of blocks inside). - """ - if not isinstance(block, dict): - return str(block) - if block.get("type") == "text": - return str(block.get("text", "")) - return f"[unsupported {block.get('type', 'content')} block]" - - class ToolInvoker: """Runs one agent's tool calls: gate, execute, record. @@ -152,11 +131,11 @@ def __init__( self._usage = usage_tracker self._system_namespace = system_namespace - async def invoke(self, tc: dict[str, Any]) -> str: + async def invoke(self, tc: dict[str, Any]) -> NormalizedToolResult: """Resolve, gate, and execute one tool call for both local and MCP tools. The pipeline short-circuits at the first step that stops the call, each - returning a model-facing string: + returning a normalized result whose text is model-facing: 1. resolve the tool (unknown name → error); 2. enforce execution limits (blocked → controlled message); @@ -167,25 +146,21 @@ async def invoke(self, tc: dict[str, Any]) -> str: 5. execute with the ``before/after/on_error`` lifecycle hooks. Every outcome that reached a decision — executed, failed, or denied — is - recorded against the run before the text is returned. + recorded against the run before the result is returned. """ tool = self._binding.get(tc["name"]) if tool is None: - return f"Unknown tool: {tc['name']}" + return NormalizedToolResult.text_only(f"Unknown tool: {tc['name']}") call = self._describe_call(tool, tc) blocked = self._enforce_limits(call) if blocked is not None: - return blocked + return NormalizedToolResult.text_only(blocked) gate = await self._gate_tool_call(call) if isinstance(gate, DenyTool): await self._usage.record_denied(call.identity) - return gate.message - - cached = await self._cached_result(call) - if cached is not None: - return cached + return NormalizedToolResult.text_only(gate.message) return await self._execute(call) @@ -235,17 +210,10 @@ def _enforce_limits(self, call: _ToolCall) -> str | None: return blocked_message(exc) return None - async def _cached_result(self, call: _ToolCall) -> str | None: - """Return a prior successful result for this exact call, if any. - - Guards against a second side effect when a graph re-entry after resume - replays the node with the same ``tool_call_id`` (the primary - duplicate-execution protection). Re-recording the same identity is an - upsert, so the replay does not add a second usage record either. - """ - cached = await self._execution_manager.already_executed(call.exec_id) - if cached is None or cached.result is None: - return None + async def _replay(self, call: _ToolCall, claim: ToolExecutionClaim) -> NormalizedToolResult: + """Return the terminal result published by another execution owner.""" + if claim.result is None or claim.status is None: + raise RuntimeError("replayed tool execution has no terminal result") log( logger, logging.INFO, @@ -255,33 +223,41 @@ async def _cached_result(self, call: _ToolCall) -> str | None: tool_call_id=call.tool_call_id, execution_id=call.exec_id, ) - await self._usage.record_success(call.identity) - return cached.result - - async def _execute(self, call: _ToolCall) -> str: - """Run the provider exactly once, wrapped in the idempotency ledger and the - ``before_tool_call`` hook, dispatching to the success or error recorder. - """ - await self._execution_manager.begin_execution( + if claim.status == ToolExecutionStatus.SUCCEEDED: + await self._usage.record_success(call.identity) + else: + await self._usage.record_failure(call.identity, error=claim.result.text) + return claim.result + + async def _execute(self, call: _ToolCall) -> NormalizedToolResult: + """Claim, run, normalize, and publish one logical tool invocation.""" + claim = await self._execution_manager.claim_execution( call.exec_id, tool_call_id=call.tool_call_id, run_id=call.run_id, tool_name=call.name, ) + if not claim.should_execute: + return await self._replay(call, claim) + self._log_call(logging.INFO, "tool call started", call) - await self._hook_manager.run_before_tool_call( - current_run_context.get(), - ToolRequestContext( - agent_id=self._spec.id, - tool_name=call.name, - provider=call.provider, - server_id=call.server_id, - ), - ) + try: + await self._hook_manager.run_before_tool_call( + current_run_context.get(), + ToolRequestContext( + agent_id=self._spec.id, + tool_name=call.name, + provider=call.provider, + server_id=call.server_id, + ), + ) + except BaseException: + await self._publish_processing_failure(call, "Tool execution blocked before invocation") + raise start = time.perf_counter() try: - result = await call.tool.ainvoke( + provider_result = await call.tool.ainvoke( { "type": "tool_call", "id": call.tool_call_id, @@ -289,16 +265,53 @@ async def _execute(self, call: _ToolCall) -> str: "args": call.args, } ) + except asyncio.CancelledError: + await self._publish_processing_failure(call, "Tool execution was cancelled") + raise except Exception as exc: return await self._record_error(call, exc, _elapsed_ms(start)) - return await self._record_success(call, result, _elapsed_ms(start)) + except BaseException: + await self._publish_processing_failure(call, "Tool execution was interrupted") + raise - async def _record_error(self, call: _ToolCall, exc: Exception, latency_ms: int) -> str: + try: + normalized = normalize_tool_result(provider_result) + except ProviderToolResultError as exc: + return await self._record_error( + call, + exc, + _elapsed_ms(start), + model_text=exc.model_text, + ) + except ToolResultNormalizationError as exc: + return await self._record_error( + call, + exc, + _elapsed_ms(start), + model_text="Tool error: invalid tool result", + ) + + return await self._record_success(call, normalized, _elapsed_ms(start)) + + async def _record_error( + self, + call: _ToolCall, + exc: Exception, + latency_ms: int, + *, + model_text: str | None = None, + ) -> NormalizedToolResult: """Record a failed call, fire ``on_tool_error``, and return the error text. The failure is returned (not raised) so the model can read it and recover. """ error = str(exc)[:200] + result = NormalizedToolResult.text_only(model_text or f"Tool error: {exc}") + await self._execution_manager.finish_execution( + call.exec_id, + status=ToolExecutionStatus.FAILED, + result=result, + ) await self._usage.record_failure(call.identity, error=error) self._log_call(logging.WARNING, "tool call failed", call, ms=latency_ms, error=error) @@ -306,45 +319,73 @@ async def _record_error(self, call: _ToolCall, exc: Exception, latency_ms: int) current_run_context.get(), self._call_context(call, "failed", latency_ms, error=error), ) - await self._execution_manager.finish_execution( - call.exec_id, status="failed", result=f"Tool error: {exc}" - ) - return f"Tool error: {exc}" + return result - async def _record_success(self, call: _ToolCall, result: object, latency_ms: int) -> str: + async def _record_success( + self, + call: _ToolCall, + result: NormalizedToolResult, + latency_ms: int, + ) -> NormalizedToolResult: """Record a successful call, fire ``after_tool_call`` and result-transform hooks, persist the result to the idempotency ledger, and return it. """ - await self._usage.record_success(call.identity) - self._log_call(logging.INFO, "tool call ended", call, ms=latency_ms) - await self._hook_manager.run_after_tool_call( - current_run_context.get(), - self._call_context(call, "succeeded", latency_ms), + try: + await self._usage.record_success(call.identity) + self._log_call(logging.INFO, "tool call ended", call, ms=latency_ms) + await self._hook_manager.run_after_tool_call( + current_run_context.get(), + self._call_context(call, "succeeded", latency_ms), + ) + normalized = await self._transform_result(call, result, latency_ms) + except BaseException: + await self._publish_processing_failure(call, "Tool result processing failed") + raise + + await self._execution_manager.finish_execution( + call.exec_id, + status=ToolExecutionStatus.SUCCEEDED, + result=normalized, ) - result_text = await self._transform_result(call, _extract_result_text(result), latency_ms) + return normalized + + async def _publish_processing_failure(self, call: _ToolCall, message: str) -> None: + """Unblock duplicate callers without exposing exception payloads.""" await self._execution_manager.finish_execution( - call.exec_id, status="succeeded", result=result_text + call.exec_id, + status=ToolExecutionStatus.FAILED, + result=NormalizedToolResult.text_only(message), ) - return result_text - async def _transform_result(self, call: _ToolCall, result_text: str, latency_ms: int) -> str: + async def _transform_result( + self, + call: _ToolCall, + result: NormalizedToolResult, + latency_ms: int, + ) -> NormalizedToolResult: """Let ``transform_tool_result`` hooks reshape the result (e.g. truncate oversized MCP output). The context is only built when such a hook exists. """ if not self._hook_manager.has("transform_tool_result"): - return result_text + return result transformed = await self._hook_manager.run_transform_tool_result( current_run_context.get(), ToolResultContext( agent_id=self._spec.id, tool_name=call.name, provider=call.provider, - result=result_text, + result=result.text, + structured_result=result.structured, + artifact=result.artifact, server_id=call.server_id, latency_ms=latency_ms, ), ) - return transformed.result + return NormalizedToolResult( + text=transformed.result, + structured=transformed.structured_result, + artifact=transformed.artifact, + ) def _call_context( self, diff --git a/src/agent_engine/engine/langgraph/tools/tool_result_normalizer.py b/src/agent_engine/engine/langgraph/tools/tool_result_normalizer.py new file mode 100644 index 00000000..d40bf2b0 --- /dev/null +++ b/src/agent_engine/engine/langgraph/tools/tool_result_normalizer.py @@ -0,0 +1,225 @@ +"""LangChain tool-result adapter into Extra's runtime contract.""" + +from __future__ import annotations + +import json +from collections.abc import Mapping +from dataclasses import dataclass + +from langchain_core.messages import ToolMessage + +from agent_engine.runtime.tool_results import ( + NormalizedToolResult, + ToolResultValidationError, + deterministic_json, +) + +_STRUCTURED_KEYS = ("structured_content", "structuredContent") +_EMPTY_RESULT_TEXT = "[tool returned no content]" +_MAX_ARTIFACT_DEPTH = 8 +_MAX_ARTIFACT_ITEMS = 128 +_MAX_ARTIFACT_TOTAL_VALUES = 1024 +_MAX_ARTIFACT_STRING_CHARS = 8192 +_MAX_ARTIFACT_TOTAL_STRING_CHARS = 32768 +_MAX_ARTIFACT_KEY_CHARS = 256 +_MAX_ARTIFACT_JSON_CHARS = 65536 + + +class ToolResultNormalizationError(RuntimeError): + """A provider result does not satisfy Extra's tool-result contract.""" + + +class ProviderToolResultError(ToolResultNormalizationError): + """The provider returned an explicit error result instead of raising.""" + + def __init__(self, model_text: str) -> None: + self.model_text = model_text + super().__init__("tool provider returned an error result") + + +@dataclass +class _ArtifactBudget: + values_left: int = _MAX_ARTIFACT_TOTAL_VALUES + string_chars_left: int = _MAX_ARTIFACT_TOTAL_STRING_CHARS + + def take_value(self) -> bool: + if self.values_left <= 0: + return False + self.values_left -= 1 + return True + + def take_string(self, size: int) -> bool: + if size > self.string_chars_left: + return False + self.string_chars_left -= size + return True + + +def normalize_tool_result(result: object) -> NormalizedToolResult: + """Normalize one LangChain or direct local-tool result exactly once.""" + try: + if isinstance(result, ToolMessage): + model_text, content_structured = _tool_message_content(result.content) + artifact_structured, artifact = _split_artifact(result.artifact) + structured = ( + artifact_structured if artifact_structured is not None else content_structured + ) + normalized = NormalizedToolResult(model_text, structured, artifact) + normalized = _with_model_fallback(normalized) + if result.status == "error": + raise ProviderToolResultError(normalized.text) + return normalized + + model_text, structured = _direct_result(result) + return _with_model_fallback(NormalizedToolResult(model_text, structured)) + except ProviderToolResultError: + raise + except ToolResultValidationError as exc: + raise ToolResultNormalizationError(str(exc)) from exc + + +def _tool_message_content(content: object) -> tuple[str, object | None]: + if isinstance(content, str): + return content, _json_container(content) + if isinstance(content, list): + return "\n".join(_content_block_text(block) for block in content), None + raise ToolResultNormalizationError( + f"LangChain ToolMessage content has unsupported type {type(content).__name__}" + ) + + +def _direct_result(result: object) -> tuple[str, object | None]: + if isinstance(result, str): + return result, _json_container(result) + if isinstance(result, (Mapping, list, tuple)): + return deterministic_json(result), result + if result is None or isinstance(result, (bool, int, float)): + return deterministic_json(result), result + raise ToolResultNormalizationError( + f"tool returned unsupported result type {type(result).__name__}" + ) + + +def _content_block_text(block: object) -> str: + if isinstance(block, str): + return block + if not isinstance(block, Mapping): + raise ToolResultNormalizationError( + f"tool content block has unsupported type {type(block).__name__}" + ) + block_type = block.get("type") + if not isinstance(block_type, str) or not block_type: + raise ToolResultNormalizationError("tool content block must declare a string type") + if block_type == "text": + text = block.get("text") + if not isinstance(text, str): + raise ToolResultNormalizationError("text content block must contain string text") + return text + return f"[unsupported {block_type} block]" + + +def _json_container(content: str) -> object | None: + """Recover dict/list local results serialized by LangChain ToolMessage.""" + try: + parsed = json.loads(content) + except json.JSONDecodeError: + return None + return parsed if isinstance(parsed, (dict, list)) else None + + +def _with_model_fallback(result: NormalizedToolResult) -> NormalizedToolResult: + if result.text: + return result + structured_text = result.structured_text() + return result.with_text(structured_text or _EMPTY_RESULT_TEXT) + + +def _split_artifact(artifact: object) -> tuple[object | None, object | None]: + if artifact is None: + return None, None + if not isinstance(artifact, Mapping): + raise ToolResultNormalizationError("tool artifact must be a mapping") + + present_keys = [key for key in _STRUCTURED_KEYS if key in artifact] + if len(present_keys) > 1: + raise ToolResultNormalizationError("tool artifact declares structured output twice") + structured = artifact[present_keys[0]] if present_keys else None + metadata: dict[str, object] = {} + budget = _ArtifactBudget() + for key, value in artifact.items(): + if key in _STRUCTURED_KEYS: + continue + name = _string_key(key, path="artifact") + metadata[name] = _bounded_artifact_value( + value, + path="artifact.", + depth=0, + budget=budget, + ) + if metadata: + encoded = deterministic_json(metadata, field_name="artifact metadata") + if len(encoded) > _MAX_ARTIFACT_JSON_CHARS: + metadata = { + "type": "artifact_metadata", + "size": len(encoded), + "omitted": True, + } + return structured, metadata or None + + +def _string_key(key: object, *, path: str) -> str: + if not isinstance(key, str): + raise ToolResultNormalizationError(f"{path} contains a non-string key") + if len(key) > _MAX_ARTIFACT_KEY_CHARS: + raise ToolResultNormalizationError(f"{path} contains an oversized key") + return key + + +def _bounded_artifact_value( + value: object, + *, + path: str, + depth: int, + budget: _ArtifactBudget, +) -> object: + """Keep bounded JSON metadata and replace known large payload shapes.""" + if not budget.take_value(): + return {"type": "value", "omitted": True} + if isinstance(value, (bytes, bytearray, memoryview)): + return {"type": "binary", "size": len(value), "omitted": True} + if isinstance(value, str): + if len(value) > _MAX_ARTIFACT_STRING_CHARS or not budget.take_string(len(value)): + return {"type": "text", "size": len(value), "omitted": True} + return value + if value is None or isinstance(value, (bool, int, float)): + return value + if depth >= _MAX_ARTIFACT_DEPTH: + return {"type": "nested", "omitted": True} + if isinstance(value, Mapping): + if len(value) > _MAX_ARTIFACT_ITEMS: + return {"type": "object", "size": len(value), "omitted": True} + nested: dict[str, object] = {} + for key, item in value.items(): + name = _string_key(key, path=path) + nested[name] = _bounded_artifact_value( + item, + path=f"{path}.", + depth=depth + 1, + budget=budget, + ) + return nested + if isinstance(value, (list, tuple)): + if len(value) > _MAX_ARTIFACT_ITEMS: + return {"type": "array", "size": len(value), "omitted": True} + return [ + _bounded_artifact_value( + item, + path=f"{path}[{index}]", + depth=depth + 1, + budget=budget, + ) + for index, item in enumerate(value) + ] + raise ToolResultNormalizationError( + f"{path} contains unsupported value type {type(value).__name__}" + ) diff --git a/src/agent_engine/observability/providers/logging/provider.py b/src/agent_engine/observability/providers/logging/provider.py index f7b83b18..bef752b4 100644 --- a/src/agent_engine/observability/providers/logging/provider.py +++ b/src/agent_engine/observability/providers/logging/provider.py @@ -30,11 +30,11 @@ def on_llm_error(self, error: BaseException, **kwargs: Any) -> None: def on_tool_start(self, serialized: dict[str, Any], input_str: str, **kwargs: Any) -> None: log(logger, logging.INFO, "tool start", name=(serialized or {}).get("name", "?")) - log(logger, logging.DEBUG, "tool input", value=input_str[:300]) + log(logger, logging.DEBUG, "tool input", chars=len(input_str)) def on_tool_end(self, output: Any, **kwargs: Any) -> None: log(logger, logging.INFO, "tool end", status="ok") - log(logger, logging.DEBUG, "tool output", value=str(output)[:300]) + log(logger, logging.DEBUG, "tool output", output_type=type(output).__name__) def on_tool_error(self, error: BaseException, **kwargs: Any) -> None: log(logger, logging.WARNING, "tool end", status="error", error=str(error)) diff --git a/src/agent_engine/runtime/hooks/models.py b/src/agent_engine/runtime/hooks/models.py index 8df0146a..50cd5e30 100644 --- a/src/agent_engine/runtime/hooks/models.py +++ b/src/agent_engine/runtime/hooks/models.py @@ -13,6 +13,8 @@ from dataclasses import dataclass, field from typing import Any, Literal, TypeVar +from agent_engine.runtime.tool_results import JsonValue + # The supported lifecycle points. Adding a point here is the one place that # enables a new hook kind across schema validation, loading, and execution. HookPoint = Literal[ @@ -236,10 +238,12 @@ class ToolResultContext: Unlike the observe-only ``after_tool_call`` (which never sees results), this hook is given the tool's ``result`` text precisely so it can shape it — - truncating oversized MCP output, redacting, normalizing — and **must return a - ``ToolResultContext``** carrying the original or modified ``result``. Because it - handles raw tool output it is trusted code that may see sensitive content; - never log the ``result`` body, only safe metadata (sizes, names, ids). + truncating oversized MCP output, redacting, normalizing — alongside optional + machine-readable ``structured_result`` and safe ``artifact`` metadata. It + **must return a ``ToolResultContext``** carrying the original or modified + values. Because it handles raw tool output it is trusted code that may see + sensitive content; never log result bodies, only safe metadata (sizes, + names, ids). """ agent_id: str @@ -249,7 +253,17 @@ class ToolResultContext: server_id: str | None = None latency_ms: int | None = None metadata: dict[str, object] = field(default_factory=dict) + structured_result: JsonValue | None = None + artifact: JsonValue | None = None def with_result(self, result: str) -> ToolResultContext: - """Return a copy with ``result`` replaced (immutable update).""" + """Replace text while preserving structured output and artifacts.""" return dataclasses.replace(self, result=result) + + def with_structured_result(self, structured_result: JsonValue | None) -> ToolResultContext: + """Return a copy with the machine-readable result replaced.""" + return dataclasses.replace(self, structured_result=structured_result) + + def with_artifact(self, artifact: JsonValue | None) -> ToolResultContext: + """Return a copy with safe artifact metadata replaced.""" + return dataclasses.replace(self, artifact=artifact) diff --git a/src/agent_engine/runtime/tool_results.py b/src/agent_engine/runtime/tool_results.py new file mode 100644 index 00000000..005d737b --- /dev/null +++ b/src/agent_engine/runtime/tool_results.py @@ -0,0 +1,239 @@ +"""Provider-independent, persistable tool-result value objects.""" + +from __future__ import annotations + +import json +import math +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import TypeAlias, cast + +JsonScalar: TypeAlias = str | int | float | bool | None +JsonValue: TypeAlias = JsonScalar | list["JsonValue"] | dict[str, "JsonValue"] +PersistedToolResult: TypeAlias = str | dict[str, object] + +_PERSISTENCE_VERSION = 1 +_MAX_JSON_DEPTH = 64 +_MAX_JSON_VALUES = 10_000 +_MAX_JSON_STRING_BYTES = 1_000_000 +_MAX_JSON_KEY_BYTES = 1024 +_MAX_JSON_TOTAL_KEY_BYTES = 256_000 +_MAX_JSON_BYTES = 1_048_576 +_MAX_JSON_INTEGER_BITS = 4096 +_PERSISTED_KEYS = frozenset({"version", "text", "structured", "artifact"}) + + +class ToolResultValidationError(ValueError): + """A tool result cannot satisfy Extra's JSON-safe runtime contract.""" + + +@dataclass +class _JsonBudget: + values_left: int = _MAX_JSON_VALUES + string_bytes_left: int = _MAX_JSON_STRING_BYTES + key_bytes_left: int = _MAX_JSON_TOTAL_KEY_BYTES + + def take_value(self, *, path: str) -> None: + if self.values_left <= 0: + raise ToolResultValidationError(f"{path} exceeds the JSON value budget") + self.values_left -= 1 + + def take_string(self, value: str, *, path: str) -> None: + size = _utf8_size(value, path=path) + if size > self.string_bytes_left: + raise ToolResultValidationError(f"{path} exceeds the JSON string budget") + self.string_bytes_left -= size + + def take_key(self, key: str, *, path: str) -> None: + size = _utf8_size(key, path=path) + if size > _MAX_JSON_KEY_BYTES: + raise ToolResultValidationError(f"{path} contains an oversized object key") + if size > self.key_bytes_left: + raise ToolResultValidationError(f"{path} exceeds the JSON object-key budget") + self.key_bytes_left -= size + + +@dataclass(frozen=True, init=False) +class NormalizedToolResult: + """One authoritative result for hooks, replay, and model text. + + Nested values are validated and stored as canonical JSON rather than as + mutable provider-owned objects. Accessors decode fresh values, so callers + cannot mutate the result held by the execution ledger. + """ + + text: str + _structured_json: str | None = field(repr=False) + _artifact_json: str | None = field(repr=False) + + def __init__( + self, + text: str, + structured: object | None = None, + artifact: object | None = None, + ) -> None: + if not isinstance(text, str): + raise ToolResultValidationError("tool result text must be a string") + object.__setattr__(self, "text", text) + object.__setattr__( + self, + "_structured_json", + _canonical_json(structured, field_name="structured result") + if structured is not None + else None, + ) + object.__setattr__( + self, + "_artifact_json", + _canonical_json(artifact, field_name="artifact metadata") + if artifact is not None + else None, + ) + + @classmethod + def text_only(cls, text: str) -> NormalizedToolResult: + return cls(text=text) + + @property + def structured(self) -> JsonValue | None: + return _decode_json(self._structured_json) + + @property + def artifact(self) -> JsonValue | None: + return _decode_json(self._artifact_json) + + @property + def has_structured(self) -> bool: + return self._structured_json is not None + + @property + def has_artifact(self) -> bool: + return self._artifact_json is not None + + def structured_text(self) -> str | None: + """Return deterministic JSON for a structured-only model fallback.""" + return self._structured_json + + def with_text(self, text: str) -> NormalizedToolResult: + return NormalizedToolResult(text, self.structured, self.artifact) + + def to_persisted(self) -> dict[str, object]: + """Return a versioned value containing JSON-serializable primitives.""" + return { + "version": _PERSISTENCE_VERSION, + "text": self.text, + "structured": self.structured, + "artifact": self.artifact, + } + + @classmethod + def from_persisted( + cls, value: PersistedToolResult | Mapping[str, object] + ) -> NormalizedToolResult: + """Restore current payloads and legacy text-only ledger entries.""" + if isinstance(value, str): + return cls.text_only(value) + if not isinstance(value, Mapping): + raise ToolResultValidationError("persisted tool result must be text or a mapping") + unknown = set(value) - _PERSISTED_KEYS + if unknown: + names = ", ".join(sorted(str(key) for key in unknown)) + raise ToolResultValidationError(f"persisted tool result has unknown fields: {names}") + if value.get("version") != _PERSISTENCE_VERSION: + raise ToolResultValidationError("unsupported persisted tool result version") + text = value.get("text") + if not isinstance(text, str): + raise ToolResultValidationError("persisted tool result text must be a string") + return cls( + text=text, + structured=value.get("structured"), + artifact=value.get("artifact"), + ) + + +def deterministic_json(value: object, *, field_name: str = "tool result") -> str: + """Serialize one JSON-like value canonically or raise a typed error.""" + return _canonical_json(value, field_name=field_name) + + +def _canonical_json(value: object, *, field_name: str) -> str: + normalized = _json_value( + value, + path=field_name, + depth=0, + budget=_JsonBudget(), + ) + encoded = json.dumps( + normalized, + ensure_ascii=False, + allow_nan=False, + sort_keys=True, + separators=(",", ":"), + ) + if _utf8_size(encoded, path=field_name) > _MAX_JSON_BYTES: + raise ToolResultValidationError(f"{field_name} exceeds the encoded JSON budget") + return encoded + + +def _json_value( + value: object, + *, + path: str, + depth: int, + budget: _JsonBudget, +) -> JsonValue: + if depth > _MAX_JSON_DEPTH: + raise ToolResultValidationError(f"{path} exceeds maximum nesting depth") + budget.take_value(path=path) + if value is None or isinstance(value, bool): + return value + if isinstance(value, str): + budget.take_string(value, path=path) + return value + if isinstance(value, int): + if value.bit_length() > _MAX_JSON_INTEGER_BITS: + raise ToolResultValidationError(f"{path} contains an oversized integer") + return value + if isinstance(value, float): + if not math.isfinite(value): + raise ToolResultValidationError(f"{path} contains a non-finite number") + return value + if isinstance(value, Mapping): + normalized: dict[str, JsonValue] = {} + for key, item in value.items(): + if not isinstance(key, str): + raise ToolResultValidationError(f"{path} contains a non-string object key") + budget.take_key(key, path=path) + normalized[key] = _json_value( + item, + path=f"{path}.", + depth=depth + 1, + budget=budget, + ) + return normalized + if isinstance(value, (list, tuple)): + return [ + _json_value( + item, + path=f"{path}[{index}]", + depth=depth + 1, + budget=budget, + ) + for index, item in enumerate(value) + ] + raise ToolResultValidationError( + f"{path} contains unsupported value type {type(value).__name__}" + ) + + +def _decode_json(value: str | None) -> JsonValue | None: + if value is None: + return None + return cast(JsonValue, json.loads(value)) + + +def _utf8_size(value: str, *, path: str) -> int: + try: + return len(value.encode("utf-8")) + except UnicodeEncodeError as exc: + raise ToolResultValidationError(f"{path} contains invalid Unicode") from exc diff --git a/tests/approvals/test_engine_hitl.py b/tests/approvals/test_engine_hitl.py index 10be6266..dedb9661 100644 --- a/tests/approvals/test_engine_hitl.py +++ b/tests/approvals/test_engine_hitl.py @@ -28,6 +28,13 @@ from agent_engine.approvals.in_memory_session_approval_repository import ( InMemorySessionApprovalRepository, ) +from agent_engine.approvals.in_memory_tool_execution_repository import ( + InMemoryToolExecutionRepository, +) +from agent_engine.approvals.tool_execution_manager import ( + ToolExecutionManager, + execution_id_for, +) from agent_engine.core.spec import ( AgentSpec, BasePromptSet, @@ -40,6 +47,7 @@ from agent_engine.engine.langgraph.engine import LangGraphEngine from agent_engine.runs.in_memory import InMemoryRunRepository from agent_engine.runtime.hooks import RunContext +from agent_engine.tool_usage.models import stable_tool_call_id _MODEL = ModelConfig(provider="fake", name="fake", temperature=None) @@ -238,6 +246,18 @@ def _write_counting_tool(base_dir: Path, tool_id: str, counter: Path) -> None: ) +def _write_structured_counting_tool(base_dir: Path, tool_id: str, counter: Path) -> None: + tools_dir = base_dir / "plugins" / "tools" + tools_dir.mkdir(parents=True, exist_ok=True) + (tools_dir / f"{tool_id}.py").write_text( + f"def {tool_id}(message: str) -> dict:\n" + f" with open({str(counter)!r}, 'a') as f:\n" + " f.write('x')\n" + " return {'status': 'sent', 'message': message}\n", + encoding="utf-8", + ) + + def _spec(tool_id: str, *, auto_mode: bool = False) -> SystemSpec: agent = AgentSpec( id="writer", @@ -320,6 +340,41 @@ async def test_allow_once_resumes_same_run_and_executes_once(tmp_path: Path) -> assert recovered == resumed +async def test_hitl_resume_persists_structured_local_result(tmp_path: Path) -> None: + counter = tmp_path / "calls.log" + tool_name = "send_structured" + _write_structured_counting_tool(tmp_path, tool_name, counter) + manager = ToolExecutionManager(execution_repository=InMemoryToolExecutionRepository()) + async with LangGraphEngine( + tmp_path, + model_factory=_factory, + execution_manager=manager, + ) as engine: + await engine.build(_spec(tool_name)) + pending = await engine.run("hi", context=RunContext(run_id="run-structured")) + assert pending.pending_approval is not None + + resumed = await engine.resume( + "run-structured", + pending.pending_approval.approval_id, + "allow once", + ) + + tool_call_id = stable_tool_call_id( + "run-structured", + "writer", + "local", + None, + tool_name, + {"message": "go"}, + ) + persisted = await manager.restored_result(execution_id_for(tool_call_id)) + assert resumed.status == "completed" + assert _executions(counter) == 1 + assert persisted is not None + assert persisted.structured == {"message": "go", "status": "sent"} + + async def test_retrying_first_decision_recovers_second_pending_approval(tmp_path: Path) -> None: first_counter = tmp_path / "first.log" second_counter = tmp_path / "second.log" diff --git a/tests/approvals/test_manager.py b/tests/approvals/test_manager.py index 959b177b..4bdccd67 100644 --- a/tests/approvals/test_manager.py +++ b/tests/approvals/test_manager.py @@ -20,12 +20,19 @@ from agent_engine.approvals.in_memory_tool_execution_repository import ( InMemoryToolExecutionRepository, ) -from agent_engine.approvals.models import ApprovalStatus, RunRecord, RunStatus +from agent_engine.approvals.models import ( + ApprovalStatus, + RunRecord, + RunStatus, + ToolExecutionRecord, + ToolExecutionStatus, +) from agent_engine.approvals.tool_execution_manager import ( ToolExecutionManager, execution_id_for, ) from agent_engine.runs.in_memory import InMemoryRunRepository +from agent_engine.runtime.tool_results import NormalizedToolResult def _manager() -> ToolExecutionManager: @@ -35,27 +42,54 @@ def _manager() -> ToolExecutionManager: async def test_idempotency_reports_prior_success() -> None: mgr = _manager() exec_id = execution_id_for("tc1") - assert await mgr.already_executed(exec_id) is None - assert ( - await mgr.begin_execution(exec_id, tool_call_id="tc1", run_id="r1", tool_name="t") is True + assert await mgr.restored_result(exec_id) is None + claim = await mgr.claim_execution(exec_id, tool_call_id="tc1", run_id="r1", tool_name="t") + assert claim.should_execute is True + result = NormalizedToolResult( + text="Found 2 invoices", + structured={"count": 2}, + artifact={"source": "billing"}, ) - # A second begin for the same key is a duplicate. + await mgr.finish_execution(exec_id, status=ToolExecutionStatus.SUCCEEDED, result=result) + prior = await mgr.restored_result(exec_id) + assert prior == result + replay = await mgr.claim_execution(exec_id, tool_call_id="tc1", run_id="r1", tool_name="t") + assert replay.should_execute is False + assert replay.status == ToolExecutionStatus.SUCCEEDED + assert replay.result == result assert ( await mgr.begin_execution(exec_id, tool_call_id="tc1", run_id="r1", tool_name="t") is False ) - await mgr.finish_execution(exec_id, status="succeeded", result="R") - prior = await mgr.already_executed(exec_id) - assert prior is not None and prior.result == "R" async def test_idempotency_no_repository_never_dedupes() -> None: mgr = ToolExecutionManager() # no repository exec_id = execution_id_for("tc1") - assert await mgr.already_executed(exec_id) is None - assert ( - await mgr.begin_execution(exec_id, tool_call_id="tc1", run_id="r1", tool_name="t") is True + assert await mgr.restored_result(exec_id) is None + claim = await mgr.claim_execution(exec_id, tool_call_id="tc1", run_id="r1", tool_name="t") + assert claim.should_execute is True + + +async def test_legacy_text_result_replays_as_text_only() -> None: + repository = InMemoryToolExecutionRepository() + manager = ToolExecutionManager(execution_repository=repository) + exec_id = execution_id_for("legacy") + await repository.start( + ToolExecutionRecord( + execution_id=exec_id, + tool_call_id="legacy", + run_id="r1", + tool_name="t", + ) + ) + await repository.complete( + exec_id, + status=ToolExecutionStatus.SUCCEEDED, + result="legacy text", ) + assert await manager.restored_result(exec_id) == NormalizedToolResult.text_only("legacy text") + # ------------------------------- ApprovalManager ------------------------------ # diff --git a/tests/approvals/test_repository.py b/tests/approvals/test_repository.py index 02fbdd92..b1afb278 100644 --- a/tests/approvals/test_repository.py +++ b/tests/approvals/test_repository.py @@ -15,7 +15,9 @@ ApprovalRecord, ApprovalStatus, ToolExecutionRecord, + ToolExecutionStatus, ) +from agent_engine.runtime.tool_results import NormalizedToolResult pytestmark = pytest.mark.asyncio @@ -97,7 +99,8 @@ async def test_execution_idempotency_start_is_create_if_absent() -> None: second, created2 = await repo.start(rec) assert created1 is True assert created2 is False # duplicate attempt detected - assert first is second + assert first == second + assert first is not second # callers cannot mutate the repository's record async def test_execution_complete_records_result() -> None: @@ -105,6 +108,39 @@ async def test_execution_complete_records_result() -> None: await repo.start( ToolExecutionRecord(execution_id="e1", tool_call_id="tc1", run_id="r1", tool_name="t") ) - await repo.complete("e1", status="succeeded", result="done") + result = NormalizedToolResult(text="done", structured={"ok": True}) + await repo.complete( + "e1", + status=ToolExecutionStatus.SUCCEEDED, + result=result.to_persisted(), + ) rec = await repo.get("e1") - assert rec is not None and rec.status == "succeeded" and rec.result == "done" + assert rec is not None and rec.status == ToolExecutionStatus.SUCCEEDED + assert rec.result == result.to_persisted() + + with pytest.raises(ValueError, match="already terminal"): + await repo.complete( + "e1", + status=ToolExecutionStatus.FAILED, + result=NormalizedToolResult.text_only("replacement").to_persisted(), + ) + + unchanged = await repo.get("e1") + assert unchanged == rec + + +async def test_execution_waiter_receives_the_terminal_snapshot() -> None: + repo = InMemoryToolExecutionRepository() + await repo.start( + ToolExecutionRecord(execution_id="e1", tool_call_id="tc1", run_id="r1", tool_name="t") + ) + waiter = asyncio.create_task(repo.wait_for_completion("e1")) + await asyncio.sleep(0) + assert waiter.done() is False + + result = NormalizedToolResult("done", structured={"ok": True}).to_persisted() + await repo.complete("e1", status=ToolExecutionStatus.SUCCEEDED, result=result) + + completed = await waiter + assert completed.status == ToolExecutionStatus.SUCCEEDED + assert completed.result == result diff --git a/tests/runtime/hooks/test_hook_manager.py b/tests/runtime/hooks/test_hook_manager.py index d0cdaed4..64921e46 100644 --- a/tests/runtime/hooks/test_hook_manager.py +++ b/tests/runtime/hooks/test_hook_manager.py @@ -398,9 +398,19 @@ def test_has_reports_declared_points() -> None: async def test_transform_tool_result_returns_modified_result() -> None: mgr = _manager(HookSpec("transform_tool_result", f"{_FIX}:truncate_tool_result")) out = await mgr.run_transform_tool_result( - None, ToolResultContext("a", "t", "mcp", result="abcdefgh") + None, + ToolResultContext( + "a", + "t", + "mcp", + result="abcdefgh", + structured_result={"count": 2}, + artifact={"source": "billing"}, + ), ) assert out.result == "abc" # truncated to the configured limit + assert out.structured_result == {"count": 2} + assert out.artifact == {"source": "billing"} async def test_transform_tool_result_warn_failure_keeps_original() -> None: diff --git a/tests/runtime/test_mcp_tool_result_extraction.py b/tests/runtime/test_mcp_tool_result_extraction.py index 9cb87621..787e687a 100644 --- a/tests/runtime/test_mcp_tool_result_extraction.py +++ b/tests/runtime/test_mcp_tool_result_extraction.py @@ -13,15 +13,26 @@ from __future__ import annotations +import asyncio +import logging from collections.abc import AsyncIterator from pathlib import Path from typing import Any, cast +import pytest from langchain_core.language_models import BaseChatModel from langchain_core.messages import AIMessage, ToolMessage from langchain_core.messages.tool import ToolCall from langchain_core.tools import StructuredTool +from agent_engine.approvals.approval_provider import ApprovalProvider, ApprovalRequest +from agent_engine.approvals.coordinator import ApprovalCoordinator +from agent_engine.approvals.decision import ApprovalDecision +from agent_engine.approvals.in_memory_tool_execution_repository import ( + InMemoryToolExecutionRepository, +) +from agent_engine.approvals.models import ToolExecutionStatus +from agent_engine.approvals.tool_execution_manager import ToolExecutionManager from agent_engine.core.spec import ( AgentSpec, BasePromptSet, @@ -34,23 +45,70 @@ ToolSpec, ) from agent_engine.engine.langgraph.engine import LangGraphEngine +from agent_engine.engine.langgraph.tools.agent_tool_binding import AgentToolBinding +from agent_engine.engine.langgraph.tools.tool_invoker import ToolInvoker +from agent_engine.runtime.hooks import HookManager, RunContext, current_run_context +from agent_engine.runtime.tool_results import NormalizedToolResult, PersistedToolResult +from agent_engine.tool_usage.in_memory import InMemoryToolUsageRepository +from agent_engine.tool_usage.tracker import ToolUsageTracker _MODEL = ModelConfig(provider="fake", name="fake", temperature=None) +class _UnusedApprovalProvider(ApprovalProvider): + async def request_decision(self, request: ApprovalRequest) -> ApprovalDecision: + raise AssertionError("auto-mode tool execution must not request approval") + + +class _CapturingExecutionRepository(InMemoryToolExecutionRepository): + def __init__(self) -> None: + super().__init__() + self.completed_result: PersistedToolResult | None = None + + async def complete( + self, + execution_id: str, + status: ToolExecutionStatus, + result: PersistedToolResult, + ) -> None: + await super().complete(execution_id, status, result) + self.completed_result = result + + +class _BlockingCompletionRepository(InMemoryToolExecutionRepository): + def __init__(self) -> None: + super().__init__() + self.completing = asyncio.Event() + self.release = asyncio.Event() + + async def complete( + self, + execution_id: str, + status: ToolExecutionStatus, + result: PersistedToolResult, + ) -> None: + self.completing.set() + await self.release.wait() + await super().complete(execution_id, status, result) + + class EchoToolResultModel: """Calls the tool once, then answers with exactly the tool result text it was given — lets a test see precisely what reached the conversation. """ def __init__( - self, tool_names: list[str] | None = None, tool_args: dict[str, Any] | None = None + self, + tool_names: list[str] | None = None, + tool_args: dict[str, Any] | None = None, + captured: list[ToolMessage] | None = None, ) -> None: self._tool_names = tool_names or [] self._tool_args = tool_args or {} + self._captured = captured if captured is not None else [] def bind_tools(self, tools: list[Any]) -> EchoToolResultModel: - return EchoToolResultModel([t.name for t in tools], self._tool_args) + return EchoToolResultModel([t.name for t in tools], self._tool_args, self._captured) async def ainvoke(self, messages: list[Any]) -> AIMessage: return self._respond(messages) @@ -60,6 +118,8 @@ async def astream(self, messages: list[Any]) -> AsyncIterator[AIMessage]: def _respond(self, messages: list[Any]) -> AIMessage: tool_msgs = [m for m in messages if isinstance(m, ToolMessage)] + if tool_msgs: + self._captured.append(tool_msgs[-1]) if self._tool_names and not tool_msgs: return AIMessage( content="", @@ -70,9 +130,13 @@ def _respond(self, messages: list[Any]) -> AIMessage: def _model_factory( tool_args: dict[str, Any] | None = None, + captured: list[ToolMessage] | None = None, ) -> Any: def factory(provider: str, name: str, temperature: float | None) -> BaseChatModel: - return cast(BaseChatModel, EchoToolResultModel(tool_args=tool_args)) + return cast( + BaseChatModel, + EchoToolResultModel(tool_args=tool_args, captured=captured), + ) return factory @@ -100,14 +164,25 @@ def _system(graph: GraphNode) -> SystemSpec: ) -async def _run_with_mcp_tool(tmp_path: Path, mcp_tool: StructuredTool) -> str: +async def _run_with_mcp_tool( + tmp_path: Path, mcp_tool: StructuredTool +) -> tuple[str, NormalizedToolResult, ToolMessage]: spec = _system(_agent("research", mcps=(MCPSpec(id="wiki", url="https://wiki.test/mcp"),))) - async with LangGraphEngine(tmp_path, model_factory=_model_factory()) as engine: + captured: list[ToolMessage] = [] + repository = _CapturingExecutionRepository() + manager = ToolExecutionManager(execution_repository=repository) + async with LangGraphEngine( + tmp_path, + model_factory=_model_factory(captured=captured), + execution_manager=manager, + ) as engine: await engine.build(spec) engine._mcp_tools["wiki"] = [mcp_tool] engine._app = engine._build_graph(spec) result = await engine.run("search please") - return result.answer + assert repository.completed_result is not None + normalized = NormalizedToolResult.from_persisted(repository.completed_result) + return result.answer, normalized, captured[-1] async def test_mcp_text_result_is_not_garbled(tmp_path: Path) -> None: @@ -123,9 +198,12 @@ def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: response_format="content_and_artifact", ) - answer = await _run_with_mcp_tool(tmp_path, mcp_tool) + answer, runtime_result, message = await _run_with_mcp_tool(tmp_path, mcp_tool) assert answer == "clean text" + assert runtime_result.text == "clean text" + assert runtime_result.structured == {"value": "clean text"} + assert message.artifact is None assert "'type':" not in answer assert "'text':" not in answer @@ -135,7 +213,7 @@ def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: return [ {"type": "text", "text": "first block", "id": "lc_1"}, {"type": "text", "text": "second block", "id": "lc_2"}, - ], {} + ], {"structured_content": {"count": 2}} mcp_tool = StructuredTool.from_function( fake_mcp_tool, @@ -144,11 +222,191 @@ def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: response_format="content_and_artifact", ) - answer = await _run_with_mcp_tool(tmp_path, mcp_tool) + answer, runtime_result, _ = await _run_with_mcp_tool(tmp_path, mcp_tool) assert "first block" in answer assert "second block" in answer assert "'type':" not in answer + assert runtime_result.structured == {"count": 2} + + +async def test_mcp_structured_only_result_is_preserved(tmp_path: Path) -> None: + def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + return [], {"structured_content": {"balance": 1250}} + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="account_balance", + description="balance", + response_format="content_and_artifact", + ) + + answer, runtime_result, message = await _run_with_mcp_tool(tmp_path, mcp_tool) + + assert answer == '{"balance":1250}' + assert message.content == '{"balance":1250}' + assert message.artifact is None + assert runtime_result.structured == {"balance": 1250} + + +async def test_structured_only_fallback_is_deterministic_json(tmp_path: Path) -> None: + def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + return [], {"structured_content": {"z": 1, "a": [2, 1]}} + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="deterministic_result", + description="result", + response_format="content_and_artifact", + ) + + answer, runtime_result, _ = await _run_with_mcp_tool(tmp_path, mcp_tool) + + assert answer == '{"a":[2,1],"z":1}' + assert runtime_result.text == answer + assert runtime_result.structured == {"a": [2, 1], "z": 1} + + +async def test_unsupported_mcp_block_does_not_discard_structured_result( + tmp_path: Path, +) -> None: + def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + return [{"type": "image", "url": "https://files.test/chart.png"}], { + "structured_content": {"chart_id": "chart-1"} + } + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="chart", + description="chart", + response_format="content_and_artifact", + ) + + answer, runtime_result, _ = await _run_with_mcp_tool(tmp_path, mcp_tool) + + assert answer == "[unsupported image block]" + assert runtime_result.structured == {"chart_id": "chart-1"} + + +async def test_artifact_metadata_is_preserved_without_binary_body(tmp_path: Path) -> None: + def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + return [{"type": "text", "text": "report ready"}], { + "structured_content": {"report_id": "r-1"}, + "file": {"uri": "s3://reports/r-1.pdf", "mime_type": "application/pdf"}, + "preview": b"binary-preview", + } + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="report", + description="report", + response_format="content_and_artifact", + ) + + answer, runtime_result, message = await _run_with_mcp_tool(tmp_path, mcp_tool) + + assert answer == "report ready" + assert message.artifact is None + assert runtime_result.structured == {"report_id": "r-1"} + assert runtime_result.artifact == { + "file": {"uri": "s3://reports/r-1.pdf", "mime_type": "application/pdf"}, + "preview": {"type": "binary", "size": 14, "omitted": True}, + } + + +async def test_structured_result_values_are_not_automatically_logged( + tmp_path: Path, + caplog: pytest.LogCaptureFixture, +) -> None: + sensitive_value = "structured-value-must-stay-private" + + def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + return [{"type": "text", "text": "complete"}], { + "structured_content": {"private": sensitive_value} + } + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="private_report", + description="report", + response_format="content_and_artifact", + ) + + with caplog.at_level(logging.DEBUG): + _, runtime_result, message = await _run_with_mcp_tool(tmp_path, mcp_tool) + + assert runtime_result.structured == {"private": sensitive_value} + assert message.artifact is None + assert sensitive_value not in caplog.text + assert all( + sensitive_value not in repr(getattr(record, "fields", {})) for record in caplog.records + ) + + +async def test_malformed_artifact_fails_with_controlled_model_text(tmp_path: Path) -> None: + def fake_mcp_tool() -> tuple[list[dict[str, str]], str]: + return [{"type": "text", "text": "valid text"}], "not-a-mapping" + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="malformed", + description="malformed", + response_format="content_and_artifact", + ) + + answer, runtime_result, _ = await _run_with_mcp_tool(tmp_path, mcp_tool) + + assert answer == "Tool error: invalid tool result" + assert runtime_result == NormalizedToolResult.text_only(answer) + + +async def test_non_json_structured_value_fails_without_using_repr_or_logging_data( + tmp_path: Path, + caplog: pytest.LogCaptureFixture, +) -> None: + unexpected = object() + sensitive_key = "private-account-identifier" + + def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + return [{"type": "text", "text": "valid text"}], { + "structured_content": {sensitive_key: unexpected} + } + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="non_json", + description="non-json", + response_format="content_and_artifact", + ) + + with caplog.at_level(logging.DEBUG): + answer, runtime_result, _ = await _run_with_mcp_tool(tmp_path, mcp_tool) + + assert answer == "Tool error: invalid tool result" + assert hex(id(unexpected)) not in answer + assert sensitive_key not in answer + assert hex(id(unexpected)) not in caplog.text + assert sensitive_key not in caplog.text + assert runtime_result == NormalizedToolResult.text_only(answer) + + +async def test_oversized_mcp_structured_result_fails_with_controlled_text( + tmp_path: Path, +) -> None: + def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + return [{"type": "text", "text": "valid text"}], {"structured_content": list(range(10_001))} + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="oversized", + description="oversized", + response_format="content_and_artifact", + ) + + answer, runtime_result, _ = await _run_with_mcp_tool(tmp_path, mcp_tool) + + assert answer == "Tool error: invalid tool result" + assert runtime_result == NormalizedToolResult.text_only(answer) def _write_tool(base_dir: Path, tool_id: str) -> None: @@ -171,3 +429,207 @@ async def test_local_tool_result_unaffected(tmp_path: Path) -> None: result = await engine.run("book please") assert result.answer == "did: go" + + +async def test_local_structured_tool_preserves_data_and_existing_text(tmp_path: Path) -> None: + tool_id = "get_orders" + body = f"def {tool_id}(message: str) -> dict:\n return {{'orders': [{{'id': 'ORDER-1'}}]}}\n" + tools_dir = tmp_path / "plugins" / "tools" + tools_dir.mkdir(parents=True, exist_ok=True) + (tools_dir / f"{tool_id}.py").write_text(body, encoding="utf-8") + spec = _system(_agent("orders", tools=(ToolSpec(tool_id, "orders"),))) + captured: list[ToolMessage] = [] + repository = _CapturingExecutionRepository() + + factory = _model_factory(tool_args={"message": "go"}, captured=captured) + async with LangGraphEngine( + tmp_path, + model_factory=factory, + execution_manager=ToolExecutionManager(execution_repository=repository), + ) as engine: + await engine.build(spec) + result = await engine.run("orders please") + + assert repository.completed_result is not None + runtime_result = NormalizedToolResult.from_persisted(repository.completed_result) + assert result.answer == '{"orders": [{"id": "ORDER-1"}]}' + assert captured[-1].artifact is None + assert runtime_result.structured == {"orders": [{"id": "ORDER-1"}]} + + +async def test_idempotent_replay_restores_identical_normalized_result() -> None: + calls = 0 + + def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + nonlocal calls + calls += 1 + return [{"type": "text", "text": "Found 2 invoices"}], { + "structured_content": {"count": 2}, + "source": "billing", + } + + tool = StructuredTool.from_function( + fake_mcp_tool, + name="invoice_search", + description="search invoices", + response_format="content_and_artifact", + ) + execution_manager = ToolExecutionManager(execution_repository=InMemoryToolExecutionRepository()) + invoker = _invoker(tool, execution_manager) + tool_call = {"id": "call-1", "name": tool.name, "args": {}} + token = current_run_context.set(RunContext(run_id="run-1")) + try: + first = await invoker.invoke(tool_call) + replayed = await invoker.invoke(tool_call) + finally: + current_run_context.reset(token) + + assert calls == 1 + assert replayed == first + assert replayed.text == "Found 2 invoices" + assert replayed.structured == {"count": 2} + assert replayed.artifact == {"source": "billing"} + + +async def test_concurrent_duplicate_calls_share_one_provider_execution() -> None: + calls = 0 + started = asyncio.Event() + release = asyncio.Event() + + async def fake_mcp_tool() -> tuple[list[dict[str, str]], dict[str, Any]]: + nonlocal calls + calls += 1 + started.set() + await release.wait() + return [{"type": "text", "text": "done"}], {"structured_content": {"ok": True}} + + tool = StructuredTool.from_function( + coroutine=fake_mcp_tool, + name="slow_tool", + description="slow", + response_format="content_and_artifact", + ) + invoker = _invoker( + tool, + ToolExecutionManager(execution_repository=InMemoryToolExecutionRepository()), + ) + tool_call = {"id": "call-1", "name": tool.name, "args": {}} + token = current_run_context.set(RunContext(run_id="run-concurrent")) + try: + owner = asyncio.create_task(invoker.invoke(tool_call)) + await started.wait() + duplicate = asyncio.create_task(invoker.invoke(tool_call)) + await asyncio.sleep(0) + release.set() + first, second = await asyncio.gather(owner, duplicate) + finally: + current_run_context.reset(token) + + assert calls == 1 + assert first == second == NormalizedToolResult("done", structured={"ok": True}) + + +async def test_cancellation_during_failed_result_write_does_not_strand_duplicate() -> None: + calls = 0 + + async def failing_tool() -> str: + nonlocal calls + calls += 1 + raise RuntimeError("provider failed") + + tool = StructuredTool.from_function( + coroutine=failing_tool, + name="failing_tool", + description="fails", + ) + repository = _BlockingCompletionRepository() + invoker = _invoker( + tool, + ToolExecutionManager(execution_repository=repository), + ) + tool_call = {"id": "call-1", "name": tool.name, "args": {}} + token = current_run_context.set(RunContext(run_id="run-cancelled-failure")) + owner = asyncio.create_task(invoker.invoke(tool_call)) + try: + await repository.completing.wait() + duplicate = asyncio.create_task(invoker.invoke(tool_call)) + owner.cancel() + await asyncio.sleep(0) + assert duplicate.done() is False + + repository.release.set() + with pytest.raises(asyncio.CancelledError): + await owner + replayed = await asyncio.wait_for(duplicate, timeout=1) + finally: + repository.release.set() + if not owner.done(): + owner.cancel() + current_run_context.reset(token) + + assert calls == 1 + assert replayed == NormalizedToolResult.text_only("Tool error: provider failed") + + +async def test_cancellation_during_successful_result_write_preserves_success() -> None: + calls = 0 + + async def successful_tool() -> str: + nonlocal calls + calls += 1 + return "completed" + + tool = StructuredTool.from_function( + coroutine=successful_tool, + name="successful_tool", + description="succeeds", + ) + repository = _BlockingCompletionRepository() + invoker = _invoker( + tool, + ToolExecutionManager(execution_repository=repository), + ) + tool_call = {"id": "call-1", "name": tool.name, "args": {}} + token = current_run_context.set(RunContext(run_id="run-cancelled-success")) + owner = asyncio.create_task(invoker.invoke(tool_call)) + try: + await repository.completing.wait() + duplicate = asyncio.create_task(invoker.invoke(tool_call)) + owner.cancel() + await asyncio.sleep(0) + assert duplicate.done() is False + + repository.release.set() + with pytest.raises(asyncio.CancelledError): + await owner + replayed = await asyncio.wait_for(duplicate, timeout=1) + finally: + repository.release.set() + if not owner.done(): + owner.cancel() + current_run_context.reset(token) + + assert calls == 1 + assert replayed == NormalizedToolResult.text_only("completed") + + +def _invoker(tool: StructuredTool, execution_manager: ToolExecutionManager) -> ToolInvoker: + return ToolInvoker( + spec=AgentSpec( + id="billing", + name="billing", + description="billing", + model=_MODEL, + auto_mode=True, + ), + node_path="billing", + binding=AgentToolBinding( + tools={tool.name: tool}, + mcp_tool_names=frozenset({tool.name}), + mcp_server_by_tool={tool.name: "billing-mcp"}, + ), + hook_manager=HookManager.empty(), + execution_manager=execution_manager, + approval_coordinator=ApprovalCoordinator(_UnusedApprovalProvider()), + usage_tracker=ToolUsageTracker(InMemoryToolUsageRepository()), + ) diff --git a/tests/runtime/test_tool_hooks.py b/tests/runtime/test_tool_hooks.py index dd3f947d..d9ac3e8f 100644 --- a/tests/runtime/test_tool_hooks.py +++ b/tests/runtime/test_tool_hooks.py @@ -229,6 +229,38 @@ def factory(provider: str, name: str, temperature: float | None) -> BaseChatMode assert result.answer == "Y" * 3 +async def test_transform_hook_receives_mcp_structured_result_without_losing_it( + tmp_path: Path, +) -> None: + spec = _system( + _agent("research", mcps=(MCPSpec(id="wiki", url="https://wiki.test/mcp"),)), + HookSpec("transform_tool_result", f"{_FIX}:truncate_tool_result"), + ) + + def fake_mcp_tool(message: str) -> tuple[list[dict[str, str]], dict[str, Any]]: + return [{"type": "text", "text": "abcdef"}], {"structured_content": {"count": 2}} + + mcp_tool = StructuredTool.from_function( + fake_mcp_tool, + name="wiki_search", + description="search", + response_format="content_and_artifact", + ) + + def factory(provider: str, name: str, temperature: float | None) -> BaseChatModel: + return cast(BaseChatModel, EchoToolResultModel()) + + async with LangGraphEngine(tmp_path, model_factory=factory) as engine: + await engine.build(spec) + engine._mcp_tools["wiki"] = [mcp_tool] + engine._app = engine._build_graph(spec) + result = await engine.run("go") + + transformed = next(c[1] for c in fixtures.CALLS if c[0] == "transform_tool_result") + assert transformed.structured_result == {"count": 2} + assert result.answer == "abc" + + async def test_after_tool_call_receives_provider_and_server_id( tmp_path: Path, model_factory: Any ) -> None: diff --git a/tests/runtime/test_tool_results.py b/tests/runtime/test_tool_results.py new file mode 100644 index 00000000..6989e7f0 --- /dev/null +++ b/tests/runtime/test_tool_results.py @@ -0,0 +1,120 @@ +"""Normalized tool-result invariants and persistence compatibility.""" + +from __future__ import annotations + +import json + +import pytest +from langchain_core.messages import ToolMessage + +from agent_engine.engine.langgraph.tools.tool_result_normalizer import ( + ProviderToolResultError, + normalize_tool_result, +) +from agent_engine.runtime.tool_results import ( + NormalizedToolResult, + ToolResultValidationError, +) + + +def test_persisted_round_trip_is_json_serializable_and_semantically_equal() -> None: + result = NormalizedToolResult( + "Found invoices", + structured={"z": 1, "invoices": [{"id": "INV-1"}]}, + artifact={"source": "billing"}, + ) + + payload = result.to_persisted() + encoded = json.dumps(payload, sort_keys=True, allow_nan=False) + restored = NormalizedToolResult.from_persisted(json.loads(encoded)) + + assert restored == result + assert restored.structured_text() == '{"invoices":[{"id":"INV-1"}],"z":1}' + + +def test_nested_values_cannot_mutate_the_owned_result() -> None: + source = {"items": [{"id": "one"}]} + result = NormalizedToolResult("ok", structured=source) + source["items"][0]["id"] = "changed" + exposed = result.structured + assert isinstance(exposed, dict) + exposed["items"] = [] + + assert result.structured == {"items": [{"id": "one"}]} + + +@pytest.mark.parametrize( + "value", + [ + {1: "non-string key"}, + {"value": object()}, + {"value": float("nan")}, + ], +) +def test_non_json_structured_values_are_rejected(value: object) -> None: + with pytest.raises(ToolResultValidationError): + NormalizedToolResult("ok", structured=value) + + +@pytest.mark.parametrize( + "value", + [ + list(range(10_001)), + "x" * 1_000_001, + ], +) +def test_oversized_structured_values_are_rejected(value: object) -> None: + with pytest.raises(ToolResultValidationError, match="budget"): + NormalizedToolResult("ok", structured=value) + + +def test_legacy_text_ledger_value_restores_as_text_only() -> None: + assert NormalizedToolResult.from_persisted("legacy") == NormalizedToolResult.text_only("legacy") + + +def test_unknown_persistence_version_is_rejected() -> None: + with pytest.raises(ToolResultValidationError, match="unsupported"): + NormalizedToolResult.from_persisted( + {"version": 99, "text": "future", "structured": None, "artifact": None} + ) + + +def test_empty_provider_result_has_stable_model_text() -> None: + result = normalize_tool_result(ToolMessage(content=[], tool_call_id="call-1")) + + assert result == NormalizedToolResult.text_only("[tool returned no content]") + + +def test_explicit_provider_error_preserves_only_model_safe_text() -> None: + message = ToolMessage( + content=[{"type": "text", "text": "account not found"}], + tool_call_id="call-1", + status="error", + ) + + with pytest.raises(ProviderToolResultError) as raised: + normalize_tool_result(message) + + assert raised.value.model_text == "account not found" + assert "account not found" not in str(raised.value) + + +def test_artifact_metadata_has_an_aggregate_size_budget() -> None: + result = normalize_tool_result( + ToolMessage( + content="ready", + tool_call_id="call-1", + artifact={ + "structured_content": {"report_id": "r-1"}, + "metadata": {f"field-{index}": "x" * 8000 for index in range(20)}, + }, + ) + ) + + assert result.text == "ready" + assert result.structured == {"report_id": "r-1"} + assert result.artifact is not None + encoded = json.dumps(result.artifact) + assert len(encoded) < 65536 + assert encoded.count("x" * 8000) <= 4 + assert '"omitted": true' in encoded diff --git a/tests/test_logging.py b/tests/test_logging.py index b4a222e6..b441a871 100644 --- a/tests/test_logging.py +++ b/tests/test_logging.py @@ -2,6 +2,8 @@ import logging +from langchain_core.messages import ToolMessage + from agent_engine.logging_config import ( StructuredFormatter, configure_logging, @@ -72,6 +74,26 @@ def test_tool_event_emits_structured_record(caplog): assert record.fields == {"name": "search"} +def test_tool_callback_logs_only_safe_result_metadata(caplog): + sensitive = "structured-value-must-stay-private" + handler = LoggingCallbackHandler() + + with caplog.at_level(logging.DEBUG, logger="agent_engine.trace"): + handler.on_tool_start({"name": "search"}, sensitive) + handler.on_tool_end( + ToolMessage( + content="complete", + tool_call_id="call-1", + artifact={"structured_content": {"private": sensitive}}, + ) + ) + + fields = [getattr(record, "fields", {}) for record in caplog.records] + assert all(sensitive not in repr(value) for value in fields) + assert {"chars": len(sensitive)} in fields + assert {"output_type": "ToolMessage"} in fields + + def test_log_helper_attaches_fields(caplog): with caplog.at_level(logging.INFO, logger="agent_engine.test"): log(logging.getLogger("agent_engine.test"), logging.INFO, "ping", n=1)