diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index f0b33e8..98356eb 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -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 diff --git a/EvoScientist/config/settings.py b/EvoScientist/config/settings.py index 421ce8e..8188f2d 100644 --- a/EvoScientist/config/settings.py +++ b/EvoScientist/config/settings.py @@ -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", diff --git a/EvoScientist/llm/errors.py b/EvoScientist/llm/errors.py index 1140f12..d953d79 100644 --- a/EvoScientist/llm/errors.py +++ b/EvoScientist/llm/errors.py @@ -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. diff --git a/EvoScientist/llm/patches.py b/EvoScientist/llm/patches.py index 9c30044..ca3b46e 100644 --- a/EvoScientist/llm/patches.py +++ b/EvoScientist/llm/patches.py @@ -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 diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py index bb8afcb..d79e0ed 100644 --- a/EvoScientist/middleware/__init__.py +++ b/EvoScientist/middleware/__init__.py @@ -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", diff --git a/EvoScientist/middleware/error_normalization.py b/EvoScientist/middleware/error_normalization.py index 3ec816d..992445c 100644 --- a/EvoScientist/middleware/error_normalization.py +++ b/EvoScientist/middleware/error_normalization.py @@ -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): diff --git a/EvoScientist/middleware/model_fallback.py b/EvoScientist/middleware/model_fallback.py index 29a0de7..42c2222 100644 --- a/EvoScientist/middleware/model_fallback.py +++ b/EvoScientist/middleware/model_fallback.py @@ -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" diff --git a/EvoScientist/middleware/repetitive_tool_guard.py b/EvoScientist/middleware/repetitive_tool_guard.py new file mode 100644 index 0000000..5572167 --- /dev/null +++ b/EvoScientist/middleware/repetitive_tool_guard.py @@ -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)) diff --git a/EvoScientist/middleware/tool_protocol_guard.py b/EvoScientist/middleware/tool_protocol_guard.py new file mode 100644 index 0000000..eb149d4 --- /dev/null +++ b/EvoScientist/middleware/tool_protocol_guard.py @@ -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 "", + "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 diff --git a/EvoScientist/middleware/tool_selector.py b/EvoScientist/middleware/tool_selector.py index b498d0f..787b9d2 100644 --- a/EvoScientist/middleware/tool_selector.py +++ b/EvoScientist/middleware/tool_selector.py @@ -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. " diff --git a/EvoScientist/stream/emitter.py b/EvoScientist/stream/emitter.py index cfa268a..714d23c 100644 --- a/EvoScientist/stream/emitter.py +++ b/EvoScientist/stream/emitter.py @@ -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) diff --git a/EvoScientist/stream/events.py b/EvoScientist/stream/events.py index 309953b..2e63eef 100644 --- a/EvoScientist/stream/events.py +++ b/EvoScientist/stream/events.py @@ -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: diff --git a/tests/test_agent_factory_extensions.py b/tests/test_agent_factory_extensions.py index b0f7301..e05dce5 100644 --- a/tests/test_agent_factory_extensions.py +++ b/tests/test_agent_factory_extensions.py @@ -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", + ] diff --git a/tests/test_config.py b/tests/test_config.py index b25d0d9..07a040e 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -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 diff --git a/tests/test_error_normalization_middleware.py b/tests/test_error_normalization_middleware.py index 62649bf..efd1799 100644 --- a/tests/test_error_normalization_middleware.py +++ b/tests/test_error_normalization_middleware.py @@ -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``. diff --git a/tests/test_graph_gateway.py b/tests/test_graph_gateway.py index ebd2859..05daaf5 100644 --- a/tests/test_graph_gateway.py +++ b/tests/test_graph_gateway.py @@ -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", diff --git a/tests/test_host_metering_extensions.py b/tests/test_host_metering_extensions.py new file mode 100644 index 0000000..7123e04 --- /dev/null +++ b/tests/test_host_metering_extensions.py @@ -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" diff --git a/tests/test_llm.py b/tests/test_llm.py index 16d0c81..c4ec1a9 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -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 diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py index 30f1a03..36ea9d6 100644 --- a/tests/test_model_fallback.py +++ b/tests/test_model_fallback.py @@ -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): diff --git a/tests/test_repetitive_tool_guard.py b/tests/test_repetitive_tool_guard.py new file mode 100644 index 0000000..008037b --- /dev/null +++ b/tests/test_repetitive_tool_guard.py @@ -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) diff --git a/tests/test_stream_events.py b/tests/test_stream_events.py index c1d3491..dfb8f1b 100644 --- a/tests/test_stream_events.py +++ b/tests/test_stream_events.py @@ -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.""" diff --git a/tests/test_stream_recovery.py b/tests/test_stream_recovery.py index de57eda..0893859 100644 --- a/tests/test_stream_recovery.py +++ b/tests/test_stream_recovery.py @@ -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", + ] diff --git a/tests/test_tool_protocol_guard.py b/tests/test_tool_protocol_guard.py new file mode 100644 index 0000000..a82e5a5 --- /dev/null +++ b/tests/test_tool_protocol_guard.py @@ -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:")