fix: harden tool-call protocol and fallback handling
This commit is contained in:
@@ -19,6 +19,7 @@ Usage:
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from collections.abc import Sequence
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
@@ -304,9 +305,12 @@ def _inject_subagent_middleware(
|
||||
path doesn't fall back to the global-writing ``_ensure_chat_model()``.
|
||||
"""
|
||||
from .middleware import (
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
ContextOverflowMapperMiddleware,
|
||||
ErrorNormalizationMiddleware,
|
||||
RepetitiveToolCallGuardMiddleware,
|
||||
ToolErrorHandlerMiddleware,
|
||||
ToolProtocolGuardMiddleware,
|
||||
create_context_editing_middleware,
|
||||
create_memory_lifecycle_middleware,
|
||||
create_memory_middleware,
|
||||
@@ -315,6 +319,16 @@ def _inject_subagent_middleware(
|
||||
)
|
||||
|
||||
cfg = cfg if cfg is not None else _ensure_config()
|
||||
repetitive_tool_call_threshold = getattr(
|
||||
cfg,
|
||||
"repetitive_tool_call_threshold",
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
)
|
||||
if not isinstance(repetitive_tool_call_threshold, int):
|
||||
repetitive_tool_call_threshold = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD
|
||||
max_consecutive_tool_errors = getattr(cfg, "max_consecutive_tool_errors", 3)
|
||||
if not isinstance(max_consecutive_tool_errors, int):
|
||||
max_consecutive_tool_errors = 3
|
||||
memory_controls = MemoryControls.from_config(cfg)
|
||||
memory_dir = str(_paths_mod.MEMORIES_DIR)
|
||||
memory_scheduler = default_memory_scheduler()
|
||||
@@ -339,6 +353,11 @@ def _inject_subagent_middleware(
|
||||
# them into a non-dataclass envelope wrapper before
|
||||
# anything downstream sees them.
|
||||
ErrorNormalizationMiddleware(),
|
||||
RepetitiveToolCallGuardMiddleware(
|
||||
threshold=repetitive_tool_call_threshold,
|
||||
max_consecutive_errors=max_consecutive_tool_errors,
|
||||
),
|
||||
ToolProtocolGuardMiddleware(),
|
||||
# Subagents share the main agent's model: use the threaded
|
||||
# ``chat_model`` on the pure path, else defer to the factory's
|
||||
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
|
||||
@@ -654,6 +673,7 @@ def _get_default_middleware(
|
||||
tool_selector_threshold: int | None = None,
|
||||
memory_max_inline_profile_chars: int | None = None,
|
||||
enable_background_execution: bool = True,
|
||||
enable_legacy_model_fallback: bool = True,
|
||||
):
|
||||
"""Build the default middleware list.
|
||||
|
||||
@@ -675,11 +695,14 @@ def _get_default_middleware(
|
||||
Async sub-agent factories pass their deployed agent name here.
|
||||
"""
|
||||
from .middleware import (
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
ConfigurableModelMiddleware,
|
||||
ContextOverflowMapperMiddleware,
|
||||
ErrorNormalizationMiddleware,
|
||||
ModelFallbackMiddleware,
|
||||
RepetitiveToolCallGuardMiddleware,
|
||||
ToolErrorHandlerMiddleware,
|
||||
ToolProtocolGuardMiddleware,
|
||||
create_code_interpreter_middleware,
|
||||
create_context_editing_middleware,
|
||||
create_memory_lifecycle_middleware,
|
||||
@@ -692,6 +715,16 @@ def _get_default_middleware(
|
||||
)
|
||||
|
||||
cfg = cfg if cfg is not None else _ensure_config()
|
||||
repetitive_tool_call_threshold = getattr(
|
||||
cfg,
|
||||
"repetitive_tool_call_threshold",
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
)
|
||||
if not isinstance(repetitive_tool_call_threshold, int):
|
||||
repetitive_tool_call_threshold = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD
|
||||
max_consecutive_tool_errors = getattr(cfg, "max_consecutive_tool_errors", 3)
|
||||
if not isinstance(max_consecutive_tool_errors, int):
|
||||
max_consecutive_tool_errors = 3
|
||||
if cfg.model_fallbacks:
|
||||
load_fallback_chain(cfg.model_fallbacks)
|
||||
model = chat_model if chat_model is not None else _ensure_chat_model()
|
||||
@@ -741,6 +774,15 @@ def _get_default_middleware(
|
||||
from .llm import get_chat_model
|
||||
|
||||
tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider)
|
||||
selector_middlewares = create_tool_selector_middleware(
|
||||
**(
|
||||
{"threshold": tool_selector_threshold}
|
||||
if tool_selector_threshold is not None
|
||||
else {}
|
||||
),
|
||||
model=tool_selector_model,
|
||||
track_stream_selection=not for_async_subagent,
|
||||
)
|
||||
mw = [
|
||||
# Outermost — catches provider-SDK exceptions from the model
|
||||
# call (including exceptions surfaced through inner
|
||||
@@ -749,18 +791,15 @@ def _get_default_middleware(
|
||||
ErrorNormalizationMiddleware(),
|
||||
ConfigurableModelMiddleware(),
|
||||
create_context_editing_middleware(model),
|
||||
ModelFallbackMiddleware(),
|
||||
*([ModelFallbackMiddleware()] if enable_legacy_model_fallback else []),
|
||||
RepetitiveToolCallGuardMiddleware(
|
||||
threshold=repetitive_tool_call_threshold,
|
||||
max_consecutive_errors=max_consecutive_tool_errors,
|
||||
),
|
||||
ContextOverflowMapperMiddleware(),
|
||||
ToolErrorHandlerMiddleware(),
|
||||
*create_tool_selector_middleware(
|
||||
**(
|
||||
{"threshold": tool_selector_threshold}
|
||||
if tool_selector_threshold is not None
|
||||
else {}
|
||||
),
|
||||
model=tool_selector_model,
|
||||
track_stream_selection=not for_async_subagent,
|
||||
),
|
||||
*selector_middlewares,
|
||||
ToolProtocolGuardMiddleware(),
|
||||
# Interpreter prompt must land before runtime/memory context, so this
|
||||
# middleware sits ahead of runtime_context in the stack.
|
||||
create_code_interpreter_middleware(
|
||||
@@ -897,6 +936,8 @@ def create_cli_agent(
|
||||
memory_max_inline_profile_chars: int | None = None,
|
||||
enable_subagents: bool = True,
|
||||
enable_background_execution: bool = True,
|
||||
main_agent_outer_middlewares: Sequence[AgentMiddleware] | None = None,
|
||||
main_agent_route_middleware: AgentMiddleware | None = None,
|
||||
) -> "CompiledStateGraph":
|
||||
"""Create agent with checkpointer for CLI multi-turn support.
|
||||
|
||||
@@ -933,6 +974,12 @@ def create_cli_agent(
|
||||
enable_background_execution: Whether local background-process tools are
|
||||
installed. Embedding hosts should disable this when process execution
|
||||
is provided by an external backend.
|
||||
main_agent_outer_middlewares: Optional host-owned middleware installed
|
||||
only on the top-level agent, outside EvoScientist's default chain.
|
||||
main_agent_route_middleware: Optional host-owned route middleware placed
|
||||
after ConfigurableModelMiddleware and before tool selection. When
|
||||
provided, EvoScientist's legacy model fallback is disabled for the
|
||||
top-level agent so the host is the only fallback authority.
|
||||
"""
|
||||
import os as _os
|
||||
|
||||
@@ -1017,7 +1064,22 @@ def create_cli_agent(
|
||||
tool_selector_threshold=tool_selector_threshold,
|
||||
memory_max_inline_profile_chars=memory_max_inline_profile_chars,
|
||||
enable_background_execution=enable_background_execution,
|
||||
enable_legacy_model_fallback=main_agent_route_middleware is None,
|
||||
)
|
||||
if main_agent_route_middleware is not None:
|
||||
configurable_index = next(
|
||||
(
|
||||
index
|
||||
for index, middleware in enumerate(mw)
|
||||
if getattr(middleware, "name", "") == "configurable_model"
|
||||
),
|
||||
None,
|
||||
)
|
||||
if configurable_index is None:
|
||||
raise RuntimeError("ConfigurableModelMiddleware route slot is unavailable")
|
||||
mw.insert(configurable_index + 1, main_agent_route_middleware)
|
||||
if main_agent_outer_middlewares:
|
||||
mw = [*main_agent_outer_middlewares, *mw]
|
||||
|
||||
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
|
||||
# would propagate it to every subagent, breaking parallel execute calls
|
||||
|
||||
@@ -262,6 +262,13 @@ class EvoScientistConfig:
|
||||
# Lower (e.g., 5000) if you want a tighter safety net against runaway loops.
|
||||
recursion_limit: int = 1_000_000
|
||||
|
||||
# Number of consecutive model rounds with the same structured tool name and
|
||||
# arguments that activates provider-facing loop repair. Set 0 to disable.
|
||||
repetitive_tool_call_threshold: int = 2
|
||||
# Number of consecutive deterministic tool errors allowed before the next
|
||||
# model call is blocked. Transient provider/network errors are not counted.
|
||||
max_consecutive_tool_errors: int = 3
|
||||
|
||||
# Memory Settings
|
||||
# Profile memory injects and maintains `/memories/profile/...` files.
|
||||
memory_profile_enabled: bool = True
|
||||
@@ -473,6 +480,14 @@ class EvoScientistConfig:
|
||||
stt_compute_type: str = "int8" # "int8" | "float16" | "float32"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
for field_name in (
|
||||
"repetitive_tool_call_threshold",
|
||||
"max_consecutive_tool_errors",
|
||||
):
|
||||
value = getattr(self, field_name)
|
||||
if not isinstance(value, int) or isinstance(value, bool) or value < 0:
|
||||
raise ValueError(f"{field_name} must be a non-negative integer")
|
||||
|
||||
# A non-positive or non-int sandbox_execute_timeout (e.g. a hand-edited
|
||||
# config file value — load_config does not coerce file values — or a
|
||||
# 0/negative env value) would raise inside CustomSandboxBackend.__init__
|
||||
@@ -735,6 +750,11 @@ def set_config_value(key: str, value: Any) -> bool:
|
||||
|
||||
if key == "sandbox_execute_timeout" and value <= 0:
|
||||
return False
|
||||
if key in {
|
||||
"repetitive_tool_call_threshold",
|
||||
"max_consecutive_tool_errors",
|
||||
} and (isinstance(value, bool) or value < 0):
|
||||
return False
|
||||
if key == "memory_skill_synthesis_time":
|
||||
value = _normalize_hhmm(value)
|
||||
if value is None:
|
||||
@@ -813,6 +833,10 @@ _ENV_MAPPINGS = {
|
||||
"langgraph_dev_file_persistence": "EVOSCIENTIST_LANGGRAPH_DEV_FILE_PERSISTENCE",
|
||||
"langgraph_dev_jobs_per_worker": "EVOSCIENTIST_LANGGRAPH_DEV_JOBS_PER_WORKER",
|
||||
"recursion_limit": "EVOSCIENTIST_RECURSION_LIMIT",
|
||||
"repetitive_tool_call_threshold": (
|
||||
"EVOSCIENTIST_REPETITIVE_TOOL_CALL_THRESHOLD"
|
||||
),
|
||||
"max_consecutive_tool_errors": "EVOSCIENTIST_MAX_CONSECUTIVE_TOOL_ERRORS",
|
||||
"memory_profile_enabled": "EVOSCIENTIST_MEMORY_PROFILE_ENABLED",
|
||||
"memory_observations_enabled": "EVOSCIENTIST_MEMORY_OBSERVATIONS_ENABLED",
|
||||
"memory_observation_writer": "EVOSCIENTIST_MEMORY_OBSERVATION_WRITER",
|
||||
|
||||
@@ -32,6 +32,98 @@ from typing import Any
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class AgentControlError(Exception):
|
||||
"""Host-defined terminal control error that must bypass model fallback."""
|
||||
|
||||
non_fallbackable = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
code: str,
|
||||
message: str,
|
||||
*,
|
||||
status_code: int = 403,
|
||||
retryable: bool = False,
|
||||
) -> None:
|
||||
super().__init__(message)
|
||||
self.code = code
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
self.retryable = retryable
|
||||
|
||||
def model_dump(self) -> dict[str, Any]:
|
||||
return {
|
||||
"error": type(self).__name__,
|
||||
"code": self.code,
|
||||
"message": self.message,
|
||||
"status_code": self.status_code,
|
||||
"retryable": self.retryable,
|
||||
}
|
||||
|
||||
|
||||
class ModelToolProtocolError(AgentControlError):
|
||||
"""A completed model response contained an invalid tool-call protocol."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reason: str,
|
||||
*,
|
||||
provider: str | None = None,
|
||||
model: str | None = None,
|
||||
route_key: str | None = None,
|
||||
config_generation: int | None = None,
|
||||
api_mode: str | None = None,
|
||||
endpoint: str | None = None,
|
||||
tool_call_transport: str | None = None,
|
||||
call_id: str | None = None,
|
||||
call_diagnostic: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
"MODEL_TOOL_PROTOCOL_INVALID",
|
||||
"The model returned an invalid structured tool call.",
|
||||
status_code=502,
|
||||
retryable=False,
|
||||
)
|
||||
self.reason = reason
|
||||
self.provider = provider
|
||||
self.model = model
|
||||
self.route_key = route_key
|
||||
self.config_generation = config_generation
|
||||
self.api_mode = api_mode
|
||||
self.endpoint = endpoint
|
||||
self.tool_call_transport = tool_call_transport
|
||||
self.call_id = call_id
|
||||
# Internal-only, redacted structure for server logs. Deliberately omitted
|
||||
# from model_dump() so it never becomes part of the public SSE contract.
|
||||
self.call_diagnostic = dict(call_diagnostic or {})
|
||||
self.fallbackable = True
|
||||
self.recoverable = True
|
||||
|
||||
def model_dump(self) -> dict[str, Any]:
|
||||
payload = super().model_dump()
|
||||
payload.update(
|
||||
{
|
||||
"reason": self.reason,
|
||||
"fallbackable": self.fallbackable,
|
||||
"recoverable": self.recoverable,
|
||||
}
|
||||
)
|
||||
for key in (
|
||||
"provider",
|
||||
"model",
|
||||
"route_key",
|
||||
"config_generation",
|
||||
"api_mode",
|
||||
"endpoint",
|
||||
"tool_call_transport",
|
||||
"call_id",
|
||||
):
|
||||
value = getattr(self, key)
|
||||
if value is not None:
|
||||
payload[key] = value
|
||||
return payload
|
||||
|
||||
|
||||
class ProviderStreamError(Exception):
|
||||
"""Envelope-shaped wrapper for a provider SDK exception raised
|
||||
inside a chat model call.
|
||||
|
||||
+276
-55
@@ -284,71 +284,288 @@ def _stable_tool_call_id(message: Any, message_index: int, call_index: int) -> s
|
||||
return "call_" + hashlib.sha256(seed.encode("utf-8")).hexdigest()[:24]
|
||||
|
||||
|
||||
def _ensure_openai_tool_call_ids(messages: list[Any]) -> list[Any]:
|
||||
"""Copy messages and repair missing AI/ToolMessage call identifiers."""
|
||||
import copy
|
||||
from collections import deque
|
||||
def _tool_message_match_index(
|
||||
tool_messages: list[Any],
|
||||
used_indexes: set[int],
|
||||
*,
|
||||
call_id: str,
|
||||
call_name: str,
|
||||
) -> int | None:
|
||||
"""Find the best unused result for one assistant tool call."""
|
||||
|
||||
pending_call_ids: deque[str] = deque()
|
||||
normalized: list[Any] = []
|
||||
def _matches(index: int, *, require_id: bool, require_name: bool) -> bool:
|
||||
if index in used_indexes:
|
||||
return False
|
||||
message = tool_messages[index]
|
||||
result_id = str(getattr(message, "tool_call_id", "") or "")
|
||||
result_name = str(getattr(message, "name", "") or "")
|
||||
if require_id and result_id != call_id:
|
||||
return False
|
||||
if not require_id and result_id:
|
||||
return False
|
||||
return not require_name or not result_name or result_name == call_name
|
||||
|
||||
for message_index, message in enumerate(messages):
|
||||
message_type = getattr(message, "type", None)
|
||||
if message_type == "ai":
|
||||
tool_calls = list(getattr(message, "tool_calls", None) or [])
|
||||
if not tool_calls:
|
||||
normalized.append(message)
|
||||
if call_id:
|
||||
for require_name in (True, False):
|
||||
for index in range(len(tool_messages)):
|
||||
if _matches(index, require_id=True, require_name=require_name):
|
||||
return index
|
||||
for require_name in (True, False):
|
||||
for index in range(len(tool_messages)):
|
||||
if _matches(index, require_id=False, require_name=require_name):
|
||||
return index
|
||||
return None
|
||||
|
||||
# A result-side identifier is more authoritative than a generated fallback.
|
||||
for require_name in (True, False):
|
||||
for index, message in enumerate(tool_messages):
|
||||
if index in used_indexes:
|
||||
continue
|
||||
result_id = str(getattr(message, "tool_call_id", "") or "")
|
||||
result_name = str(getattr(message, "name", "") or "")
|
||||
if result_id and (
|
||||
not require_name or not result_name or result_name == call_name
|
||||
):
|
||||
return index
|
||||
for require_name in (True, False):
|
||||
for index in range(len(tool_messages)):
|
||||
if _matches(index, require_id=False, require_name=require_name):
|
||||
return index
|
||||
return None
|
||||
|
||||
copied = copy.copy(message)
|
||||
normalized_calls: list[dict[str, Any]] = []
|
||||
for call_index, original_call in enumerate(tool_calls):
|
||||
call = dict(original_call)
|
||||
call_id = str(call.get("id") or "") or _stable_tool_call_id(
|
||||
message, message_index, call_index
|
||||
)
|
||||
call["id"] = call_id
|
||||
normalized_calls.append(call)
|
||||
pending_call_ids.append(call_id)
|
||||
copied.tool_calls = normalized_calls
|
||||
|
||||
if isinstance(copied.content, list):
|
||||
call_index = 0
|
||||
blocks: list[Any] = []
|
||||
for original_block in copied.content:
|
||||
if not isinstance(original_block, dict):
|
||||
blocks.append(original_block)
|
||||
continue
|
||||
block = dict(original_block)
|
||||
if block.get("type") in {"tool_call", "function_call"}:
|
||||
if call_index < len(normalized_calls):
|
||||
block["id"] = normalized_calls[call_index]["id"]
|
||||
call_index += 1
|
||||
blocks.append(block)
|
||||
copied.content = blocks
|
||||
normalized.append(copied)
|
||||
def _copy_ai_message_with_tool_pairs(
|
||||
message: Any,
|
||||
message_index: int,
|
||||
tool_messages: list[Any],
|
||||
) -> tuple[Any | None, list[Any]]:
|
||||
"""Return a replay-safe assistant message and its matched tool results."""
|
||||
import copy
|
||||
|
||||
copied = copy.copy(message)
|
||||
additional_kwargs = dict(getattr(message, "additional_kwargs", None) or {})
|
||||
# Parsed tool_calls are canonical. Raw copies can otherwise re-introduce an
|
||||
# invalid call after invalid_tool_calls has been cleared.
|
||||
additional_kwargs.pop("tool_calls", None)
|
||||
copied.additional_kwargs = additional_kwargs
|
||||
copied.invalid_tool_calls = []
|
||||
|
||||
original_calls = list(getattr(message, "tool_calls", None) or [])
|
||||
used_results: set[int] = set()
|
||||
matched_calls: list[dict[str, Any]] = []
|
||||
matched_result_indexes: list[int] = []
|
||||
original_to_matched_call: dict[int, tuple[str, str]] = {}
|
||||
|
||||
for call_index, original_call in enumerate(original_calls):
|
||||
call = dict(original_call)
|
||||
call_id = str(call.get("id") or "")
|
||||
call_name = str(call.get("name") or "").strip()
|
||||
# A missing name is structurally unreplayable. Never infer it from
|
||||
# arguments or retain its paired ToolMessage in provider history.
|
||||
if not call_name:
|
||||
continue
|
||||
call["name"] = call_name
|
||||
result_index = _tool_message_match_index(
|
||||
tool_messages,
|
||||
used_results,
|
||||
call_id=call_id,
|
||||
call_name=call_name,
|
||||
)
|
||||
# A historical client-side function call is only replayable together
|
||||
# with its result. Incomplete calls are discarded instead of asking the
|
||||
# provider to continue a broken tool turn.
|
||||
if result_index is None:
|
||||
continue
|
||||
if not call_id:
|
||||
result_id = str(
|
||||
getattr(tool_messages[result_index], "tool_call_id", "") or ""
|
||||
)
|
||||
call_id = result_id or _stable_tool_call_id(
|
||||
message, message_index, call_index
|
||||
)
|
||||
call["id"] = call_id
|
||||
matched_calls.append(call)
|
||||
matched_result_indexes.append(result_index)
|
||||
original_to_matched_call[call_index] = (call_id, call_name)
|
||||
used_results.add(result_index)
|
||||
|
||||
copied.tool_calls = matched_calls
|
||||
if isinstance(copied.content, list):
|
||||
original_call_index = 0
|
||||
blocks: list[Any] = []
|
||||
for original_block in copied.content:
|
||||
if not isinstance(original_block, dict):
|
||||
blocks.append(original_block)
|
||||
continue
|
||||
block = dict(original_block)
|
||||
if block.get("type") in {"tool_call", "function_call"}:
|
||||
matched_call = original_to_matched_call.get(original_call_index)
|
||||
original_call_index += 1
|
||||
if matched_call is None:
|
||||
continue
|
||||
call_id, call_name = matched_call
|
||||
# LangChain content blocks use id; the Responses converter later
|
||||
# maps it to call_id.
|
||||
block["id"] = call_id
|
||||
block["name"] = call_name
|
||||
if isinstance(block.get("function"), dict):
|
||||
block["function"] = {**block["function"], "name": call_name}
|
||||
blocks.append(block)
|
||||
copied.content = blocks
|
||||
|
||||
matched_results: list[Any] = []
|
||||
result_to_call_id = {
|
||||
result_index: matched_calls[index]["id"]
|
||||
for index, result_index in enumerate(matched_result_indexes)
|
||||
}
|
||||
for result_index, result in enumerate(tool_messages):
|
||||
call_id = result_to_call_id.get(result_index)
|
||||
if call_id is None:
|
||||
continue
|
||||
copied_result = copy.copy(result)
|
||||
copied_result.tool_call_id = call_id
|
||||
matched_results.append(copied_result)
|
||||
|
||||
had_tool_protocol = bool(original_calls) or bool(
|
||||
getattr(message, "invalid_tool_calls", None)
|
||||
)
|
||||
if not matched_calls and had_tool_protocol:
|
||||
replayable_content = _flatten_message_content(copied.content)
|
||||
if not replayable_content:
|
||||
return None, matched_results
|
||||
|
||||
return copied, matched_results
|
||||
|
||||
|
||||
def _sanitize_openai_tool_history(messages: list[Any]) -> list[Any]:
|
||||
"""Copy history while retaining only complete, replayable tool turns."""
|
||||
|
||||
normalized: list[Any] = []
|
||||
index = 0
|
||||
while index < len(messages):
|
||||
message = messages[index]
|
||||
message_type = getattr(message, "type", None)
|
||||
if message_type == "tool":
|
||||
# A tool result without its immediately preceding assistant call is
|
||||
# invalid for both Chat Completions and Responses APIs.
|
||||
index += 1
|
||||
continue
|
||||
if message_type != "ai":
|
||||
normalized.append(message)
|
||||
index += 1
|
||||
continue
|
||||
|
||||
if message_type == "tool":
|
||||
tool_call_id = str(getattr(message, "tool_call_id", "") or "")
|
||||
if tool_call_id:
|
||||
try:
|
||||
pending_call_ids.remove(tool_call_id)
|
||||
except ValueError:
|
||||
pass
|
||||
normalized.append(message)
|
||||
continue
|
||||
if pending_call_ids:
|
||||
copied = copy.copy(message)
|
||||
copied.tool_call_id = pending_call_ids.popleft()
|
||||
normalized.append(copied)
|
||||
continue
|
||||
|
||||
normalized.append(message)
|
||||
next_index = index + 1
|
||||
tool_messages: list[Any] = []
|
||||
while (
|
||||
next_index < len(messages)
|
||||
and getattr(messages[next_index], "type", None) == "tool"
|
||||
):
|
||||
tool_messages.append(messages[next_index])
|
||||
next_index += 1
|
||||
copied, matched_results = _copy_ai_message_with_tool_pairs(
|
||||
message,
|
||||
index,
|
||||
tool_messages,
|
||||
)
|
||||
if copied is not None:
|
||||
normalized.append(copied)
|
||||
normalized.extend(matched_results)
|
||||
index = next_index
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def _ensure_openai_tool_call_ids(messages: list[Any]) -> list[Any]:
|
||||
"""Backward-compatible alias for replay-safe tool history normalization."""
|
||||
|
||||
return _sanitize_openai_tool_history(messages)
|
||||
|
||||
|
||||
def _has_assistant_tool_protocol(messages: list[Any]) -> bool:
|
||||
"""Return whether history contains assistant-side tool protocol state."""
|
||||
|
||||
for message in messages:
|
||||
if getattr(message, "type", None) != "ai":
|
||||
continue
|
||||
if getattr(message, "tool_calls", None) or getattr(
|
||||
message, "invalid_tool_calls", None
|
||||
):
|
||||
return True
|
||||
additional_kwargs = getattr(message, "additional_kwargs", None) or {}
|
||||
if additional_kwargs.get("tool_calls"):
|
||||
return True
|
||||
content = getattr(message, "content", None)
|
||||
if isinstance(content, list) and any(
|
||||
isinstance(block, dict)
|
||||
and block.get("type") in {"tool_call", "function_call"}
|
||||
for block in content
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _validate_openai_tool_history(messages: list[Any]) -> None:
|
||||
"""Raise when sanitized history still contains an invalid tool protocol."""
|
||||
|
||||
available_call_ids: set[str] = set()
|
||||
for message in messages:
|
||||
message_type = getattr(message, "type", None)
|
||||
if message_type == "ai":
|
||||
if getattr(message, "invalid_tool_calls", None):
|
||||
raise ValueError("invalid_tool_calls must not be replayed")
|
||||
response_call_ids: set[str] = set()
|
||||
response_calls: dict[str, str] = {}
|
||||
for call in getattr(message, "tool_calls", None) or []:
|
||||
call_name = str(call.get("name") or "").strip()
|
||||
if not call_name:
|
||||
raise ValueError("assistant tool call is missing a name")
|
||||
call_id = str(call.get("id") or "").strip()
|
||||
if not call_id:
|
||||
raise ValueError("assistant tool call is missing an id")
|
||||
if call_id in response_call_ids or call_id in available_call_ids:
|
||||
raise ValueError(
|
||||
"assistant tool call id is duplicated while outstanding"
|
||||
)
|
||||
response_call_ids.add(call_id)
|
||||
available_call_ids.add(call_id)
|
||||
response_calls[call_id] = call_name
|
||||
content = getattr(message, "content", None)
|
||||
content_call_ids: set[str] = set()
|
||||
if isinstance(content, list):
|
||||
for block in content:
|
||||
if not isinstance(block, dict) or block.get("type") not in {
|
||||
"tool_call",
|
||||
"function_call",
|
||||
}:
|
||||
continue
|
||||
block_id = str(
|
||||
block.get("id") or block.get("call_id") or ""
|
||||
).strip()
|
||||
block_name = block.get("name") or block.get("tool_name")
|
||||
function = block.get("function")
|
||||
if not block_name and isinstance(function, dict):
|
||||
block_name = function.get("name")
|
||||
block_name = str(block_name or "").strip()
|
||||
if (
|
||||
not block_id
|
||||
or block_id in content_call_ids
|
||||
or response_calls.get(block_id) != block_name
|
||||
):
|
||||
raise ValueError(
|
||||
"assistant content block does not match parsed tool call"
|
||||
)
|
||||
content_call_ids.add(block_id)
|
||||
elif message_type == "tool":
|
||||
call_id = str(getattr(message, "tool_call_id", "") or "")
|
||||
if not call_id or call_id not in available_call_ids:
|
||||
raise ValueError("tool result does not match a prior tool call")
|
||||
available_call_ids.remove(call_id)
|
||||
|
||||
if available_call_ids:
|
||||
raise ValueError("assistant tool call is missing its tool result")
|
||||
|
||||
|
||||
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
|
||||
"""Flatten list content for OpenAI-compatible APIs, preserving media.
|
||||
|
||||
@@ -364,7 +581,9 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
|
||||
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
messages = _ensure_openai_tool_call_ids(messages)
|
||||
sanitize_tool_history = _has_assistant_tool_protocol(messages)
|
||||
if sanitize_tool_history:
|
||||
messages = _sanitize_openai_tool_history(messages)
|
||||
out: list[Any] = []
|
||||
pending_media: list[Any] = [] # media hoisted out of a run of tool messages
|
||||
|
||||
@@ -403,6 +622,8 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
|
||||
msg.content = flat
|
||||
out.append(msg)
|
||||
_flush() # conversation may end with tool messages
|
||||
if sanitize_tool_history:
|
||||
_validate_openai_tool_history(out)
|
||||
return out
|
||||
|
||||
|
||||
|
||||
@@ -29,16 +29,25 @@ from .memory_lifecycle import (
|
||||
default_memory_scheduler,
|
||||
)
|
||||
from .model_fallback import ModelFallbackMiddleware, load_fallback_chain
|
||||
from .repetitive_tool_guard import (
|
||||
DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS,
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
RepetitiveToolCallGuardMiddleware,
|
||||
collapse_repetitive_tool_rounds,
|
||||
)
|
||||
from .runtime_context import RuntimeContextMiddleware, create_runtime_context_middleware
|
||||
from .scheduler import (
|
||||
SchedulerMiddleware,
|
||||
create_scheduler_middleware,
|
||||
)
|
||||
from .tool_error_handler import ToolErrorHandlerMiddleware
|
||||
from .tool_protocol_guard import ToolProtocolGuardMiddleware
|
||||
from .tool_selector import create_tool_selector_middleware
|
||||
from .utils import disable_thinking
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS",
|
||||
"DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD",
|
||||
"AskUserMiddleware",
|
||||
"AskUserRequest",
|
||||
"AskUserWidgetResult",
|
||||
@@ -50,9 +59,12 @@ __all__ = [
|
||||
"EvoMemoryMiddleware",
|
||||
"ModelFallbackMiddleware",
|
||||
"Question",
|
||||
"RepetitiveToolCallGuardMiddleware",
|
||||
"RuntimeContextMiddleware",
|
||||
"SchedulerMiddleware",
|
||||
"ToolErrorHandlerMiddleware",
|
||||
"ToolProtocolGuardMiddleware",
|
||||
"collapse_repetitive_tool_rounds",
|
||||
"compute_context_editing_trigger",
|
||||
"create_code_interpreter_middleware",
|
||||
"create_context_editing_middleware",
|
||||
|
||||
@@ -20,9 +20,9 @@ recognized provider SDK client, wraps the exception in a
|
||||
envelope pre-baked on its instance attributes.
|
||||
|
||||
Contract: the wrap decision is based on the **model**, not the
|
||||
exception. Every exception raised inside a call to a recognized
|
||||
provider model gets wrapped — SDK exceptions, httpx errors,
|
||||
langchain-wrapper failures, and even builtins like ``RuntimeError``.
|
||||
exception, after platform and graph control signals have been excluded.
|
||||
Provider SDK exceptions, httpx errors, langchain-wrapper failures, and
|
||||
even builtins like ``RuntimeError`` get wrapped for a recognized model.
|
||||
At the middleware boundary we can tell which provider was in use, but
|
||||
not the exception's precise origin; a uniform envelope is more useful
|
||||
to the WebUI than gambling on the exception class. If the model isn't
|
||||
@@ -138,11 +138,15 @@ def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError
|
||||
**outside** the user middleware stack to compress history and
|
||||
retry. Wrapping it here would change the type and break that
|
||||
self-healing fallback.
|
||||
- ``AgentControlError`` — a platform-owned typed decision. Gateway route
|
||||
fallback and canonical error mapping depend on its concrete type and
|
||||
structured fields, so it must never become a provider incident.
|
||||
- Models we don't recognize as a provider SDK.
|
||||
"""
|
||||
from langchain_core.exceptions import ContextOverflowError
|
||||
|
||||
from ..llm.errors import (
|
||||
AgentControlError,
|
||||
ProviderStreamError,
|
||||
_extract_error_type,
|
||||
_extract_provider_code,
|
||||
@@ -157,6 +161,12 @@ def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError
|
||||
if isinstance(exc, ProviderStreamError):
|
||||
return None
|
||||
|
||||
# Platform control errors are raised by inner middleware after the provider
|
||||
# response has already been interpreted. Wrapping them would erase routing,
|
||||
# retry and recovery semantics such as ModelToolProtocolError.fallbackable.
|
||||
if isinstance(exc, AgentControlError):
|
||||
return None
|
||||
|
||||
# LangGraph control-flow / structural signals must propagate
|
||||
# untouched, regardless of which caller invoked us.
|
||||
if _should_pass_through(exc):
|
||||
|
||||
@@ -48,6 +48,8 @@ _MALFORMED_REQUEST_PATTERNS: list[str] = [
|
||||
"invalid_request_error",
|
||||
"invalid request",
|
||||
"malformed",
|
||||
"repetitive tool calls",
|
||||
"identical name and arguments",
|
||||
]
|
||||
"""Substrings that identify a malformed request (client-side bug)."""
|
||||
|
||||
@@ -215,6 +217,9 @@ def _is_non_fallbackable(exc: Exception) -> str | None:
|
||||
"""
|
||||
from langchain_core.exceptions import ContextOverflowError
|
||||
|
||||
if getattr(exc, "non_fallbackable", False):
|
||||
return f"platform control error: {getattr(exc, 'code', type(exc).__name__)}"
|
||||
|
||||
if isinstance(exc, ContextOverflowError):
|
||||
return "context length exceeded"
|
||||
|
||||
|
||||
@@ -0,0 +1,350 @@
|
||||
"""Detect deterministic tool loops and compact only provider-facing history."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
ModelRequest,
|
||||
ModelResponse,
|
||||
)
|
||||
|
||||
from ..llm.errors import AgentControlError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD = 2
|
||||
DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS = 3
|
||||
|
||||
_TRANSIENT_PATTERNS = (
|
||||
"timeout",
|
||||
"timed out",
|
||||
"cancelled",
|
||||
"canceled",
|
||||
"connection",
|
||||
"rate limit",
|
||||
"too many requests",
|
||||
"temporarily unavailable",
|
||||
"service unavailable",
|
||||
"overloaded",
|
||||
"bad gateway",
|
||||
"gateway timeout",
|
||||
"http 500",
|
||||
"http 502",
|
||||
"http 503",
|
||||
"http 504",
|
||||
)
|
||||
_DETERMINISTIC_PATTERNS: tuple[tuple[str, tuple[str, ...]], ...] = (
|
||||
(
|
||||
"INVALID_ARGUMENTS",
|
||||
("invalid argument", "validation error", "schema", "bad input"),
|
||||
),
|
||||
("UNKNOWN_TOOL", ("not a valid tool", "unknown tool", "tool not found")),
|
||||
("UNSUPPORTED", ("not supported", "unsupported", "not implemented")),
|
||||
(
|
||||
"POLICY_DENIED",
|
||||
("permission denied", "forbidden", "policy denied", "not allowed"),
|
||||
),
|
||||
)
|
||||
_SAFE_CODE_RE = re.compile(r"^[A-Za-z][A-Za-z0-9_.:-]{0,95}$")
|
||||
_DETERMINISTIC_CODE_MARKERS = (
|
||||
"INVALID",
|
||||
"VALIDATION",
|
||||
"SCHEMA",
|
||||
"UNKNOWN_TOOL",
|
||||
"NOT_FOUND",
|
||||
"UNSUPPORTED",
|
||||
"NOT_IMPLEMENTED",
|
||||
"POLICY",
|
||||
"PERMISSION",
|
||||
"FORBIDDEN",
|
||||
"DENIED",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RepetitiveToolHistoryRepair:
|
||||
messages: list[Any]
|
||||
blocked_tool_names: frozenset[str]
|
||||
removed_rounds: int
|
||||
tail_repetitions: int = 0
|
||||
tail_consecutive_errors: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ToolRound:
|
||||
messages: tuple[Any, ...]
|
||||
signature: tuple[tuple[str, str, str], ...]
|
||||
tool_names: frozenset[str]
|
||||
deterministic_error: bool
|
||||
|
||||
|
||||
def _canonical_tool_args(value: Any) -> str:
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return value.strip()
|
||||
try:
|
||||
return json.dumps(
|
||||
value,
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
default=str,
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
return repr(value)
|
||||
|
||||
|
||||
def _deterministic_result_code(message: Any) -> str | None:
|
||||
additional = getattr(message, "additional_kwargs", None)
|
||||
additional = additional if isinstance(additional, Mapping) else {}
|
||||
raw_code = additional.get("error_code") or additional.get("code")
|
||||
status = str(getattr(message, "status", "") or "").lower()
|
||||
content = str(getattr(message, "content", "") or "")
|
||||
lowered = content.lower()
|
||||
|
||||
if any(pattern in lowered for pattern in _TRANSIENT_PATTERNS):
|
||||
return None
|
||||
if isinstance(raw_code, str) and _SAFE_CODE_RE.fullmatch(raw_code.strip()):
|
||||
normalized = raw_code.strip().upper()
|
||||
if any(
|
||||
pattern.replace(" ", "_") in normalized for pattern in _TRANSIENT_PATTERNS
|
||||
):
|
||||
return None
|
||||
if any(marker in normalized for marker in _DETERMINISTIC_CODE_MARKERS):
|
||||
return normalized
|
||||
return None
|
||||
is_error = status == "error" or lowered.startswith("error:")
|
||||
if not is_error:
|
||||
return None
|
||||
for code, patterns in _DETERMINISTIC_PATTERNS:
|
||||
if any(pattern in lowered for pattern in patterns):
|
||||
return code
|
||||
return None
|
||||
|
||||
|
||||
def _parse_tool_round(
|
||||
messages: Sequence[Any], start: int
|
||||
) -> tuple[_ToolRound, int] | None:
|
||||
assistant = messages[start]
|
||||
if getattr(assistant, "type", None) != "ai":
|
||||
return None
|
||||
raw_calls = list(getattr(assistant, "tool_calls", None) or [])
|
||||
calls = [call for call in raw_calls if isinstance(call, Mapping)]
|
||||
if not calls or len(calls) != len(raw_calls):
|
||||
return None
|
||||
|
||||
end = start + 1
|
||||
results: list[Any] = []
|
||||
while end < len(messages) and getattr(messages[end], "type", None) == "tool":
|
||||
results.append(messages[end])
|
||||
end += 1
|
||||
if not results:
|
||||
return None
|
||||
results_by_id = {
|
||||
str(getattr(result, "tool_call_id", "") or "").strip(): result
|
||||
for result in results
|
||||
if str(getattr(result, "tool_call_id", "") or "").strip()
|
||||
}
|
||||
|
||||
signature: list[tuple[str, str, str]] = []
|
||||
tool_names: set[str] = set()
|
||||
for index, call in enumerate(calls):
|
||||
name = str(call.get("name") or "").strip()
|
||||
call_id = str(call.get("id") or "").strip()
|
||||
if not name or not call_id:
|
||||
return None
|
||||
result = results_by_id.get(call_id)
|
||||
if result is None and index < len(results):
|
||||
candidate = results[index]
|
||||
if not str(getattr(candidate, "tool_call_id", "") or "").strip():
|
||||
result = candidate
|
||||
if result is None:
|
||||
return None
|
||||
result_code = _deterministic_result_code(result)
|
||||
if result_code is None:
|
||||
return _ToolRound(
|
||||
messages=(assistant, *results),
|
||||
signature=(),
|
||||
tool_names=frozenset(),
|
||||
deterministic_error=False,
|
||||
), end
|
||||
signature.append((name, _canonical_tool_args(call.get("args")), result_code))
|
||||
tool_names.add(name)
|
||||
|
||||
return (
|
||||
_ToolRound(
|
||||
messages=(assistant, *results),
|
||||
signature=tuple(signature),
|
||||
tool_names=frozenset(tool_names),
|
||||
deterministic_error=True,
|
||||
),
|
||||
end,
|
||||
)
|
||||
|
||||
|
||||
def collapse_repetitive_tool_rounds(
|
||||
messages: Sequence[Any],
|
||||
*,
|
||||
threshold: int = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
) -> RepetitiveToolHistoryRepair:
|
||||
"""Build a provider-only projection while preserving audit history.
|
||||
|
||||
Only the middle rounds of three-or-more identical deterministic error
|
||||
groups are omitted. The first and last observations remain, and callers
|
||||
must never persist this projection back to a checkpoint.
|
||||
"""
|
||||
if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0:
|
||||
raise ValueError("repetitive tool call threshold must be non-negative")
|
||||
original = list(messages)
|
||||
segments: list[Any | _ToolRound] = []
|
||||
index = 0
|
||||
while index < len(original):
|
||||
parsed = _parse_tool_round(original, index)
|
||||
if parsed is None:
|
||||
segments.append(original[index])
|
||||
index += 1
|
||||
continue
|
||||
tool_round, index = parsed
|
||||
segments.append(tool_round)
|
||||
|
||||
tail_repetitions = 0
|
||||
tail_consecutive_errors = 0
|
||||
if segments and isinstance(segments[-1], _ToolRound):
|
||||
tail = segments[-1]
|
||||
if tail.deterministic_error:
|
||||
cursor = len(segments) - 1
|
||||
while cursor >= 0 and isinstance(segments[cursor], _ToolRound):
|
||||
current = segments[cursor]
|
||||
if not current.deterministic_error:
|
||||
break
|
||||
tail_consecutive_errors += 1
|
||||
cursor -= 1
|
||||
cursor = len(segments) - 1
|
||||
while cursor >= 0 and isinstance(segments[cursor], _ToolRound):
|
||||
current = segments[cursor]
|
||||
if (
|
||||
not current.deterministic_error
|
||||
or current.signature != tail.signature
|
||||
):
|
||||
break
|
||||
tail_repetitions += 1
|
||||
cursor -= 1
|
||||
|
||||
projected: list[Any] = []
|
||||
removed_rounds = 0
|
||||
index = 0
|
||||
while index < len(segments):
|
||||
segment = segments[index]
|
||||
if not isinstance(segment, _ToolRound) or not segment.deterministic_error:
|
||||
if isinstance(segment, _ToolRound):
|
||||
projected.extend(segment.messages)
|
||||
else:
|
||||
projected.append(segment)
|
||||
index += 1
|
||||
continue
|
||||
end = index + 1
|
||||
while (
|
||||
end < len(segments)
|
||||
and isinstance(segments[end], _ToolRound)
|
||||
and segments[end].deterministic_error
|
||||
and segments[end].signature == segment.signature
|
||||
):
|
||||
end += 1
|
||||
group = segments[index:end]
|
||||
should_compact = threshold > 0 and len(group) >= threshold and len(group) > 2
|
||||
if should_compact:
|
||||
projected.extend(group[0].messages)
|
||||
projected.extend(group[-1].messages)
|
||||
removed_rounds += len(group) - 2
|
||||
else:
|
||||
for item in group:
|
||||
projected.extend(item.messages)
|
||||
index = end
|
||||
|
||||
blocked = (
|
||||
segments[-1].tool_names
|
||||
if tail_repetitions and isinstance(segments[-1], _ToolRound)
|
||||
else frozenset()
|
||||
)
|
||||
return RepetitiveToolHistoryRepair(
|
||||
messages=projected,
|
||||
blocked_tool_names=blocked,
|
||||
removed_rounds=removed_rounds,
|
||||
tail_repetitions=tail_repetitions,
|
||||
tail_consecutive_errors=tail_consecutive_errors,
|
||||
)
|
||||
|
||||
|
||||
class RepetitiveToolCallGuardMiddleware(AgentMiddleware):
|
||||
"""Stop deterministic loops before another model request is made."""
|
||||
|
||||
name = "repetitive_tool_call_guard"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
threshold: int = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||
max_consecutive_errors: int = DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
for name, value in {
|
||||
"threshold": threshold,
|
||||
"max_consecutive_errors": max_consecutive_errors,
|
||||
}.items():
|
||||
if not isinstance(value, int) or isinstance(value, bool) or value < 0:
|
||||
raise ValueError(f"{name} must be a non-negative integer")
|
||||
self.threshold = threshold
|
||||
self.max_consecutive_errors = max_consecutive_errors
|
||||
|
||||
def _prepare_request(self, request: ModelRequest) -> ModelRequest:
|
||||
repair = collapse_repetitive_tool_rounds(
|
||||
request.messages,
|
||||
threshold=self.threshold,
|
||||
)
|
||||
if self.threshold and repair.tail_repetitions >= self.threshold:
|
||||
raise AgentControlError(
|
||||
"MODEL_TOOL_LOOP_DETECTED",
|
||||
"A deterministic repeated tool-call loop was stopped.",
|
||||
status_code=422,
|
||||
retryable=False,
|
||||
)
|
||||
if (
|
||||
self.max_consecutive_errors
|
||||
and repair.tail_consecutive_errors >= self.max_consecutive_errors
|
||||
):
|
||||
raise AgentControlError(
|
||||
"MODEL_TOOL_ERROR_LIMIT",
|
||||
"Too many consecutive deterministic tool errors were stopped.",
|
||||
status_code=422,
|
||||
retryable=False,
|
||||
)
|
||||
if repair.removed_rounds:
|
||||
logger.info(
|
||||
"Compacted deterministic tool errors for provider projection: removed_rounds=%d",
|
||||
repair.removed_rounds,
|
||||
)
|
||||
return request.override(messages=repair.messages)
|
||||
return request
|
||||
|
||||
def wrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
return handler(self._prepare_request(request))
|
||||
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
return await handler(self._prepare_request(request))
|
||||
@@ -0,0 +1,361 @@
|
||||
"""Validate completed model tool calls before they can reach ToolNode."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
ExtendedModelResponse,
|
||||
ModelRequest,
|
||||
ModelResponse,
|
||||
)
|
||||
from langchain_core.messages import AIMessage
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
from ..llm.errors import ModelToolProtocolError, _provider_from_model
|
||||
|
||||
_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call"})
|
||||
_MAX_DIAGNOSTIC_KEYS = 16
|
||||
_MAX_DIAGNOSTIC_KEY_CHARS = 64
|
||||
|
||||
|
||||
def _tool_name(tool: BaseTool | Mapping[str, Any] | Any) -> str | None:
|
||||
if isinstance(tool, BaseTool):
|
||||
return tool.name.strip() or None
|
||||
if isinstance(tool, Mapping):
|
||||
value = tool.get("name")
|
||||
if not value and isinstance(tool.get("function"), Mapping):
|
||||
value = tool["function"].get("name")
|
||||
if isinstance(value, str) and value.strip():
|
||||
return value.strip()
|
||||
return None
|
||||
value = getattr(tool, "name", None)
|
||||
return value.strip() if isinstance(value, str) and value.strip() else None
|
||||
|
||||
|
||||
def _ai_messages(response: Any) -> list[AIMessage]:
|
||||
"""Extract final AI messages from every LangChain middleware response shape."""
|
||||
if isinstance(response, AIMessage):
|
||||
return [response]
|
||||
if isinstance(response, ExtendedModelResponse):
|
||||
response = response.model_response
|
||||
elif not isinstance(response, ModelResponse):
|
||||
nested = getattr(response, "model_response", None)
|
||||
if nested is not None:
|
||||
response = nested
|
||||
result = getattr(response, "result", None)
|
||||
if not isinstance(result, Sequence) or isinstance(result, str | bytes):
|
||||
return []
|
||||
return [message for message in result if isinstance(message, AIMessage)]
|
||||
|
||||
|
||||
def _block_identity(block: Mapping[str, Any]) -> tuple[str, str]:
|
||||
call_id = str(block.get("id") or block.get("call_id") or "").strip()
|
||||
name = block.get("name") or block.get("tool_name")
|
||||
function = block.get("function")
|
||||
if not name and isinstance(function, Mapping):
|
||||
name = function.get("name")
|
||||
return call_id, str(name or "").strip()
|
||||
|
||||
|
||||
def _value_digest(value: Any) -> str:
|
||||
try:
|
||||
encoded = json.dumps(
|
||||
value,
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
default=lambda item: f"<{type(item).__name__}>",
|
||||
).encode("utf-8")
|
||||
except (TypeError, ValueError):
|
||||
encoded = f"<{type(value).__name__}:unserializable>".encode()
|
||||
return "sha256:" + hashlib.sha256(encoded).hexdigest()[:16]
|
||||
|
||||
|
||||
def _argument_diagnostic(value: Any, *, present: bool) -> dict[str, Any]:
|
||||
if not present:
|
||||
return {"args_present": False, "args_type": "missing"}
|
||||
if isinstance(value, Mapping):
|
||||
keys = sorted(str(key)[:_MAX_DIAGNOSTIC_KEY_CHARS] for key in value)
|
||||
return {
|
||||
"args_present": True,
|
||||
"args_type": "object",
|
||||
"args_key_count": len(keys),
|
||||
"args_keys": keys[:_MAX_DIAGNOSTIC_KEYS],
|
||||
"args_keys_truncated": len(keys) > _MAX_DIAGNOSTIC_KEYS,
|
||||
"args_digest": _value_digest(value),
|
||||
}
|
||||
if isinstance(value, Sequence) and not isinstance(value, str | bytes):
|
||||
value_type = "array"
|
||||
elif isinstance(value, str):
|
||||
value_type = "string"
|
||||
elif value is None:
|
||||
value_type = "null"
|
||||
else:
|
||||
value_type = type(value).__name__
|
||||
return {
|
||||
"args_present": True,
|
||||
"args_type": value_type,
|
||||
"args_digest": _value_digest(value),
|
||||
}
|
||||
|
||||
|
||||
def _summarize_call(call: Any) -> dict[str, Any]:
|
||||
if not isinstance(call, Mapping):
|
||||
return {"call_type": type(call).__name__}
|
||||
function = call.get("function")
|
||||
function = function if isinstance(function, Mapping) else {}
|
||||
call_id = str(call.get("id") or call.get("call_id") or "").strip()
|
||||
name = call.get("name") or call.get("tool_name") or function.get("name")
|
||||
name = str(name or "").strip()
|
||||
if "args" in call:
|
||||
args = call.get("args")
|
||||
args_present = True
|
||||
elif "arguments" in call:
|
||||
args = call.get("arguments")
|
||||
args_present = True
|
||||
elif "arguments" in function:
|
||||
args = function.get("arguments")
|
||||
args_present = True
|
||||
else:
|
||||
args = None
|
||||
args_present = False
|
||||
summary = {
|
||||
"call_type": "object",
|
||||
"name": name or "<missing>",
|
||||
"id_present": bool(call_id),
|
||||
**_argument_diagnostic(args, present=args_present),
|
||||
}
|
||||
if call_id:
|
||||
summary["id_fingerprint"] = _value_digest(call_id)
|
||||
return summary
|
||||
|
||||
|
||||
def _raw_openai_call(message: AIMessage, call_index: int) -> Any | None:
|
||||
additional = getattr(message, "additional_kwargs", None)
|
||||
additional = additional if isinstance(additional, Mapping) else {}
|
||||
raw_calls = additional.get("tool_calls")
|
||||
if (
|
||||
isinstance(raw_calls, Sequence)
|
||||
and not isinstance(raw_calls, str | bytes)
|
||||
and call_index < len(raw_calls)
|
||||
):
|
||||
return raw_calls[call_index]
|
||||
return None
|
||||
|
||||
|
||||
def _call_diagnostic(
|
||||
message: AIMessage,
|
||||
call: Any,
|
||||
*,
|
||||
source: str,
|
||||
call_index: int,
|
||||
call_count: int,
|
||||
) -> dict[str, Any]:
|
||||
diagnostic = {
|
||||
"source": source,
|
||||
"call_index": call_index,
|
||||
"call_count": call_count,
|
||||
**_summarize_call(call),
|
||||
}
|
||||
raw_call = _raw_openai_call(message, call_index)
|
||||
diagnostic["raw_openai_call_available"] = raw_call is not None
|
||||
if raw_call is not None:
|
||||
diagnostic["raw_openai_call"] = _summarize_call(raw_call)
|
||||
return diagnostic
|
||||
|
||||
|
||||
def _route_metadata(request: ModelRequest) -> dict[str, Any]:
|
||||
model = request.model
|
||||
metadata = getattr(model, "metadata", None)
|
||||
metadata = metadata if isinstance(metadata, Mapping) else {}
|
||||
provider = metadata.get("route_provider") or _provider_from_model(model)
|
||||
model_id = metadata.get("route_model")
|
||||
if not model_id:
|
||||
model_id = (
|
||||
getattr(model, "model_name", None)
|
||||
or getattr(model, "model", None)
|
||||
or getattr(model, "model_id", None)
|
||||
)
|
||||
generation = metadata.get("route_config_generation")
|
||||
try:
|
||||
config_generation = int(generation) if generation is not None else None
|
||||
except (TypeError, ValueError):
|
||||
config_generation = None
|
||||
return {
|
||||
"provider": str(provider) if provider else None,
|
||||
"model": str(model_id) if model_id else None,
|
||||
"route_key": str(metadata.get("route_key"))
|
||||
if metadata.get("route_key")
|
||||
else None,
|
||||
"config_generation": config_generation,
|
||||
"api_mode": str(metadata.get("route_api_mode"))
|
||||
if metadata.get("route_api_mode")
|
||||
else None,
|
||||
"endpoint": str(metadata.get("route_endpoint"))
|
||||
if metadata.get("route_endpoint")
|
||||
else None,
|
||||
"tool_call_transport": str(metadata.get("route_tool_call_transport"))
|
||||
if metadata.get("route_tool_call_transport")
|
||||
else None,
|
||||
}
|
||||
|
||||
|
||||
def _raise_protocol_error(
|
||||
request: ModelRequest,
|
||||
reason: str,
|
||||
*,
|
||||
call_id: str | None = None,
|
||||
call_diagnostic: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
raise ModelToolProtocolError(
|
||||
reason,
|
||||
call_id=call_id or None,
|
||||
call_diagnostic=call_diagnostic,
|
||||
**_route_metadata(request),
|
||||
)
|
||||
|
||||
|
||||
def _validate_message(
|
||||
message: AIMessage,
|
||||
request: ModelRequest,
|
||||
allowed_names: frozenset[str],
|
||||
) -> None:
|
||||
invalid_calls = list(getattr(message, "invalid_tool_calls", None) or [])
|
||||
if invalid_calls:
|
||||
invalid = invalid_calls[0]
|
||||
call_id = str(invalid.get("id") or "") if isinstance(invalid, Mapping) else ""
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
"invalid_final_call",
|
||||
call_id=call_id,
|
||||
call_diagnostic=_call_diagnostic(
|
||||
message,
|
||||
invalid,
|
||||
source="invalid_tool_calls",
|
||||
call_index=0,
|
||||
call_count=len(invalid_calls),
|
||||
),
|
||||
)
|
||||
|
||||
parsed_by_id: dict[str, str] = {}
|
||||
parsed_calls = list(getattr(message, "tool_calls", None) or [])
|
||||
for call_index, raw_call in enumerate(parsed_calls):
|
||||
diagnostic = _call_diagnostic(
|
||||
message,
|
||||
raw_call,
|
||||
source="parsed_tool_calls",
|
||||
call_index=call_index,
|
||||
call_count=len(parsed_calls),
|
||||
)
|
||||
if not isinstance(raw_call, Mapping):
|
||||
_raise_protocol_error(
|
||||
request, "invalid_final_call", call_diagnostic=diagnostic
|
||||
)
|
||||
call_id = str(raw_call.get("id") or raw_call.get("call_id") or "").strip()
|
||||
name = str(raw_call.get("name") or "").strip()
|
||||
if not name:
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
"missing_name",
|
||||
call_id=call_id,
|
||||
call_diagnostic=diagnostic,
|
||||
)
|
||||
if name not in allowed_names:
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
"unknown_name",
|
||||
call_id=call_id,
|
||||
call_diagnostic=diagnostic,
|
||||
)
|
||||
if not call_id:
|
||||
_raise_protocol_error(request, "missing_id", call_diagnostic=diagnostic)
|
||||
if call_id in parsed_by_id:
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
"duplicate_id",
|
||||
call_id=call_id,
|
||||
call_diagnostic=diagnostic,
|
||||
)
|
||||
args = raw_call.get("args")
|
||||
if not isinstance(args, Mapping):
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
"invalid_args",
|
||||
call_id=call_id,
|
||||
call_diagnostic=diagnostic,
|
||||
)
|
||||
parsed_by_id[call_id] = name
|
||||
|
||||
content = getattr(message, "content", None)
|
||||
if not isinstance(content, list):
|
||||
return
|
||||
seen_block_ids: set[str] = set()
|
||||
tool_blocks = [
|
||||
block
|
||||
for block in content
|
||||
if isinstance(block, Mapping) and block.get("type") in _TOOL_BLOCK_TYPES
|
||||
]
|
||||
for block_index, block in enumerate(tool_blocks):
|
||||
diagnostic = _call_diagnostic(
|
||||
message,
|
||||
block,
|
||||
source="content_blocks",
|
||||
call_index=block_index,
|
||||
call_count=len(tool_blocks),
|
||||
)
|
||||
call_id, name = _block_identity(block)
|
||||
if not call_id:
|
||||
_raise_protocol_error(request, "missing_id", call_diagnostic=diagnostic)
|
||||
if call_id in seen_block_ids:
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
"duplicate_id",
|
||||
call_id=call_id,
|
||||
call_diagnostic=diagnostic,
|
||||
)
|
||||
seen_block_ids.add(call_id)
|
||||
parsed_name = parsed_by_id.get(call_id)
|
||||
if parsed_name is None or (name and name != parsed_name):
|
||||
_raise_protocol_error(
|
||||
request,
|
||||
"inconsistent_block",
|
||||
call_id=call_id,
|
||||
call_diagnostic=diagnostic,
|
||||
)
|
||||
|
||||
|
||||
class ToolProtocolGuardMiddleware(AgentMiddleware):
|
||||
"""Fail closed on malformed final tool calls using the actual request tools."""
|
||||
|
||||
name = "tool_protocol_guard"
|
||||
|
||||
@staticmethod
|
||||
def _validate(response: Any, request: ModelRequest) -> None:
|
||||
allowed_names = frozenset(
|
||||
name for tool in request.tools if (name := _tool_name(tool)) is not None
|
||||
)
|
||||
for message in _ai_messages(response):
|
||||
_validate_message(message, request, allowed_names)
|
||||
|
||||
def wrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], ModelResponse],
|
||||
) -> ModelResponse:
|
||||
response = handler(request)
|
||||
self._validate(response, request)
|
||||
return response
|
||||
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
request: ModelRequest,
|
||||
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
|
||||
) -> ModelResponse:
|
||||
response = await handler(request)
|
||||
self._validate(response, request)
|
||||
return response
|
||||
@@ -48,6 +48,7 @@ DEFAULT_ALWAYS_INCLUDE_TOOLS: frozenset[str] = frozenset(
|
||||
"read_memory",
|
||||
"record_observation",
|
||||
"search_observations",
|
||||
"write_todos",
|
||||
}
|
||||
)
|
||||
|
||||
@@ -276,6 +277,15 @@ def create_tool_selector_middleware(
|
||||
|
||||
model = _ensure_chat_model()
|
||||
safe_model = disable_thinking(model)
|
||||
safe_model = safe_model.model_copy(
|
||||
update={
|
||||
"tags": [*(safe_model.tags or []), "metering:tool_selector"],
|
||||
"metadata": {
|
||||
**(safe_model.metadata or {}),
|
||||
"metering_scope": "tool_selector",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
system_prompt = (
|
||||
"You are selecting tools for a scientific research agent. "
|
||||
|
||||
@@ -7,6 +7,15 @@ All events contain a type and associated data dict.
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
STREAM_PROTOCOL_CAPABILITIES = frozenset(
|
||||
{
|
||||
"task_snapshot_v1",
|
||||
"complete_tool_call_v1",
|
||||
"correlated_tool_call_id_v1",
|
||||
"final_invalid_tool_call_v1",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamEvent:
|
||||
@@ -158,6 +167,14 @@ class StreamEventEmitter:
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def task_snapshot(source: str, items: list[dict[str, Any]]) -> StreamEvent:
|
||||
"""Emit the complete root-agent task state without product-specific IDs."""
|
||||
return StreamEvent(
|
||||
"task_snapshot",
|
||||
{"type": "task_snapshot", "source": source, "items": items},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def interrupt(
|
||||
interrupt_id: str,
|
||||
@@ -213,6 +230,19 @@ class StreamEventEmitter:
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def error(message: str) -> StreamEvent:
|
||||
def error(
|
||||
message: str,
|
||||
*,
|
||||
code: str | None = None,
|
||||
recoverable: bool | None = None,
|
||||
details: dict[str, Any] | None = None,
|
||||
) -> StreamEvent:
|
||||
"""Error event."""
|
||||
return StreamEvent("error", {"type": "error", "message": message})
|
||||
data: dict[str, Any] = {"type": "error", "message": message}
|
||||
if code is not None:
|
||||
data["code"] = code
|
||||
if recoverable is not None:
|
||||
data["recoverable"] = recoverable
|
||||
if details is not None:
|
||||
data["details"] = details
|
||||
return StreamEvent("error", data)
|
||||
|
||||
+276
-45
@@ -16,7 +16,7 @@ from typing import Any, TypeAlias
|
||||
from langchain_core._api import LangChainBetaWarning
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
|
||||
from langgraph.graph import END
|
||||
from langgraph.types import Command, Interrupt
|
||||
from langgraph.types import Command, Interrupt, Overwrite
|
||||
|
||||
from ..memory.worker_activity import clear_completed_memory_activity_counts
|
||||
from .emitter import StreamEventEmitter
|
||||
@@ -95,9 +95,10 @@ async def _clear_interrupted_graph_state(
|
||||
no output and leaves the messages channel unchanged. From the user's side the
|
||||
conversation looks like it lost all history because the agent stops responding.
|
||||
|
||||
The fix: ``aupdate_state(config, None, as_node=END)`` clears all pending tasks
|
||||
and writes a checkpoint whose ``next`` is the empty tuple, without touching
|
||||
any channel values (message history is preserved).
|
||||
Recovery first removes malformed/incomplete tool protocol from the messages
|
||||
channel, then ``aupdate_state(config, None, as_node=END)`` clears pending
|
||||
tasks and writes a checkpoint whose ``next`` is the empty tuple. Completed
|
||||
tool call/result pairs and all non-tool history are preserved.
|
||||
|
||||
Critically, this only runs when the stuck state is *not* a legitimate
|
||||
human-in-the-loop interrupt. The agent pauses via ``interrupt()`` /
|
||||
@@ -114,10 +115,9 @@ async def _clear_interrupted_graph_state(
|
||||
_log = logging.getLogger(__name__)
|
||||
try:
|
||||
snapshot = await agent.aget_state(config)
|
||||
# Only act when the graph is genuinely stuck (non-empty next tuple)...
|
||||
if not snapshot or not getattr(snapshot, "next", None):
|
||||
if not snapshot:
|
||||
return
|
||||
# ...and not parked at a real human-in-the-loop interrupt.
|
||||
# Never alter a real human-in-the-loop pause.
|
||||
if _snapshot_has_pending_interrupt(snapshot):
|
||||
_log.debug(
|
||||
"Leaving interrupted graph state intact for thread %s: "
|
||||
@@ -127,6 +127,13 @@ async def _clear_interrupted_graph_state(
|
||||
)
|
||||
return
|
||||
|
||||
await _repair_malformed_tool_history(agent, config, snapshot=snapshot)
|
||||
|
||||
# Only force END when the graph is genuinely stuck. Message repair also
|
||||
# applies to failures that already left next empty.
|
||||
if not getattr(snapshot, "next", None):
|
||||
return
|
||||
|
||||
stuck_at = snapshot.next
|
||||
await agent.aupdate_state(config, None, as_node=END)
|
||||
_log.debug(
|
||||
@@ -142,6 +149,49 @@ async def _clear_interrupted_graph_state(
|
||||
)
|
||||
|
||||
|
||||
async def _repair_malformed_tool_history(
|
||||
agent: Any,
|
||||
config: dict[str, Any],
|
||||
*,
|
||||
snapshot: Any | None = None,
|
||||
) -> bool:
|
||||
"""Rewrite a checkpoint's messages to a replay-safe tool history.
|
||||
|
||||
Only structurally invalid protocol is removed. Completed tool call/result
|
||||
pairs, including repeated successes and repeated errors, are audit and
|
||||
billing facts and must remain in persistent history.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from ..llm.patches import _sanitize_openai_tool_history
|
||||
|
||||
_log = logging.getLogger(__name__)
|
||||
if snapshot is None:
|
||||
snapshot = await agent.aget_state(config)
|
||||
if not snapshot or _snapshot_has_pending_interrupt(snapshot):
|
||||
return False
|
||||
values = getattr(snapshot, "values", None)
|
||||
if not isinstance(values, Mapping):
|
||||
return False
|
||||
messages = values.get("messages")
|
||||
if not isinstance(messages, list):
|
||||
return False
|
||||
|
||||
repaired = _sanitize_openai_tool_history(messages)
|
||||
if repaired == messages:
|
||||
return False
|
||||
|
||||
await agent.aupdate_state(config, {"messages": Overwrite(repaired)})
|
||||
_log.warning(
|
||||
"Repaired structurally invalid tool history for thread %s: messages %d -> %d",
|
||||
config.get("configurable", {}).get("thread_id", "?"),
|
||||
len(messages),
|
||||
len(repaired),
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _SubagentInfo:
|
||||
path: tuple[str, ...]
|
||||
@@ -217,7 +267,12 @@ class _V3EventProcessor:
|
||||
tuple[tuple[str, ...], str], tuple[str, dict[str, Any]]
|
||||
] = {}
|
||||
self._emitted_tool_calls: set[tuple[tuple[str, ...], str]] = set()
|
||||
self._pending_tool_calls: dict[
|
||||
tuple[tuple[str, ...], str], tuple[str, dict[str, Any]]
|
||||
] = {}
|
||||
self._emitted_interrupts: set[str] = set()
|
||||
self._pending_invalid_tool_calls: dict[str, tuple[str, str]] = {}
|
||||
self._last_task_snapshot: tuple[tuple[str, str], ...] | None = None
|
||||
self._selector = _ToolSelectionSuppressor(emitter)
|
||||
|
||||
@staticmethod
|
||||
@@ -245,13 +300,26 @@ class _V3EventProcessor:
|
||||
if method == "tools":
|
||||
return self._process_tool_event(namespace, _event_data(event), subagent)
|
||||
if method == "updates":
|
||||
return self._process_update_event(_event_data(event))
|
||||
return self._process_update_event(
|
||||
_event_data(event), namespace=namespace, source="update"
|
||||
)
|
||||
if method == "values":
|
||||
events: list[dict[str, Any]] = []
|
||||
params = event.get("params") or {}
|
||||
interrupts = params.get("interrupts") or ()
|
||||
if interrupts:
|
||||
events.extend(self._process_update_event({"__interrupt__": interrupts}))
|
||||
events.extend(
|
||||
self._process_update_event(
|
||||
{"__interrupt__": interrupts},
|
||||
namespace=namespace,
|
||||
source="values",
|
||||
)
|
||||
)
|
||||
events.extend(
|
||||
self._process_update_event(
|
||||
_event_data(event), namespace=namespace, source="values"
|
||||
)
|
||||
)
|
||||
if self._process_value_message_snapshots and not namespace:
|
||||
events.extend(self._process_value_messages(_event_data(event)))
|
||||
return events
|
||||
@@ -390,6 +458,7 @@ class _V3EventProcessor:
|
||||
inp, out = _usage_counts(usage) if usage is not None else (0, 0)
|
||||
if inp or out:
|
||||
events.append(self.emitter.usage_stats(inp, out).data)
|
||||
events.extend(self._flush_invalid_tool_calls())
|
||||
return events
|
||||
return []
|
||||
|
||||
@@ -415,14 +484,12 @@ class _V3EventProcessor:
|
||||
if tool_call is None:
|
||||
return events
|
||||
tool_name, args, tool_call_id = tool_call
|
||||
events.extend(
|
||||
self._emit_tool_call_once(
|
||||
namespace=namespace,
|
||||
subagent=subagent,
|
||||
name=tool_name,
|
||||
args=args,
|
||||
tool_call_id=tool_call_id,
|
||||
)
|
||||
self._pending_invalid_tool_calls.pop(tool_call_id, None)
|
||||
self._pending_tool_calls[
|
||||
(self._tool_scope(namespace, subagent), tool_call_id)
|
||||
] = (
|
||||
tool_name,
|
||||
args,
|
||||
)
|
||||
return events
|
||||
|
||||
@@ -465,6 +532,20 @@ class _V3EventProcessor:
|
||||
]
|
||||
return [self.emitter.tool_call(name, args, tool_call_id).data]
|
||||
|
||||
def _pending_call_id(
|
||||
self,
|
||||
*,
|
||||
scope: tuple[str, ...],
|
||||
name: str,
|
||||
args: dict[str, Any],
|
||||
) -> str:
|
||||
matches = [
|
||||
call_id
|
||||
for (candidate_scope, call_id), candidate in self._pending_tool_calls.items()
|
||||
if candidate_scope == scope and candidate == (name, args)
|
||||
]
|
||||
return matches[0] if len(matches) == 1 else ""
|
||||
|
||||
def _process_whole_message(
|
||||
self,
|
||||
msg: AIMessage | AIMessageChunk,
|
||||
@@ -472,6 +553,16 @@ class _V3EventProcessor:
|
||||
namespace: tuple[str, ...],
|
||||
) -> list[dict[str, Any]]:
|
||||
events: list[dict[str, Any]] = []
|
||||
for invalid in getattr(msg, "invalid_tool_calls", ()) or ():
|
||||
invalid_map = _as_raw_map(invalid)
|
||||
if invalid_map is None:
|
||||
continue
|
||||
call_id = str(
|
||||
invalid_map.get("id") or invalid_map.get("tool_call_id") or ""
|
||||
)
|
||||
name = str(invalid_map.get("name") or invalid_map.get("tool_name") or "")
|
||||
key = call_id or f"chunk_{len(self._pending_invalid_tool_calls)}"
|
||||
self._pending_invalid_tool_calls[key] = (call_id, name)
|
||||
additional = msg.additional_kwargs
|
||||
reasoning = additional.get("reasoning_content")
|
||||
emitted_reasoning = False
|
||||
@@ -494,14 +585,12 @@ class _V3EventProcessor:
|
||||
if tool_call is None:
|
||||
continue
|
||||
tool_name, args, tool_call_id = tool_call
|
||||
events.extend(
|
||||
self._emit_tool_call_once(
|
||||
namespace=namespace,
|
||||
subagent=subagent,
|
||||
name=tool_name,
|
||||
args=args,
|
||||
tool_call_id=tool_call_id,
|
||||
)
|
||||
self._pending_invalid_tool_calls.pop(tool_call_id, None)
|
||||
self._pending_tool_calls[
|
||||
(self._tool_scope(namespace, subagent), tool_call_id)
|
||||
] = (
|
||||
tool_name,
|
||||
args,
|
||||
)
|
||||
|
||||
if subagent is None:
|
||||
@@ -534,6 +623,9 @@ class _V3EventProcessor:
|
||||
name,
|
||||
args,
|
||||
)
|
||||
self._pending_tool_calls.pop(
|
||||
(self._tool_scope(namespace, subagent), tool_call_id), None
|
||||
)
|
||||
events.extend(
|
||||
self._emit_tool_call_once(
|
||||
namespace=namespace,
|
||||
@@ -580,6 +672,7 @@ class _V3EventProcessor:
|
||||
content += "\n... (truncated)"
|
||||
success = is_success(content)
|
||||
|
||||
lifecycle_key = (self._tool_scope(namespace, subagent), tool_call_id)
|
||||
if subagent is not None:
|
||||
events.append(
|
||||
self.emitter.subagent_tool_result(
|
||||
@@ -591,22 +684,78 @@ class _V3EventProcessor:
|
||||
instance_id=subagent.instance_id,
|
||||
).data
|
||||
)
|
||||
return events
|
||||
events.append(
|
||||
self.emitter.tool_result(
|
||||
name, content, success, tool_call_id=tool_call_id
|
||||
).data
|
||||
)
|
||||
else:
|
||||
events.append(
|
||||
self.emitter.tool_result(
|
||||
name, content, success, tool_call_id=tool_call_id
|
||||
).data
|
||||
)
|
||||
self._emitted_tool_calls.discard(lifecycle_key)
|
||||
self._pending_tool_calls.pop(lifecycle_key, None)
|
||||
return events
|
||||
|
||||
return []
|
||||
|
||||
def _process_update_event(self, data: object) -> list[dict[str, Any]]:
|
||||
@staticmethod
|
||||
def _normalize_task_items(value: object) -> list[dict[str, str]] | None:
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
aliases = {
|
||||
"todo": "pending",
|
||||
"pending": "pending",
|
||||
"active": "in_progress",
|
||||
"in-progress": "in_progress",
|
||||
"in_progress": "in_progress",
|
||||
"done": "completed",
|
||||
"completed": "completed",
|
||||
}
|
||||
items: list[dict[str, str]] = []
|
||||
for raw in value:
|
||||
raw_map = _as_raw_map(raw)
|
||||
if raw_map is None:
|
||||
continue
|
||||
content = str(raw_map.get("content") or raw_map.get("task") or "").strip()
|
||||
if not content:
|
||||
continue
|
||||
status = aliases.get(str(raw_map.get("status") or "pending").lower())
|
||||
if status is None:
|
||||
continue
|
||||
items.append({"content": content, "status": status})
|
||||
return items
|
||||
|
||||
@classmethod
|
||||
def _find_task_items(cls, data: object) -> list[dict[str, str]] | None:
|
||||
data_map = _as_raw_map(data)
|
||||
if data_map is None:
|
||||
return None
|
||||
if "todos" in data_map:
|
||||
return cls._normalize_task_items(data_map["todos"])
|
||||
for value in data_map.values():
|
||||
nested = _as_raw_map(value)
|
||||
if nested is not None and "todos" in nested:
|
||||
return cls._normalize_task_items(nested["todos"])
|
||||
return None
|
||||
|
||||
def _process_update_event(
|
||||
self,
|
||||
data: object,
|
||||
*,
|
||||
namespace: tuple[str, ...] = (),
|
||||
source: str = "update",
|
||||
) -> list[dict[str, Any]]:
|
||||
events: list[dict[str, Any]] = []
|
||||
data_map = _as_raw_map(data)
|
||||
if data_map is not None and "__interrupt__" in data_map:
|
||||
events.extend(self._process_interrupts(data_map["__interrupt__"]))
|
||||
|
||||
if not namespace:
|
||||
items = self._find_task_items(data)
|
||||
if items is not None:
|
||||
signature = tuple((item["content"], item["status"]) for item in items)
|
||||
if signature != self._last_task_snapshot:
|
||||
self._last_task_snapshot = signature
|
||||
events.append(self.emitter.task_snapshot(source, items).data)
|
||||
|
||||
summarization_event = _find_summarization_event_payload(data)
|
||||
if summarization_event and not self._summarization_in_progress:
|
||||
signature = _summarization_event_signature(summarization_event)
|
||||
@@ -622,6 +771,10 @@ class _V3EventProcessor:
|
||||
events.extend(self._emit_summarization_text(summary_text))
|
||||
return events
|
||||
|
||||
def _flush_invalid_tool_calls(self) -> list[dict[str, Any]]:
|
||||
self._pending_invalid_tool_calls.clear()
|
||||
return []
|
||||
|
||||
def _process_interrupts(self, interrupts: object) -> list[dict[str, Any]]:
|
||||
events: list[dict[str, Any]] = []
|
||||
if not isinstance(interrupts, list | tuple):
|
||||
@@ -660,26 +813,74 @@ class _V3EventProcessor:
|
||||
raw_questions = interrupt_map.get("questions")
|
||||
questions = raw_questions if isinstance(raw_questions, list) else []
|
||||
tc_id = str(interrupt_map.get("tool_call_id", ""))
|
||||
return self._dedupe_interrupt_event(
|
||||
self.emitter.ask_user_interrupt(
|
||||
interrupt_id,
|
||||
questions,
|
||||
tc_id,
|
||||
).data
|
||||
events: list[dict[str, Any]] = []
|
||||
candidate = self._pending_tool_calls.get(((), tc_id)) if tc_id else None
|
||||
if candidate is not None:
|
||||
events.extend(
|
||||
self._emit_tool_call_once(
|
||||
namespace=(),
|
||||
subagent=None,
|
||||
name=candidate[0],
|
||||
args=candidate[1],
|
||||
tool_call_id=tc_id,
|
||||
)
|
||||
)
|
||||
events.extend(
|
||||
self._dedupe_interrupt_event(
|
||||
self.emitter.ask_user_interrupt(
|
||||
interrupt_id,
|
||||
questions,
|
||||
tc_id,
|
||||
).data
|
||||
)
|
||||
)
|
||||
return events
|
||||
|
||||
raw_action_reqs = interrupt_map.get("action_requests")
|
||||
action_reqs = raw_action_reqs if isinstance(raw_action_reqs, list) else []
|
||||
raw_review_cfgs = interrupt_map.get("review_configs")
|
||||
review_cfgs = raw_review_cfgs if isinstance(raw_review_cfgs, list) else None
|
||||
if action_reqs:
|
||||
return self._dedupe_interrupt_event(
|
||||
self.emitter.interrupt(
|
||||
interrupt_id,
|
||||
action_reqs,
|
||||
review_cfgs,
|
||||
).data
|
||||
events: list[dict[str, Any]] = []
|
||||
for raw_request in action_reqs:
|
||||
request_map = _as_raw_map(raw_request)
|
||||
if request_map is None:
|
||||
continue
|
||||
call_id = str(
|
||||
request_map.get("id") or request_map.get("tool_call_id") or ""
|
||||
)
|
||||
name = str(request_map.get("name") or request_map.get("tool_name") or "")
|
||||
args_map = _as_raw_map(
|
||||
request_map.get("args")
|
||||
if "args" in request_map
|
||||
else request_map.get("input")
|
||||
)
|
||||
if not call_id and name and args_map is not None:
|
||||
call_id = self._pending_call_id(
|
||||
scope=(),
|
||||
name=name,
|
||||
args=dict(args_map),
|
||||
)
|
||||
if call_id and name and args_map is not None:
|
||||
events.extend(
|
||||
self._emit_tool_call_once(
|
||||
namespace=(),
|
||||
subagent=None,
|
||||
name=name,
|
||||
args=dict(args_map),
|
||||
tool_call_id=call_id,
|
||||
)
|
||||
)
|
||||
events.extend(
|
||||
self._dedupe_interrupt_event(
|
||||
self.emitter.interrupt(
|
||||
interrupt_id,
|
||||
action_reqs,
|
||||
review_cfgs,
|
||||
).data
|
||||
)
|
||||
)
|
||||
return events
|
||||
return []
|
||||
|
||||
def _dedupe_interrupt_event(self, event: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
@@ -805,6 +1006,8 @@ async def stream_agent_events(
|
||||
thread_id: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
media: list[str] | None = None,
|
||||
callbacks: list[Any] | None = None,
|
||||
error_mode: str = "emit",
|
||||
) -> AsyncGenerator[dict[str, Any], None]:
|
||||
"""Stream events from a DeepAgents/LangGraph v3 run.
|
||||
|
||||
@@ -820,6 +1023,9 @@ async def stream_agent_events(
|
||||
metadata: Optional metadata dict merged into the LangGraph config
|
||||
(e.g. agent_name, updated_at for checkpoint persistence).
|
||||
media: Optional list of local file paths for attachments.
|
||||
callbacks: Optional Runnable callbacks propagated to all nested model calls.
|
||||
error_mode: ``emit`` preserves the generic error event; ``raise`` lets an
|
||||
embedding host produce the single terminal error envelope.
|
||||
|
||||
Yields:
|
||||
Event dicts: thinking, text, tool_call, tool_result,
|
||||
@@ -829,6 +1035,8 @@ async def stream_agent_events(
|
||||
config: dict[str, Any] = {"configurable": {"thread_id": thread_id}}
|
||||
if metadata:
|
||||
config["metadata"] = metadata
|
||||
if callbacks:
|
||||
config["callbacks"] = callbacks
|
||||
emitter = StreamEventEmitter()
|
||||
existing_summarization_event: Mapping[str, object] | None = None
|
||||
try:
|
||||
@@ -957,7 +1165,30 @@ async def stream_agent_events(
|
||||
yield item
|
||||
except Exception as e:
|
||||
_run_raised = True
|
||||
yield emitter.error(str(e)).data
|
||||
if error_mode == "emit":
|
||||
payload = e.model_dump() if hasattr(e, "model_dump") else {}
|
||||
if not isinstance(payload, Mapping):
|
||||
payload = {}
|
||||
code = str(payload.get("code") or "") or None
|
||||
details = {
|
||||
key: payload[key]
|
||||
for key in (
|
||||
"reason",
|
||||
"provider",
|
||||
"model",
|
||||
"route_key",
|
||||
"config_generation",
|
||||
"api_mode",
|
||||
"call_id",
|
||||
)
|
||||
if payload.get(key) is not None
|
||||
}
|
||||
yield emitter.error(
|
||||
str(payload.get("message") or e),
|
||||
code=code,
|
||||
recoverable=bool(payload.get("recoverable", True)) if code else None,
|
||||
details=details or None,
|
||||
).data
|
||||
raise
|
||||
finally:
|
||||
if stream is not None:
|
||||
|
||||
@@ -69,3 +69,69 @@ def test_create_cli_agent_accepts_host_backend_and_memory_options(
|
||||
assert calls["middleware_kwargs"]["memory_max_inline_profile_chars"] == 1000
|
||||
assert calls["middleware_kwargs"]["enable_background_execution"] is False
|
||||
assert calls["agent_config"] == {"recursion_limit": 321}
|
||||
|
||||
|
||||
def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, tmp_path):
|
||||
import EvoScientist.EvoScientist as agent_module
|
||||
from EvoScientist.config.settings import EvoScientistConfig
|
||||
|
||||
calls = {}
|
||||
|
||||
class _Middleware:
|
||||
def __init__(self, name):
|
||||
self.name = name
|
||||
|
||||
class _Backend:
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
class _CompositeBackend:
|
||||
def __init__(self, **_kwargs):
|
||||
pass
|
||||
|
||||
class _Agent:
|
||||
def with_config(self, _config):
|
||||
return self
|
||||
|
||||
default_chain = [
|
||||
_Middleware("error_normalization"),
|
||||
_Middleware("configurable_model"),
|
||||
_Middleware("context_editing"),
|
||||
_Middleware("tool_protocol_guard"),
|
||||
]
|
||||
route = _Middleware("gateway_route_fallback")
|
||||
|
||||
monkeypatch.setattr("deepagents.backends.CompositeBackend", _CompositeBackend)
|
||||
monkeypatch.setattr("deepagents.create_deep_agent", lambda **_kwargs: _Agent())
|
||||
monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _Backend)
|
||||
monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _Backend)
|
||||
monkeypatch.setattr(agent_module, "set_active_workspace", lambda _path: None)
|
||||
def fake_default_middleware(**kwargs):
|
||||
calls["middleware_kwargs"] = kwargs
|
||||
return list(default_chain)
|
||||
|
||||
monkeypatch.setattr(agent_module, "_get_default_middleware", fake_default_middleware)
|
||||
|
||||
def fake_load(_backend, middleware, **_kwargs):
|
||||
calls["middleware"] = middleware
|
||||
return {"subagents": []}
|
||||
|
||||
monkeypatch.setattr(agent_module, "load_mcp_and_build_kwargs", fake_load)
|
||||
|
||||
agent_module.create_cli_agent(
|
||||
workspace_dir=str(tmp_path),
|
||||
checkpointer=object(),
|
||||
config=EvoScientistConfig(auto_approve=True),
|
||||
chat_model=object(),
|
||||
workspace_backend=object(),
|
||||
main_agent_route_middleware=route,
|
||||
)
|
||||
|
||||
assert calls["middleware_kwargs"]["enable_legacy_model_fallback"] is False
|
||||
assert [middleware.name for middleware in calls["middleware"][:5]] == [
|
||||
"error_normalization",
|
||||
"configurable_model",
|
||||
"gateway_route_fallback",
|
||||
"context_editing",
|
||||
"tool_protocol_guard",
|
||||
]
|
||||
|
||||
@@ -153,6 +153,8 @@ class TestEvoScientistConfig:
|
||||
assert config.channel_debug_tracing is False
|
||||
assert config.imessage_enabled is False
|
||||
assert config.imessage_allowed_senders == ""
|
||||
assert config.repetitive_tool_call_threshold == 2
|
||||
assert config.max_consecutive_tool_errors == 3
|
||||
|
||||
def test_auth_mode_default(self):
|
||||
"""Test that anthropic_auth_mode defaults to api_key."""
|
||||
@@ -200,6 +202,18 @@ class TestEvoScientistConfig:
|
||||
assert config.dangerous_mode is True
|
||||
assert config.auto_approve is True
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"kwargs",
|
||||
[
|
||||
{"repetitive_tool_call_threshold": -1},
|
||||
{"max_consecutive_tool_errors": -1},
|
||||
{"max_consecutive_tool_errors": True},
|
||||
],
|
||||
)
|
||||
def test_tool_guard_thresholds_must_be_non_negative_integers(self, kwargs):
|
||||
with pytest.raises(ValueError, match="non-negative integer"):
|
||||
EvoScientistConfig(**kwargs)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Test config path functions
|
||||
|
||||
@@ -15,7 +15,11 @@ from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from EvoScientist.llm.errors import ProviderStreamError
|
||||
from EvoScientist.llm.errors import (
|
||||
AgentControlError,
|
||||
ModelToolProtocolError,
|
||||
ProviderStreamError,
|
||||
)
|
||||
from EvoScientist.middleware.error_normalization import (
|
||||
ErrorNormalizationMiddleware,
|
||||
_normalize,
|
||||
@@ -148,6 +152,23 @@ class TestNormalize:
|
||||
)
|
||||
assert _normalize(req, pre_wrapped) is None
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error",
|
||||
[
|
||||
AgentControlError("MODEL_TOOL_LOOP_DETECTED", "loop stopped"),
|
||||
ModelToolProtocolError(
|
||||
"missing_name",
|
||||
provider="openai",
|
||||
model="gpt-example",
|
||||
route_key="route-1",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_platform_control_error_passes_through(self, error):
|
||||
req = _request(_openai_model())
|
||||
|
||||
assert _normalize(req, error) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _is_provider_error — used by tool selector to distinguish provider
|
||||
@@ -355,6 +376,25 @@ class TestMiddleware:
|
||||
self._run_awrap(mw, req, handler)
|
||||
assert excinfo.value is raised
|
||||
|
||||
def test_awrap_preserves_model_tool_protocol_error_identity(self):
|
||||
raised = ModelToolProtocolError(
|
||||
"missing_name",
|
||||
provider="openai",
|
||||
model="gpt-example",
|
||||
route_key="route-1",
|
||||
)
|
||||
|
||||
async def handler(_req):
|
||||
raise raised
|
||||
|
||||
req = _request(_openai_model())
|
||||
with pytest.raises(ModelToolProtocolError) as excinfo:
|
||||
self._run_awrap(ErrorNormalizationMiddleware(), req, handler)
|
||||
|
||||
assert excinfo.value is raised
|
||||
assert excinfo.value.code == "MODEL_TOOL_PROTOCOL_INVALID"
|
||||
assert excinfo.value.fallbackable is True
|
||||
|
||||
def test_awrap_wraps_any_exception_from_recognized_model(self):
|
||||
"""Any exception raised inside a call to a provider-recognized
|
||||
model gets wrapped — including builtins like ``RuntimeError``.
|
||||
|
||||
@@ -924,6 +924,12 @@ async def test_langgraph_server_gateway_emits_state_interrupt_before_done():
|
||||
events = await _collect()
|
||||
|
||||
assert events == [
|
||||
{
|
||||
"type": "tool_call",
|
||||
"name": "execute",
|
||||
"args": {"command": "echo hello"},
|
||||
"id": "tool-1",
|
||||
},
|
||||
{
|
||||
"type": "interrupt",
|
||||
"interrupt_id": "interrupt-1",
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
from EvoScientist.llm.errors import AgentControlError
|
||||
from EvoScientist.middleware.model_fallback import _is_non_fallbackable
|
||||
|
||||
|
||||
def test_agent_control_error_is_non_fallbackable():
|
||||
error = AgentControlError(
|
||||
"INSUFFICIENT_BALANCE",
|
||||
"balance unavailable",
|
||||
status_code=403,
|
||||
)
|
||||
|
||||
assert "platform control error" in (_is_non_fallbackable(error) or "")
|
||||
assert error.model_dump()["code"] == "INSUFFICIENT_BALANCE"
|
||||
@@ -1293,6 +1293,107 @@ class TestPatchOpenAICompatContent:
|
||||
assert len(set(call_ids)) == 2
|
||||
assert [message.tool_call_id for message in first[1:]] == call_ids
|
||||
|
||||
def test_content_tool_block_is_normalized_to_parsed_call(self):
|
||||
from langchain_core.messages import AIMessage, ToolMessage
|
||||
|
||||
from EvoScientist.llm.patches import _ensure_openai_tool_call_ids
|
||||
|
||||
normalized = _ensure_openai_tool_call_ids(
|
||||
[
|
||||
AIMessage(
|
||||
content=[
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "wrong-id",
|
||||
"name": "wrong-name",
|
||||
"args": {},
|
||||
}
|
||||
],
|
||||
tool_calls=[{"id": "call-1", "name": "execute", "args": {}}],
|
||||
),
|
||||
ToolMessage(content="ok", tool_call_id="call-1"),
|
||||
]
|
||||
)
|
||||
|
||||
assert normalized[0].content[0]["id"] == "call-1"
|
||||
assert normalized[0].content[0]["name"] == "execute"
|
||||
|
||||
def test_invalid_tool_call_is_not_replayed_to_responses_api(self):
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from langchain_openai.chat_models.base import _construct_responses_api_input
|
||||
|
||||
from EvoScientist.llm.patches import _sanitize_messages
|
||||
|
||||
invalid = AIMessage(
|
||||
content=[
|
||||
{"type": "reasoning", "reasoning": "partial"},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": None,
|
||||
"name": "execute",
|
||||
"args": '{"command":',
|
||||
},
|
||||
],
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"type": "invalid_tool_call",
|
||||
"id": None,
|
||||
"name": "execute",
|
||||
"args": '{"command":',
|
||||
"error": "Failed to parse tool call arguments as JSON",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
normalized = _sanitize_messages([invalid, HumanMessage(content="retry")])
|
||||
payload = _construct_responses_api_input(normalized)
|
||||
|
||||
assert all(item.get("type") != "function_call" for item in payload)
|
||||
assert [message.type for message in normalized] == ["human"]
|
||||
|
||||
def test_invalid_tool_call_preserves_replayable_assistant_text(self):
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
from EvoScientist.llm.patches import _sanitize_messages
|
||||
|
||||
invalid = AIMessage(
|
||||
content="I could not finish the tool request.",
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"type": "invalid_tool_call",
|
||||
"id": None,
|
||||
"name": "execute",
|
||||
"args": "{",
|
||||
"error": "bad json",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
normalized = _sanitize_messages([invalid])
|
||||
|
||||
assert len(normalized) == 1
|
||||
assert normalized[0].content == "I could not finish the tool request."
|
||||
assert normalized[0].invalid_tool_calls == []
|
||||
|
||||
def test_orphan_tool_results_and_unanswered_calls_are_removed(self):
|
||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||
|
||||
from EvoScientist.llm.patches import _sanitize_messages
|
||||
|
||||
messages = [
|
||||
ToolMessage(content="orphan", tool_call_id="missing"),
|
||||
AIMessage(
|
||||
content="waiting",
|
||||
tool_calls=[{"id": "call_unanswered", "name": "execute", "args": {}}],
|
||||
),
|
||||
HumanMessage(content="continue"),
|
||||
]
|
||||
|
||||
normalized = _sanitize_messages(messages)
|
||||
|
||||
assert [message.type for message in normalized] == ["ai", "human"]
|
||||
assert normalized[0].tool_calls == []
|
||||
|
||||
def test_generate_flattened(self):
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
|
||||
@@ -84,6 +84,7 @@ class TestIsNonFallbackable:
|
||||
"Error 400: invalid_request_error",
|
||||
"400 Bad Request: invalid request body",
|
||||
"400: malformed JSON in request",
|
||||
"<400> InvalidParameter: Repetitive tool calls detected in history",
|
||||
],
|
||||
)
|
||||
def test_malformed_request_400_patterns(self, msg):
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
"""Deterministic tool-loop guard and provider projection tests."""
|
||||
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from langchain.agents.middleware.types import ModelResponse
|
||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||
|
||||
from EvoScientist.llm.errors import AgentControlError
|
||||
from EvoScientist.middleware.repetitive_tool_guard import (
|
||||
RepetitiveToolCallGuardMiddleware,
|
||||
collapse_repetitive_tool_rounds,
|
||||
)
|
||||
|
||||
|
||||
def _round(
|
||||
call_id: str,
|
||||
*,
|
||||
name: str = "execute",
|
||||
command: str = "pwd",
|
||||
content: str = "Error: invalid argument: command rejected by schema",
|
||||
status: str = "error",
|
||||
) -> list[Any]:
|
||||
return [
|
||||
AIMessage(
|
||||
content="",
|
||||
tool_calls=[{"id": call_id, "name": name, "args": {"command": command}}],
|
||||
),
|
||||
ToolMessage(
|
||||
content=content,
|
||||
tool_call_id=call_id,
|
||||
name=name,
|
||||
status=status,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _Request:
|
||||
messages: list[Any]
|
||||
tools: list[Any]
|
||||
|
||||
def override(self, **updates: Any):
|
||||
return replace(self, **updates)
|
||||
|
||||
|
||||
def test_provider_projection_keeps_first_and_last_deterministic_error_rounds():
|
||||
messages = [HumanMessage(content="inspect")]
|
||||
for index in range(4):
|
||||
messages.extend(_round(f"call-{index}"))
|
||||
messages.append(HumanMessage(content="continue"))
|
||||
|
||||
repair = collapse_repetitive_tool_rounds(messages, threshold=2)
|
||||
|
||||
assert repair.removed_rounds == 2
|
||||
assert [m.type for m in repair.messages] == [
|
||||
"human",
|
||||
"ai",
|
||||
"tool",
|
||||
"ai",
|
||||
"tool",
|
||||
"human",
|
||||
]
|
||||
assert repair.messages[1].tool_calls[0]["id"] == "call-0"
|
||||
assert repair.messages[3].tool_calls[0]["id"] == "call-3"
|
||||
|
||||
|
||||
def test_successful_repeated_calls_are_never_projected_away():
|
||||
messages = [
|
||||
*_round("call-1", content="ok", status="success"),
|
||||
*_round("call-2", content="ok", status="success"),
|
||||
*_round("call-3", content="ok", status="success"),
|
||||
]
|
||||
repair = collapse_repetitive_tool_rounds(messages)
|
||||
assert repair.messages == messages
|
||||
assert repair.removed_rounds == 0
|
||||
assert repair.tail_repetitions == 0
|
||||
|
||||
|
||||
def test_transient_and_unknown_errors_do_not_count_as_semantic_loop():
|
||||
transient = [
|
||||
*_round("call-1", content="Error: connection timeout"),
|
||||
*_round("call-2", content="Error: connection timeout"),
|
||||
]
|
||||
unknown = [
|
||||
*_round("call-3", content="Error: something unusual"),
|
||||
*_round("call-4", content="Error: something unusual"),
|
||||
]
|
||||
assert collapse_repetitive_tool_rounds(transient).tail_repetitions == 0
|
||||
assert collapse_repetitive_tool_rounds(unknown).tail_consecutive_errors == 0
|
||||
|
||||
|
||||
def test_generic_raw_execution_error_code_remains_unknown():
|
||||
messages = _round("call-1", content="Error: something unusual")
|
||||
messages[1].additional_kwargs["error_code"] = "TOOL_EXECUTION_FAILED"
|
||||
|
||||
repair = collapse_repetitive_tool_rounds(messages)
|
||||
|
||||
assert repair.tail_consecutive_errors == 0
|
||||
|
||||
|
||||
def test_identical_tail_loop_stops_before_next_model_call():
|
||||
request = _Request(
|
||||
messages=[*_round("call-1"), *_round("call-2")],
|
||||
tools=[{"name": "execute"}],
|
||||
)
|
||||
called = False
|
||||
|
||||
def handler(_request):
|
||||
nonlocal called
|
||||
called = True
|
||||
return ModelResponse(result=[AIMessage(content="should not run")])
|
||||
|
||||
with pytest.raises(AgentControlError) as caught:
|
||||
RepetitiveToolCallGuardMiddleware(threshold=2).wrap_model_call(request, handler)
|
||||
|
||||
assert caught.value.code == "MODEL_TOOL_LOOP_DETECTED"
|
||||
assert called is False
|
||||
|
||||
|
||||
def test_different_deterministic_errors_hit_consecutive_limit():
|
||||
request = _Request(
|
||||
messages=[
|
||||
*_round("one", name="execute"),
|
||||
*_round("two", name="read_file"),
|
||||
*_round("three", name="search"),
|
||||
],
|
||||
tools=[],
|
||||
)
|
||||
|
||||
with pytest.raises(AgentControlError) as caught:
|
||||
RepetitiveToolCallGuardMiddleware(
|
||||
threshold=0, max_consecutive_errors=3
|
||||
).wrap_model_call(request, lambda _request: None)
|
||||
|
||||
assert caught.value.code == "MODEL_TOOL_ERROR_LIMIT"
|
||||
|
||||
|
||||
def test_user_message_breaks_tail_loop_but_historical_projection_is_temporary():
|
||||
original = [
|
||||
*_round("call-1"),
|
||||
*_round("call-2"),
|
||||
*_round("call-3"),
|
||||
HumanMessage(content="try a new approach"),
|
||||
]
|
||||
request = _Request(messages=original, tools=[])
|
||||
captured = []
|
||||
|
||||
def handler(prepared):
|
||||
captured.append(prepared)
|
||||
return ModelResponse(result=[AIMessage(content="continued")])
|
||||
|
||||
RepetitiveToolCallGuardMiddleware().wrap_model_call(request, handler)
|
||||
assert len(captured[0].messages) == 5
|
||||
assert len(original) == 7
|
||||
|
||||
|
||||
def test_zero_thresholds_disable_only_semantic_loop_guards():
|
||||
request = _Request(messages=[*_round("one"), *_round("two")], tools=[])
|
||||
captured = []
|
||||
middleware = RepetitiveToolCallGuardMiddleware(
|
||||
threshold=0, max_consecutive_errors=0
|
||||
)
|
||||
middleware.wrap_model_call(
|
||||
request,
|
||||
lambda prepared: (
|
||||
captured.append(prepared) or ModelResponse(result=[AIMessage(content="ok")])
|
||||
),
|
||||
)
|
||||
assert captured == [request]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kwargs", [{"threshold": -1}, {"max_consecutive_errors": -1}])
|
||||
def test_negative_threshold_is_rejected(kwargs):
|
||||
with pytest.raises(ValueError, match="non-negative"):
|
||||
RepetitiveToolCallGuardMiddleware(**kwargs)
|
||||
@@ -11,6 +11,7 @@ from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.types import Command, Interrupt
|
||||
|
||||
from EvoScientist.middleware.ask_user import AskUserMiddleware
|
||||
from EvoScientist.stream.emitter import STREAM_PROTOCOL_CAPABILITIES
|
||||
from EvoScientist.stream.events import stream_agent_events
|
||||
from EvoScientist.stream.summarization import (
|
||||
_extract_summary_message_text,
|
||||
@@ -1121,6 +1122,74 @@ class TestUsageStatsExtraction:
|
||||
assert len(usage_events) == 0
|
||||
|
||||
|
||||
class TestCanonicalSourceCapabilities:
|
||||
async def test_root_update_emits_full_task_snapshot_and_empty_clear(self):
|
||||
agent = FakeV3Agent(
|
||||
[
|
||||
protocol_event(
|
||||
"updates",
|
||||
{"model": {"todos": [{"content": "Inspect", "status": "active"}]}},
|
||||
),
|
||||
protocol_event("updates", {"model": {"todos": []}}),
|
||||
]
|
||||
)
|
||||
events = await collect_events(agent)
|
||||
snapshots = [event for event in events if event.get("type") == "task_snapshot"]
|
||||
assert snapshots == [
|
||||
{
|
||||
"type": "task_snapshot",
|
||||
"source": "update",
|
||||
"items": [{"content": "Inspect", "status": "in_progress"}],
|
||||
},
|
||||
{"type": "task_snapshot", "source": "update", "items": []},
|
||||
]
|
||||
|
||||
async def test_subagent_todos_do_not_replace_root_snapshot(self):
|
||||
agent = FakeV3Agent(
|
||||
[
|
||||
protocol_event(
|
||||
"updates",
|
||||
{"todos": [{"content": "Nested", "status": "pending"}]},
|
||||
namespace=("subagent",),
|
||||
)
|
||||
]
|
||||
)
|
||||
events = await collect_events(agent)
|
||||
assert not any(event.get("type") == "task_snapshot" for event in events)
|
||||
|
||||
async def test_invalid_tool_call_candidate_is_not_committed_by_stream_processor(self):
|
||||
invalid = AIMessage(
|
||||
content="",
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"name": "write_todos",
|
||||
"args": "{bad",
|
||||
"id": "call-invalid",
|
||||
"error": "invalid json",
|
||||
"type": "invalid_tool_call",
|
||||
}
|
||||
],
|
||||
)
|
||||
agent = FakeV3Agent(
|
||||
[
|
||||
protocol_event("messages", (invalid, {})),
|
||||
message_finish(),
|
||||
]
|
||||
)
|
||||
events = await collect_events(agent)
|
||||
assert not any(event.get("type") in {"tool_call", "error"} for event in events)
|
||||
|
||||
def test_stream_capabilities_are_explicit(self):
|
||||
assert STREAM_PROTOCOL_CAPABILITIES == frozenset(
|
||||
{
|
||||
"task_snapshot_v1",
|
||||
"complete_tool_call_v1",
|
||||
"correlated_tool_call_id_v1",
|
||||
"final_invalid_tool_call_v1",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class TestSummarizationHelpers:
|
||||
"""Summarization extraction helpers."""
|
||||
|
||||
|
||||
@@ -9,10 +9,12 @@ so they actually verify the two claims the recovery rests on:
|
||||
left intact, so a pending question is never silently discarded.
|
||||
"""
|
||||
|
||||
from typing import TypedDict
|
||||
from typing import Annotated, TypedDict
|
||||
|
||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from langgraph.graph.message import add_messages
|
||||
from langgraph.types import interrupt
|
||||
|
||||
from EvoScientist.stream.events import _clear_interrupted_graph_state
|
||||
@@ -22,6 +24,10 @@ class _S(TypedDict):
|
||||
x: int
|
||||
|
||||
|
||||
class _MessageState(TypedDict):
|
||||
messages: Annotated[list, add_messages]
|
||||
|
||||
|
||||
def _crashing_app():
|
||||
# Node 'b' crashes once, then succeeds — so a post-recovery run can complete
|
||||
# and prove the graph is genuinely unstuck (not replaying the dead step).
|
||||
@@ -57,6 +63,75 @@ def _interrupting_app():
|
||||
return g.compile(checkpointer=InMemorySaver())
|
||||
|
||||
|
||||
def _invalid_tool_call_app():
|
||||
def write_invalid_call(state):
|
||||
return {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
content="",
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"type": "invalid_tool_call",
|
||||
"id": None,
|
||||
"name": "execute",
|
||||
"args": '{"command":',
|
||||
"error": "bad json",
|
||||
}
|
||||
],
|
||||
)
|
||||
]
|
||||
}
|
||||
|
||||
def crash(state):
|
||||
raise RuntimeError("provider stream failed")
|
||||
|
||||
g = StateGraph(_MessageState)
|
||||
g.add_node("write_invalid_call", write_invalid_call)
|
||||
g.add_node("crash", crash)
|
||||
g.add_edge(START, "write_invalid_call")
|
||||
g.add_edge("write_invalid_call", "crash")
|
||||
g.add_edge("crash", END)
|
||||
return g.compile(checkpointer=InMemorySaver())
|
||||
|
||||
|
||||
def _repetitive_tool_call_app():
|
||||
messages = [HumanMessage(content="inspect")]
|
||||
for call_id in ("call-1", "call-2"):
|
||||
messages.extend(
|
||||
[
|
||||
AIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": call_id,
|
||||
"name": "execute",
|
||||
"args": {"command": "pwd"},
|
||||
}
|
||||
],
|
||||
),
|
||||
ToolMessage(
|
||||
content="/workspace",
|
||||
tool_call_id=call_id,
|
||||
name="execute",
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
def write_repetitive_history(state):
|
||||
return {"messages": messages}
|
||||
|
||||
def crash(state):
|
||||
raise RuntimeError("provider rejected repetitive tool history")
|
||||
|
||||
g = StateGraph(_MessageState)
|
||||
g.add_node("write_repetitive_history", write_repetitive_history)
|
||||
g.add_node("crash", crash)
|
||||
g.add_edge(START, "write_repetitive_history")
|
||||
g.add_edge("write_repetitive_history", "crash")
|
||||
g.add_edge("crash", END)
|
||||
return g.compile(checkpointer=InMemorySaver())
|
||||
|
||||
|
||||
async def test_recovery_clears_stuck_state_after_crash():
|
||||
app = _crashing_app()
|
||||
cfg = {"configurable": {"thread_id": "t1"}}
|
||||
@@ -91,3 +166,52 @@ async def test_recovery_preserves_pending_hitl_interrupt():
|
||||
after = app.get_state(cfg)
|
||||
assert after.next == ("ask",) # interrupt left intact, still resumable
|
||||
assert after.interrupts
|
||||
|
||||
|
||||
async def test_recovery_removes_invalid_tool_call_from_checkpoint():
|
||||
app = _invalid_tool_call_app()
|
||||
cfg = {"configurable": {"thread_id": "tool-history"}}
|
||||
try:
|
||||
await app.ainvoke(
|
||||
{"messages": [HumanMessage(content="run the command")]},
|
||||
cfg,
|
||||
)
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
before = await app.aget_state(cfg)
|
||||
assert before.next == ("crash",)
|
||||
assert any(
|
||||
isinstance(message, AIMessage) and message.invalid_tool_calls
|
||||
for message in before.values["messages"]
|
||||
)
|
||||
|
||||
await _clear_interrupted_graph_state(app, cfg)
|
||||
|
||||
after = await app.aget_state(cfg)
|
||||
assert after.next == ()
|
||||
assert [message.type for message in after.values["messages"]] == ["human"]
|
||||
|
||||
|
||||
async def test_recovery_preserves_complete_repetitive_tool_rounds_in_checkpoint():
|
||||
app = _repetitive_tool_call_app()
|
||||
cfg = {"configurable": {"thread_id": "tool-loop-history"}}
|
||||
try:
|
||||
await app.ainvoke({"messages": []}, cfg)
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
before = await app.aget_state(cfg)
|
||||
assert before.next == ("crash",)
|
||||
assert len(before.values["messages"]) == 5
|
||||
|
||||
await _clear_interrupted_graph_state(app, cfg)
|
||||
|
||||
after = await app.aget_state(cfg)
|
||||
messages = after.values["messages"]
|
||||
assert after.next == ()
|
||||
assert [message.type for message in messages] == ["human", "ai", "tool", "ai", "tool"]
|
||||
assert [messages[1].tool_calls[0]["id"], messages[3].tool_calls[0]["id"]] == [
|
||||
"call-1",
|
||||
"call-2",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,246 @@
|
||||
"""Final model tool protocol validation tests."""
|
||||
|
||||
from dataclasses import dataclass, field, replace
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from langchain.agents.middleware.types import ExtendedModelResponse, ModelResponse
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
from EvoScientist.llm.errors import ModelToolProtocolError
|
||||
from EvoScientist.middleware.tool_protocol_guard import ToolProtocolGuardMiddleware
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _Request:
|
||||
tools: list[Any]
|
||||
model: Any = field(default_factory=lambda: SimpleNamespace(metadata={}))
|
||||
|
||||
def override(self, **updates: Any):
|
||||
return replace(self, **updates)
|
||||
|
||||
|
||||
def _response(*calls: dict[str, Any], content: Any = "") -> ModelResponse:
|
||||
return ModelResponse(result=[AIMessage(content=content, tool_calls=list(calls))])
|
||||
|
||||
|
||||
def _call(call_id: str = "call-1", name: str = "search", args: Any = None):
|
||||
return {"id": call_id, "name": name, "args": {} if args is None else args}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("call", "reason"),
|
||||
[
|
||||
(_call(name=""), "missing_name"),
|
||||
(_call(name=" "), "missing_name"),
|
||||
(_call(name="missing"), "unknown_name"),
|
||||
(_call(call_id=""), "missing_id"),
|
||||
],
|
||||
)
|
||||
def test_invalid_final_tool_call_fails_closed(call, reason):
|
||||
middleware = ToolProtocolGuardMiddleware()
|
||||
request = _Request(tools=[{"name": "search"}])
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
middleware.wrap_model_call(request, lambda _request: _response(call))
|
||||
|
||||
assert caught.value.reason == reason
|
||||
assert caught.value.retryable is False
|
||||
assert caught.value.fallbackable is True
|
||||
|
||||
|
||||
def test_non_mapping_args_are_rejected_if_adapter_bypasses_message_validation():
|
||||
message = AIMessage(content="", tool_calls=[_call()])
|
||||
message.tool_calls[0]["args"] = "{}"
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
|
||||
assert caught.value.reason == "invalid_args"
|
||||
|
||||
|
||||
def test_duplicate_parallel_call_id_rejects_whole_response():
|
||||
request = _Request(tools=[{"name": "search"}, {"name": "read_file"}])
|
||||
response = _response(_call(name="search"), _call(name="read_file"))
|
||||
|
||||
with pytest.raises(
|
||||
ModelToolProtocolError, match="invalid structured tool call"
|
||||
) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
request, lambda _request: response
|
||||
)
|
||||
|
||||
assert caught.value.reason == "duplicate_id"
|
||||
|
||||
|
||||
def test_one_invalid_parallel_call_rejects_atomically():
|
||||
request = _Request(tools=[{"name": "search"}])
|
||||
response = _response(_call(call_id="one"), _call(call_id="two", name="missing"))
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
request, lambda _request: response
|
||||
)
|
||||
|
||||
assert caught.value.reason == "unknown_name"
|
||||
|
||||
|
||||
def test_final_invalid_tool_calls_are_rejected():
|
||||
message = AIMessage(
|
||||
content="",
|
||||
invalid_tool_calls=[
|
||||
{"id": "bad", "name": "search", "args": "{", "error": "bad json"}
|
||||
],
|
||||
)
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
|
||||
assert caught.value.reason == "invalid_final_call"
|
||||
assert caught.value.call_id == "bad"
|
||||
|
||||
|
||||
def test_content_block_must_match_parsed_call():
|
||||
response = _response(
|
||||
_call(),
|
||||
content=[
|
||||
{"type": "tool_call", "id": "call-1", "name": "read_file", "args": {}}
|
||||
],
|
||||
)
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}, {"name": "read_file"}]),
|
||||
lambda _request: response,
|
||||
)
|
||||
|
||||
assert caught.value.reason == "inconsistent_block"
|
||||
|
||||
|
||||
def test_parsed_only_valid_call_and_extended_response_pass():
|
||||
response = ExtendedModelResponse(model_response=_response(_call()))
|
||||
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"type": "function", "function": {"name": "search"}}]),
|
||||
lambda _request: response,
|
||||
)
|
||||
assert result is response
|
||||
|
||||
|
||||
async def test_async_direct_ai_message_shape_passes():
|
||||
response = AIMessage(content="", tool_calls=[_call()])
|
||||
|
||||
async def handler(_request):
|
||||
return response
|
||||
|
||||
result = await ToolProtocolGuardMiddleware().awrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]), handler
|
||||
)
|
||||
assert result is response
|
||||
|
||||
|
||||
def test_error_carries_safe_route_metadata():
|
||||
model = SimpleNamespace(
|
||||
metadata={
|
||||
"route_provider": "openai",
|
||||
"route_model": "gpt-example",
|
||||
"route_key": "route-safe",
|
||||
"route_config_generation": 12,
|
||||
"route_api_mode": "chat_completions",
|
||||
"route_endpoint": "primary",
|
||||
"route_tool_call_transport": "streaming",
|
||||
}
|
||||
)
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}], model=model),
|
||||
lambda _request: _response(_call(name="")),
|
||||
)
|
||||
|
||||
payload = caught.value.model_dump()
|
||||
assert payload["route_key"] == "route-safe"
|
||||
assert payload["config_generation"] == 12
|
||||
assert payload["endpoint"] == "primary"
|
||||
assert payload["tool_call_transport"] == "streaming"
|
||||
assert "args" not in payload
|
||||
|
||||
|
||||
def test_missing_id_carries_redacted_call_diagnostic_only_for_internal_logging():
|
||||
call = _call(call_id="", name="search", args={"query": "private search text"})
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: _response(call),
|
||||
)
|
||||
|
||||
diagnostic = caught.value.call_diagnostic
|
||||
assert diagnostic == {
|
||||
"source": "parsed_tool_calls",
|
||||
"call_index": 0,
|
||||
"call_count": 1,
|
||||
"call_type": "object",
|
||||
"name": "search",
|
||||
"id_present": False,
|
||||
"args_present": True,
|
||||
"args_type": "object",
|
||||
"args_key_count": 1,
|
||||
"args_keys": ["query"],
|
||||
"args_keys_truncated": False,
|
||||
"args_digest": diagnostic["args_digest"],
|
||||
"raw_openai_call_available": False,
|
||||
}
|
||||
assert diagnostic["args_digest"].startswith("sha256:")
|
||||
assert "private search text" not in str(diagnostic)
|
||||
assert "call_diagnostic" not in caught.value.model_dump()
|
||||
|
||||
|
||||
def test_diagnostic_compares_parsed_and_preserved_raw_openai_call_shapes():
|
||||
parsed = _call(call_id="", name="search", args={"query": "secret"})
|
||||
raw = {
|
||||
"id": "provider-call-id",
|
||||
"type": "function",
|
||||
"function": {"name": "search", "arguments": '{"query":"secret"}'},
|
||||
}
|
||||
message = AIMessage(
|
||||
content="",
|
||||
tool_calls=[parsed],
|
||||
additional_kwargs={"tool_calls": [raw]},
|
||||
)
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
|
||||
diagnostic = caught.value.call_diagnostic
|
||||
assert diagnostic["id_present"] is False
|
||||
assert diagnostic["raw_openai_call_available"] is True
|
||||
assert diagnostic["raw_openai_call"]["id_present"] is True
|
||||
assert diagnostic["raw_openai_call"]["name"] == "search"
|
||||
assert "provider-call-id" not in str(diagnostic)
|
||||
assert "secret" not in str(diagnostic)
|
||||
|
||||
|
||||
def test_diagnostic_failure_cannot_mask_the_protocol_error():
|
||||
circular: dict[str, Any] = {}
|
||||
circular["self"] = circular
|
||||
message = AIMessage(content="", tool_calls=[_call(call_id="", args={})])
|
||||
message.tool_calls[0]["args"] = circular
|
||||
|
||||
with pytest.raises(ModelToolProtocolError) as caught:
|
||||
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||
_Request(tools=[{"name": "search"}]),
|
||||
lambda _request: ModelResponse(result=[message]),
|
||||
)
|
||||
|
||||
assert caught.value.reason == "missing_id"
|
||||
assert caught.value.call_diagnostic["args_digest"].startswith("sha256:")
|
||||
Reference in New Issue
Block a user