diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index 326e1e1..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,8 +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, @@ -314,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() @@ -333,6 +348,16 @@ def _inject_subagent_middleware( memory_scheduler=memory_scheduler, ) middleware = [ + # Outermost — catches provider-SDK exceptions from the + # model call (including inner middlewares) and normalizes + # 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``). @@ -641,9 +666,14 @@ def _get_default_middleware( *, for_async_subagent: bool = False, workspace_dir: str | Path | None = None, + memory_dir: str | Path | None = None, cfg=None, chat_model=None, memory_source_agent: str = "EvoScientist", + 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. @@ -665,10 +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, @@ -681,10 +715,20 @@ 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() - memory_dir = str(_paths_mod.MEMORIES_DIR) + memory_dir = str(memory_dir or _paths_mod.MEMORIES_DIR) source_type = ( MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN ) @@ -699,18 +743,20 @@ def _get_default_middleware( # ``ModelFallbackMiddleware``: a configurable.model override sets the # PRIMARY model only, leaving the fallback chain free to try its own # alternatives instead of re-overriding every retry to the same model. - memory_middleware = create_memory_middleware( - memory_dir, - workspace_dir=workspace_dir, - source_type=source_type, - source_agent=memory_source_agent, - enable_profile_memory=memory_controls.profile_enabled, - enable_observation_memory=memory_controls.observations_enabled, - enable_observation_tool=memory_controls.observation_tool_enabled( + memory_kwargs = { + "workspace_dir": workspace_dir, + "source_type": source_type, + "source_agent": memory_source_agent, + "enable_profile_memory": memory_controls.profile_enabled, + "enable_observation_memory": memory_controls.observations_enabled, + "enable_observation_tool": memory_controls.observation_tool_enabled( MemoryObservationTarget.AGENT ), - memory_scheduler=memory_scheduler, - ) + "memory_scheduler": memory_scheduler, + } + if memory_max_inline_profile_chars is not None: + memory_kwargs["max_inline_profile_chars"] = memory_max_inline_profile_chars + memory_middleware = create_memory_middleware(memory_dir, **memory_kwargs) # Main-agent tool selection may use the auxiliary model; async sub-agents # keep the main model (they do real work, not a one-off helper call). # context_editing stays on the main model — its model only sizes the @@ -728,16 +774,32 @@ 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 + # middlewares) and normalizes them into a non-dataclass + # envelope wrapper before anything downstream sees them. + 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( - 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( @@ -770,7 +832,7 @@ def _get_default_middleware( # Background-process tools (run_in_background / check_process / stop_process / # list_processes) — main agent only. Async sub-agents run on langgraph-dev and # must not spawn local OS processes. - if not for_async_subagent: + if not for_async_subagent and enable_background_execution: from .middleware.background import BackgroundExecutionMiddleware mw.append(BackgroundExecutionMiddleware()) @@ -868,6 +930,14 @@ def create_cli_agent( chat_model=None, *, on_mcp_progress=None, + workspace_backend=None, + memory_dir: str | Path | None = None, + tool_selector_threshold: int | None = None, + 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. @@ -894,6 +964,22 @@ def create_cli_agent( chat_model: Optional pre-built chat model. Only triggers the pure path when ``config`` is also explicit; otherwise it is ignored in favor of the ``_ensure_chat_model()`` fallback. + workspace_backend: Optional host-provided backend for the workspace + route. The default remains ``CustomSandboxBackend``. + memory_dir: Optional memory root used by both the backend route and + memory middleware. + tool_selector_threshold: Optional adaptive tool-selection threshold. + memory_max_inline_profile_chars: Optional memory profile injection cap. + enable_subagents: Whether configured subagents are available to the 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 @@ -935,19 +1021,21 @@ def create_cli_agent( workspace_dir = str(_paths.WORKSPACE_ROOT) # Read paths dynamically so runtime set_workspace_root() changes are picked up - _mem_dir = str(_paths.MEMORIES_DIR) + _mem_dir = str(memory_dir or _paths.MEMORIES_DIR) _usr_skills_dir = str(_paths.USER_SKILLS_DIR) _global_skills_dir = str(_paths.GLOBAL_SKILLS_DIR) # Always construct fresh backends from current paths (avoids stale # module-level backend when workspace root changed at runtime). set_active_workspace(workspace_dir) - ws_backend = CustomSandboxBackend( - root_dir=workspace_dir, - virtual_mode=True, - timeout=cfg.sandbox_execute_timeout, - dangerous=cfg.dangerous_mode, - ) + ws_backend = workspace_backend + if ws_backend is None: + ws_backend = CustomSandboxBackend( + root_dir=workspace_dir, + virtual_mode=True, + timeout=cfg.sandbox_execute_timeout, + dangerous=cfg.dangerous_mode, + ) sk_backend = MergedSkillsBackend( primary_dir=_usr_skills_dir, global_dir=_global_skills_dir, @@ -969,8 +1057,29 @@ def create_cli_agent( # CLI agent never drifts from the default chain. Anything CLI-specific # (e.g. ``HumanInTheLoopMiddleware``) is appended below. mw: list[AgentMiddleware] = _get_default_middleware( - workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model + workspace_dir=workspace_dir, + memory_dir=_mem_dir, + cfg=cfg, + chat_model=chat_model, + 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 @@ -995,6 +1104,8 @@ def create_cli_agent( chat_model=chat_model, workspace_dir=workspace_dir, ) + if not enable_subagents: + kwargs = {**kwargs, "subagents": []} return create_deep_agent( **kwargs, diff --git a/EvoScientist/__init__.py b/EvoScientist/__init__.py index cabee0a..a58e6aa 100644 --- a/EvoScientist/__init__.py +++ b/EvoScientist/__init__.py @@ -9,6 +9,8 @@ from __future__ import annotations from importlib import import_module +__version__ = "0.2.2" + _EXPORTS: dict[str, tuple[str, str]] = { # Agent graph (lazy to avoid expensive initialization at import time) "EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"), diff --git a/EvoScientist/ccproxy_manager.py b/EvoScientist/ccproxy_manager.py index 1dfc4ec..865d407 100644 --- a/EvoScientist/ccproxy_manager.py +++ b/EvoScientist/ccproxy_manager.py @@ -20,6 +20,9 @@ from EvoScientist.config import EvoScientistConfig logger = logging.getLogger(__name__) +_CCPROXY_AUTH_TIMEOUT_SECONDS = 30 +_CCPROXY_HEALTH_TIMEOUT_SECONDS = 180 + # ============================================================================= # Availability & auth checks @@ -127,7 +130,11 @@ def check_ccproxy_auth(provider: str = "claude_api") -> tuple[bool, str]: [exe, "auth", "status", provider], capture_output=True, text=True, - timeout=10, + # ccproxy's CLI initializes its full plugin system on every + # invocation — a cold start takes ~10s on Apple Silicon, so a + # 10s timeout made OAuth startup fail intermittently with + # "Auth check timed out". + timeout=_CCPROXY_AUTH_TIMEOUT_SECONDS, ) import re as _re @@ -176,6 +183,33 @@ def is_ccproxy_running(port: int) -> bool: return False +def write_ccproxy_config() -> str: + """Write the ccproxy config file EvoScientist passes to ``serve --config``. + + Disables ccproxy's default Codex model mappings, which rewrite any + ``gpt-*``/``o1-*``/``o3-*``/``claude-*`` model to ``gpt-5.3-codex`` + before forwarding — silently overriding the model the user configured + (and failing outright on accounts where ``gpt-5.3-codex`` is not + served). With no mappings, the requested model reaches the Codex + backend unmodified. + + Returns: + Absolute path to the generated config file. + """ + from EvoScientist.config import get_config_dir + + path = get_config_dir() / "ccproxy.toml" + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text( + "# Generated by EvoScientist (ccproxy_manager) — do not edit;\n" + "# regenerated on every ccproxy start.\n" + "[plugins.codex]\n" + "model_mappings = []\n", + encoding="utf-8", + ) + return str(path) + + def start_ccproxy(port: int) -> subprocess.Popen: """Start ccproxy serve as a background process. @@ -186,18 +220,32 @@ def start_ccproxy(port: int) -> subprocess.Popen: The Popen handle for the ccproxy process. Raises: - RuntimeError: If ccproxy fails to become healthy within 30 seconds. + RuntimeError: If ccproxy fails to become healthy within + ``_CCPROXY_HEALTH_TIMEOUT_SECONDS``. FileNotFoundError: If ccproxy binary is not found. """ exe = _ccproxy_exe() or "ccproxy" + cmd = [exe, "serve", "--port", str(port)] + try: + cmd += ["--config", write_ccproxy_config()] + except (OSError, UnicodeError) as exc: + logger.warning( + "Could not write ccproxy config (%s); starting with defaults — " + "Codex model mappings will rewrite gpt-* models to gpt-5.3-codex", + exc, + ) + logger.warning( + "Starting ccproxy on port %d; first startup may take up to %d seconds", + port, + _CCPROXY_HEALTH_TIMEOUT_SECONDS, + ) proc = subprocess.Popen( - [exe, "serve", "--port", str(port)], + cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) - # Wait for health (ccproxy can take up to ~11s on first start) - deadline = time.monotonic() + 30 + deadline = time.monotonic() + _CCPROXY_HEALTH_TIMEOUT_SECONDS while time.monotonic() < deadline: if proc.poll() is not None: raise RuntimeError( @@ -213,7 +261,10 @@ def start_ccproxy(port: int) -> subprocess.Popen: proc.wait(timeout=3) except subprocess.TimeoutExpired: proc.kill() - raise RuntimeError("ccproxy did not become healthy within 30 seconds") + raise RuntimeError( + "ccproxy did not become healthy within " + f"{_CCPROXY_HEALTH_TIMEOUT_SECONDS} seconds" + ) def stop_ccproxy(proc: subprocess.Popen | None) -> None: diff --git a/EvoScientist/config/settings.py b/EvoScientist/config/settings.py index 0793a9c..8188f2d 100644 --- a/EvoScientist/config/settings.py +++ b/EvoScientist/config/settings.py @@ -106,11 +106,23 @@ def _normalize_hhmm(value: Any) -> str | None: def get_config_dir() -> Path: """Get the configuration directory path. - Uses XDG_CONFIG_HOME if set, otherwise ~/.config/evoscientist/ + Priority: + 1. EVOSCIENTIST_CONFIG_DIR + 2. EVOSCIENTIST_HOME/config + 3. XDG_CONFIG_HOME/evoscientist + 4. ~/.config/evoscientist """ + configured = os.environ.get("EVOSCIENTIST_CONFIG_DIR") + if configured: + return Path(configured).expanduser().resolve() + + home = os.environ.get("EVOSCIENTIST_HOME") + if home: + return Path(home).expanduser().resolve() / "config" + xdg_config = os.environ.get("XDG_CONFIG_HOME") if xdg_config: - return Path(xdg_config) / "evoscientist" + return Path(xdg_config).expanduser() / "evoscientist" return Path.home() / ".config" / "evoscientist" @@ -123,6 +135,16 @@ def get_config_path() -> Path: # Configuration dataclass # ============================================================================= +# OpenRouter app-attribution defaults (issue #339). Single source of truth: the +# EvoScientistConfig fields below default to these, and llm/models.py imports +# them for its env-fallback, so the values never drift across the two layers. +OPENROUTER_DEFAULT_HTTP_REFERER = "https://github.com/EvoScientist/EvoScientist" +OPENROUTER_DEFAULT_APP_TITLE = "EvoScientist" +# OpenRouter honors only the first 2 categories per request (server-side limit) +# and silently ignores the rest, so keep the two most relevant ones. Chosen per +# maintainer review — creative-writing is a less competitive marketplace group. +OPENROUTER_DEFAULT_APP_CATEGORIES = "creative-writing,personal-agent" + @dataclass class EvoScientistConfig: @@ -240,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 @@ -278,10 +307,21 @@ class EvoScientistConfig: # a deploy-style langgraph server instead of the in-terminal CLI/TUI. ui_backend: Literal["cli", "tui", "webui"] = "tui" log_level: str = "warning" - reasoning_effort: str = "high" + # Empty means use the provider/model default. A non-empty value is an + # explicit user override exported as EVOSCIENTIST_REASONING_EFFORT. + reasoning_effort: str = "" # Anthropic prompt caching for OpenRouter anthropic/* models. Opt out if # cache-write costs outweigh the benefit for a workflow. openrouter_anthropic_prompt_cache: bool = True + # OpenRouter app attribution (issue #339). Sent only for the openrouter + # provider; identifies EvoScientist in OpenRouter's app rankings/analytics. + # Override (e.g. a private fork) via these fields or their env vars. + # Defaults live in the module constants above (also imported by llm/models.py). + openrouter_http_referer: str = OPENROUTER_DEFAULT_HTTP_REFERER + openrouter_app_title: str = OPENROUTER_DEFAULT_APP_TITLE + # Comma-separated; split into a list before being passed to + # langchain-openrouter (its app_categories kwarg expects list[str]). + openrouter_app_categories: str = OPENROUTER_DEFAULT_APP_CATEGORIES # Channel Settings channel_enabled: str = "" # "imessage" | "telegram" | "discord" | "slack" | "wechat" | "dingtalk" | "feishu" | "email" | "qq" | "signal" | "" (comma-separated for multiple) @@ -440,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__ @@ -702,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: @@ -761,6 +814,9 @@ _ENV_MAPPINGS = { "openrouter_anthropic_prompt_cache": ( "EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE" ), + "openrouter_http_referer": "EVOSCIENTIST_OPENROUTER_HTTP_REFERER", + "openrouter_app_title": "EVOSCIENTIST_OPENROUTER_APP_TITLE", + "openrouter_app_categories": "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE", "channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING", "ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT", @@ -777,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", @@ -892,6 +952,22 @@ def apply_config_to_env(config: EvoScientistConfig) -> None: os.environ["TAVILY_API_KEY"] = config.tavily_api_key if config.reasoning_effort and not os.environ.get("EVOSCIENTIST_REASONING_EFFORT"): os.environ["EVOSCIENTIST_REASONING_EFFORT"] = config.reasoning_effort + if config.openrouter_http_referer and not os.environ.get( + "EVOSCIENTIST_OPENROUTER_HTTP_REFERER" + ): + os.environ["EVOSCIENTIST_OPENROUTER_HTTP_REFERER"] = ( + config.openrouter_http_referer + ) + if config.openrouter_app_title and not os.environ.get( + "EVOSCIENTIST_OPENROUTER_APP_TITLE" + ): + os.environ["EVOSCIENTIST_OPENROUTER_APP_TITLE"] = config.openrouter_app_title + if config.openrouter_app_categories and not os.environ.get( + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES" + ): + os.environ["EVOSCIENTIST_OPENROUTER_APP_CATEGORIES"] = ( + config.openrouter_app_categories + ) if not config.openrouter_anthropic_prompt_cache and not os.environ.get( "EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE" ): diff --git a/EvoScientist/langgraph_dev/http.py b/EvoScientist/langgraph_dev/http.py index 6b83bd7..0d9c6f6 100644 --- a/EvoScientist/langgraph_dev/http.py +++ b/EvoScientist/langgraph_dev/http.py @@ -43,7 +43,7 @@ async def get_models(_request: Request) -> JSONResponse: ``discover_ollama_models()`` call, same 1.5-s timeout, same fail-soft semantics (the probe returns ``[]`` on any error, never raises). The TUI's "Custom Ollama model…" sentinel is intentionally - omitted: that's a widget-specific input affordance, not part of + omitted — that's a widget-specific input affordance, not part of the registry surface. ``default`` reflects the deployment's currently-configured fallback diff --git a/EvoScientist/langgraph_dev/main_graph.py b/EvoScientist/langgraph_dev/main_graph.py index fc1bc41..7921621 100644 --- a/EvoScientist/langgraph_dev/main_graph.py +++ b/EvoScientist/langgraph_dev/main_graph.py @@ -5,8 +5,244 @@ in ``EvoScientist/EvoScientist.py`` so it doesn't construct on plain ``import EvoScientist``. ``langgraph dev`` 's symbol resolver inspects module attributes directly and doesn't trigger ``__getattr__``, so we re-export here to make it visible. + +Before re-export we upgrade the compiled graph's class in place to +``_EvoFilteredGraph``, which strips ``PrivateStateAttr``-marked fields +(currently just ``_quickjs_snapshot_payload``) from ``get_state`` / +``get_state_history`` responses. Upstream ``langchain_quickjs`` annotates +the field ``PrivateStateAttr = OmitFromSchema(input=True, output=True)``, +but LangGraph's ``_prepare_state_snapshot`` doesn't honor that on +checkpoint reads — every ``getState`` materializes the delta chain back +into a full ~1.4 MB blob, which the WebUI then downloads. The subclass +closes the gap without touching the middleware's write path, preserving +cross-turn REPL persistence as ``langchain-ai/deepagents#3064`` shipped it. """ -from EvoScientist.EvoScientist import EvoScientist_agent +from langgraph.graph.state import CompiledStateGraph +from langgraph.types import PregelTask, StateSnapshot + +from EvoScientist.EvoScientist import EvoScientist_agent as _agent + +_PRIVATE_STATE_FIELDS = frozenset({"_quickjs_snapshot_payload"}) + +# Sanity check on the LangGraph internals ``_strip_private`` scrubs. If any +# of these attributes disappear or get renamed in a future upstream bump, +# the assertion fires at import time — the deployment refuses to start, +# instead of silently degrading (the filter would ``.get()`` its way to a +# no-op and the private-field payload would come back on the wire without +# anyone noticing until a user reports slow thread switches again). +# +# Doesn't cover every internal we depend on — ``metadata["writes"]`` / +# ``metadata["counters_since_delta_snapshot"]`` dict keys aren't a canary +# target because ``dict.get`` already tolerates their absence. What we +# canary here is the ``NamedTuple`` field set: renames there would be the +# highest-impact silent regression. +_EXPECTED_SNAPSHOT_FIELDS = frozenset({"values", "metadata", "tasks"}) +_EXPECTED_TASK_FIELDS = frozenset({"result", "state"}) + +_missing_snap = _EXPECTED_SNAPSHOT_FIELDS - set(StateSnapshot._fields) +_missing_task = _EXPECTED_TASK_FIELDS - set(PregelTask._fields) +if _missing_snap or _missing_task: + raise RuntimeError( + "LangGraph state shape drifted from the version _strip_private was " + f"written against. Missing StateSnapshot fields: {_missing_snap or set()}. " + f"Missing PregelTask fields: {_missing_task or set()}. Review " + "_strip_private and re-verify against the current upstream shape " + "before removing this assertion." + ) + + +def _strip_private(snap): + """Strip ``PrivateStateAttr``-marked fields from a ``StateSnapshot``. + + Empirically verified against a live history response for a thread with + a single touched turn: the private field leaks on four surfaces — three + trivial, one heavy: + + * ``snap.values`` — the materialized channel state exposed as the main + payload. For DeltaChannels this is the delta chain replayed into full + bytes (~1.4 MB for the quickjs snapshot). ``get_state`` and every + history entry. + * ``snap.metadata['writes']`` — ``{node_name: {channel: value}}`` map of + the raw writes that produced each checkpoint. On the ``after_agent`` + step that first snapshots the REPL, ``value`` is the encoded write + record ``("snap", full_bytes)`` ≈ 1.4 MB. + * ``snap.tasks[*].result`` — the return dict of each completed + ``PregelTask``. ``after_agent`` returns + ``{"_quickjs_snapshot_payload": ("snap", bytes)}``; this dict becomes + the task's ``result`` field, which the API surfaces verbatim under + ``tasks[*].result`` (``langgraph_api.state:106``). This is the + dominant leak: 1.7 MB in the last history entry of any thread whose + most-recent-in-window checkpoint had a snapshot anchor. + * ``snap.metadata['counters_since_delta_snapshot']`` — DeltaChannel's + snapshot cadence bookkeeping ``{channel: [count, superstep]}``. Tiny + (~20 B) but exposes the private field name; strip for cleanliness. + * ``snap.tasks[*].state`` (nested ``StateSnapshot``) — populated when the + caller passes ``subgraphs=True``. Repeats all of the above surfaces + for each subgraph task, so recurse into it. Not exercised by the + current WebUI (which doesn't pass ``subgraphs=True`` on REST reads), + but SDK / curl / gRPC callers can. + """ + if snap is None: + return snap + values = {k: v for k, v in snap.values.items() if k not in _PRIVATE_STATE_FIELDS} + metadata = snap.metadata + if metadata: + new_metadata = metadata + if new_metadata.get("writes"): + scrubbed_writes = { + node: { + k: v for k, v in ch_writes.items() if k not in _PRIVATE_STATE_FIELDS + } + for node, ch_writes in new_metadata["writes"].items() + } + new_metadata = {**new_metadata, "writes": scrubbed_writes} + if new_metadata.get("counters_since_delta_snapshot"): + scrubbed_counters = { + k: v + for k, v in new_metadata["counters_since_delta_snapshot"].items() + if k not in _PRIVATE_STATE_FIELDS + } + new_metadata = { + **new_metadata, + "counters_since_delta_snapshot": scrubbed_counters, + } + metadata = new_metadata + tasks = snap.tasks + if tasks: + new_tasks = [] + changed = False + for t in tasks: + replace_kwargs: dict = {} + result = getattr(t, "result", None) + if isinstance(result, dict) and any( + k in result for k in _PRIVATE_STATE_FIELDS + ): + replace_kwargs["result"] = { + k: v for k, v in result.items() if k not in _PRIVATE_STATE_FIELDS + } + # ``t.state`` is a ``RunnableConfig | StateSnapshot | None`` per + # ``PregelTask``'s typing. When ``subgraphs=True`` on the caller, + # this holds the subgraph's fully-materialized ``StateSnapshot`` — + # which repeats the same four leak surfaces (``values``, + # ``metadata.writes``, ``metadata.counters_since_delta_snapshot``, + # ``tasks[*].result/state``). Recurse so the whole tree is clean. + nested_state = getattr(t, "state", None) + if isinstance(nested_state, StateSnapshot): + scrubbed_state = _strip_private(nested_state) + if scrubbed_state is not nested_state: + replace_kwargs["state"] = scrubbed_state + if replace_kwargs: + new_tasks.append(t._replace(**replace_kwargs)) + changed = True + else: + new_tasks.append(t) + if changed: + tasks = tuple(new_tasks) + return snap._replace(values=values, metadata=metadata, tasks=tasks) + + +class _EvoFilteredGraph(CompiledStateGraph): + """Filters ``PrivateStateAttr``-marked state fields from checkpoint reads. + + ``Pregel.copy`` uses ``self.__class__(**attrs)`` so this subclass + survives the ``graph_obj.copy(update=...)`` call in + ``langgraph_api.graph.get_graph`` that binds the checkpointer / store + before yielding to endpoint handlers. + + **Known gap — streaming paths.** The overrides only cover ``get_state`` + / ``get_state_history``. On this compiled graph, + ``self.output_channels`` correctly excludes ``_quickjs_snapshot_payload`` + (respects ``OmitFromSchema(output=True)``), but + ``self.stream_channels_asis`` includes it alongside other private + fields (``jump_to``, ``_summarization_event``) — the two lists are + built by ``langgraph.graph.state``'s graph builder and only the first + checks the output schema. So a client streaming with + ``stream_mode="values"`` or ``stream_mode="events"`` (which fall back + to ``stream_channels_asis`` when ``output_keys`` is ``None``) can pull + the anchor blob in per-run event data. Empirically the WebUI's + ``stream_mode=["updates"]`` path is clean, so this is transient per-run + rather than the persistent per-getState download this PR targets. + Filter here first; extend into the stream layer if a client relying on + ``values`` / ``events`` reports it. + """ + + async def aget_state(self, config, *, subgraphs=False): + return _strip_private(await super().aget_state(config, subgraphs=subgraphs)) + + def get_state(self, config, *, subgraphs=False): + return _strip_private(super().get_state(config, subgraphs=subgraphs)) + + async def aget_state_history(self, config, **kw): + async for snap in super().aget_state_history(config, **kw): + yield _strip_private(snap) + + def get_state_history(self, config, **kw): + for snap in super().get_state_history(config, **kw): + yield _strip_private(snap) + + +# In-place ``__class__`` swap: the subclass adds only methods (no new +# instance attributes) so the memory layout is identical and the swap is +# safe. Constructing a fresh ``_EvoFilteredGraph`` via ``.copy()`` would +# require reproducing the deep-agent build pipeline; the swap avoids that. +_agent.__class__ = _EvoFilteredGraph +EvoScientist_agent = _agent + + +def _apply_filter_to_all_registered_graphs() -> None: + """Extend the class swap to every graph registered in ``langgraph.json``. + + ``EvoScientist.py:_build_middleware_stack`` installs + ``create_code_interpreter_middleware`` unconditionally — it's not gated + on the ``for_async_subagent`` flag — so every subagent (sync ``task`` + dispatch and async ``start_async_task``) carries the QuickJS REPL and + can produce ``_quickjs_snapshot_payload`` writes on its own checkpoint + namespace. + + Async subagents get their own ``thread_id`` and their ``/threads/{id}/state`` + endpoint is served by their own compiled graph. Without swapping the + class on those graphs, the filter we applied to ``EvoScientist_agent`` + doesn't reach that endpoint and any real code_interpreter touch inside + a subagent leaks the anchor snapshot verbatim. + + Reads the graph registry straight from ``langgraph.json`` so a new + subagent added to the config picks up the swap automatically — no + hardcoded list to keep in sync. + + Idempotent (skips graphs already swapped) and safe on graphs that don't + use the middleware — ``_strip_private`` returns snapshots unchanged when + the private field is absent. Best-effort: if the config is unreadable + or an entry can't be resolved, the deployment still starts — only the + unresolvable subagents remain unfiltered. + """ + import json + from importlib import import_module + from pathlib import Path + + config_path = Path(__file__).parent / "langgraph.json" + try: + config = json.loads(config_path.read_text()) + except (OSError, json.JSONDecodeError): + return + + for path in config.get("graphs", {}).values(): + # Format: "module.dotted.path:attr_name" + if ":" not in path: + continue + module_path, attr = path.rsplit(":", 1) + try: + module = import_module(module_path) + except ImportError: + continue + graph = getattr(module, attr, None) + if isinstance(graph, CompiledStateGraph) and not isinstance( + graph, _EvoFilteredGraph + ): + graph.__class__ = _EvoFilteredGraph + + +_apply_filter_to_all_registered_graphs() + __all__ = ["EvoScientist_agent"] diff --git a/EvoScientist/llm/errors.py b/EvoScientist/llm/errors.py new file mode 100644 index 0000000..d953d79 --- /dev/null +++ b/EvoScientist/llm/errors.py @@ -0,0 +1,386 @@ +"""Provider-error surface for langgraph SSE frames. + +Provides :class:`ProviderStreamError` — a normalized, non-dataclass +exception raised by ``ErrorNormalizationMiddleware`` in place of the +provider SDK exception that a chat model call raised. Non-dataclass on +purpose: since orjson 3.0, dataclass instances are serialized natively +via their field enumeration, skipping the ``default=`` hook that +would otherwise build our SSE envelope. Some provider SDKs (openrouter +today) decorate their exceptions with ``@dataclass``, so their errors +emerge on the wire as raw dataclass fields — no envelope, no way for +the WebUI to distinguish quota / auth / rate-limit. Wrapping them in +a plain ``Exception`` subclass here keeps orjson on the ``default=`` +path, which then calls :meth:`ProviderStreamError.model_dump` +(upstream ``langgraph_api.serde.default`` checks that hook before its +``BaseException`` branch) — no serde monkey-patch needed. + +Also lives here: the pure-function helpers the middleware uses to +build the envelope (provider tag from ``ModelRequest.model``, SDK +field extractors, env-driven API-key redaction). They stay next to +:class:`ProviderStreamError` because the middleware is their only +consumer. +""" + +from __future__ import annotations + +import os +import re +from typing import Any + +# --------------------------------------------------------------------------- +# ProviderStreamError +# --------------------------------------------------------------------------- + + +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. + + Attributes mirror the SSE envelope one-for-one: + + - ``provider`` — concrete provider tag (``openai`` / ``anthropic`` + / ``deepseek`` / ``openrouter`` / ``openai_compat`` / …) + - ``class_qualname`` — fully qualified name of the underlying + exception's class (e.g. ``openrouter.errors.…``) + - ``message`` — API-key-redacted ``str(exc)`` + - ``status_code`` — HTTP status if the SDK exposed one + - ``code`` — provider error code (``insufficient_quota``, …) + - ``err_type`` — provider error type label (openai's ``.type``) + - ``request_id`` — SDK-provided correlation id + + The underlying exception is available via ``__cause__`` (set by + ``raise ProviderStreamError(...) from exc`` in the middleware). + """ + + def __init__( + self, + provider: str, + class_qualname: str, + message: str, + *, + status_code: int | None = None, + code: str | None = None, + err_type: str | None = None, + request_id: str | None = None, + ) -> None: + super().__init__(message) + self.provider = provider + self.class_qualname = class_qualname + self.message = message + self.status_code = status_code + self.code = code + self.err_type = err_type + self.request_id = request_id + + def as_envelope(self) -> dict[str, Any]: + """Return the SSE envelope dict — the shape the WebUI consumes.""" + payload: dict[str, Any] = { + "error": self.class_qualname.rsplit(".", 1)[-1], + "class": self.class_qualname, + "message": self.message, + "provider": self.provider, + } + if self.status_code is not None: + payload["status_code"] = self.status_code + if self.code is not None: + payload["code"] = self.code + if self.err_type is not None: + payload["type"] = self.err_type + if self.request_id: + payload["request_id"] = self.request_id + return payload + + def model_dump(self) -> dict[str, Any]: + """Serialization hook consumed by ``langgraph_api.serde.default``. + + Upstream's dispatch checks ``hasattr(obj, 'model_dump')`` BEFORE + the ``isinstance(obj, BaseException)`` branch, so exposing this + method lets upstream emit our envelope with no monkey-patch on + its ``default`` callable. The name matches Pydantic's + convention deliberately — it's the hook upstream is looking + for. + """ + return self.as_envelope() + + +# --------------------------------------------------------------------------- +# API-key redaction — env-driven, prefix-only +# --------------------------------------------------------------------------- +# +# Redaction is built from credentials actually deployed via env vars, +# not from generic key shapes. Rationale: (a) zero false positives — +# we only scrub strings we know are secrets, (b) defense-in-depth — +# the compiled regex holds only the first 8 chars of each key, so a +# leak of the regex object itself (traceback locals, process dump) +# can't expose the secret. Suffix-greedy match consumes the rest of +# the key shape at runtime. The table is rebuilt on every +# ``_redact_api_keys`` call so credentials loaded after import +# (typically ``load_dotenv`` in a main entry point) still get +# scrubbed. ``re.compile`` caches by source string internally, so an +# unchanged env costs a dict lookup. + +_API_KEY_ENV_SUFFIXES = ("_API_KEY", "_TOKEN", "_SECRET") +_API_KEY_MIN_LEN = 12 +_API_KEY_PREFIX_LEN = 8 + + +def _build_env_key_redaction_re() -> re.Pattern[str] | None: + prefixes: list[str] = [] + for k, v in os.environ.items(): + if not k.endswith(_API_KEY_ENV_SUFFIXES): + continue + if not isinstance(v, str) or len(v) < _API_KEY_MIN_LEN: + continue + prefixes.append(re.escape(v[:_API_KEY_PREFIX_LEN])) + if not prefixes: + return None + alternation = "|".join(f"{p}[A-Za-z0-9_+/=.-]*" for p in prefixes) + return re.compile(alternation) + + +def _redact_api_keys(message: str) -> str: + """Replace any deployed key prefix in *message* with ````. + + Defensive; provider error messages occasionally echo the + authorization header back. Rebuilt per call so credentials loaded + after import (typical ``load_dotenv`` pattern) are still redacted. + """ + pattern = _build_env_key_redaction_re() + if pattern is None: + return message + return pattern.sub("", message) + + +# --------------------------------------------------------------------------- +# Provider inference from ModelRequest.model +# --------------------------------------------------------------------------- +# +# Host → concrete provider. Hand-maintained snapshot mirroring the +# routed-provider tables in ``llm/models.py`` +# (``_OPENAI_ROUTED_PROVIDERS`` + ``_ANTHROPIC_ROUTED_PROVIDERS``). +# Kept here rather than imported from ``models.py`` to keep the +# import surface of ``errors.py`` minimal — importing ``models.py`` +# would pull in every langchain chat-model client at first +# middleware access. Consumed by ``_lookup_host_or_compat``; unknown +# hosts fall back to ``_compat`` so the WebUI knows +# "openai/anthropic SDK, but not native" instead of getting a +# misleading concrete tag. Update when a new routed provider is +# added to ``models.py``. +# +# Related sibling: ``_PROVIDER_EXC_MODULE_PREFIXES`` in +# ``middleware/error_normalization.py`` — the exception-side +# provider allow-list. Adding a whole new provider SDK (not just a +# new base_url routed through an existing one) means updating that +# list too. + +_HOST_TO_PROVIDER: dict[str, str] = { + "api.openai.com": "openai", + "api.anthropic.com": "anthropic", + "api.deepseek.com": "deepseek", + "api.moonshot.cn": "moonshot", + "api.siliconflow.cn": "siliconflow", + "open.bigmodel.cn": "zhipu", # zhipu + zhipu-code share this host + "ark.cn-beijing.volces.com": "volcengine", + "dashscope.aliyuncs.com": "dashscope", + "coding.dashscope.aliyuncs.com": "dashscope", + "api.minimaxi.com": "minimax", + "api.kimi.com": "kimi", # kimi-coding shares this host + "openrouter.ai": "openrouter", +} + + +def _provider_from_model(model: Any) -> str | None: + """Derive the concrete provider tag from a chat model instance. + + Class-based dispatch for unambiguous providers (``ChatOpenRouter``, + ``ChatGoogleGenerativeAI``); ``openai_api_base`` / + ``anthropic_api_url`` looked up in ``_HOST_TO_PROVIDER`` for + openai/anthropic-shape clients (native + routed). Returns ``None`` + when the model isn't from a recognized provider SDK — the caller + (``ErrorNormalizationMiddleware``) then passes the exception + through unchanged. + """ + cls_module = type(model).__module__ or "" + if cls_module.startswith("langchain_openrouter"): + return "openrouter" + if cls_module.startswith("langchain_google_genai"): + return "google_genai" + if cls_module.startswith("langchain_openai"): + return _lookup_host_or_compat( + getattr(model, "openai_api_base", None), module_tag="openai" + ) + if cls_module.startswith("langchain_anthropic"): + return _lookup_host_or_compat( + getattr(model, "anthropic_api_url", None), module_tag="anthropic" + ) + return None + + +def _lookup_host_or_compat(base_url: str | None, module_tag: str) -> str: + """Extract host from *base_url* and look up in ``_HOST_TO_PROVIDER``. + + Falls back to *module_tag* when no ``base_url`` is set (native SDK + default endpoint) or ``_compat`` for an unrecognized + host — the honest "openai SDK shape but unknown upstream" tag. + """ + if not base_url: + return module_tag + try: + from urllib.parse import urlparse + + host = urlparse(base_url).hostname + except Exception: + host = None + if not host: + return module_tag + return _HOST_TO_PROVIDER.get(host.lower(), f"{module_tag}_compat") + + +# --------------------------------------------------------------------------- +# SDK-field extractors — populate the envelope's optional fields +# --------------------------------------------------------------------------- + + +def _extract_status_code(exc: BaseException) -> int | None: + """Best-effort HTTP status code from a provider SDK exception. + + Order matters: openai/anthropic store it on ``.status_code``; + httpx-wrappers expose it via ``.response.status_code``; + ``google.genai.errors.APIError`` (unusually) stores it as an + integer ``.code`` — type-disambiguated from openai/anthropic's + string ``.code`` (provider error code, surfaced separately). + """ + status_code = getattr(exc, "status_code", None) + if isinstance(status_code, int): + return status_code + response = getattr(exc, "response", None) + if response is not None: + rsc = getattr(response, "status_code", None) + if isinstance(rsc, int): + return rsc + code = getattr(exc, "code", None) + if isinstance(code, int): + return code + return None + + +def _extract_provider_code(exc: BaseException) -> str | None: + """Provider error code (e.g. ``insufficient_quota``, + ``invalid_api_key``). Distinct from HTTP status; higher signal for + a WebUI toast than the integer alone. + """ + code = getattr(exc, "code", None) + if isinstance(code, str) and code: + return code + return None + + +def _extract_error_type(exc: BaseException) -> str | None: + """Provider error type label. + + - openai exposes this as ``.type`` (``rate_limit_error`` etc.) + - ``google.genai.errors.APIError`` stores a string label at + ``.status`` (``"NOT_FOUND"``, ``"RESOURCE_EXHAUSTED"``, …) — a + good fit for the same field. + + ``.type`` takes precedence when both are set. + """ + err_type = getattr(exc, "type", None) + if isinstance(err_type, str) and err_type: + return err_type + status = getattr(exc, "status", None) + if isinstance(status, str) and status: + return status + return None diff --git a/EvoScientist/llm/models.py b/EvoScientist/llm/models.py index 5484676..f9d6f27 100644 --- a/EvoScientist/llm/models.py +++ b/EvoScientist/llm/models.py @@ -10,11 +10,19 @@ endpoints) and convenient short names for common models. from __future__ import annotations import os +import re +import subprocess import warnings +from functools import lru_cache from typing import Any from langchain.chat_models import init_chat_model +from ..config.settings import ( + OPENROUTER_DEFAULT_APP_CATEGORIES, + OPENROUTER_DEFAULT_APP_TITLE, + OPENROUTER_DEFAULT_HTTP_REFERER, +) from .context_window import apply_known_context_window from .patches import ( _is_ccproxy_codex, @@ -37,6 +45,50 @@ _DEEPSEEK_BASE_URL = "https://api.deepseek.com" _MOONSHOT_BASE_URL = "https://api.moonshot.cn/v1" _KIMI_CODING_BASE_URL = "https://api.kimi.com/coding/" +# Minimum Codex CLI version advertised when no explicit override is set. Newer +# installed versions are advertised automatically. +_CODEX_CLIENT_VERSION_FALLBACK = "0.144.1" + + +@lru_cache(maxsize=1) +def _installed_codex_client_version() -> str: + """Return the installed Codex CLI version, or an empty string.""" + try: + result = subprocess.run( + ["codex", "--version"], + capture_output=True, + text=True, + timeout=2, + check=False, + ) + except (OSError, subprocess.TimeoutExpired): + return "" + + if result.returncode != 0: + return "" + match = re.search(r"\b(\d+\.\d+\.\d+)\b", result.stdout + result.stderr) + return match.group(1) if match else "" + + +def _resolve_codex_client_version() -> str: + """Resolve an explicit override or the newer of installed and minimum versions.""" + override = os.environ.get("EVOSCIENTIST_CODEX_CLIENT_VERSION", "").strip() + if override: + return override + + installed = _installed_codex_client_version() + if installed and tuple(map(int, installed.split("."))) >= tuple( + map(int, _CODEX_CLIENT_VERSION_FALLBACK.split(".")) + ): + return installed + return _CODEX_CLIENT_VERSION_FALLBACK + + +def _resolve_reasoning_effort(default: str) -> str: + """Return the configured reasoning effort or a provider-specific default.""" + return os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or default + + # Providers routed through the OpenAI provider with a custom base_url. # Maps provider name → (base_url or None, env var for API key). _OPENAI_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = { @@ -68,6 +120,19 @@ _THINKING_CAPABLE_PROVIDERS: set[str] = {"minimax"} _TRUTHY_ENV_VALUES = {"1", "true", "yes", "on"} _FALSEY_ENV_VALUES = {"0", "false", "no", "off"} +# OpenRouter app attribution (issue #339). Default values are the single source +# of truth in config/settings.py (imported above); langchain-openrouter maps +# app_url → HTTP-Referer, app_title → X-Title, app_categories → +# X-OpenRouter-Categories. OpenRouter honors at most this many categories per +# request (server-side limit) and silently ignores the rest, so the sent list is +# capped to this many below. https://openrouter.ai/docs/app-attribution +_OPENROUTER_MAX_CATEGORIES_PER_REQUEST = 2 + +# Legacy/provider-specific options that are not accepted by the installed +# LangChain chat model constructors. Leaving them at the top level makes +# LangChain move them into model_kwargs and can later leak them into SDK calls. +_UNSUPPORTED_CHAT_MODEL_KWARGS = frozenset({"sanitize_openai_sdk_headers"}) + # Model registry: list of (short_name, model_id, provider) # Allows same short_name across different providers. _MODEL_ENTRIES: list[tuple[str, str, str]] = [ @@ -264,6 +329,15 @@ def _env_flag_disabled(name: str) -> bool: return value is not None and value.strip().lower() in _FALSEY_ENV_VALUES +def _drop_unsupported_chat_model_kwargs(kwargs: dict[str, Any]) -> None: + for key in _UNSUPPORTED_CHAT_MODEL_KWARGS: + kwargs.pop(key, None) + model_kwargs = kwargs.get("model_kwargs") + if isinstance(model_kwargs, dict): + for key in _UNSUPPORTED_CHAT_MODEL_KWARGS: + model_kwargs.pop(key, None) + + def _supports_openrouter_anthropic_prompt_cache(provider: str, model_id: str) -> bool: """Return whether EvoScientist should declare OpenRouter Claude caching.""" return provider == "openrouter" and model_id.startswith( @@ -321,8 +395,16 @@ def _apply_auto_config( Mutates *kwargs* in place. Only sets keys that the caller hasn't already provided, so explicit user settings are never overridden. """ + disable_reasoning = bool(kwargs.pop("_disable_reasoning", False)) + disable_thinking = bool(kwargs.pop("_disable_thinking", False)) + if disable_reasoning: + kwargs.pop("reasoning", None) + kwargs.pop("include_thoughts", None) + if disable_thinking: + kwargs.pop("thinking", None) + # Anthropic: extended thinking - if provider == "anthropic" and "thinking" not in kwargs: + if provider == "anthropic" and not disable_thinking and "thinking" not in kwargs: _supports_thinking = original_provider in _THINKING_CAPABLE_PROVIDERS # Detect local proxy (e.g. ccproxy): thinking blocks in conversation # history cause 422 errors because the proxy doesn't accept 'thinking' @@ -341,24 +423,31 @@ def _apply_auto_config( kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000} # OpenAI (native, not third-party routed): reasoning - if provider == "openai" and not is_third_party and "reasoning" not in kwargs: - if _is_ccproxy_codex(): - # ccproxy uses Chat Completions which doesn't support reasoning. - pass - else: - _eff = ( - "xhigh" - if ("5.4" in model_id or "5.5" in model_id or "codex" in model_id) - else "high" + if ( + provider == "openai" + and not is_third_party + and not disable_reasoning + and "reasoning" not in kwargs + ): + _default_effort = ( + "xhigh" + if ( + "5.4" in model_id + or "5.5" in model_id + or "5.6" in model_id + or "codex" in model_id ) - kwargs["reasoning"] = {"effort": _eff, "summary": "auto"} + else "high" + ) + _eff = _resolve_reasoning_effort(_default_effort) + kwargs["reasoning"] = {"effort": _eff, "summary": "auto"} # Google GenAI: surface thinking traces - if provider == "google-genai": + if provider == "google-genai" and not disable_reasoning: kwargs.setdefault("include_thoughts", True) # Ollama: separate reasoning content from response for thinking models - if provider == "ollama" and "reasoning" not in kwargs: + if provider == "ollama" and not disable_reasoning and "reasoning" not in kwargs: kwargs["reasoning"] = True @@ -385,7 +474,46 @@ def get_chat_model( >>> model = get_chat_model("gpt-4o") # OpenAI model >>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID """ - model = model or DEFAULT_MODEL + skip_runtime_resolver = bool(kwargs.pop("_skip_runtime_model_resolver", False)) + runtime_provider_name: str | None = None + runtime_supports_reasoning: bool | None = None + runtime_resolved = None + if not skip_runtime_resolver: + from EvoScientist.runtime_integrations import resolve_runtime_model + + runtime_resolved = resolve_runtime_model(model, provider) + + if runtime_resolved is not None: + resolved_params = dict(getattr(runtime_resolved, "params", {}) or {}) + extra_body = resolved_params.pop("_extra_body", None) + default_headers = resolved_params.pop("_default_headers", None) + if extra_body: + resolved_params["extra_body"] = extra_body + if default_headers: + resolved_params["default_headers"] = default_headers + resolved_params.update(kwargs) + kwargs = resolved_params + + resolved_api_key = str(getattr(runtime_resolved, "api_key", "") or "") + resolved_base_url = str(getattr(runtime_resolved, "base_url", "") or "") + if resolved_api_key: + kwargs.setdefault("api_key", resolved_api_key) + if resolved_base_url: + kwargs.setdefault("base_url", resolved_base_url.rstrip("/")) + + runtime_provider_name = str( + getattr(runtime_resolved, "provider_name", "") or "" + ) + runtime_supports_reasoning = bool( + getattr(runtime_resolved, "supports_reasoning", False) + ) + if not runtime_supports_reasoning: + kwargs.setdefault("_disable_reasoning", True) + kwargs.setdefault("_disable_thinking", True) + model = str(runtime_resolved.model_id) + provider = str(runtime_resolved.protocol) + else: + model = model or DEFAULT_MODEL # Look up short name in registry (provider-aware) model_id = None @@ -420,22 +548,35 @@ def get_chat_model( _is_third_party = ( provider in _OPENAI_ROUTED_PROVIDERS or provider in _ANTHROPIC_ROUTED_PROVIDERS ) + if runtime_provider_name and runtime_provider_name != provider: + _is_third_party = True + if ( + runtime_resolved is not None + and provider == "openai" + and resolved_base_url + and "api.openai.com" not in resolved_base_url.lower() + ): + _is_third_party = True _is_openai_proxy = False - _original_provider: str | None = None + _original_provider: str | None = ( + runtime_provider_name if runtime_provider_name != provider else None + ) if provider == "anthropic": base_url = os.environ.get("ANTHROPIC_BASE_URL", "") if base_url: - kwargs["base_url"] = base_url + kwargs.setdefault("base_url", base_url) api_key = os.environ.get("ANTHROPIC_API_KEY", "") if api_key: - kwargs["api_key"] = api_key + kwargs.setdefault("api_key", api_key) # Native OpenAI base_url override (e.g. ccproxy Codex at localhost:8000/codex/v1) elif provider == "openai": base_url = os.environ.get("OPENAI_BASE_URL", "") if base_url: - kwargs["base_url"] = base_url - _is_openai_proxy = _is_ccproxy_codex() + kwargs.setdefault("base_url", base_url) + _is_openai_proxy = _is_ccproxy_codex( + kwargs.get("base_url"), kwargs.get("api_key") + ) if _is_openai_proxy: # Use Responses API for ccproxy: bypasses the format chain # converter (Chat→Responses→Chat) which returns 502 on @@ -448,9 +589,23 @@ def get_chat_model( # for Chat Completions tool_call duplication — not an issue # with the Responses API SSE format.) kwargs.pop("streaming", None) # remove if set elsewhere + # ccproxy forwards client headers upstream and only + # gap-fills its own, so the Codex backend sees this + # client's identity. Without Codex-CLI-shaped headers it + # rejects current models ("The '' model requires + # a newer version of Codex"). + _codex_ver = _resolve_codex_client_version() + _headers = kwargs.get("default_headers") or {} + kwargs["default_headers"] = _headers + _headers.setdefault("originator", "codex_cli_rs") + _headers.setdefault("version", _codex_ver) + _headers.setdefault( + "User-Agent", + f"codex_cli_rs/{_headers['version']} (EvoScientist)", + ) api_key = os.environ.get("OPENAI_API_KEY", "") if api_key: - kwargs["api_key"] = api_key + kwargs.setdefault("api_key", api_key) # OpenAI-routed providers → route through OpenAI provider with base_url elif provider in _OPENAI_ROUTED_PROVIDERS: @@ -468,10 +623,10 @@ def get_chat_model( else: base_url = base_url_default if base_url: - kwargs["base_url"] = base_url + kwargs.setdefault("base_url", base_url) api_key = os.environ.get(api_key_env, "") if api_key: - kwargs["api_key"] = api_key + kwargs.setdefault("api_key", api_key) # SiliconFlow: disable thinking — LangChain drops reasoning_content # from history, causing error 20015 on multi-turn requests. if provider == "siliconflow": @@ -488,15 +643,61 @@ def get_chat_model( _is_third_party = True api_key = os.environ.get("OPENROUTER_API_KEY", "") if api_key: - kwargs["api_key"] = api_key + kwargs.setdefault("api_key", api_key) # Reasoning via `effort` + `summary: "auto"` so a readable reasoning # summary is returned for display. OpenAI-Responses also emits encrypted # reasoning items (`rs_*` id) that can't be replayed on multi-turn # passback (OpenRouter's `/responses` beta is stateless, store=false — # "Item with id 'rs_...' not found"); the patch strips them on passback, # so enabling `summary` is safe. See langchain-ai/langchain#37777. - effort = os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or "high" + effort = _resolve_reasoning_effort("high") kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"}) + # App attribution (issue #339): identify EvoScientist to OpenRouter so + # usage is credited to the project (app rankings, model app tabs, + # analytics) rather than langchain-openrouter's LangChain-branded + # defaults. setdefault so an explicit caller kwarg wins; values are + # configurable via EVOSCIENTIST_OPENROUTER_* env (fed from the config + # file by apply_config_to_env). Applied only here, so no other provider + # ever receives these kwargs. + kwargs.setdefault( + "app_url", + os.environ.get("EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "").strip() + or OPENROUTER_DEFAULT_HTTP_REFERER, + ) + kwargs.setdefault( + "app_title", + os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE", "").strip() + or OPENROUTER_DEFAULT_APP_TITLE, + ) + # app_categories must be a list[str] (langchain-openrouter joins it into + # the X-OpenRouter-Categories header); split the comma-separated config + # value and drop blanks so a stray comma/space can't emit an empty one. + _app_categories_raw = ( + os.environ.get("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "").strip() + or OPENROUTER_DEFAULT_APP_CATEGORIES + ) + _app_categories = [ + c.strip() for c in _app_categories_raw.split(",") if c.strip() + ] + # Cap to the per-request limit and warn, so a misconfigured extra is + # dropped predictably here (and surfaced to the user) rather than being + # silently truncated server-side. + _limit = _OPENROUTER_MAX_CATEGORIES_PER_REQUEST + if len(_app_categories) > _limit: + warnings.warn( + f"OpenRouter accepts at most {_limit} app categories per " + f"request, so only the first {_limit} are sent: " + f"{_app_categories[:_limit]}. Ignoring the rest: " + f"{_app_categories[_limit:]}. Set " + f"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES (or the " + f"openrouter_app_categories config) to at most {_limit} " + f"categories to silence this warning.", + UserWarning, + stacklevel=2, + ) + _app_categories = _app_categories[:_limit] + if _app_categories: + kwargs.setdefault("app_categories", _app_categories) _patch_openrouter_strip_responses_reasoning() # Anthropic-routed providers → route through Anthropic provider with base_url @@ -517,10 +718,10 @@ def get_chat_model( else: base_url = base_url_default if base_url: - kwargs["base_url"] = base_url + kwargs.setdefault("base_url", base_url) api_key = os.environ.get(api_key_env, "") if api_key: - kwargs["api_key"] = api_key + kwargs.setdefault("api_key", api_key) # Kimi Coding Plan requires claude-code User-Agent header if provider == "kimi-coding": kwargs.setdefault("default_headers", {})["User-Agent"] = "claude-code/0.1.0" @@ -529,8 +730,9 @@ def get_chat_model( elif provider == "ollama": base_url = os.environ.get("OLLAMA_BASE_URL", "") if base_url: - kwargs["base_url"] = base_url + kwargs.setdefault("base_url", base_url) + _drop_unsupported_chat_model_kwargs(kwargs) _apply_auto_config(provider, model_id, _is_third_party, kwargs, _original_provider) _apply_openrouter_anthropic_prompt_cache(provider, model_id, kwargs) @@ -547,7 +749,14 @@ def get_chat_model( elif _responses_api_setting == "true": kwargs["use_responses_api"] = True - chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs) + anthropic_auth_token = None + if provider == "anthropic" and kwargs.get("api_key"): + anthropic_auth_token = os.environ.pop("ANTHROPIC_AUTH_TOKEN", None) + try: + chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs) + finally: + if anthropic_auth_token is not None: + os.environ["ANTHROPIC_AUTH_TOKEN"] = anthropic_auth_token # Flatten list content to strings for strict OpenAI-compatible providers # (DeepSeek, SiliconFlow, OpenRouter, custom-openai, etc.) and diff --git a/EvoScientist/llm/patches.py b/EvoScientist/llm/patches.py index db786d3..ca3b46e 100644 --- a/EvoScientist/llm/patches.py +++ b/EvoScientist/llm/patches.py @@ -25,6 +25,7 @@ Utilities: from __future__ import annotations +import hashlib import os from typing import Any @@ -178,15 +179,20 @@ _patch_ccproxy_codex_compat() # --------------------------------------------------------------------------- # Utility: detect ccproxy's Codex adapter (as opposed to generic localhost). # --------------------------------------------------------------------------- -def _is_ccproxy_codex() -> bool: +def _is_ccproxy_codex( + base_url: str | None = None, + api_key: str | None = None, +) -> bool: """Return True if the OpenAI endpoint is ccproxy's Codex adapter. Checks for the ccproxy-specific markers set by ``setup_codex_env()`` in ``ccproxy_manager.py``: the sentinel API key and the ``/codex/v1`` path. Plain localhost endpoints (vLLM, Ollama, etc.) are not affected. """ - base_url = os.environ.get("OPENAI_BASE_URL", "") - api_key = os.environ.get("OPENAI_API_KEY", "") + if base_url is None: + base_url = os.environ.get("OPENAI_BASE_URL", "") + if api_key is None: + api_key = os.environ.get("OPENAI_API_KEY", "") return ( ("127.0.0.1" in base_url or "localhost" in base_url) and api_key == "ccproxy-oauth" @@ -267,6 +273,299 @@ def _flatten_message_content(content: Any) -> str | list[Any] | Any: return "\n\n".join(parts) if parts else "" +def _stable_tool_call_id(message: Any, message_index: int, call_index: int) -> str: + seed = ":".join( + ( + str(getattr(message, "id", "") or "message"), + str(message_index), + str(call_index), + ) + ) + return "call_" + hashlib.sha256(seed.encode("utf-8")).hexdigest()[:24] + + +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.""" + + 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 + + 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 + + +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 + + 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. @@ -282,6 +581,9 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li from langchain_core.messages import HumanMessage + 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 @@ -320,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 @@ -729,6 +1033,66 @@ def _patch_openai_capture_reasoning_content() -> None: _patch_openai_capture_reasoning_content() +# --------------------------------------------------------------------------- +# Patch (module-level): silence langgraph_api's OpenAPI schema-generation +# warnings for endpoints whose docstrings aren't valid YAML. +# +# Upstream ``langgraph_api.utils.SchemaGenerator.get_schema`` calls +# ``parse_docstring`` (inherited from Starlette's ``BaseSchemaGenerator``) +# on every registered endpoint. When the docstring is prose with stray +# ``:`` characters, ``yaml.safe_load`` raises and upstream logs the +# failure + full traceback at WARNING level. It then falls back to +# ``{"description": docstring}`` — the endpoint still ends up in the +# schema with its prose as the description, just without structured +# ``parameters``/``responses``/``tags`` fields. +# +# The fallback path is fine; the warning + traceback is just noise. And +# it's only triggered for our deploy because mounting any custom Starlette +# app (``EvoScientist/langgraph_dev/http.py``) makes upstream call +# ``update_openapi_spec`` at startup — which iterates EVERY route, +# including upstream's own endpoints whose prose docstrings predate the +# YAML convention. +# +# Fix: wrap ``parse_docstring`` itself and absorb ``yaml.YAMLError`` by +# returning the same fallback shape upstream's except branch produces. +# Non-YAML exceptions are deliberately left to propagate — upstream's +# ``get_schema`` already catches them and logs WARNING + traceback, so +# unexpected failures remain debuggable. Patching ``parse_docstring`` (a +# small, stable method) instead of ``get_schema`` (the larger loop body) +# minimizes our exposure to upstream churn. +# --------------------------------------------------------------------------- +_langgraph_schema_silenced_patched = False + + +def _patch_langgraph_schema_generator_silence_warnings() -> None: + global _langgraph_schema_silenced_patched + if _langgraph_schema_silenced_patched: + return + try: + import langgraph_api.utils as _lgapi_utils + import yaml + + _SchemaGenerator = _lgapi_utils.SchemaGenerator + _orig_parse_docstring = _SchemaGenerator.parse_docstring + + def _patched_parse_docstring(self: Any, func: Any) -> dict[str, Any]: + try: + return _orig_parse_docstring(self, func) + except yaml.YAMLError: + return {"description": getattr(func, "__doc__", None) or ""} + + _SchemaGenerator.parse_docstring = _patched_parse_docstring + _langgraph_schema_silenced_patched = True + except Exception: + # Patches are loader-safe: never crash the import. Silent failure + # here just leaves the upstream warnings visible in deploy logs, + # which is a benign fallback. + pass + + +_patch_langgraph_schema_generator_silence_warnings() + + # --------------------------------------------------------------------------- # Patch (lazy, OpenRouter only): strip OpenAI-Responses encrypted reasoning # items from outgoing assistant messages. diff --git a/EvoScientist/logging_config.py b/EvoScientist/logging_config.py new file mode 100644 index 0000000..a8abdfd --- /dev/null +++ b/EvoScientist/logging_config.py @@ -0,0 +1,302 @@ +"""Shared logging configuration helpers.""" + +from __future__ import annotations + +import logging +import os +import sys +from datetime import UTC, datetime +from pathlib import Path +from typing import Any, TextIO + +DEFAULT_LOG_RETENTION_DAYS = 30 +DEFAULT_LOG_FORMAT = "%(asctime)s [%(levelname)s] %(name)s: %(message)s" +DEFAULT_LOG_DATE_FORMAT = "%Y-%m-%d %H:%M:%S" +MANAGED_HANDLER_ATTR = "_evoscientist_managed_handler" + + +def resolve_log_level(level: int | str | None, default: int = logging.INFO) -> int: + """Resolve a logging level from config or environment input.""" + if isinstance(level, int): + return level + raw = str(level or "").strip() + if not raw: + return default + if raw.isdigit(): + return int(raw) + normalized = raw.upper() + if normalized == "WARN": + normalized = "WARNING" + resolved = logging.getLevelNamesMapping().get(normalized) + return resolved if isinstance(resolved, int) else default + + +def _mark_managed(handler: logging.Handler, kind: str) -> logging.Handler: + setattr(handler, MANAGED_HANDLER_ATTR, kind) + return handler + + +def _managed_kind(handler: logging.Handler) -> str | None: + kind = getattr(handler, MANAGED_HANDLER_ATTR, None) + return kind if isinstance(kind, str) else None + + +def remove_managed_handlers( + logger: logging.Logger | None = None, + *, + kinds: set[str] | None = None, +) -> None: + """Remove handlers installed by this module without touching external ones.""" + target = logger or logging.getLogger() + for handler in target.handlers[:]: + kind = _managed_kind(handler) + if kind and (kinds is None or kind in kinds): + target.removeHandler(handler) + handler.close() + + +def _standard_formatter() -> logging.Formatter: + return logging.Formatter(DEFAULT_LOG_FORMAT, datefmt=DEFAULT_LOG_DATE_FORMAT) + + +class DailyLogFileHandler(logging.FileHandler): + """File handler that writes the active log to a date-based filename.""" + + def __init__( + self, + log_dir: str | Path, + *, + prefix: str = "evoscientist", + retention_days: int = DEFAULT_LOG_RETENTION_DAYS, + encoding: str = "utf-8", + utc: bool = False, + ) -> None: + self.log_dir = Path(log_dir).expanduser() + self.prefix = prefix + self.retention_days = max(1, retention_days) + self.utc = utc + self.log_dir.mkdir(parents=True, exist_ok=True) + super().__init__(self._dated_log_path(), encoding=encoding, delay=True) + + @property + def active_log_path(self) -> Path: + """Return the active log path for the current date.""" + return self._dated_log_path() + + def _dated_log_path(self) -> Path: + now = datetime.now(UTC if self.utc else None) + return self.log_dir / f"{self.prefix}-{now:%Y-%m-%d}.log" + + def emit(self, record: logging.LogRecord) -> None: + try: + expected = str(self.active_log_path) + if self.baseFilename != expected: + if self.stream: + self.stream.close() + self.stream = None + self.baseFilename = expected + self._delete_expired_logs() + super().emit(record) + except OSError: + self.handleError(record) + + def getFilesToDelete(self) -> list[str]: + candidates = sorted(self.log_dir.glob(f"{self.prefix}-????-??-??.log")) + if len(candidates) <= self.retention_days: + return [] + return [str(path) for path in candidates[: -self.retention_days]] + + def _delete_expired_logs(self) -> None: + for path in self.getFilesToDelete(): + try: + os.remove(path) + except OSError: + pass + + +def default_log_dir() -> Path: + """Return the default runtime log directory.""" + env_dir = os.environ.get("EVOSCIENTIST_LOG_DIR") + if env_dir: + return Path(env_dir).expanduser() + + from EvoScientist.paths import DATA_DIR + + return DATA_DIR / "logs" + + +def configure_daily_file_logging( + logger: logging.Logger | None = None, + *, + log_dir: str | Path | None = None, + prefix: str = "evoscientist", + level: int | str = logging.INFO, + retention_days: int = DEFAULT_LOG_RETENTION_DAYS, +) -> DailyLogFileHandler: + """Attach a daily file handler, replacing older matching handlers.""" + target = logger or logging.getLogger() + resolved_level = resolve_log_level(level, default=logging.INFO) + retention_days = max(1, int(retention_days)) + resolved_dir = Path(log_dir).expanduser() if log_dir else default_log_dir() + + for handler in target.handlers[:]: + if ( + isinstance(handler, DailyLogFileHandler) + and handler.prefix == prefix + and handler.log_dir == resolved_dir + ): + target.removeHandler(handler) + handler.close() + + handler = DailyLogFileHandler( + resolved_dir, + prefix=prefix, + retention_days=retention_days, + ) + _mark_managed(handler, "file") + handler.setLevel(resolved_level) + handler.setFormatter(_standard_formatter()) + target.addHandler(handler) + if target.level == logging.NOTSET or target.level > resolved_level: + target.setLevel(resolved_level) + return handler + + +def configure_console_logging( + logger: logging.Logger | None = None, + *, + level: int | str | None = logging.INFO, + stream: TextIO | None = None, + replace: bool = True, +) -> logging.StreamHandler: + """Attach a standard console handler for non-interactive entry points.""" + target = logger or logging.getLogger() + resolved_level = resolve_log_level(level, default=logging.INFO) + if replace: + remove_managed_handlers(target, kinds={"console", "rich"}) + + handler = logging.StreamHandler(stream or sys.stderr) + _mark_managed(handler, "console") + handler.setLevel(resolved_level) + handler.setFormatter(_standard_formatter()) + target.addHandler(handler) + target.setLevel(resolved_level) + return handler + + +def configure_rich_console_logging( + logger: logging.Logger | None = None, + *, + level: int | str | None = logging.INFO, + console: Any = None, + replace: bool = True, + dim_warnings: bool = False, + show_time: bool | None = None, + show_path: bool | None = None, + show_level: bool | None = None, +) -> logging.Handler: + """Attach a Rich console handler for interactive CLI output.""" + from rich.logging import RichHandler + from rich.markup import escape + + target = logger or logging.getLogger() + resolved_level = resolve_log_level(level, default=logging.INFO) + verbose = resolved_level <= logging.DEBUG + if replace: + remove_managed_handlers(target, kinds={"console", "rich"}) + + class DimWarningHandler(RichHandler): + def emit(self, record: logging.LogRecord) -> None: + if dim_warnings and record.levelno == logging.WARNING and console is not None: + console.print( + "[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] " + f"[dim]{escape(record.getMessage())}[/dim]" + ) + return + super().emit(record) + + handler = DimWarningHandler( + console=console, + show_time=verbose if show_time is None else show_time, + show_path=verbose if show_path is None else show_path, + show_level=verbose if show_level is None else show_level, + ) + _mark_managed(handler, "rich") + handler.setLevel(resolved_level) + target.addHandler(handler) + target.setLevel(resolved_level) + return handler + + +def configure_logging( + logger: logging.Logger | None = None, + *, + level: int | str | None = logging.INFO, + log_dir: str | Path | None = None, + retention_days: int = DEFAULT_LOG_RETENTION_DAYS, + prefix: str = "evoscientist", + console: bool = True, + file: bool = True, + replace_managed: bool = True, +) -> list[logging.Handler]: + """Configure standard EvoScientist console and daily file logging.""" + target = logger or logging.getLogger() + resolved_level = resolve_log_level(level, default=logging.INFO) + if replace_managed: + remove_managed_handlers(target, kinds={"console", "rich", "file"}) + + handlers: list[logging.Handler] = [] + if console: + handlers.append( + configure_console_logging(target, level=resolved_level, replace=False) + ) + if file: + handlers.append( + configure_daily_file_logging( + target, + log_dir=log_dir, + prefix=prefix, + level=resolved_level, + retention_days=retention_days, + ) + ) + target.setLevel(resolved_level) + return handlers + + +def configure_logging_from_settings( + logger: logging.Logger | None = None, + *, + default_level: int = logging.INFO, + prefix: str = "evoscientist", + console: bool = True, + file: bool = True, +) -> list[logging.Handler]: + """Configure logging from EvoScientist settings and environment overrides.""" + level: int | str | None = os.environ.get("EVOSCIENTIST_LOG_LEVEL") + log_dir: str | Path | None = os.environ.get("EVOSCIENTIST_LOG_DIR") or None + retention_days = int( + os.environ.get("EVOSCIENTIST_LOG_RETENTION_DAYS", DEFAULT_LOG_RETENTION_DAYS) + ) + + try: + from EvoScientist.config import get_effective_config + + cfg = get_effective_config() + level = level or getattr(cfg, "log_level", None) + log_dir = log_dir or getattr(cfg, "log_dir", None) or None + retention_days = int( + getattr(cfg, "log_retention_days", DEFAULT_LOG_RETENTION_DAYS) + ) + except Exception: + level = level or default_level + + return configure_logging( + logger, + level=resolve_log_level(level, default=default_level), + log_dir=log_dir, + retention_days=retention_days, + prefix=prefix, + console=console, + file=file, + ) diff --git a/EvoScientist/mcp/__init__.py b/EvoScientist/mcp/__init__.py index 7747f55..1187f19 100644 --- a/EvoScientist/mcp/__init__.py +++ b/EvoScientist/mcp/__init__.py @@ -10,6 +10,7 @@ from .client import ( build_mcp_add_kwargs, build_mcp_edit_fields, edit_mcp_server, + get_mcp_server_errors, load_mcp_config, load_mcp_tools, parse_mcp_add_args, @@ -38,6 +39,7 @@ __all__ = [ "find_server_by_name", "get_all_tags", "get_installed_names", + "get_mcp_server_errors", "install_mcp_server", "install_mcp_servers", "load_mcp_config", diff --git a/EvoScientist/mcp/client.py b/EvoScientist/mcp/client.py index bf931d7..ca984df 100644 --- a/EvoScientist/mcp/client.py +++ b/EvoScientist/mcp/client.py @@ -114,6 +114,10 @@ _URL_TRANSPORTS = {"http", "streamable_http", "sse", "websocket"} # still parallelizing the common 3–7 server case to completion. _MAX_CONCURRENT_CONNECTIONS = 8 +# Last connection error per configured server. This is process-local runtime +# diagnostics for the Web/CLI status surfaces, not persisted configuration. +_MCP_SERVER_ERRORS: dict[str, str] = {} + # Env vars forwarded to stdio MCP subprocesses on top of the MCP SDK's # minimal default set (HOME/PATH/USER/…). Without this, servers behind # a proxy or with a custom CA bundle silently fail with long timeouts. @@ -764,6 +768,9 @@ async def _load_tools( if not connections: return {} + for stale_name in set(_MCP_SERVER_ERRORS) - set(connections): + _MCP_SERVER_ERRORS.pop(stale_name, None) + client = MultiServerMCPClient(connections) # type: ignore[invalid-argument-type] def _report(event: str, name: str, detail: str = "") -> None: @@ -787,10 +794,13 @@ async def _load_tools( _report("start", name) try: tools = await client.get_tools(server_name=name) + _MCP_SERVER_ERRORS.pop(name, None) logger.info("MCP server %r: loaded %d tool(s)", name, len(tools)) _report("success", name, str(len(tools))) return name, tools except Exception as exc: + detail = str(exc) or type(exc).__name__ + _MCP_SERVER_ERRORS[name] = detail # When the caller wired up ``on_progress`` they own the # user-facing display; downgrade the logger so we don't # double-print. @@ -798,7 +808,7 @@ async def _load_tools( logger.warning("MCP server %r: failed to load tools: %s", name, exc) else: logger.debug("MCP server %r: failed to load tools: %s", name, exc) - _report("error", name, str(exc)) + _report("error", name, detail) return name, [] # ``return_exceptions=False`` is fine because ``_fetch`` already @@ -807,6 +817,11 @@ async def _load_tools( return dict(results) +def get_mcp_server_errors() -> dict[str, str]: + """Return a snapshot of the most recent per-server connection errors.""" + return dict(_MCP_SERVER_ERRORS) + + async def aload_mcp_tools( config: dict[str, Any] | None = None, *, diff --git a/EvoScientist/memory/agents/memory_worker.py b/EvoScientist/memory/agents/memory_worker.py index 9fa0e19..5009d72 100644 --- a/EvoScientist/memory/agents/memory_worker.py +++ b/EvoScientist/memory/agents/memory_worker.py @@ -428,6 +428,7 @@ def _memory_worker_middleware( enable_observation_memory: bool = True, ): """Build middleware for memory workers, excluding task execution tools.""" + from ...middleware.error_normalization import ErrorNormalizationMiddleware from ...middleware.memory import create_memory_middleware memory_controls = MemoryControls( @@ -439,18 +440,23 @@ def _memory_worker_middleware( enable_observation_tool = memory_controls.observation_tool_enabled( _memory_worker_observation_target(source_type) ) - return memory_agent_middleware( - create_memory_middleware( - str(memory_dir), - workspace_dir=workspace_dir, - source_type=source_type, - source_agent=_memory_worker_agent_name(source_type), - enable_profile_memory=enable_profile_memory, - enable_observation_memory=enable_observation_memory, - enable_observation_tool=enable_observation_tool, + return [ + # Outermost — normalize provider-SDK exceptions from the + # auxiliary model call before any inner middleware sees them. + ErrorNormalizationMiddleware(), + *memory_agent_middleware( + create_memory_middleware( + str(memory_dir), + workspace_dir=workspace_dir, + source_type=source_type, + source_agent=_memory_worker_agent_name(source_type), + enable_profile_memory=enable_profile_memory, + enable_observation_memory=enable_observation_memory, + enable_observation_tool=enable_observation_tool, + ), + excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS, ), - excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS, - ) + ] def _build_memory_worker_agent( diff --git a/EvoScientist/memory/agents/observation_linker.py b/EvoScientist/memory/agents/observation_linker.py index fb170c0..044158b 100644 --- a/EvoScientist/memory/agents/observation_linker.py +++ b/EvoScientist/memory/agents/observation_linker.py @@ -71,6 +71,8 @@ def build_observation_linker_graph( workspace_dir: str | Path | None = None, ) -> CompiledStateGraph: """Build the registered LangGraph observation linker.""" + from ...middleware.error_normalization import ErrorNormalizationMiddleware + agent_paths = resolve_memory_agent_paths( memory_dir=memory_dir, workspace_dir=workspace_dir, @@ -85,5 +87,7 @@ def build_observation_linker_graph( tools=tools, memory_dir=agent_paths.memory_dir, workspace_dir=agent_paths.workspace_dir, - middleware=memory_agent_middleware(), + # Outermost — normalize provider-SDK exceptions from the + # auxiliary model call before any inner middleware sees them. + middleware=[ErrorNormalizationMiddleware(), *memory_agent_middleware()], ) diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py index cd28de6..d79e0ed 100644 --- a/EvoScientist/middleware/__init__.py +++ b/EvoScientist/middleware/__init__.py @@ -18,6 +18,7 @@ from .context_editing import ( create_context_editing_middleware, ) from .context_overflow import ContextOverflowMapperMiddleware +from .error_normalization import ErrorNormalizationMiddleware from .memory import ( EvoMemoryMiddleware, create_memory_middleware, @@ -28,29 +29,42 @@ 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", "Choice", "ConfigurableModelMiddleware", "ContextOverflowMapperMiddleware", + "ErrorNormalizationMiddleware", "EvoMemoryLifecycleMiddleware", "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/code_interpreter.py b/EvoScientist/middleware/code_interpreter.py index 1f54c50..467cb77 100644 --- a/EvoScientist/middleware/code_interpreter.py +++ b/EvoScientist/middleware/code_interpreter.py @@ -45,7 +45,21 @@ _MEMORY_FIRST_INTERPRETER_PROMPT = ( class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware): - """Code interpreter middleware with EvoScientist's memory preflight hint.""" + """Code interpreter middleware with EvoScientist's memory preflight hint. + + ``after_agent`` / ``aafter_agent`` are intentionally NOT overridden. An + earlier "conditional snapshot" gate that skipped ``after_agent`` on turns + where ``code_interpreter`` wasn't called saved ~50 ms/turn of + ``create_snapshot()`` work, but also skipped the slot eviction upstream + performs in the same hook (``finally: self._registry.evict(thread_id)`` + in ``langchain_quickjs.middleware.CodeInterpreterMiddleware.after_agent``). + ``before_agent`` restores the REPL on every turn that follows a touched + one via ``self._registry.get(thread_id)`` (get-or-create), so skipping + eviction leaked one ``ThreadWorker`` + QuickJS Runtime per persistent + ``thread_id`` that ever went touched → quiet. The regression test + ``test_after_agent_evicts_slot_on_untouched_turn`` guards against + reintroducing the gate. + """ def _prepare_for_call(self, request: ModelRequest) -> str: return super()._prepare_for_call(request) + _MEMORY_FIRST_INTERPRETER_PROMPT diff --git a/EvoScientist/middleware/error_normalization.py b/EvoScientist/middleware/error_normalization.py new file mode 100644 index 0000000..992445c --- /dev/null +++ b/EvoScientist/middleware/error_normalization.py @@ -0,0 +1,240 @@ +"""ErrorNormalizationMiddleware — catch provider-SDK exceptions at the +model boundary and re-raise as a normalized non-dataclass wrapper. + +Some provider SDKs (openrouter.errors.* today) decorate their exception +classes with ``@dataclass``. When langgraph_api emits an SSE error +frame via ``json_dumpb`` → ``orjson.dumps(obj, default=default, +option=OPT_SERIALIZE_DATACLASS)``, orjson's dataclass fast-path +enumerates the fields directly and skips the ``default=`` hook that +builds our envelope. The wire payload comes out as +``{"message": …, "status_code": …, "body": …, "headers": null, +"raw_response": null, "data": {…}}`` with no ``error`` / ``class`` / +``provider`` envelope and no way for the WebUI to distinguish quota / +auth / rate-limit / model-not-found. + +This middleware sits at the model-call boundary. It catches +``BaseException`` from ``handler()``, and if ``request.model`` is a +recognized provider SDK client, wraps the exception in a +:class:`~EvoScientist.llm.errors.ProviderStreamError` (a plain +``Exception`` subclass, not a dataclass). The wrapper carries the SSE +envelope pre-baked on its instance attributes. + +Contract: the wrap decision is based on the **model**, not the +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 +from a recognized provider, or the request carries no ``.model``, the +exception re-raises unchanged and upstream's whitelist / catch-all +behavior takes over. +""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from typing import TYPE_CHECKING + +from langchain.agents.middleware.types import ( + AgentMiddleware, + ModelRequest, + ModelResponse, +) + +if TYPE_CHECKING: + from ..llm.errors import ProviderStreamError + + +def _should_pass_through(exc: BaseException) -> bool: + """True if *exc* is a LangGraph-level signal that must propagate + untouched — either a control-flow signal or a structural error + that isn't a provider failure. + + Covers everything in ``langgraph.errors.*``: + + - **Control flow** (breaking these would corrupt the interrupt / + resume protocol): ``GraphBubbleUp`` and its subclasses + ``GraphInterrupt``, ``NodeInterrupt``, ``ParentCommand``, + ``GraphDrained``. + - **Structural** (wrapping would mis-attribute a graph-level + issue as a provider failure): ``InvalidUpdateError``, + ``EmptyInputError``, ``EmptyChannelError``, ``TaskNotFound``, + ``GraphRecursionError``, ``NodeCancelledError``, + ``NodeTimeoutError``. + + Symmetric with upstream ``langgraph_api.serde.default``'s + whitelist, which also exposes these classes' ``str(exc)`` untouched + rather than swallowing them behind a provider envelope. + + ``KeyboardInterrupt``, ``SystemExit``, and ``asyncio.CancelledError`` + are handled implicitly by catching ``Exception`` — they inherit + from ``BaseException``. + """ + return (type(exc).__module__ or "").startswith("langgraph.errors") + + +# Module prefixes for provider SDK exceptions. Consumed by +# ``_is_provider_error`` to decide whether an exception raised inside +# a model call should surface as a provider incident or gracefully +# degrade (used by ``_ConditionalToolSelectorMiddleware``). +# +# Related sibling: ``_HOST_TO_PROVIDER`` in ``llm/errors.py`` — the +# host-side allow-list. Adding a whole new provider SDK means updating +# both; adding a new routed provider (new base_url through an existing +# SDK) only touches ``_HOST_TO_PROVIDER``. +_PROVIDER_EXC_MODULE_PREFIXES: tuple[str, ...] = ( + "openai", + "anthropic", + "google.genai", + "google.api_core", + "openrouter", + "langchain_openai", + "langchain_anthropic", + "langchain_google_genai", + "langchain_openrouter", + "httpx", +) + + +def _is_provider_error(exc: BaseException) -> bool: + """True if *exc* looks like it originated inside a provider SDK + (openai, anthropic, google.genai, openrouter, httpx, or their + langchain wrappers), as opposed to a shape / config error (structured + output not supported, malformed schema, missing tool, …). + + Used by callers that need to decide whether an exception from the + model call is worth surfacing to the user (provider errors) or + can be silently degraded around (shape errors). Cheap alternative + to inspecting ``status_code`` / ``request`` because some provider + errors — connection errors, timeouts — don't carry those attributes. + """ + module = type(exc).__module__ or "" + return any(module.startswith(p) for p in _PROVIDER_EXC_MODULE_PREFIXES) + + +def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError | None: + """Return a :class:`ProviderStreamError` wrapping *exc* if the model + on *request* comes from a recognized provider SDK, or ``None`` if + the caller should re-raise *exc* unchanged. + + Provider is read from ``request.model`` — the definitive config + the exception was raised under, not inferred from the exception + class / URL. Status / code / redaction still come from the raised + exception because those fields are populated by the SDK at raise + time. + + Returns ``None`` (caller re-raises unchanged) for: + + - Already-normalized wrappers (would double-attribute). + - LangGraph control-flow / structural errors — see + ``_should_pass_through``. This gate lives here so every caller + of ``_normalize`` (not just the wrap sites of this middleware) + gets the protection automatically. Notably + ``ModelFallbackMiddleware`` also calls ``_normalize`` at the + raise point of its fallback chain. + - ``ContextOverflowError`` — a cross-layer control signal that + deepagents' ``SummarizationMiddleware`` catches by type from + **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, + _extract_status_code, + _provider_from_model, + _redact_api_keys, + ) + + # Already normalized (e.g. by ModelFallbackMiddleware wrapping against + # the actual failing model rather than the original request's model). + # Pass through — re-wrapping would double-attribute. + 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): + return None + + # SummarizationMiddleware sits outside our stack and catches this + # by exact type to trigger reactive history compression + retry. + if isinstance(exc, ContextOverflowError): + return None + + provider = _provider_from_model(getattr(request, "model", None)) + if provider is None: + return None + cls = type(exc) + mod = cls.__module__ or "" + class_qualname = f"{mod}.{cls.__qualname__}" if mod else cls.__qualname__ + + request_id_attr = getattr(exc, "request_id", None) + request_id = ( + request_id_attr + if isinstance(request_id_attr, str) and request_id_attr + else None + ) + + return ProviderStreamError( + provider=provider, + class_qualname=class_qualname, + message=_redact_api_keys(str(exc)), + status_code=_extract_status_code(exc), + code=_extract_provider_code(exc), + err_type=_extract_error_type(exc), + request_id=request_id, + ) + + +class ErrorNormalizationMiddleware(AgentMiddleware): + """Wrap the model call in try/except and normalize provider SDK + exceptions into a non-dataclass envelope wrapper. + + Place this middleware **outermost** in the chain (first in the + middleware list) so it catches exceptions raised by inner + middlewares as well as the model handler itself. + """ + + name = "error_normalization" + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + try: + return handler(request) + except Exception as exc: + normalized = _normalize(request, exc) + if normalized is None: + raise + raise normalized from exc + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + try: + return await handler(request) + except Exception as exc: + normalized = _normalize(request, exc) + if normalized is None: + raise + raise normalized from exc diff --git a/EvoScientist/middleware/model_fallback.py b/EvoScientist/middleware/model_fallback.py index 0dee6ff..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" @@ -263,7 +268,15 @@ async def _try_fallbacks( "Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc ) + # Track the request whose model actually raised ``last_exc`` so we + # can attribute the exception to the failing model, not the + # original ``request.model``. Without this, a fallback chain + # ``deepseek → moonshot`` where moonshot exhausts its quota would + # surface as ``provider: deepseek`` — the model the user never + # actually saw fail. last_exc = primary_exc + last_failing_request = request + for model_name, provider in get_fallback_chain(): _emit( f" -> Falling back to {model_name} ({provider}) " @@ -288,8 +301,9 @@ async def _try_fallbacks( f"-- aborting fallback chain", style="red", ) - raise + _raise_normalized(fb_request, fb_exc) last_exc = fb_exc + last_failing_request = fb_request _emit( f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}", style="red", @@ -303,7 +317,24 @@ async def _try_fallbacks( ) _emit(" All fallbacks exhausted -- re-raising last error", style="red") - raise last_exc + _raise_normalized(last_failing_request, last_exc) + + +def _raise_normalized(request: ModelRequest, exc: Exception) -> None: + """Wrap *exc* in a ``ProviderStreamError`` attributed to + ``request.model`` and raise, so the outer chain sees the failure + tagged with the model that actually raised. + + Falls back to a plain ``raise`` when the model isn't from a + recognized provider (``_normalize`` returns None) — nothing useful + to add. + """ + from .error_normalization import _normalize + + normalized = _normalize(request, exc) + if normalized is not None: + raise normalized from exc + raise exc def _guard_and_fallback( @@ -330,7 +361,7 @@ def _guard_and_fallback( f"Model error ({reason}) -- not eligible for fallback, re-raising", style="red", ) - raise primary_exc + _raise_normalized(request, primary_exc) return _try_fallbacks(request, invoke, primary_exc) 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 3e28432..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", } ) @@ -132,10 +133,21 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware): return self._build_selector(request).wrap_model_call( request, _handler_after_selection ) - except Exception: + except Exception as exc: if _handler_called: raise # Error from downstream model — don't retry - # Selector itself failed (e.g., structured output not supported). + from ..llm.errors import ProviderStreamError + from .error_normalization import _is_provider_error + + if isinstance(exc, ProviderStreamError) or _is_provider_error(exc): + # Auth / quota / connection failures on the selector's + # own model. Falling back to "use all tools" would hit + # the same provider anyway (same client, likely same + # credentials). Surface it instead so the user sees + # the real cause. + raise + # Structured-output shape / config failure — gracefully + # degrade to using all tools. logger.debug("Tool selector failed, using all tools", exc_info=True) if self._track_stream_selection: _selector_active = False @@ -171,9 +183,16 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware): return await self._build_selector(request).awrap_model_call( request, _handler_after_selection ) - except Exception: + except Exception as exc: if _handler_called: raise + from ..llm.errors import ProviderStreamError + from .error_normalization import _is_provider_error + + if isinstance(exc, ProviderStreamError) or _is_provider_error(exc): + # See sync path — surface provider errors, degrade only + # on shape / config failures. + raise logger.debug("Tool selector failed, using all tools", exc_info=True) if self._track_stream_selection: _selector_active = False @@ -258,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/paths.py b/EvoScientist/paths.py index c813f77..cfd247e 100644 --- a/EvoScientist/paths.py +++ b/EvoScientist/paths.py @@ -5,6 +5,7 @@ from __future__ import annotations import logging import os import shutil +from collections.abc import Iterator from datetime import datetime from pathlib import Path @@ -215,3 +216,65 @@ def resolve_virtual_path(virtual_path: str) -> Path: """Resolve a virtual workspace path (e.g. /image.png) to a real filesystem path.""" vpath = virtual_path if virtual_path.startswith("/") else "/" + virtual_path return (_active_workspace / vpath.lstrip("/")).resolve() + + +def evoscientist_root() -> Path: + """Return the application root used by Gateway-managed runtime data.""" + env_root = os.environ.get("EVOSCIENTIST_HOME") + if env_root: + return Path(env_root).expanduser().resolve() + return DATA_DIR.expanduser().resolve() + + +_EVOSCIENTIST_DATA_ROOT: Path | None = None + + +def _data_root() -> Path: + """Return the root directory for isolated Web user workspaces.""" + global _EVOSCIENTIST_DATA_ROOT + if _EVOSCIENTIST_DATA_ROOT is not None: + return _EVOSCIENTIST_DATA_ROOT + + env_root = os.environ.get("EVOSCIENTIST_DATA_ROOT") + if env_root: + root = Path(env_root).expanduser().resolve() + else: + root = evoscientist_root() / "data" + _EVOSCIENTIST_DATA_ROOT = root + return root + + +def user_data_dir(user_id: str) -> Path: + """Return and create the isolated data directory for a Web user.""" + path = _data_root() / user_id + path.mkdir(parents=True, exist_ok=True) + return path + + +def iter_user_data_dirs() -> Iterator[Path]: + """Yield existing Web user directories without creating the data root.""" + root = _data_root() + if not root.exists(): + return + for path in root.iterdir(): + if path.is_dir(): + yield path + + +def thread_data_dir(user_id: str, thread_id: str) -> Path: + """Return and create a user's isolated thread workspace.""" + path = user_data_dir(user_id) / thread_id + path.mkdir(parents=True, exist_ok=True) + return path + + +def global_data_dir(user_id: str) -> Path: + """Return and create a user's directory shared across all threads.""" + path = user_data_dir(user_id) / "__global__" + path.mkdir(parents=True, exist_ok=True) + return path + + +def uploads_dir() -> Path: + """Return the Gateway upload staging directory.""" + return evoscientist_root() / "uploads" diff --git a/EvoScientist/runtime_integrations.py b/EvoScientist/runtime_integrations.py new file mode 100644 index 0000000..bc0ae2f --- /dev/null +++ b/EvoScientist/runtime_integrations.py @@ -0,0 +1,116 @@ +"""Optional runtime services supplied by an application embedding EvoScientist. + +The CLI package must not import a concrete web gateway. Applications such as +Ai4Sci-Web can register their database, storage, metering, and media services +at process startup through this module. +""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from dataclasses import dataclass, replace +from datetime import date +from pathlib import Path +from typing import Any + +AsyncProvider = Callable[[], Awaitable[Any]] +AsyncFileHandler = Callable[[Path], Awaitable[Any]] +AsyncUsageRecorder = Callable[[str, str], Awaitable[Any]] +ModelResolver = Callable[[str | None, str | None], Any | None] + + +class RuntimeIntegrationUnavailable(RuntimeError): + """Raised when an optional host-provided service is not configured.""" + + +@dataclass(frozen=True) +class RuntimeIntegrations: + app_connection_provider: AsyncProvider | None = None + session_connection_provider: AsyncProvider | None = None + session_dsn_provider: Callable[[], str | None] | None = None + current_date_provider: Callable[[], date] | None = None + user_storage_root_provider: Callable[[str], Path] | None = None + knowledge_file_handler: AsyncFileHandler | None = None + usage_recorder: AsyncUsageRecorder | None = None + image_backend_factory: Callable[[], Any] | None = None + model_resolver: ModelResolver | None = None + + +_integrations = RuntimeIntegrations() + + +def configure_runtime_integrations(**services: Any) -> RuntimeIntegrations: + """Register host-provided services and return the resulting configuration.""" + global _integrations + _integrations = replace(_integrations, **services) + return _integrations + + +def reset_runtime_integrations() -> None: + """Clear all host-provided services, primarily for tests.""" + global _integrations + _integrations = RuntimeIntegrations() + + +def has_session_connection_provider() -> bool: + return _integrations.session_connection_provider is not None + + +def get_session_dsn() -> str | None: + provider = _integrations.session_dsn_provider + return provider() if provider is not None else None + + +async def get_session_connection() -> Any: + provider = _integrations.session_connection_provider + if provider is None: + raise RuntimeIntegrationUnavailable( + "No session connection provider is configured" + ) + return await provider() + + +async def get_app_connection() -> Any: + provider = _integrations.app_connection_provider + if provider is None: + raise RuntimeIntegrationUnavailable( + "No application connection provider is configured" + ) + return await provider() + + +def current_date() -> date: + provider = _integrations.current_date_provider + return provider() if provider is not None else date.today() + + +def resolve_user_storage_root(user_id: str) -> Path | None: + provider = _integrations.user_storage_root_provider + return provider(user_id) if provider is not None else None + + +def resolve_runtime_model(model: str | None, provider: str | None = None) -> Any | None: + """Resolve a host-managed model configuration when one is registered.""" + resolver = _integrations.model_resolver + return resolver(model, provider) if resolver is not None else None + + +async def handle_knowledge_file(path: Path) -> None: + handler = _integrations.knowledge_file_handler + if handler is not None: + await handler(path) + + +async def record_service_usage(service: str, action: str) -> None: + recorder = _integrations.usage_recorder + if recorder is not None: + await recorder(service, action) + + +def get_image_backend() -> Any: + factory = _integrations.image_backend_factory + if factory is None: + raise RuntimeIntegrationUnavailable( + "Image generation is unavailable in this runtime. Configure an image backend first." + ) + return factory() 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 f22c161..2e63eef 100644 --- a/EvoScientist/stream/events.py +++ b/EvoScientist/stream/events.py @@ -8,13 +8,15 @@ import base64 import inspect import mimetypes import os +import warnings from collections.abc import AsyncGenerator, AsyncIterator, Mapping from dataclasses import dataclass 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 @@ -43,6 +45,12 @@ GraphRunInput: TypeAlias = str | Command LangGraphStreamInput: TypeAlias = dict[str, list[dict[str, object]]] | Command _ValueMessageKey: TypeAlias = tuple[str, ...] +warnings.filterwarnings( + "ignore", + message=r"The v3 streaming protocol on Pregel is experimental\.", + category=LangChainBetaWarning, +) + @dataclass(frozen=True, slots=True) class _AssistantValueMessage: @@ -87,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()`` / @@ -106,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: " @@ -119,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( @@ -134,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, ...] @@ -209,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 @@ -237,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 @@ -382,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 [] @@ -407,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 @@ -457,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, @@ -464,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 @@ -486,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: @@ -526,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, @@ -572,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( @@ -583,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) @@ -614,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): @@ -652,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]]: @@ -797,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. @@ -812,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, @@ -821,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: @@ -949,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/conftest.py b/tests/conftest.py index a6458d7..1e46d61 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,7 +1,11 @@ """Shared fixtures for EvoScientist tests.""" +from pathlib import Path + import pytest +_NONEXISTENT_DOTENV = str(Path(__file__).with_name(".pytest-dotenv-does-not-exist")) + @pytest.fixture(autouse=True) def _reset_tool_selection_state(): @@ -164,3 +168,24 @@ def restore_model_passthrough_patch(): yield finally: _reset() + + +@pytest.fixture(autouse=True) +def _isolate_dotenv(monkeypatch): + """Keep the developer's real .env out of the test environment. + + ``get_effective_config`` runs ``load_dotenv(find_dotenv(usecwd=True), + override=True)``, so any test that loads config injects the repo's + real .env into ``os.environ`` for the rest of the pytest process. + An empty-valued line like ``MINIMAX_BASE_URL=`` then makes + ``os.environ.get(key, default)`` return "" instead of the default, + breaking unrelated tests later in the run (see issue #322). + + Pointing ``find_dotenv`` at a fixed path that does not exist makes + ``load_dotenv`` a no-op without creating a temporary directory for + every test. + """ + monkeypatch.setattr( + "EvoScientist.config.settings.find_dotenv", + lambda *args, **kwargs: _NONEXISTENT_DOTENV, + ) diff --git a/tests/test_agent_factory_extensions.py b/tests/test_agent_factory_extensions.py new file mode 100644 index 0000000..e05dce5 --- /dev/null +++ b/tests/test_agent_factory_extensions.py @@ -0,0 +1,137 @@ +def test_create_cli_agent_accepts_host_backend_and_memory_options( + monkeypatch, tmp_path +): + import EvoScientist.EvoScientist as agent_module + from EvoScientist.config.settings import EvoScientistConfig + + calls = {} + workspace_backend = object() + chat_model = object() + + class _CompositeBackend: + def __init__(self, *, default, routes): + calls["default_backend"] = default + calls["routes"] = routes + + class _MemoryBackend: + def __init__(self, **kwargs): + calls["memory_backend_kwargs"] = kwargs + + class _SkillsBackend: + def __init__(self, **kwargs): + calls["skills_backend_kwargs"] = kwargs + + class _Agent: + def with_config(self, config): + calls["agent_config"] = config + return self + + cfg = EvoScientistConfig(auto_approve=True, recursion_limit=321) + + monkeypatch.setattr("deepagents.backends.CompositeBackend", _CompositeBackend) + monkeypatch.setattr("deepagents.create_deep_agent", lambda **kwargs: _Agent()) + monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _MemoryBackend) + monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _SkillsBackend) + monkeypatch.setattr(agent_module, "set_active_workspace", lambda path: None) + monkeypatch.setattr( + agent_module, + "_get_default_middleware", + lambda **kwargs: calls.setdefault("middleware_kwargs", kwargs) or [], + ) + monkeypatch.setattr( + agent_module, + "load_mcp_and_build_kwargs", + lambda *args, **kwargs: {"subagents": [{"name": "research"}]}, + ) + + memory_dir = tmp_path / "memory" + result = agent_module.create_cli_agent( + workspace_dir=str(tmp_path / "workspace"), + checkpointer=object(), + config=cfg, + chat_model=chat_model, + workspace_backend=workspace_backend, + memory_dir=memory_dir, + tool_selector_threshold=8, + memory_max_inline_profile_chars=1000, + enable_subagents=False, + enable_background_execution=False, + ) + + assert isinstance(result, _Agent) + assert calls["default_backend"] is workspace_backend + assert calls["memory_backend_kwargs"] == { + "root_dir": str(memory_dir), + "virtual_mode": True, + } + assert calls["middleware_kwargs"]["memory_dir"] == str(memory_dir) + assert calls["middleware_kwargs"]["tool_selector_threshold"] == 8 + 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_ccproxy_manager.py b/tests/test_ccproxy_manager.py index 7b0c7f9..7b9547e 100644 --- a/tests/test_ccproxy_manager.py +++ b/tests/test_ccproxy_manager.py @@ -6,6 +6,8 @@ from unittest.mock import MagicMock, patch import pytest from EvoScientist.ccproxy_manager import ( + _CCPROXY_AUTH_TIMEOUT_SECONDS, + _CCPROXY_HEALTH_TIMEOUT_SECONDS, check_ccproxy_auth, ensure_ccproxy, is_ccproxy_available, @@ -15,6 +17,7 @@ from EvoScientist.ccproxy_manager import ( setup_codex_env, start_ccproxy, stop_ccproxy, + write_ccproxy_config, ) # ============================================================================= @@ -52,6 +55,8 @@ class TestCheckCcproxyAuth: mock_run.assert_called_once() cmd = mock_run.call_args[0][0] assert cmd[1:] == ["auth", "status", "claude_api"] + # ccproxy CLI cold start takes ~10s; timeout must leave headroom + assert mock_run.call_args[1]["timeout"] == _CCPROXY_AUTH_TIMEOUT_SECONDS @patch("subprocess.run") def test_valid_auth_codex(self, mock_run): @@ -123,9 +128,10 @@ class TestIsCcproxyRunning: class TestStartCcproxy: + @patch("EvoScientist.ccproxy_manager.logger.warning") @patch("EvoScientist.ccproxy_manager.is_ccproxy_running") @patch("subprocess.Popen") - def test_success(self, mock_popen, mock_running): + def test_success(self, mock_popen, mock_running, mock_warning): proc = MagicMock() proc.poll.return_value = None mock_popen.return_value = proc @@ -134,6 +140,11 @@ class TestStartCcproxy: result = start_ccproxy(8000) assert result is proc + mock_warning.assert_called_once_with( + "Starting ccproxy on port %d; first startup may take up to %d seconds", + 8000, + _CCPROXY_HEALTH_TIMEOUT_SECONDS, + ) @patch("EvoScientist.ccproxy_manager.is_ccproxy_running", return_value=False) @patch("EvoScientist.ccproxy_manager.time") @@ -143,7 +154,11 @@ class TestStartCcproxy: proc.poll.return_value = None mock_popen.return_value = proc # Simulate time passing beyond deadline - mock_time.monotonic.side_effect = [0, 0, 31] + mock_time.monotonic.side_effect = [ + 0, + 0, + _CCPROXY_HEALTH_TIMEOUT_SECONDS + 1, + ] mock_time.sleep = MagicMock() with pytest.raises(RuntimeError, match="did not become healthy"): @@ -154,6 +169,54 @@ class TestStartCcproxy: with pytest.raises(FileNotFoundError): start_ccproxy(8000) + @patch("EvoScientist.ccproxy_manager.is_ccproxy_running") + @patch("subprocess.Popen") + def test_passes_generated_config(self, mock_popen, mock_running, tmp_path): + proc = MagicMock() + proc.poll.return_value = None + mock_popen.return_value = proc + mock_running.side_effect = [True] + + with patch("EvoScientist.config.get_config_dir", return_value=tmp_path): + start_ccproxy(8000) + + cmd = mock_popen.call_args[0][0] + assert "--config" in cmd + assert cmd[cmd.index("--config") + 1] == str(tmp_path / "ccproxy.toml") + + @patch("EvoScientist.ccproxy_manager.is_ccproxy_running") + @patch("EvoScientist.ccproxy_manager.write_ccproxy_config", side_effect=OSError) + @patch("subprocess.Popen") + def test_config_write_failure_starts_without_config( + self, mock_popen, mock_write, mock_running + ): + proc = MagicMock() + proc.poll.return_value = None + mock_popen.return_value = proc + mock_running.side_effect = [True] + + start_ccproxy(8000) + + cmd = mock_popen.call_args[0][0] + assert "--config" not in cmd + + +# ============================================================================= +# write_ccproxy_config +# ============================================================================= + + +class TestWriteCcproxyConfig: + def test_writes_codex_mapping_override(self, tmp_path): + config_dir = tmp_path / "missing" / "config" + with patch("EvoScientist.config.get_config_dir", return_value=config_dir): + path = write_ccproxy_config() + + assert path == str(config_dir / "ccproxy.toml") + content = (config_dir / "ccproxy.toml").read_text(encoding="utf-8") + assert "[plugins.codex]" in content + assert "model_mappings = []" in content + # ============================================================================= # ensure_ccproxy diff --git a/tests/test_code_interpreter_middleware.py b/tests/test_code_interpreter_middleware.py index 8d5103f..8baa89d 100644 --- a/tests/test_code_interpreter_middleware.py +++ b/tests/test_code_interpreter_middleware.py @@ -1,4 +1,5 @@ -"""Regression tests for the code_interpreter PTC allowlist. +"""Regression tests for the code_interpreter PTC allowlist and the +``EvoCodeInterpreterMiddleware`` subclass shape. langchain-quickjs >=0.3 reserves the ``task`` sub-agent dispatch tool as the top-level REPL global and raises ``ValueError`` if ``task`` appears in the @@ -8,7 +9,10 @@ allowlist (``task()`` stays reachable as the REPL global, with responseSchema). from __future__ import annotations +from unittest.mock import MagicMock + import pytest +from langchain_core.messages import AIMessage, HumanMessage from EvoScientist.middleware.code_interpreter import ( _DEFAULT_PTC_ALLOWLIST, @@ -45,3 +49,290 @@ def test_filter_tools_for_ptc_accepts_default_allowlist(): def test_create_code_interpreter_middleware_builds(): assert create_code_interpreter_middleware() is not None + + +def test_middleware_uses_thread_mode(): + """Upstream ``mode="thread"`` (the default) preserves cross-turn REPL + state as ``langchain-ai/deepagents#3064`` shipped it. The wire-cost + bloat that motivated the earlier ``mode="turn"`` regression guard is + fixed at the API serialization layer (``EvoFilteredGraph`` in + ``EvoScientist/langgraph_dev/main_graph.py``), not by revoking the + persistence feature. + """ + mw = create_code_interpreter_middleware() + assert mw._mode == "thread" + + +def test_after_agent_evicts_slot_on_untouched_turn(): + """Regression guard against reintroducing a conditional-snapshot gate + that skips ``after_agent`` on untouched turns. + + Upstream ``after_agent`` in ``langchain_quickjs/middleware.py`` performs + two things: snapshot the REPL AND evict the slot (``finally: + self._registry.evict(thread_id)``). ``before_agent`` restores the REPL + on any turn that follows a touched one via ``self._registry.get`` — + which is get-or-create. So if ``after_agent`` returns early without + evicting, one ``ThreadWorker`` + QuickJS Runtime leaks per persistent + ``thread_id`` that ever went touched → quiet. + + Fix: don't override ``after_agent`` / ``aafter_agent`` at all — inherit + upstream's unconditional snapshot+evict behavior. This test creates a + slot the way ``before_agent`` would, calls ``after_agent`` with an + untouched-state input, and asserts the slot was evicted. + """ + mw = create_code_interpreter_middleware() + tid = mw._fallback_thread_id + + # Simulate the slot creation that ``before_agent`` performs when it sees + # a prior turn's snapshot payload in state. + mw._registry.get(tid) + assert len(mw._registry._slots) == 1 + + # Untouched-turn state: no ``code_interpreter`` tool call between the + # last ``HumanMessage`` and end. Under the earlier buggy gate this + # returned ``{}`` without evicting — leaking the slot created above. + untouched_state = { + "_quickjs_snapshot_payload": b"payload-from-prior-turn", + "messages": [ + HumanMessage(content="thanks"), + AIMessage(content="you're welcome"), + ], + } + mw.after_agent(untouched_state, runtime=None) + + assert len(mw._registry._slots) == 0, ( + "after_agent must evict the slot even on untouched turns, because " + "before_agent already restored a REPL that owns a ThreadWorker + " + "QuickJS Runtime. Skipping eviction leaks those resources." + ) + + +def test_evo_filtered_graph_strips_private_snapshot_field(): + """The ``StateSnapshot`` returned by ``EvoScientist_agent.get_state`` must + not contain ``_quickjs_snapshot_payload`` in either ``values`` (the + materialized channel payload) or ``metadata['writes']`` (the raw write + records surfaced by ``get_state_history``). + """ + from EvoScientist.langgraph_dev.main_graph import _EvoFilteredGraph, _strip_private + + snap = MagicMock() + snap.values = { + "messages": ["m1"], + "_quickjs_snapshot_payload": b"x" * 100, + "skills_metadata": [], + } + snap.metadata = { + "source": "loop", + "step": 42, + "writes": { + "CodeInterpreterMiddleware.after_agent": { + "_quickjs_snapshot_payload": ("snap", b"y" * 1_400_000), + "messages": [], + }, + "model": {"messages": ["m1"]}, + }, + "parents": {}, + } + _strip_private(snap) + snap._replace.assert_called_once() + kwargs = snap._replace.call_args.kwargs + assert "_quickjs_snapshot_payload" not in kwargs["values"] + assert "messages" in kwargs["values"] + assert "skills_metadata" in kwargs["values"] + scrubbed_writes = kwargs["metadata"]["writes"] + assert ( + "_quickjs_snapshot_payload" + not in scrubbed_writes["CodeInterpreterMiddleware.after_agent"] + ) + assert "messages" in scrubbed_writes["CodeInterpreterMiddleware.after_agent"] + assert scrubbed_writes["model"] == {"messages": ["m1"]} + # Non-writes metadata keys are preserved. + assert kwargs["metadata"]["source"] == "loop" + assert kwargs["metadata"]["step"] == 42 + # Sanity: the class exists and inherits from CompiledStateGraph. + from langgraph.graph.state import CompiledStateGraph + + assert issubclass(_EvoFilteredGraph, CompiledStateGraph) + + +def test_strip_private_handles_missing_metadata_writes(): + """``metadata['writes']`` can be missing or ``None`` on some snapshots + (e.g. initial state). The filter must not crash and must still strip + values. + """ + from EvoScientist.langgraph_dev.main_graph import _strip_private + + snap = MagicMock() + snap.values = {"_quickjs_snapshot_payload": b"x", "messages": []} + snap.metadata = {"source": "input", "step": -1, "writes": None} + snap.tasks = () + _strip_private(snap) + kwargs = snap._replace.call_args.kwargs + assert "_quickjs_snapshot_payload" not in kwargs["values"] + # writes was None, metadata passes through unchanged. + assert kwargs["metadata"]["writes"] is None + + +def test_strip_private_scrubs_task_result_snapshot_blob(): + """``tasks[*].result`` is where ``after_agent``'s return dict lands. + When the middleware snapshots, ``result`` carries + ``{"_quickjs_snapshot_payload": ("snap", ~1.4 MB bytes)}``. Verified + on live history: this is the dominant per-response leak, larger than + ``values`` and ``metadata.writes`` combined for anchor checkpoints. + """ + from EvoScientist.langgraph_dev.main_graph import _strip_private + + class FakeTask: + def __init__(self, id_, result): + self.id = id_ + self.name = "CodeInterpreterMiddleware.after_agent" + self.result = result + + def _replace(self, **kwargs): + for k, v in kwargs.items(): + setattr(self, k, v) + return self + + leaking_task = FakeTask( + "t1", {"_quickjs_snapshot_payload": ("snap", b"z" * 1_400_000), "messages": []} + ) + clean_task = FakeTask("t2", {"messages": ["hi"]}) + snap = MagicMock() + snap.values = {} + snap.metadata = {"source": "loop", "step": 5} + snap.tasks = (leaking_task, clean_task) + _strip_private(snap) + kwargs = snap._replace.call_args.kwargs + tasks_after = kwargs["tasks"] + assert "_quickjs_snapshot_payload" not in tasks_after[0].result + assert "messages" in tasks_after[0].result + # Clean task is passed through untouched. + assert tasks_after[1] is clean_task + + +def test_agent_uses_filtered_graph_class(): + """The ``__class__`` swap in ``main_graph.py`` is the load-bearing wiring + that makes ``_strip_private`` reach the langgraph-api endpoints. + ``_strip_private`` and ``_EvoFilteredGraph`` in isolation don't prove the + swap ran; every other test in this file passes even if someone drops the + swap line. This asserts the compiled agent is actually the filtered + subclass at module-load time, and that the subclass survives + ``Pregel.copy(update=...)`` — the call langgraph-api makes in + ``get_graph`` before yielding the graph to endpoint handlers. + """ + from EvoScientist.langgraph_dev.main_graph import ( + EvoScientist_agent, + _EvoFilteredGraph, + ) + + assert isinstance(EvoScientist_agent, _EvoFilteredGraph) + assert isinstance(EvoScientist_agent.copy(update={}), _EvoFilteredGraph) + + +def test_all_registered_graphs_use_filtered_graph_class(): + """Every graph registered in ``langgraph.json`` (main + all subagents) + gets the ``__class__`` swap via ``_apply_filter_to_all_registered_graphs``. + Iterating the config directly matches the auto-detect refactor: adding + a new subagent to ``langgraph.json`` should not require a corresponding + test update. + + Subagents get ``create_code_interpreter_middleware`` unconditionally + (``EvoScientist.py:_build_middleware_stack``), so they can touch the + QuickJS REPL and write ``_quickjs_snapshot_payload`` on their own + checkpoint namespace. Async subagents also get their own ``thread_id`` + and their ``/threads/{id}/state`` endpoint runs on their own compiled + graph — without the swap on those graphs, our filter would miss that + endpoint entirely. + """ + import json + from importlib import import_module + from pathlib import Path + + # Import triggers ``main_graph``'s swap loop. + from EvoScientist.langgraph_dev import main_graph + from EvoScientist.langgraph_dev.main_graph import _EvoFilteredGraph + + config_path = Path(main_graph.__file__).parent / "langgraph.json" + config = json.loads(config_path.read_text()) + for name, path in config["graphs"].items(): + module_path, attr = path.rsplit(":", 1) + graph = getattr(import_module(module_path), attr) + assert isinstance(graph, _EvoFilteredGraph), ( + f"graph {name!r} ({path}) did not receive the class swap" + ) + + +def test_strip_private_recurses_into_nested_subgraph_state(): + """When ``subgraphs=True``, ``PregelTask.state`` holds a nested + ``StateSnapshot`` for the subgraph. Its ``values`` (and its own nested + tasks) can carry ``_quickjs_snapshot_payload`` just like the parent. + Recursion covers the compound leak path CodeRabbit flagged. + """ + from langgraph.types import StateSnapshot + + from EvoScientist.langgraph_dev.main_graph import _strip_private + + nested_snap = StateSnapshot( + values={"_quickjs_snapshot_payload": b"n" * 1_400_000, "messages": []}, + next=(), + config={}, + metadata={"source": "loop", "step": 3}, + created_at="2026-07-01T12:00:00Z", + parent_config=None, + tasks=(), + interrupts=(), + ) + + class FakeTask: + def __init__(self, state): + self.id = "sub-1" + self.name = "subgraph" + self.result = None + self.state = state + + def _replace(self, **kwargs): + for k, v in kwargs.items(): + setattr(self, k, v) + return self + + task_with_nested = FakeTask(nested_snap) + task_with_config_state = FakeTask({"configurable": {"thread_id": "t"}}) + snap = MagicMock() + snap.values = {} + snap.metadata = {"source": "loop", "step": 5} + snap.tasks = (task_with_nested, task_with_config_state) + _strip_private(snap) + kwargs = snap._replace.call_args.kwargs + tasks_after = kwargs["tasks"] + # Nested StateSnapshot got recursively scrubbed. + assert "_quickjs_snapshot_payload" not in tasks_after[0].state.values + assert "messages" in tasks_after[0].state.values + # A dict (RunnableConfig-shaped) state passes through unchanged — we only + # recurse into ``StateSnapshot`` instances. + assert tasks_after[1].state == {"configurable": {"thread_id": "t"}} + + +def test_strip_private_scrubs_delta_counters(): + """``metadata['counters_since_delta_snapshot']`` is a small + ``{channel: [count, superstep]}`` bookkeeping map. Not a size problem, + but leaks the channel name — strip for consistency with the private + annotation. + """ + from EvoScientist.langgraph_dev.main_graph import _strip_private + + snap = MagicMock() + snap.values = {} + snap.metadata = { + "source": "loop", + "step": 5, + "counters_since_delta_snapshot": { + "_quickjs_snapshot_payload": [1, 14], + "messages": [3, 14], + }, + } + snap.tasks = () + _strip_private(snap) + kwargs = snap._replace.call_args.kwargs + counters = kwargs["metadata"]["counters_since_delta_snapshot"] + assert "_quickjs_snapshot_payload" not in counters + assert "messages" in counters diff --git a/tests/test_config.py b/tests/test_config.py index 8799da5..07a040e 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -52,12 +52,9 @@ def _restore_dangerous_env(): def temp_config_dir(tmp_path, monkeypatch): """Use a temporary directory for config during tests.""" config_dir = tmp_path / "evoscientist" + monkeypatch.delenv("EVOSCIENTIST_CONFIG_DIR", raising=False) + monkeypatch.delenv("EVOSCIENTIST_HOME", raising=False) monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) - # Prevent load_dotenv from loading the project's real .env file - monkeypatch.setattr( - "EvoScientist.config.settings.find_dotenv", - lambda *a, **k: str(tmp_path / ".env"), - ) # Also clear any API keys from environment for key in [ "ANTHROPIC_API_KEY", @@ -77,6 +74,9 @@ def temp_config_dir(tmp_path, monkeypatch): "EVOSCIENTIST_AUXILIARY_MODEL", "EVOSCIENTIST_AUXILIARY_PROVIDER", "EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE", + "EVOSCIENTIST_OPENROUTER_HTTP_REFERER", + "EVOSCIENTIST_OPENROUTER_APP_TITLE", + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "EVOSCIENTIST_DANGEROUS_MODE", ]: monkeypatch.delenv(key, raising=False) @@ -104,6 +104,9 @@ def clean_env(monkeypatch): "EVOSCIENTIST_AUXILIARY_MODEL", "EVOSCIENTIST_AUXILIARY_PROVIDER", "EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE", + "EVOSCIENTIST_OPENROUTER_HTTP_REFERER", + "EVOSCIENTIST_OPENROUTER_APP_TITLE", + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "EVOSCIENTIST_DANGEROUS_MODE", ]: monkeypatch.delenv(key, raising=False) @@ -129,8 +132,13 @@ class TestEvoScientistConfig: assert config.show_thinking is True assert config.ui_backend == "tui" assert config.log_level == "warning" - assert config.reasoning_effort == "high" + assert config.reasoning_effort == "" assert config.openrouter_anthropic_prompt_cache is True + assert config.openrouter_http_referer == ( + "https://github.com/EvoScientist/EvoScientist" + ) + assert config.openrouter_app_title == "EvoScientist" + assert config.openrouter_app_categories == "creative-writing,personal-agent" assert config.memory_profile_enabled is True assert config.memory_observations_enabled is True assert config.memory_observation_writer == MemoryObservationWriter.ALL @@ -145,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.""" @@ -192,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 @@ -199,14 +221,34 @@ class TestEvoScientistConfig: class TestConfigPaths: + def test_get_config_dir_with_explicit_override(self, monkeypatch, tmp_path): + """An explicit config directory has the highest priority.""" + config_dir = tmp_path / "gateway-config" + monkeypatch.setenv("EVOSCIENTIST_CONFIG_DIR", str(config_dir)) + monkeypatch.setenv("EVOSCIENTIST_HOME", str(tmp_path / "runtime-home")) + + assert get_config_dir() == config_dir.resolve() + + def test_get_config_dir_with_evoscientist_home(self, monkeypatch, tmp_path): + """Runtime home keeps configuration and data under one root.""" + home = tmp_path / "runtime-home" + monkeypatch.delenv("EVOSCIENTIST_CONFIG_DIR", raising=False) + monkeypatch.setenv("EVOSCIENTIST_HOME", str(home)) + + assert get_config_dir() == home.resolve() / "config" + def test_get_config_dir_with_xdg(self, monkeypatch, tmp_path): """Test config dir uses XDG_CONFIG_HOME when set.""" + monkeypatch.delenv("EVOSCIENTIST_CONFIG_DIR", raising=False) + monkeypatch.delenv("EVOSCIENTIST_HOME", raising=False) monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) config_dir = get_config_dir() assert config_dir == tmp_path / "evoscientist" def test_get_config_dir_default(self, monkeypatch): """Test config dir defaults to ~/.config/evoscientist.""" + monkeypatch.delenv("EVOSCIENTIST_CONFIG_DIR", raising=False) + monkeypatch.delenv("EVOSCIENTIST_HOME", raising=False) monkeypatch.delenv("XDG_CONFIG_HOME", raising=False) config_dir = get_config_dir() assert config_dir == Path.home() / ".config" / "evoscientist" @@ -672,6 +714,22 @@ class TestPriorityChain: config = get_effective_config() assert config.openrouter_anthropic_prompt_cache is False + def test_env_openrouter_app_attribution_override( + self, temp_config_dir, monkeypatch + ): + """OpenRouter app-attribution env vars should override file config.""" + save_config(EvoScientistConfig()) + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "https://acme.test") + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_TITLE", "Acme") + monkeypatch.setenv( + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "cli-agent,programming-app" + ) + + config = get_effective_config() + assert config.openrouter_http_referer == "https://acme.test" + assert config.openrouter_app_title == "Acme" + assert config.openrouter_app_categories == "cli-agent,programming-app" + def test_set_openrouter_anthropic_prompt_cache(self, temp_config_dir, clean_env): """Test OpenRouter Anthropic prompt cache can be set through config.""" save_config(EvoScientistConfig()) @@ -732,6 +790,60 @@ class TestApplyConfigToEnv: "false" ) + def test_openrouter_app_attribution_applied_to_env(self, clean_env, monkeypatch): + """Config app-attribution values are exported to env for models.py.""" + for env in ( + "EVOSCIENTIST_OPENROUTER_HTTP_REFERER", + "EVOSCIENTIST_OPENROUTER_APP_TITLE", + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", + ): + monkeypatch.delenv(env, raising=False) + config = EvoScientistConfig( + openrouter_http_referer="https://acme.test", + openrouter_app_title="Acme", + openrouter_app_categories="cli-agent,programming-app", + ) + apply_config_to_env(config) + + assert os.environ.get("EVOSCIENTIST_OPENROUTER_HTTP_REFERER") == ( + "https://acme.test" + ) + assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE") == "Acme" + assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES") == ( + "cli-agent,programming-app" + ) + + def test_openrouter_app_attribution_env_not_overwritten( + self, clean_env, monkeypatch + ): + """apply_config_to_env must not clobber an already-set attribution env var.""" + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_TITLE", "existing-title") + config = EvoScientistConfig(openrouter_app_title="config-title") + apply_config_to_env(config) + + assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE") == "existing-title" + + def test_openrouter_app_attribution_empty_config_not_applied( + self, clean_env, monkeypatch + ): + """Empty-string attribution config must not create env vars.""" + for env in ( + "EVOSCIENTIST_OPENROUTER_HTTP_REFERER", + "EVOSCIENTIST_OPENROUTER_APP_TITLE", + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", + ): + monkeypatch.delenv(env, raising=False) + config = EvoScientistConfig( + openrouter_http_referer="", + openrouter_app_title="", + openrouter_app_categories="", + ) + apply_config_to_env(config) + + assert os.environ.get("EVOSCIENTIST_OPENROUTER_HTTP_REFERER") is None + assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE") is None + assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES") is None + def test_dangerous_mode_round_trips_to_env(self, clean_env, monkeypatch): """dangerous_mode set via CLI override must survive a fresh re-read. @@ -849,3 +961,30 @@ def test_scheduler_config_defaults_and_env(monkeypatch): assert eff2.memory_skill_synthesis_mode == MemorySkillSynthesisMode.AUTO assert eff2.memory_skill_synthesis_cadence == MemorySkillSynthesisCadence.MONTHLY assert eff2.memory_skill_synthesis_time == "04:30" + + +# ============================================================================= +# Dotenv isolation (issue #322) +# ============================================================================= + + +class TestDotenvIsolation: + def test_env_file_not_leaked_into_process_env(self, tmp_path, monkeypatch): + """A .env in cwd must not leak into os.environ during tests. + + Without the suite-wide ``_isolate_dotenv`` fixture, + ``get_effective_config`` loads the developer's real .env with + ``override=True``; an empty-valued line like ``MINIMAX_BASE_URL=`` + then poisons ``os.environ.get(key, default)`` lookups for every + test that runs afterwards in the same process. + """ + monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) + repro_dir = tmp_path / "repro" + repro_dir.mkdir() + (repro_dir / ".env").write_text("MINIMAX_BASE_URL=\n") + monkeypatch.chdir(repro_dir) + monkeypatch.delenv("MINIMAX_BASE_URL", raising=False) + + get_effective_config() + + assert "MINIMAX_BASE_URL" not in os.environ diff --git a/tests/test_error_normalization_middleware.py b/tests/test_error_normalization_middleware.py new file mode 100644 index 0000000..efd1799 --- /dev/null +++ b/tests/test_error_normalization_middleware.py @@ -0,0 +1,440 @@ +"""Tests for ErrorNormalizationMiddleware + ProviderStreamError. + +Verifies that provider-SDK exceptions from a chat model call get +wrapped into a non-dataclass ``ProviderStreamError`` at the model +boundary, and that non-provider exceptions pass through unchanged. +The provider tag is derived from ``request.model`` (class + base_url), +not from the raised exception. +""" + +from __future__ import annotations + +import asyncio +import dataclasses +from types import SimpleNamespace + +import pytest + +from EvoScientist.llm.errors import ( + AgentControlError, + ModelToolProtocolError, + ProviderStreamError, +) +from EvoScientist.middleware.error_normalization import ( + ErrorNormalizationMiddleware, + _normalize, +) + +# --------------------------------------------------------------------------- +# Test fixtures — fake chat model instances + requests +# --------------------------------------------------------------------------- + + +def _fake_model(module: str, cls_name: str, **attrs): + """Build a fake chat model instance whose ``type(model).__module__`` + matches *module*, carrying arbitrary attributes for ``base_url`` / + ``openai_api_base`` / ``anthropic_api_url`` lookup. + """ + cls = type(cls_name, (), {"__module__": module}) + inst = cls() + for k, v in attrs.items(): + setattr(inst, k, v) + return inst + + +def _request(model): + """Fake ``ModelRequest`` with just the ``.model`` attribute the + middleware reads. + """ + return SimpleNamespace(model=model) + + +def _openai_model(base_url: str | None = None): + return _fake_model( + "langchain_openai.chat_models.base", + "ChatOpenAI", + openai_api_base=base_url, + ) + + +def _anthropic_model(base_url: str | None = None): + return _fake_model( + "langchain_anthropic.chat_models", + "ChatAnthropic", + anthropic_api_url=base_url, + ) + + +def _openrouter_model(): + return _fake_model("langchain_openrouter.chat_models", "ChatOpenRouter") + + +def _google_model(): + return _fake_model("langchain_google_genai.chat_models", "ChatGoogleGenerativeAI") + + +def _make_exc(cls_name: str = "APIError", message: str = "boom", **attrs): + """Build a plain-Exception subclass carrying arbitrary attributes + (``status_code``, ``code``, ``type``, ``request_id`` …). + """ + cls = type(cls_name, (Exception,), attrs) + return cls(message) + + +# --------------------------------------------------------------------------- +# _normalize — provider inference from ModelRequest.model +# --------------------------------------------------------------------------- + + +class TestNormalize: + def test_openai_native_model_tags_openai(self): + req = _request(_openai_model()) + exc = _make_exc(message="rate limited", status_code=429) + wrapped = _normalize(req, exc) + assert isinstance(wrapped, ProviderStreamError) + assert wrapped.provider == "openai" + assert wrapped.status_code == 429 + + def test_openai_routed_deepseek_tagged_by_base_url(self): + req = _request(_openai_model(base_url="https://api.deepseek.com")) + wrapped = _normalize(req, _make_exc(message="quota exceeded")) + assert wrapped.provider == "deepseek" + + def test_openai_routed_moonshot_tagged_by_base_url(self): + req = _request(_openai_model(base_url="https://api.moonshot.cn/v1")) + assert _normalize(req, _make_exc()).provider == "moonshot" + + def test_unknown_openai_compat_host_tagged_openai_compat(self): + req = _request(_openai_model(base_url="https://internal.corp/v1")) + assert _normalize(req, _make_exc()).provider == "openai_compat" + + def test_anthropic_native_model_tags_anthropic(self): + req = _request(_anthropic_model(base_url="https://api.anthropic.com")) + assert _normalize(req, _make_exc()).provider == "anthropic" + + def test_anthropic_routed_minimax_tagged_by_base_url(self): + req = _request(_anthropic_model(base_url="https://api.minimaxi.com/anthropic")) + assert _normalize(req, _make_exc()).provider == "minimax" + + def test_unknown_anthropic_compat_host_tagged_anthropic_compat(self): + req = _request(_anthropic_model(base_url="https://internal.corp/v1")) + assert _normalize(req, _make_exc()).provider == "anthropic_compat" + + def test_openrouter_tagged_from_class_alone(self): + req = _request(_openrouter_model()) + wrapped = _normalize(req, _make_exc(cls_name="UnauthorizedResponseError")) + assert wrapped.provider == "openrouter" + assert wrapped.class_qualname.endswith(".UnauthorizedResponseError") + + def test_google_genai_tagged_from_class_alone(self): + req = _request(_google_model()) + assert _normalize(req, _make_exc()).provider == "google_genai" + + def test_unrecognized_model_class_returns_none(self): + req = _request(_fake_model("some.other.pkg", "SomeModel")) + assert _normalize(req, _make_exc()) is None + + def test_missing_model_on_request_returns_none(self): + """If the request has no ``.model`` at all (defensive).""" + assert _normalize(SimpleNamespace(), _make_exc()) is None + + def test_already_normalized_exception_passes_through(self): + """``ModelFallbackMiddleware`` wraps against the failing model + before re-raising. The outer chain's ``_normalize`` must NOT + double-wrap — otherwise attribution flips back to the original + request's model. + """ + req = _request(_openrouter_model()) + pre_wrapped = ProviderStreamError( + provider="moonshot", + class_qualname="openai.RateLimitError", + message="quota exceeded", + ) + 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 +# failures (surface) from shape / config failures (degrade) +# --------------------------------------------------------------------------- + + +class TestIsProviderError: + def test_openai_module_is_provider_error(self): + from EvoScientist.middleware.error_normalization import _is_provider_error + + assert _is_provider_error(_make_exc(__module__="openai")) + + def test_httpx_timeout_is_provider_error(self): + from EvoScientist.middleware.error_normalization import _is_provider_error + + assert _is_provider_error( + _make_exc(cls_name="TimeoutException", __module__="httpx") + ) + + def test_langchain_wrapper_module_is_provider_error(self): + from EvoScientist.middleware.error_normalization import _is_provider_error + + assert _is_provider_error( + _make_exc( + cls_name="BadRequestError", + __module__="langchain_openai.chat_models", + ) + ) + + def test_pydantic_validation_is_not_provider_error(self): + """Structured-output shape failures come from pydantic / + langchain, NOT from a provider SDK — the tool selector's + graceful-degrade path is right for these. + """ + from EvoScientist.middleware.error_normalization import _is_provider_error + + assert not _is_provider_error( + _make_exc(cls_name="ValidationError", __module__="pydantic") + ) + + def test_builtin_is_not_provider_error(self): + from EvoScientist.middleware.error_normalization import _is_provider_error + + assert not _is_provider_error(RuntimeError("x")) + + +# --------------------------------------------------------------------------- +# ProviderStreamError envelope +# --------------------------------------------------------------------------- + + +class TestProviderStreamErrorEnvelope: + def test_envelope_contains_required_fields(self): + err = ProviderStreamError( + provider="deepseek", + class_qualname="openai.RateLimitError", + message="quota exceeded", + status_code=429, + code="insufficient_quota", + ) + env = err.as_envelope() + assert env["error"] == "RateLimitError" + assert env["class"] == "openai.RateLimitError" + assert env["message"] == "quota exceeded" + assert env["provider"] == "deepseek" + assert env["status_code"] == 429 + assert env["code"] == "insufficient_quota" + + def test_envelope_omits_absent_optional_fields(self): + err = ProviderStreamError( + provider="openrouter", + class_qualname="openrouter.errors.foo.UnauthorizedResponseError", + message="User not found.", + ) + env = err.as_envelope() + assert "status_code" not in env + assert "code" not in env + assert "type" not in env + assert "request_id" not in env + + def test_provider_stream_error_is_not_a_dataclass(self): + """The whole point of the wrapper — must not be a dataclass so + orjson's OPT_SERIALIZE_DATACLASS fast-path doesn't fire. + """ + err = ProviderStreamError("x", "y.Z", "msg") + assert not dataclasses.is_dataclass(err) + assert not dataclasses.is_dataclass(type(err)) + + def test_model_dump_returns_envelope(self): + """Upstream ``serde.default`` calls ``model_dump()`` before its + exception branch — the hook that lets us skip the serde patch. + """ + err = ProviderStreamError( + provider="openrouter", + class_qualname="openrouter.errors.foo.UnauthorizedResponseError", + message="User not found.", + status_code=401, + ) + assert err.model_dump() == err.as_envelope() + + +# --------------------------------------------------------------------------- +# Middleware behavior +# --------------------------------------------------------------------------- + + +class TestMiddleware: + def _run_awrap(self, mw, request, handler): + async def _go(): + return await mw.awrap_model_call(request=request, handler=handler) + + return asyncio.run(_go()) + + def test_awrap_normalizes_provider_exception(self): + raised = _make_exc(cls_name="UnauthorizedResponseError", message="boom") + + async def handler(_req): + raise raised + + req = _request(_openrouter_model()) + mw = ErrorNormalizationMiddleware() + with pytest.raises(ProviderStreamError) as excinfo: + self._run_awrap(mw, req, handler) + assert excinfo.value.provider == "openrouter" + assert excinfo.value.__cause__ is raised + + def test_awrap_passes_through_non_provider_model_exception(self): + """If the model isn't a recognized provider SDK, the exception + passes through unwrapped — same as any non-model exception. + """ + raised = _make_exc(message="boom") + + async def handler(_req): + raise raised + + req = _request(_fake_model("some.other.pkg", "SomeModel")) + mw = ErrorNormalizationMiddleware() + with pytest.raises(Exception, match="boom") as excinfo: + self._run_awrap(mw, req, handler) + assert excinfo.value is raised + + def _langgraph_error_samples(self): + """Instances covering both branches of ``_should_pass_through``: + control-flow (``GraphBubbleUp`` + subclasses) and structural + errors. Constructor signatures vary — some need positional + args — so build each explicitly. + """ + from langgraph.errors import ( + EmptyInputError, + GraphBubbleUp, + GraphInterrupt, + InvalidUpdateError, + NodeTimeoutError, + TaskNotFound, + ) + + return [ + GraphBubbleUp(), + GraphInterrupt(), + InvalidUpdateError("bad update"), + EmptyInputError("no input"), + TaskNotFound(), + NodeTimeoutError("node-x", 1.5, kind="run", run_timeout=1.0), + ] + + def test_awrap_passes_through_langgraph_errors(self): + """Exceptions from ``langgraph.errors.*`` must propagate + untouched even when the model is a recognized provider — + they're either control-flow signals (interrupts, HITL) or + graph-level structural errors, neither is a provider incident. + """ + req = _request(_openrouter_model()) # recognized — would normally wrap + mw = ErrorNormalizationMiddleware() + + for raised in self._langgraph_error_samples(): + + async def handler(_req, _r=raised): + raise _r + + with pytest.raises(type(raised)) as excinfo: + self._run_awrap(mw, req, handler) + assert excinfo.value is raised, ( + f"{type(raised).__name__} got wrapped instead of propagated" + ) + + def test_awrap_passes_through_context_overflow_error(self): + """``ContextOverflowError`` is a cross-layer control signal: + deepagents' ``SummarizationMiddleware`` sits outside our stack + and catches it by type to compress history and retry. Wrapping + it here would change the type and break that self-healing + fallback — regressing to a user-visible ``ProviderStreamError`` + on any long conversation. + """ + from langchain_core.exceptions import ContextOverflowError + + raised = ContextOverflowError("context length exceeded") + + async def handler(_req): + raise raised + + req = _request(_openrouter_model()) # recognized — would normally wrap + mw = ErrorNormalizationMiddleware() + with pytest.raises(ContextOverflowError) as excinfo: + 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``. + Rationale: at the middleware boundary we can tell the model is + a provider, but not the exception's origin (SDK vs + langchain-wrapper vs httpx vs our code). Wrapping uniformly + gives the WebUI a consistent envelope; upstream's + ``RuntimeError``-whitelist would emit ``{"error": + "RuntimeError", "message": str(exc)}`` which isn't more + useful. + """ + raised = RuntimeError("internal glitch") + + async def handler(_req): + raise raised + + req = _request(_openai_model()) + mw = ErrorNormalizationMiddleware() + with pytest.raises(ProviderStreamError) as excinfo: + self._run_awrap(mw, req, handler) + assert excinfo.value.provider == "openai" + assert excinfo.value.__cause__ is raised + assert excinfo.value.class_qualname == "builtins.RuntimeError" + + def test_sync_wrap_normalizes_provider_exception(self): + raised = _make_exc(message="boom") + + def handler(_req): + raise raised + + req = _request(_openrouter_model()) + mw = ErrorNormalizationMiddleware() + with pytest.raises(ProviderStreamError) as excinfo: + mw.wrap_model_call(request=req, handler=handler) + assert excinfo.value.provider == "openrouter" + + def test_success_path_returns_handler_result(self): + async def handler(_req): + return "ok" + + req = _request(_openrouter_model()) + mw = ErrorNormalizationMiddleware() + assert self._run_awrap(mw, req, handler) == "ok" 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_langgraph_schema_generator_silenced.py b/tests/test_langgraph_schema_generator_silenced.py new file mode 100644 index 0000000..0b856c5 --- /dev/null +++ b/tests/test_langgraph_schema_generator_silenced.py @@ -0,0 +1,152 @@ +"""Regression tests for the langgraph_api SchemaGenerator silencing patch. + +Reproducer: mounting our ``/api/models`` custom Starlette app makes +langgraph_api call ``update_openapi_spec`` at startup, which iterates +EVERY route (ours + upstream's). Endpoints whose docstrings aren't +valid YAML hit a warning + traceback in the deploy log — purely noise, +since the existing fallback path already produces a usable schema +entry. The patch keeps the fallback shape but silences the log spam. +""" + +from __future__ import annotations + +import os + +# ``langgraph_api.config`` reads several required env vars at import +# time via starlette's ``Config(...)`` helper. We don't actually use the +# DB or Redis here — any non-empty string keeps the loader happy. +os.environ.setdefault("DATABASE_URI", "sqlite:///:memory:") +os.environ.setdefault("REDIS_URI", "redis://localhost:6379") + +# Importing patches.py applies the eager module-level monkey-patch. +import langgraph_api.utils as _lgapi_utils + +import EvoScientist.llm.patches as _patches + +# Re-invoke the patch after env vars are set. Required because earlier test +# modules (e.g. test_llm.py) import patches.py *without* DATABASE_URI/ +# REDIS_URI, which makes ``langgraph_api.utils`` fail to import inside the +# patch's bare ``except``; the loader swallows it and the flag stays False +# forever (Python won't re-run module-level code on subsequent imports). +# The patch function is idempotent (early-return on the flag), so calling +# it here is a no-op when the patch already landed and a successful retry +# when the prior import failed. +_patches._patch_langgraph_schema_generator_silence_warnings() + + +class _FakeEndpoint: + """Minimal Starlette-like endpoint info for the schema generator.""" + + def __init__(self, path: str, method: str, func): + self.path = path + self.http_method = method + self.func = func + + +class _DocstringFixture: + """The kinds of docstrings the patched generator must handle.""" + + @staticmethod + def prose_with_colon(): + """Endpoint summary. + + Query params: + id: The thing you want. + """ + + @staticmethod + def valid_yaml(): + """ + summary: A valid YAML docstring. + description: Stays structured. + """ + + @staticmethod + def no_docstring(): + pass + + +def _generator(): + return _lgapi_utils.SchemaGenerator( + {"openapi": "3.1.0", "info": {"title": "test", "version": "0"}} + ) + + +def test_prose_docstring_no_longer_logs_warning(): + """The patched ``parse_docstring`` must silence upstream's structlog + WARNING when ``yaml.safe_load`` fails on a prose docstring. + + Inverts the patch first to prove the fixture actually trips + ``yaml.safe_load`` — without this baseline assertion the test would + pass vacuously if the fixture stopped triggering the failure path + (e.g. if upstream changed how docstrings are pre-processed). + """ + from structlog.testing import capture_logs + + gen = _generator() + endpoint = _FakeEndpoint("/x", "get", _DocstringFixture.prose_with_colon) + gen.get_endpoints = lambda _routes: [endpoint] + + patched_parse = _lgapi_utils.SchemaGenerator.parse_docstring + # Phase 1: baseline. Drop the subclass override so MRO falls through + # to Starlette's BaseSchemaGenerator.parse_docstring, which is what + # production hits before our patch installs. + del _lgapi_utils.SchemaGenerator.parse_docstring + try: + with capture_logs() as baseline_records: + gen.get_schema([]) + finally: + _lgapi_utils.SchemaGenerator.parse_docstring = patched_parse + + baseline_warnings = [r for r in baseline_records if r.get("log_level") == "warning"] + assert any( + "Unable to parse docstring" in r.get("event", "") for r in baseline_warnings + ), "fixture no longer trips parse_docstring — test would pass vacuously" + + # Phase 2: with the patch reinstated, the same call must emit no + # warning records. + with capture_logs() as patched_records: + schema = gen.get_schema([]) + + assert [r for r in patched_records if r.get("log_level") == "warning"] == [] + + # Schema still has the fallback shape — fixture's prose becomes the + # description verbatim (with leading/trailing whitespace from the + # docstring preserved by upstream's fallback path). + entry = schema["paths"]["/x"]["get"] + assert "description" in entry + assert "Query params" in entry["description"] + + +def test_valid_yaml_docstring_keeps_structured_parse(): + """Endpoints with parseable YAML keep their structured metadata — + we only changed the failure branch, not the success path. + """ + gen = _generator() + endpoint = _FakeEndpoint("/y", "get", _DocstringFixture.valid_yaml) + gen.get_endpoints = lambda _routes: [endpoint] + schema = gen.get_schema([]) + entry = schema["paths"]["/y"]["get"] + assert entry.get("summary") == "A valid YAML docstring." + assert entry.get("description") == "Stays structured." + + +def test_no_docstring_still_handled(): + """Endpoints with ``__doc__ = None`` must not raise — fallback uses + empty string for ``description``. + """ + gen = _generator() + endpoint = _FakeEndpoint("/z", "get", _DocstringFixture.no_docstring) + gen.get_endpoints = lambda _routes: [endpoint] + schema = gen.get_schema([]) + entry = schema["paths"]["/z"]["get"] + # Either description="" (fallback path) or structured (if YAML parse + # of None happens to succeed somehow — implementation detail). + # The contract is just "no exception, entry exists". + assert isinstance(entry, dict) + + +def test_patch_flag_set(): + from EvoScientist.llm.patches import _langgraph_schema_silenced_patched + + assert _langgraph_schema_silenced_patched is True diff --git a/tests/test_llm.py b/tests/test_llm.py index 23168d0..c4ec1a9 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -1,5 +1,6 @@ """Tests for EvoScientist LLM module.""" +from types import SimpleNamespace from unittest.mock import patch import pytest @@ -160,6 +161,68 @@ class TestGetModelInfo: class TestGetChatModel: + @patch("EvoScientist.llm.models.init_chat_model") + def test_uses_host_model_resolver(self, mock_init): + """An embedding host can provide model routing without a core dependency.""" + from EvoScientist.runtime_integrations import ( + configure_runtime_integrations, + reset_runtime_integrations, + ) + + mock_init.return_value = "mock_model" + resolved = SimpleNamespace( + provider_name="relay-a", + model_id="model-a", + protocol="openai", + api_key="sk-host", + base_url="https://relay.example/v1/", + params={"max_tokens": 8192, "_default_headers": {"X-Relay": "a"}}, + supports_reasoning=False, + ) + configure_runtime_integrations(model_resolver=lambda model, provider: resolved) + try: + assert get_chat_model("alias-a") == "mock_model" + finally: + reset_runtime_integrations() + + call_kwargs = mock_init.call_args.kwargs + assert call_kwargs["model"] == "model-a" + assert call_kwargs["model_provider"] == "openai" + assert call_kwargs["api_key"] == "sk-host" + assert call_kwargs["base_url"] == "https://relay.example/v1" + assert call_kwargs["max_tokens"] == 8192 + assert call_kwargs["default_headers"] == {"X-Relay": "a"} + assert "reasoning" not in call_kwargs + + @patch("EvoScientist.llm.models._patch_openai_compat_content") + @patch("EvoScientist.llm.models.init_chat_model") + def test_host_openai_provider_with_custom_base_uses_compat_patch( + self, mock_init, mock_compat + ): + from EvoScientist.runtime_integrations import ( + configure_runtime_integrations, + reset_runtime_integrations, + ) + + model_instance = object() + mock_init.return_value = model_instance + resolved = SimpleNamespace( + provider_name="openai", + model_id="gpt-5.5", + protocol="openai", + api_key="sk-host", + base_url="https://relay.example/v1", + params={}, + supports_reasoning=True, + ) + configure_runtime_integrations(model_resolver=lambda model, provider: resolved) + try: + get_chat_model("gpt-5.5", provider="openai") + finally: + reset_runtime_integrations() + + mock_compat.assert_called_once_with(model_instance, hoist_tool_media=True) + @patch("EvoScientist.llm.models.init_chat_model") def test_uses_default_model_when_none(self, mock_init): """Test that get_chat_model uses default model when model=None.""" @@ -229,6 +292,39 @@ class TestGetChatModel: assert call_kwargs["temperature"] == 0.7 assert call_kwargs["max_tokens"] == 1000 + @patch("EvoScientist.llm.models.init_chat_model") + def test_drops_unsupported_legacy_model_kwargs(self, mock_init): + mock_init.return_value = "mock_model" + + get_chat_model( + "gpt-5-nano", + provider="openai", + sanitize_openai_sdk_headers=True, + model_kwargs={"sanitize_openai_sdk_headers": False, "custom": "value"}, + ) + + call_kwargs = mock_init.call_args.kwargs + assert "sanitize_openai_sdk_headers" not in call_kwargs + assert call_kwargs["model_kwargs"] == {"custom": "value"} + + @patch("EvoScientist.llm.models.init_chat_model") + def test_explicit_credentials_override_environment(self, mock_init, monkeypatch): + """Host-provided credentials take precedence over process defaults.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENAI_API_KEY", "sk-environment") + monkeypatch.setenv("OPENAI_BASE_URL", "https://environment.example/v1") + + get_chat_model( + "gpt-5-nano", + provider="openai", + api_key="sk-explicit", + base_url="https://explicit.example/v1", + ) + + call_kwargs = mock_init.call_args.kwargs + assert call_kwargs["api_key"] == "sk-explicit" + assert call_kwargs["base_url"] == "https://explicit.example/v1" + @patch("EvoScientist.llm.models.init_chat_model") def test_infers_openai_from_gpt_prefix(self, mock_init): """Test that OpenAI is inferred from gpt- prefix.""" @@ -439,6 +535,200 @@ class TestThirdPartyRouting: call_kwargs = mock_init.call_args[1] assert call_kwargs["reasoning"] == {"effort": "medium", "summary": "auto"} + # --- OpenRouter app attribution (issue #339) --- + + _APP_ATTR_ENV = ( + "EVOSCIENTIST_OPENROUTER_HTTP_REFERER", + "EVOSCIENTIST_OPENROUTER_APP_TITLE", + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", + ) + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openrouter_app_attribution_defaults(self, mock_init, monkeypatch): + """OpenRouter init should carry EvoScientist's default app attribution.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") + # Isolate from any leaked env overrides so we assert the built-in defaults. + for _env in self._APP_ATTR_ENV: + monkeypatch.delenv(_env, raising=False) + + get_chat_model("x-ai/grok-4.3", provider="openrouter") + + call_kwargs = mock_init.call_args[1] + assert call_kwargs["app_url"] == "https://github.com/EvoScientist/EvoScientist" + assert call_kwargs["app_title"] == "EvoScientist" + # Must be a list[str] (not the comma string) — langchain-openrouter joins it. + assert call_kwargs["app_categories"] == ["creative-writing", "personal-agent"] + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openrouter_app_attribution_from_env(self, mock_init, monkeypatch): + """Env vars should override the default app attribution values.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "https://acme.test") + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_TITLE", "Acme") + # Include a space to prove each category is stripped. + monkeypatch.setenv( + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "cli-agent, programming-app" + ) + + get_chat_model("x-ai/grok-4.3", provider="openrouter") + + call_kwargs = mock_init.call_args[1] + assert call_kwargs["app_url"] == "https://acme.test" + assert call_kwargs["app_title"] == "Acme" + assert call_kwargs["app_categories"] == ["cli-agent", "programming-app"] + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openrouter_app_attribution_user_override_not_clobbered( + self, mock_init, monkeypatch + ): + """Caller-supplied attribution kwargs must beat both env and defaults.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") + # Env is also set, to prove an explicit kwarg outranks the env override + # (not just the built-in default). + monkeypatch.setenv( + "EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "https://env.example" + ) + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_TITLE", "EnvTitle") + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "env-cat") + + get_chat_model( + "x-ai/grok-4.3", + provider="openrouter", + app_url="https://mine.example", + app_title="MyApp", + app_categories=["only-this"], + ) + + call_kwargs = mock_init.call_args[1] + assert call_kwargs["app_url"] == "https://mine.example" + assert call_kwargs["app_title"] == "MyApp" + # An explicit list is preserved verbatim, not re-split. + assert call_kwargs["app_categories"] == ["only-this"] + + @patch("EvoScientist.llm.models.init_chat_model") + def test_non_openrouter_providers_get_no_app_attribution( + self, mock_init, monkeypatch + ): + """Only the openrouter provider should receive app-attribution kwargs.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-real") + monkeypatch.setenv("OLLAMA_BASE_URL", "http://localhost:11434") + + for model, provider in ( + ("claude-sonnet-4-6", "anthropic"), + ("llama3.1:8b", "ollama"), + ): + get_chat_model(model, provider=provider) + call_kwargs = mock_init.call_args[1] + assert "app_url" not in call_kwargs + assert "app_title" not in call_kwargs + assert "app_categories" not in call_kwargs + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openrouter_app_attribution_coexists_with_reasoning_and_cache( + self, mock_init, monkeypatch + ): + """Attribution must not disturb reasoning or Anthropic prompt caching.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") + monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False) + monkeypatch.delenv( + "EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE", raising=False + ) + for _env in self._APP_ATTR_ENV: + monkeypatch.delenv(_env, raising=False) + + get_chat_model("claude-sonnet-4.6", provider="openrouter") + + call_kwargs = mock_init.call_args[1] + # Existing behavior intact. + assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"} + assert call_kwargs["model_kwargs"]["cache_control"] == {"type": "ephemeral"} + # Attribution added alongside. + assert call_kwargs["app_url"] == "https://github.com/EvoScientist/EvoScientist" + assert call_kwargs["app_title"] == "EvoScientist" + assert call_kwargs["app_categories"] == ["creative-writing", "personal-agent"] + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openrouter_app_categories_env_strips_blank_items( + self, mock_init, monkeypatch + ): + """A messy comma value (stray commas / spaces) yields a clean list.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "a,, b ") + + get_chat_model("x-ai/grok-4.3", provider="openrouter") + + assert mock_init.call_args[1]["app_categories"] == ["a", "b"] + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openrouter_app_categories_capped_to_per_request_limit( + self, mock_init, monkeypatch + ): + """Over-configuring categories caps to the first N and warns the user.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") + monkeypatch.setenv( + "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", + "cli-agent,programming-app,personal-agent,writing-assistant", + ) + + with pytest.warns(UserWarning, match="at most 2 app categories"): + get_chat_model("x-ai/grok-4.3", provider="openrouter") + + # OpenRouter honors at most 2 per request, so only the first 2 are sent. + assert mock_init.call_args[1]["app_categories"] == [ + "cli-agent", + "programming-app", + ] + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openrouter_app_categories_all_separators_omit_kwarg( + self, mock_init, monkeypatch + ): + """A categories value with no real items omits the kwarg entirely.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") + monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", " , , ") + + get_chat_model("x-ai/grok-4.3", provider="openrouter") + + # No app_categories kwarg at all — not an empty list (which the library + # would reject / send as an empty header). + assert "app_categories" not in mock_init.call_args[1] + + def test_openrouter_app_attribution_lands_on_real_model(self, monkeypatch): + """Build a REAL ChatOpenRouter (no mock) and assert the attribution + values land on the instance rather than being silently dumped into + model_kwargs. + + The mocked tests above assert on the kwargs handed to init_chat_model, + so they cannot catch a param-name typo or a langchain-openrouter version + that accepts these only as passthrough model params (which the library + does with a warning, not an error). This test is the guard for both. + """ + from langchain_openrouter import ChatOpenRouter + + monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") + for _env in self._APP_ATTR_ENV: + monkeypatch.delenv(_env, raising=False) + + model = get_chat_model("x-ai/grok-4.3", provider="openrouter") + + assert isinstance(model, ChatOpenRouter) + assert model.app_url == "https://github.com/EvoScientist/EvoScientist" + assert model.app_title == "EvoScientist" + assert model.app_categories == ["creative-writing", "personal-agent"] + # Not silently swallowed into model_kwargs (the passthrough failure mode). + model_kwargs = model.model_kwargs or {} + assert "app_url" not in model_kwargs + assert "app_title" not in model_kwargs + assert "app_categories" not in model_kwargs + @patch("EvoScientist.llm.models.init_chat_model") def test_openrouter_anthropic_prompt_cache_enabled_by_default( self, mock_init, monkeypatch @@ -957,6 +1247,153 @@ class TestPatchOpenAICompatContent: model._astream = AsyncMock() return model + def test_missing_tool_call_ids_are_repaired_without_mutating_history(self): + from langchain_core.messages import AIMessage, ToolMessage + + from EvoScientist.llm.patches import _ensure_openai_tool_call_ids + + ai = AIMessage( + content=[{"type": "tool_call", "id": "", "name": "execute", "args": {}}], + tool_calls=[{"id": "", "name": "execute", "args": {}}], + ) + tool = ToolMessage(content="ok", tool_call_id="") + + normalized = _ensure_openai_tool_call_ids([ai, tool]) + + call_id = normalized[0].tool_calls[0]["id"] + assert call_id.startswith("call_") + assert normalized[0].content[0]["id"] == call_id + assert normalized[1].tool_call_id == call_id + assert ai.tool_calls[0]["id"] == "" + assert tool.tool_call_id == "" + + def test_missing_parallel_tool_call_ids_are_stable_and_ordered(self): + from langchain_core.messages import AIMessage, ToolMessage + + from EvoScientist.llm.patches import _ensure_openai_tool_call_ids + + messages = [ + AIMessage( + id="assistant-1", + content="", + tool_calls=[ + {"id": "", "name": "read_file", "args": {}}, + {"id": "", "name": "execute", "args": {}}, + ], + ), + ToolMessage(content="file", tool_call_id=""), + ToolMessage(content="command", tool_call_id=""), + ] + + first = _ensure_openai_tool_call_ids(messages) + second = _ensure_openai_tool_call_ids(messages) + call_ids = [call["id"] for call in first[0].tool_calls] + + assert call_ids == [call["id"] for call in second[0].tool_calls] + 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 @@ -2238,6 +2675,38 @@ class TestPatchOpenrouterStripResponsesReasoning: class TestAutoConfig: + @pytest.fixture(autouse=True) + def _clear_reasoning_effort_env(self, monkeypatch): + """Keep auto-config tests independent of the developer environment.""" + monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False) + + @patch("EvoScientist.llm.models.init_chat_model") + def test_internal_sentinels_disable_auto_reasoning(self, mock_init, monkeypatch): + """Internal callers can disable reasoning without leaking sentinels.""" + mock_init.return_value = "mock_model" + monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + + for model, provider in ( + ("claude-sonnet-4-6", "anthropic"), + ("gpt-5-nano", "openai"), + ("gemini-2.5-flash", "google-genai"), + ("llama3.1:8b", "ollama"), + ): + mock_init.reset_mock() + get_chat_model( + model, + provider=provider, + _disable_reasoning=True, + _disable_thinking=True, + ) + call_kwargs = mock_init.call_args.kwargs + assert "_disable_reasoning" not in call_kwargs + assert "_disable_thinking" not in call_kwargs + assert "reasoning" not in call_kwargs + assert "thinking" not in call_kwargs + assert "include_thoughts" not in call_kwargs + @patch("EvoScientist.llm.models.init_chat_model") def test_anthropic_4_5_thinking(self, mock_init, monkeypatch): """Anthropic 4-5 models get enabled thinking with budget.""" @@ -2327,6 +2796,7 @@ class TestAutoConfig: """gpt-5.4+ and codex models get xhigh reasoning.""" mock_init.return_value = "mock_model" monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False) get_chat_model("gpt-5.4", provider="openai") assert mock_init.call_args[1]["reasoning"] == { @@ -2346,6 +2816,26 @@ class TestAutoConfig: "summary": "auto", } + get_chat_model("gpt-5.6-sol", provider="openai") + assert mock_init.call_args[1]["reasoning"] == { + "effort": "xhigh", + "summary": "auto", + } + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openai_reasoning_effort_from_env(self, mock_init, monkeypatch): + """Native OpenAI reasoning effort should be configurable via env var.""" + mock_init.return_value = "mock_model" + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.setenv("EVOSCIENTIST_REASONING_EFFORT", "high") + + get_chat_model("gpt-5.5", provider="openai") + + assert mock_init.call_args[1]["reasoning"] == { + "effort": "high", + "summary": "auto", + } + @patch("EvoScientist.llm.models.init_chat_model") def test_openai_reasoning_high_fallback(self, mock_init, monkeypatch): """Other OpenAI models get high reasoning effort.""" @@ -2377,8 +2867,8 @@ class TestAutoConfig: assert call_kwargs["model_provider"] == "openai" assert call_kwargs["base_url"] == "http://127.0.0.1:8000/codex/v1" assert call_kwargs["api_key"] == "ccproxy-oauth" - # Proxy mode: reasoning skipped (ccproxy untested) - assert "reasoning" not in call_kwargs + # ccproxy uses the Responses API, so reasoning configuration is valid. + assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"} # Proxy mode: Responses API (bypasses format chain), streaming ON assert call_kwargs["use_responses_api"] is True assert "streaming" not in call_kwargs @@ -2411,6 +2901,120 @@ class TestAutoConfig: call_kwargs = mock_init.call_args[1] assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"} assert "use_responses_api" not in call_kwargs + assert "default_headers" not in call_kwargs + + @patch( + "EvoScientist.llm.models._installed_codex_client_version", + return_value="0.144.1", + ) + @patch("EvoScientist.llm.models.init_chat_model") + def test_openai_ccproxy_codex_client_headers( + self, mock_init, mock_installed_version, monkeypatch + ): + """ccproxy Codex mode sends Codex-CLI-shaped client headers.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1") + monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth") + monkeypatch.delenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", raising=False) + + get_chat_model("gpt-5.5", provider="openai") + + headers = mock_init.call_args[1]["default_headers"] + assert headers["originator"] == "codex_cli_rs" + assert headers["version"] == "0.144.1" + assert headers["User-Agent"].startswith("codex_cli_rs/0.144.1") + mock_installed_version.assert_called_once_with() + assert mock_init.call_args[1]["reasoning"]["effort"] == "xhigh" + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openai_ccproxy_codex_client_version_env(self, mock_init, monkeypatch): + """EVOSCIENTIST_CODEX_CLIENT_VERSION overrides the pinned version.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1") + monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth") + monkeypatch.setenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", "9.9.9") + + get_chat_model("gpt-5.5", provider="openai") + + headers = mock_init.call_args[1]["default_headers"] + assert headers["version"] == "9.9.9" + assert headers["User-Agent"].startswith("codex_cli_rs/9.9.9") + + @patch("EvoScientist.llm.models.subprocess.run") + def test_installed_codex_client_version(self, mock_run): + """The advertised version follows the installed Codex CLI.""" + from EvoScientist.llm.models import _installed_codex_client_version + + mock_run.return_value.returncode = 0 + mock_run.return_value.stdout = "codex-cli 0.144.1\n" + mock_run.return_value.stderr = "" + _installed_codex_client_version.cache_clear() + try: + assert _installed_codex_client_version() == "0.144.1" + assert _installed_codex_client_version() == "0.144.1" + finally: + _installed_codex_client_version.cache_clear() + mock_run.assert_called_once_with( + ["codex", "--version"], + capture_output=True, + text=True, + timeout=2, + check=False, + ) + + @patch( + "EvoScientist.llm.models._installed_codex_client_version", + return_value="0.140.0", + ) + def test_older_installed_codex_uses_fallback( + self, mock_installed_version, monkeypatch + ): + """An outdated installed CLI must not undercut the safe fallback.""" + from EvoScientist.llm.models import ( + _CODEX_CLIENT_VERSION_FALLBACK, + _resolve_codex_client_version, + ) + + monkeypatch.delenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", raising=False) + + assert _resolve_codex_client_version() == _CODEX_CLIENT_VERSION_FALLBACK + mock_installed_version.assert_called_once_with() + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openai_ccproxy_codex_headers_respect_caller(self, mock_init, monkeypatch): + """Caller-supplied default_headers keys are not overridden.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1") + monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth") + + get_chat_model( + "gpt-5.5", + provider="openai", + default_headers={"originator": "codex_vscode", "version": "9.9.9"}, + ) + + headers = mock_init.call_args[1]["default_headers"] + assert headers["originator"] == "codex_vscode" + assert headers["version"] == "9.9.9" + assert headers["User-Agent"].startswith("codex_cli_rs/9.9.9") + + @patch("EvoScientist.llm.models.init_chat_model") + def test_openai_ccproxy_codex_none_headers(self, mock_init, monkeypatch): + """An explicit default_headers=None is normalized before gap-filling.""" + mock_init.return_value = "mock_model" + monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1") + monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth") + monkeypatch.setenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", "9.9.9") + + get_chat_model( + "gpt-5.5", + provider="openai", + default_headers=None, + ) + + headers = mock_init.call_args[1]["default_headers"] + assert headers["originator"] == "codex_cli_rs" + assert headers["version"] == "9.9.9" @patch("EvoScientist.llm.models.init_chat_model") def test_openai_ccproxy_key_but_wrong_path_not_ccproxy( diff --git a/tests/test_logging_config.py b/tests/test_logging_config.py new file mode 100644 index 0000000..eda8997 --- /dev/null +++ b/tests/test_logging_config.py @@ -0,0 +1,129 @@ +import logging +from datetime import UTC, datetime +from io import StringIO + +from EvoScientist.logging_config import ( + DailyLogFileHandler, + configure_console_logging, + configure_daily_file_logging, + configure_logging, + default_log_dir, + resolve_log_level, +) + + +def test_daily_log_file_handler_uses_dated_active_file(tmp_path): + handler = DailyLogFileHandler(tmp_path, prefix="evoscientist", retention_days=30) + logger = logging.getLogger("tests.daily_log_file_handler") + logger.handlers.clear() + logger.propagate = False + logger.setLevel(logging.INFO) + logger.addHandler(handler) + + logger.info("hello") + handler.close() + + today = datetime.now().strftime("%Y-%m-%d") + assert (tmp_path / f"evoscientist-{today}.log").read_text(encoding="utf-8").strip() + + +def test_daily_log_file_handler_keeps_latest_retention_days(tmp_path): + for day in range(1, 33): + (tmp_path / f"evoscientist-2026-01-{day:02d}.log").write_text( + "x", encoding="utf-8" + ) + + handler = DailyLogFileHandler(tmp_path, prefix="evoscientist", retention_days=30) + handler._delete_expired_logs() + handler.close() + + remaining = sorted(path.name for path in tmp_path.glob("evoscientist-*.log")) + assert len(remaining) == 30 + assert remaining[0] == "evoscientist-2026-01-03.log" + + +def test_configure_daily_file_logging_replaces_matching_handler(tmp_path): + logger = logging.getLogger("tests.configure_daily_file_logging") + logger.handlers.clear() + logger.propagate = False + + first = configure_daily_file_logging(logger, log_dir=tmp_path) + second = configure_daily_file_logging(logger, log_dir=tmp_path) + + try: + handlers = [h for h in logger.handlers if isinstance(h, DailyLogFileHandler)] + assert handlers == [second] + assert first.stream is None + finally: + for handler in logger.handlers[:]: + logger.removeHandler(handler) + handler.close() + + +def test_configure_logging_replaces_only_managed_handlers(tmp_path): + logger = logging.getLogger("tests.configure_logging") + logger.handlers.clear() + logger.propagate = False + external = logging.NullHandler() + logger.addHandler(external) + + configure_logging(logger, log_dir=tmp_path, level="debug") + configure_logging(logger, log_dir=tmp_path, level="info") + + try: + daily_handlers = [h for h in logger.handlers if isinstance(h, DailyLogFileHandler)] + stream_handlers = [ + h + for h in logger.handlers + if isinstance(h, logging.StreamHandler) + and not isinstance(h, DailyLogFileHandler) + ] + assert external in logger.handlers + assert len(daily_handlers) == 1 + assert len(stream_handlers) == 1 + assert logger.level == logging.INFO + finally: + for handler in logger.handlers[:]: + logger.removeHandler(handler) + handler.close() + + +def test_configure_console_logging_emits_to_stream(): + logger = logging.getLogger("tests.configure_console_logging") + logger.handlers.clear() + logger.propagate = False + stream = StringIO() + + configure_console_logging(logger, level="INFO", stream=stream) + try: + logger.info("hello") + assert "tests.configure_console_logging: hello" in stream.getvalue() + finally: + for handler in logger.handlers[:]: + logger.removeHandler(handler) + handler.close() + + +def test_resolve_log_level_accepts_alias_numeric_and_fallback(): + assert resolve_log_level("warn") == logging.WARNING + assert resolve_log_level("10") == logging.DEBUG + assert resolve_log_level("", default=logging.ERROR) == logging.ERROR + assert resolve_log_level("not-a-level", default=logging.CRITICAL) == logging.CRITICAL + + +def test_daily_log_file_handler_supports_utc(tmp_path): + handler = DailyLogFileHandler(tmp_path, utc=True) + try: + today_utc = datetime.now(UTC).strftime("%Y-%m-%d") + assert handler.active_log_path.name == f"evoscientist-{today_utc}.log" + finally: + handler.close() + + +def test_default_log_dir_uses_current_data_dir(monkeypatch, tmp_path): + import EvoScientist.paths as paths + + monkeypatch.delenv("EVOSCIENTIST_LOG_DIR", raising=False) + monkeypatch.setattr(paths, "DATA_DIR", tmp_path / "data") + + assert default_log_dir() == tmp_path / "data" / "logs" diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 5e1bb62..4302372 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -1421,22 +1421,35 @@ class TestLoadToolsProgressCallback: ] async def test_failure_emits_start_then_error_with_detail(self, monkeypatch): - from EvoScientist.mcp.client import _load_tools + from EvoScientist.mcp import client as mcp_client events: list[tuple[str, str, str]] = [] self._patch_client(monkeypatch, {"srv": RuntimeError("boom")}) + monkeypatch.setattr(mcp_client, "_MCP_SERVER_ERRORS", {}) config = {"srv": {"transport": "stdio", "command": "demo"}} def record(event, name, detail): events.append((event, name, detail)) - await _load_tools(config, on_progress=record) + await mcp_client._load_tools(config, on_progress=record) assert events == [ ("start", "srv", ""), ("error", "srv", "boom"), ] + assert mcp_client.get_mcp_server_errors() == {"srv": "boom"} + + async def test_success_clears_previous_server_error(self, monkeypatch): + from EvoScientist.mcp import client as mcp_client + + self._patch_client(monkeypatch, {"srv": ["tool1"]}) + monkeypatch.setattr(mcp_client, "_MCP_SERVER_ERRORS", {"srv": "old error"}) + + config = {"srv": {"transport": "stdio", "command": "demo"}} + await mcp_client._load_tools(config) + + assert mcp_client.get_mcp_server_errors() == {} async def test_mixed_fleet_reports_each_server_independently(self, monkeypatch): from EvoScientist.mcp.client import _load_tools diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py index fa9e346..36ea9d6 100644 --- a/tests/test_model_fallback.py +++ b/tests/test_model_fallback.py @@ -6,6 +6,7 @@ fallback chain behaviour via _try_fallbacks / _guard_and_fallback. from __future__ import annotations +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -83,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): @@ -223,6 +225,89 @@ class TestTryFallbacks: # fb-b should never be reached. assert mock_gcm.call_count == 1 + async def test_exhausted_fallbacks_attribute_to_last_failing_model(self): + """Regression: when every fallback fails, the raised + ``ProviderStreamError`` must be attributed to the model that + ACTUALLY failed last, not the original ``request.model``. + Prevents a ``deepseek → moonshot`` chain from surfacing as + ``provider: deepseek`` after moonshot exhausts its quota. + """ + from EvoScientist.llm.errors import ProviderStreamError + + add_fallback("moonshot-model", "moonshot") + # Original request's model is openai-shape. Fallback's model + # will be openai-shape with a moonshot base_url. + req = _fake_request() + + # ChatOpenAI-shape model instance so ``_provider_from_model`` + # returns a recognized provider. + def _make_openai_model(base_url=None): + cls = type( + "ChatOpenAI", + (), + {"__module__": "langchain_openai.chat_models.base"}, + ) + inst = cls() + inst.openai_api_base = base_url + return inst + + req.model = _make_openai_model() # primary + fallback_model = _make_openai_model(base_url="https://api.moonshot.cn/v1") + # ``request.override(model=...)`` must return the request with the + # new model so ``_try_fallbacks`` tracks the failing model. + req.override = MagicMock( + side_effect=lambda **kw: SimpleNamespace(model=kw.get("model", req.model)) + ) + + async def _invoke(_r): + raise Exception("429 quota exceeded") + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = fallback_model + with pytest.raises(ProviderStreamError) as exc_info: + await _try_fallbacks(req, _invoke, Exception("openai primary failed")) + + # Attribution flipped to moonshot (the failing fallback), not + # openai (the original request's model). + assert exc_info.value.provider == "moonshot" + assert "quota exceeded" in exc_info.value.message + + async def test_langgraph_error_at_fallback_raise_point_passes_through(self): + """Regression: ``_raise_normalized`` calls ``_normalize`` + directly, so its ``_should_pass_through`` gate must fire even + without the ``ErrorNormalizationMiddleware`` wrap sites' own + check. Prevents a ``langgraph.errors.*`` exception hitting the + fallback chain from being wrapped as a provider incident. + """ + from langgraph.errors import InvalidUpdateError + + add_fallback("fb-a", "prov-a") + req = _fake_request() + + # Use a recognized-provider model so ``_provider_from_model`` + # wouldn't short-circuit — the guard has to come from + # ``_should_pass_through``, not the provider check. + cls = type( + "ChatOpenAI", (), {"__module__": "langchain_openai.chat_models.base"} + ) + model = cls() + model.openai_api_base = None + req.model = model + req.override = MagicMock( + side_effect=lambda **kw: SimpleNamespace(model=kw.get("model", req.model)) + ) + + raised = InvalidUpdateError("state mismatch") + + async def _invoke(_r): + raise raised + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = model + with pytest.raises(InvalidUpdateError) as exc_info: + await _try_fallbacks(req, _invoke, Exception("primary failed")) + assert exc_info.value is raised + # ═════════════════════════════════════════════════════════════════ # 3. _guard_and_fallback — pre-check before chain walk @@ -242,6 +327,34 @@ class TestGuardAndFallback: invoke.assert_not_awaited() + async def test_context_overflow_with_provider_model_passes_through_unwrapped(self): + """Regression: a ``ContextOverflowError`` entering + ``_guard_and_fallback`` under a recognized-provider model must + come out unwrapped. Otherwise ``_raise_normalized`` → + ``_normalize`` would wrap it as a ``ProviderStreamError`` and + deepagents' ``SummarizationMiddleware`` (which sits outside + the user middleware stack and catches by exact type) would + stop compressing history and retrying. + """ + add_fallback("fb", "prov") + req = _fake_request() + # Recognized provider — without the gate in ``_normalize`` this + # would wrap. With the gate, the raw type propagates. + cls = type( + "ChatOpenAI", (), {"__module__": "langchain_openai.chat_models.base"} + ) + model = cls() + model.openai_api_base = None + req.model = model + invoke = AsyncMock() + + raised = ContextOverflowError("context length exceeded") + with pytest.raises(ContextOverflowError) as exc_info: + await _guard_and_fallback(raised, req, invoke) + + assert exc_info.value is raised + invoke.assert_not_awaited() + async def test_malformed_400_raises_immediately(self): add_fallback("fb", "prov") req = _fake_request() diff --git a/tests/test_observation_memory.py b/tests/test_observation_memory.py index 5989034..3db0d79 100644 --- a/tests/test_observation_memory.py +++ b/tests/test_observation_memory.py @@ -2294,7 +2294,11 @@ def test_memory_worker_observation_writer_modes( observation_writer=observation_writer, ) - assert type(middleware[0]).__name__ == "ToolErrorHandlerMiddleware" + # ErrorNormalizationMiddleware wraps outermost so provider-SDK + # exceptions from the auxiliary model call get normalized before + # the tool-error handler sees them. + assert type(middleware[0]).__name__ == "ErrorNormalizationMiddleware" + assert type(middleware[1]).__name__ == "ToolErrorHandlerMiddleware" assert _memory_tool_names(middleware) == expected_tools diff --git a/tests/test_paths.py b/tests/test_paths.py index 1b9d8c4..c29e94d 100644 --- a/tests/test_paths.py +++ b/tests/test_paths.py @@ -21,6 +21,7 @@ def _restore_paths(): "GLOBAL_MEMORIES_DIR": paths.GLOBAL_MEMORIES_DIR, "USER_SKILLS_DIR": paths.USER_SKILLS_DIR, "_active_workspace": paths._active_workspace, + "_EVOSCIENTIST_DATA_ROOT": paths._EVOSCIENTIST_DATA_ROOT, } yield paths.WORKSPACE_ROOT = orig["WORKSPACE_ROOT"] @@ -32,6 +33,7 @@ def _restore_paths(): paths.GLOBAL_MEMORIES_DIR = orig["GLOBAL_MEMORIES_DIR"] paths.USER_SKILLS_DIR = orig["USER_SKILLS_DIR"] paths._active_workspace = orig["_active_workspace"] + paths._EVOSCIENTIST_DATA_ROOT = orig["_EVOSCIENTIST_DATA_ROOT"] class TestSetWorkspaceRoot: @@ -140,6 +142,63 @@ class TestDataDir: assert paths.GLOBAL_MEMORIES_DIR == paths.DATA_DIR / "memories" +class TestGatewayDataDirs: + def test_evoscientist_root_prefers_home_override(self, tmp_path, monkeypatch): + home = tmp_path / "runtime-home" + monkeypatch.setenv("EVOSCIENTIST_HOME", str(home)) + + assert paths.evoscientist_root() == home.resolve() + + def test_evoscientist_root_falls_back_to_data_dir(self, tmp_path, monkeypatch): + data_dir = tmp_path / "app-data" + monkeypatch.delenv("EVOSCIENTIST_HOME", raising=False) + monkeypatch.setattr(paths, "DATA_DIR", data_dir) + + assert paths.evoscientist_root() == data_dir.resolve() + + def test_data_root_respects_environment_override(self, tmp_path, monkeypatch): + data_root = tmp_path / "web-data" + monkeypatch.setenv("EVOSCIENTIST_DATA_ROOT", str(data_root)) + paths._EVOSCIENTIST_DATA_ROOT = None + + assert paths._data_root() == data_root.resolve() + + def test_user_thread_and_global_dirs_are_created(self, tmp_path, monkeypatch): + data_root = tmp_path / "web-data" + monkeypatch.setenv("EVOSCIENTIST_DATA_ROOT", str(data_root)) + paths._EVOSCIENTIST_DATA_ROOT = None + + user_dir = paths.user_data_dir("user-a") + thread_dir = paths.thread_data_dir("user-a", "thread-1") + shared_dir = paths.global_data_dir("user-a") + + assert user_dir == data_root / "user-a" + assert thread_dir == user_dir / "thread-1" + assert shared_dir == user_dir / "__global__" + assert user_dir.is_dir() + assert thread_dir.is_dir() + assert shared_dir.is_dir() + + def test_iter_user_data_dirs_yields_directories_only(self, tmp_path, monkeypatch): + data_root = tmp_path / "web-data" + monkeypatch.setenv("EVOSCIENTIST_DATA_ROOT", str(data_root)) + paths._EVOSCIENTIST_DATA_ROOT = None + paths.user_data_dir("user-a") + paths.user_data_dir("user-b") + (data_root / "metadata.json").write_text("{}", encoding="utf-8") + + assert {path.name for path in paths.iter_user_data_dirs()} == { + "user-a", + "user-b", + } + + def test_uploads_dir_uses_evoscientist_root(self, tmp_path, monkeypatch): + home = tmp_path / "runtime-home" + monkeypatch.setenv("EVOSCIENTIST_HOME", str(home)) + + assert paths.uploads_dir() == home.resolve() / "uploads" + + class TestLegacySessionsDbMigration: """Tests for migrate_legacy_sessions_db() — transitional helper. 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_runtime_integrations.py b/tests/test_runtime_integrations.py new file mode 100644 index 0000000..9e3ffa7 --- /dev/null +++ b/tests/test_runtime_integrations.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +import ast +from pathlib import Path + +import pytest + +from EvoScientist.runtime_integrations import ( + RuntimeIntegrationUnavailable, + configure_runtime_integrations, + get_app_connection, + get_image_backend, + get_session_connection, + get_session_dsn, + handle_knowledge_file, + record_service_usage, + reset_runtime_integrations, + resolve_runtime_model, +) + + +@pytest.fixture(autouse=True) +def reset_integrations(): + reset_runtime_integrations() + yield + reset_runtime_integrations() + + +def test_core_package_does_not_import_gateway(): + package_root = Path(__file__).resolve().parents[1] / "EvoScientist" + violations = [] + for source_file in package_root.rglob("*.py"): + tree = ast.parse( + source_file.read_text(encoding="utf-8"), filename=str(source_file) + ) + for node in ast.walk(tree): + if isinstance(node, ast.Import): + names = [alias.name for alias in node.names] + elif isinstance(node, ast.ImportFrom): + if node.level: + continue + names = [node.module or ""] + else: + continue + if any(name == "gateway" or name.startswith("gateway.") for name in names): + violations.append( + f"{source_file.relative_to(package_root)}:{node.lineno}" + ) + assert violations == [] + + +@pytest.mark.anyio +async def test_optional_integrations_are_safe_without_web_runtime(tmp_path): + assert get_session_dsn() is None + await handle_knowledge_file(tmp_path / "result.md") + await record_service_usage("search", "query") + with pytest.raises(RuntimeIntegrationUnavailable): + await get_app_connection() + with pytest.raises(RuntimeIntegrationUnavailable): + await get_session_connection() + with pytest.raises(RuntimeIntegrationUnavailable): + get_image_backend() + + +@pytest.mark.anyio +async def test_host_can_register_runtime_integrations(tmp_path): + app_connection = object() + session_connection = object() + knowledge_paths = [] + usage = [] + image_backend = object() + + async def provide_app_connection(): + return app_connection + + async def provide_session_connection(): + return session_connection + + async def handle_file(path): + knowledge_paths.append(path) + + async def record_usage(service, action): + usage.append((service, action)) + + configure_runtime_integrations( + app_connection_provider=provide_app_connection, + session_connection_provider=provide_session_connection, + session_dsn_provider=lambda: "postgresql://example/session", + knowledge_file_handler=handle_file, + usage_recorder=record_usage, + image_backend_factory=lambda: image_backend, + ) + + path = tmp_path / "result.md" + await handle_knowledge_file(path) + await record_service_usage("mineru", "parse") + + assert await get_app_connection() is app_connection + assert await get_session_connection() is session_connection + assert get_session_dsn() == "postgresql://example/session" + assert get_image_backend() is image_backend + assert knowledge_paths == [path] + assert usage == [("mineru", "parse")] + + +def test_host_can_register_model_resolver(): + resolved = object() + calls = [] + + def resolve_model(model, provider): + calls.append((model, provider)) + return resolved + + configure_runtime_integrations(model_resolver=resolve_model) + + assert resolve_runtime_model("model-a", "provider-a") is resolved + assert calls == [("model-a", "provider-a")] diff --git a/tests/test_serde_default_rich_exception.py b/tests/test_serde_default_rich_exception.py new file mode 100644 index 0000000..aa2fdbb --- /dev/null +++ b/tests/test_serde_default_rich_exception.py @@ -0,0 +1,242 @@ +"""Regression tests for the helpers ``ErrorNormalizationMiddleware`` +uses to build the SSE error envelope. + +- ``_redact_api_keys`` + ``_build_env_key_redaction_re`` — scrubs + deployed credentials that the SDK might echo back. +- ``_extract_status_code`` / ``_extract_provider_code`` / + ``_extract_error_type`` — read SDK-specific fields off the raised + exception. + +Middleware wire behavior + ``_provider_from_model`` live in +``test_error_normalization_middleware.py``. One end-to-end orjson test +at the bottom guards that a ``ProviderStreamError`` survives +langgraph_api's UNPATCHED ``serde.default`` under +``OPT_SERIALIZE_DATACLASS`` — the whole reason the wrapper exists. +""" + +from __future__ import annotations + +import os + +import langgraph_api.serde as _serde_mod + +from EvoScientist.llm.errors import ( + _API_KEY_ENV_SUFFIXES, + _build_env_key_redaction_re, + _extract_error_type, + _extract_provider_code, + _extract_status_code, + _redact_api_keys, +) + +# --------------------------------------------------------------------------- +# Redaction +# --------------------------------------------------------------------------- + + +def test_env_deployed_key_redacted_in_message(monkeypatch): + """A credential exported via env var is scrubbed by + ``_redact_api_keys``. The redaction table is rebuilt per call — + ``monkeypatch.setenv`` alone is enough, no attribute reassignment. + """ + key = "sk-proj-aBcDeFgHiJkLmNoPqRsTuVwXyZ1234567890" + monkeypatch.setenv("OPENAI_API_KEY", key) + + msg = ( + f"Invalid API key: {key}. Get a new one at https://platform.openai.com/api-keys" + ) + redacted = _redact_api_keys(msg) + assert key not in redacted + assert "" in redacted + assert "Invalid API key" in redacted + assert "platform.openai.com" in redacted + + +def test_multiple_env_keys_redacted_independently(monkeypatch): + """Each ``*_API_KEY`` / ``*_TOKEN`` / ``*_SECRET`` env var + contributes its own prefix to the alternation. + """ + k1 = "sk-or-aBcDeFg012345678901234" + k2 = "AIzaABCDEFGHIJ0123456789" + k3 = "ghp_p4t70k3n0123456789abcdef" + monkeypatch.setenv("OPENROUTER_API_KEY", k1) + monkeypatch.setenv("GOOGLE_API_KEY", k2) + monkeypatch.setenv("GITHUB_TOKEN", k3) + + msg = _redact_api_keys(f"Failures: {k1}, {k2}, {k3}") + assert k1 not in msg + assert k2 not in msg + assert k3 not in msg + assert msg.count("") == 3 + + +def test_base64_suffix_secret_fully_redacted(monkeypatch): + """A base64-style secret (``/`` ``+`` ``=``) must redact end-to-end, + not leak its tail past the first padding char. + """ + key = "AbCdEfGh/secret+tail==" + monkeypatch.setenv("SOME_SECRET", key) + + msg = _redact_api_keys(f"auth failed with token={key} on retry") + assert "secret" not in msg + assert "tail" not in msg + assert "" in msg + assert "auth failed" in msg + assert "on retry" in msg + + +def test_unknown_shape_not_redacted_without_env(monkeypatch): + """Env-only redaction: a key-shaped string not deployed via env is + left alone. Tradeoff — we only scrub what we know is a secret. + """ + for k in list(os.environ): + if k.endswith(_API_KEY_ENV_SUFFIXES): + monkeypatch.delenv(k, raising=False) + + msg = _redact_api_keys("Unknown key seen: sk-or-aBcDeFg012345678901234") + assert "sk-or-aBcDeFg012345678901234" in msg + assert "" not in msg + + +def test_env_key_loaded_after_first_call_is_redacted(monkeypatch): + """The pattern rebuilds every call so keys loaded after + ``patches.py`` imports (typical ``load_dotenv`` sequence) are + still scrubbed on the next call. + """ + for k in list(os.environ): + if k.endswith(_API_KEY_ENV_SUFFIXES): + monkeypatch.delenv(k, raising=False) + + key = "sk-proj-loaded_after_import_1234567890abcdef" + # Pass 1: env empty — key leaks. + assert key in _redact_api_keys(f"leak: {key}") + + # Pass 2: after simulated load_dotenv. + monkeypatch.setenv("OPENAI_API_KEY", key) + redacted = _redact_api_keys(f"leak: {key}") + assert key not in redacted + assert "" in redacted + + +def test_redaction_regex_holds_only_prefix(monkeypatch): + """Defense-in-depth: the compiled regex must not embed the full key. + A process-memory leak (traceback locals, debugger) exposes at most + the first 8 chars — not the secret. + """ + key = "sk-proj-aBcDeFgHiJkLmNoPqRsTuVwXyZ1234567890_secret_suffix" + monkeypatch.setenv("OPENAI_API_KEY", key) + pattern = _build_env_key_redaction_re() + assert pattern is not None + assert key not in pattern.pattern + assert "aBcDeFgHiJkLmNoPqRs" not in pattern.pattern + # Sanity: still matches the full key at runtime via prefix + suffix + # greedy. + m = pattern.search(f"err: {key}") + assert m is not None + assert m.group(0) == key + + +# --------------------------------------------------------------------------- +# Field extractors +# --------------------------------------------------------------------------- + + +def _fake_exc(**attrs): + return type("APIError", (Exception,), attrs)("boom") + + +def test_status_code_read_from_direct_attribute(): + """openai / anthropic ``APIStatusError`` carries integer + ``.status_code`` — the primary path. + """ + assert _extract_status_code(_fake_exc(status_code=429)) == 429 + + +def test_status_code_read_via_response_attribute(): + """Wrappers that don't promote status to top level expose it via + ``.response.status_code`` (httpx pattern). + """ + + class FakeResponse: + status_code = 504 + + assert _extract_status_code(_fake_exc(response=FakeResponse())) == 504 + + +def test_status_code_read_via_integer_code_attribute(): + """``google.genai.errors.APIError`` stores HTTP status as integer + ``.code`` — type-disambiguated from openai/anthropic's string + ``.code`` (provider error code). + """ + assert _extract_status_code(_fake_exc(code=400)) == 400 + + +def test_provider_code_read_from_string_code_attribute(): + """Provider error code (``insufficient_quota`` etc.) is a string + ``.code`` — higher signal than the integer HTTP status alone. + """ + assert ( + _extract_provider_code(_fake_exc(code="insufficient_quota")) + == "insufficient_quota" + ) + + +def test_provider_code_ignores_integer_code(): + """An integer ``.code`` is HTTP status (see above); must not bleed + into the provider-code path. + """ + assert _extract_provider_code(_fake_exc(code=429)) is None + + +def test_error_type_read_from_type_attribute(): + """openai exposes a ``.type`` label (``rate_limit_error``).""" + assert _extract_error_type(_fake_exc(type="rate_limit_error")) == "rate_limit_error" + + +def test_extractors_return_none_when_attributes_absent(): + """A bare exception with no SDK-shape attributes — every extractor + returns None so the envelope drops the optional fields. + """ + exc = _fake_exc() + assert _extract_status_code(exc) is None + assert _extract_provider_code(exc) is None + assert _extract_error_type(exc) is None + + +# --------------------------------------------------------------------------- +# End-to-end: ProviderStreamError survives orjson under +# OPT_SERIALIZE_DATACLASS via upstream's UNPATCHED serde.default. +# --------------------------------------------------------------------------- + + +def test_provider_stream_error_survives_orjson_dataclass_option(): + """Guard: ``ProviderStreamError`` — a plain Exception subclass with + a ``model_dump()`` hook — must emerge as the envelope on the wire + even under ``OPT_SERIALIZE_DATACLASS``, using ONLY upstream's + stock ``serde.default``. Proof that we no longer need to patch + the serde module. + """ + import orjson + + from EvoScientist.llm.errors import ProviderStreamError + + err = ProviderStreamError( + provider="openrouter", + class_qualname="openrouter.errors.foo.UnauthorizedResponseError", + message="User not found.", + status_code=401, + ) + + wire = orjson.dumps( + err, + default=_serde_mod.default, # upstream, unpatched + option=orjson.OPT_SERIALIZE_DATACLASS, + ) + decoded = orjson.loads(wire) + assert decoded == { + "error": "UnauthorizedResponseError", + "class": "openrouter.errors.foo.UnauthorizedResponseError", + "message": "User not found.", + "provider": "openrouter", + "status_code": 401, + } 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:") diff --git a/uv.lock b/uv.lock index e15f1d8..947e9fa 100644 --- a/uv.lock +++ b/uv.lock @@ -944,7 +944,7 @@ wheels = [ [[package]] name = "evoscientist" -version = "0.2.1" +version = "0.2.2" source = { editable = "." } dependencies = [ { name = "deepagents", extra = ["quickjs"] },