10 Commits

Author SHA1 Message Date
m4 3ce5614254 fix: harden tool-call protocol and fallback handling
Test / pytest (ubuntu-latest, 3.11) (pull_request) Has been cancelled
Build / build (pull_request) Has been cancelled
Docker / build (pull_request) Has been cancelled
Lint / ruff (pull_request) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (pull_request) Has been cancelled
Test / pytest (windows-latest, 3.11) (pull_request) Has been cancelled
Test / pytest (windows-latest, 3.12) (pull_request) Has been cancelled
2026-07-19 12:05:56 +08:00
m4 4fc74e7da7 EvoScientist Ai4Sci 2026-07-14 22:07:14 +08:00
jfilipiuk 753c745405 fix: silence YAML-docstring noise from custom-app OpenAPI scan (#317) 2026-07-13 15:38:28 +01:00
jfilipiuk 88ac9f5ba1 fix: surface real exception class+message in SSE error events (#315)
* fix: surface real exception class+message in SSE error events

* fix: tighten SSE error patch scope and key redaction

* fix: redact base64-style secret suffixes fully

* style: remove notes/ reference from the dosctring

* fix: rebuild env cache on each error call

* fix: route BaseException through serde.default on SSE/webhook paths

* fix: distinguish routed providers by request URL host

* feat: normalize provider-SDK exceptions via ErrorNormalizationMiddleware

* refactor: drop json_dumpb dataclass-bypass wrappers, superseded by middleware

* fix: guard _extract_host against SDK properties that raise

* refactor: derive provider tag from ModelRequest.model, not the exception

* refactor: drop serde.default patch and exception-based inference; ProviderStreamError.model_dump handles the emit

* refactor: move envelope helpers from patches.py to errors.py

* feat: extend ErrorNormalizationMiddleware coverage to every model-call path

* chore: clean up review findings from middleware pivot

* fix: pass through all langgraph.errors

* fix: move langgraph.errors pass-through into _normalize

* fix: pass through ContextOverflowError in _normalize
2026-07-13 14:17:56 +01:00
jfilipiuk 952e68efe3 fix: scope quickjs snapshot to turn to keep checkpoints small (#316)
* fix: scope quickjs snapshot to turn to keep checkpoints small

* fix: strip _quickjs_snapshot_payload from state/history responses instead of dropping mode=thread

* fix: recurse strip into nested subgraph StateSnapshot

* fix: drop conditional-snapshot gate that leaked repl slots

* fix: assert LangGraph state-shape invariants at import time

* refactor: discover graphs to filter from langgraph.json

* test: assert copy() preserves subclass; iterate langgraph.json for subagent coverage

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 12:24:52 +00:00
Mani Saint-Victor 2b28c46caf fix(llm): make gpt-5.x usable through ccproxy Codex OAuth (#324)
* fix(llm): make gpt-5.x usable through ccproxy Codex OAuth

Two independent blockers made current OpenAI models fail when routed
through ccproxy's Codex OAuth endpoint:

1. ccproxy's default Codex model mappings rewrite any gpt-*/o1-*/o3-*/
   claude-* model to gpt-5.3-codex before forwarding, silently overriding
   the configured model and failing outright on accounts where
   gpt-5.3-codex is not served ("The 'gpt-5.3-codex' model is not
   supported when using Codex with a ChatGPT account").
   start_ccproxy() now generates a config with empty codex model
   mappings and passes it via 'ccproxy serve --config'.

2. ccproxy forwards the client's own User-Agent upstream and only
   gap-fills its Codex headers, so the backend gates current models on
   the client identity ("The '<model>' model requires a newer version
   of Codex"). get_chat_model() now sends Codex-CLI-shaped
   originator/version/User-Agent headers when the ccproxy Codex adapter
   is detected, overridable via EVOSCIENTIST_CODEX_CLIENT_VERSION.

Verified live: gpt-5.5 and gpt-5.4 complete successfully through
ccproxy Codex OAuth on a ChatGPT Plus account with both fixes; each
fails without them.

* fix(ccproxy): harden Codex client routing

* fix(llm): keep Codex client identity consistent

* docs: clarify Codex version floor

* style: ruff format models.py after merge

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
Co-authored-by: X-iZhang <zacharyzhang2022@gmail.com>
2026-07-13 11:04:47 +00:00
Mani Saint-Victor da6ca38d53 fix(llm): respect reasoning_effort setting on native OpenAI path (#321)
* fix(llm): respect reasoning_effort setting on native OpenAI path

The native OpenAI provider path hardcoded reasoning effort to xhigh for
gpt-5.4/5.5/codex models, silently ignoring the user's reasoning_effort
config setting. The OpenRouter path already honors the
EVOSCIENTIST_REASONING_EFFORT env var that settings.py exports from that
setting; this applies the same lookup on the native path, falling back
to the previous defaults when unset.

Adds a regression test and isolates the existing xhigh test from the
env var.

* fix(llm): preserve model reasoning defaults

* fix(llm): preserve GPT-5.6 reasoning default

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 10:42:55 +00:00
Mani Saint-Victor 19888c2db6 fix(ccproxy): raise startup timeouts (auth check 10s→30s, serve health 30s→180s) (#328)
* fix(ccproxy): raise auth status check timeout to 30s

ccproxy's CLI initializes its full plugin system on every invocation;
a cold 'ccproxy auth status' takes ~10s wall time on Apple Silicon,
so the 10s subprocess timeout made OAuth startup fail intermittently
with 'Auth check timed out' even when credentials were valid.

* fix(ccproxy): raise serve health deadline to 120s

ccproxy boot includes plugin init plus Codex CLI detection; measured
~76s to first healthy response on an Apple Silicon Mac (ccproxy-api
0.2.9). The 30s deadline in start_ccproxy() killed the process before
it could come up, failing OAuth startup with 'ccproxy did not become
healthy within 30 seconds'.

* fix(ccproxy): widen serve health deadline to 180s

Full startup measured at ~111s on a second cold run (Apple Silicon,
ccproxy-api 0.2.9); 120s left too little headroom for boot variance.

* fix(ccproxy): centralize startup timeouts

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 10:36:05 +00:00
Mani Saint-Victor f3e65a446f fix(tests): isolate the developer's real .env from the test suite (#329)
get_effective_config() runs load_dotenv(find_dotenv(usecwd=True),
override=True), so any test that loads config injected the repo's real
.env into os.environ for the rest of the pytest process. An
empty-valued line like MINIMAX_BASE_URL= then made
os.environ.get(key, default) return '' instead of the default,
failing the MiniMax routing tests in full-suite runs while they
passed in isolation.

Generalizes the find_dotenv redirect that test_config.py's
temp_config_dir fixture already applied locally into a suite-wide
autouse fixture, pointing at a never-created path so tests writing
their own tmp_path/.env cannot collide with it. Adds a regression
test reproducing the leak.

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 10:27:40 +00:00
Zixin Dong f72f7b93d5 feat(llm): add OpenRouter app attribution headers (#339) (#344)
* feat(llm): add OpenRouter app attribution headers (#339)

Attach EvoScientist app-attribution at the shared model-init layer so all
OpenRouter calls are credited to the project. langchain-openrouter maps
app_url/app_title/app_categories -> HTTP-Referer / X-Title /
X-OpenRouter-Categories. Applied only for the openrouter provider, via
setdefault so explicit caller kwargs win. Configurable through new
openrouter_http_referer / openrouter_app_title / openrouter_app_categories
settings and their EVOSCIENTIST_OPENROUTER_* env vars.

Closes #339

* refactor(llm): centralize OpenRouter attribution defaults + cap categories

Address PR #344 review:
- Define the app-attribution default constants once in config/settings.py
  (the config fields and llm/models.py both use them) instead of duplicating
  the literals across the two modules.
- Reduce the default categories to creative-writing,personal-agent and cap the
  sent list to OpenRouter's 2-per-request limit, warning when a configured list
  exceeds it, so extras are dropped predictably (and surfaced) here rather than
  being silently truncated server-side.

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 11:08:59 +01:00
36 changed files with 5269 additions and 163 deletions
+84 -10
View File
@@ -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``).
@@ -648,6 +673,7 @@ def _get_default_middleware(
tool_selector_threshold: int | None = None,
memory_max_inline_profile_chars: int | None = None,
enable_background_execution: bool = True,
enable_legacy_model_fallback: bool = True,
):
"""Build the default middleware list.
@@ -669,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,
@@ -685,6 +715,16 @@ def _get_default_middleware(
)
cfg = cfg if cfg is not None else _ensure_config()
repetitive_tool_call_threshold = getattr(
cfg,
"repetitive_tool_call_threshold",
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
)
if not isinstance(repetitive_tool_call_threshold, int):
repetitive_tool_call_threshold = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD
max_consecutive_tool_errors = getattr(cfg, "max_consecutive_tool_errors", 3)
if not isinstance(max_consecutive_tool_errors, int):
max_consecutive_tool_errors = 3
if cfg.model_fallbacks:
load_fallback_chain(cfg.model_fallbacks)
model = chat_model if chat_model is not None else _ensure_chat_model()
@@ -734,21 +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(
**(
{"threshold": tool_selector_threshold}
if tool_selector_threshold is not None
else {}
),
model=tool_selector_model,
track_stream_selection=not for_async_subagent,
),
*selector_middlewares,
ToolProtocolGuardMiddleware(),
# Interpreter prompt must land before runtime/memory context, so this
# middleware sits ahead of runtime_context in the stack.
create_code_interpreter_middleware(
@@ -885,6 +936,8 @@ def create_cli_agent(
memory_max_inline_profile_chars: int | None = None,
enable_subagents: bool = True,
enable_background_execution: bool = True,
main_agent_outer_middlewares: Sequence[AgentMiddleware] | None = None,
main_agent_route_middleware: AgentMiddleware | None = None,
) -> "CompiledStateGraph":
"""Create agent with checkpointer for CLI multi-turn support.
@@ -921,6 +974,12 @@ def create_cli_agent(
enable_background_execution: Whether local background-process tools are
installed. Embedding hosts should disable this when process execution
is provided by an external backend.
main_agent_outer_middlewares: Optional host-owned middleware installed
only on the top-level agent, outside EvoScientist's default chain.
main_agent_route_middleware: Optional host-owned route middleware placed
after ConfigurableModelMiddleware and before tool selection. When
provided, EvoScientist's legacy model fallback is disabled for the
top-level agent so the host is the only fallback authority.
"""
import os as _os
@@ -1005,7 +1064,22 @@ def create_cli_agent(
tool_selector_threshold=tool_selector_threshold,
memory_max_inline_profile_chars=memory_max_inline_profile_chars,
enable_background_execution=enable_background_execution,
enable_legacy_model_fallback=main_agent_route_middleware is None,
)
if main_agent_route_middleware is not None:
configurable_index = next(
(
index
for index, middleware in enumerate(mw)
if getattr(middleware, "name", "") == "configurable_model"
),
None,
)
if configurable_index is None:
raise RuntimeError("ConfigurableModelMiddleware route slot is unavailable")
mw.insert(configurable_index + 1, main_agent_route_middleware)
if main_agent_outer_middlewares:
mw = [*main_agent_outer_middlewares, *mw]
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
# would propagate it to every subagent, breaking parallel execute calls
+57 -6
View File
@@ -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:
+65 -1
View File
@@ -135,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:
@@ -252,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
@@ -290,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)
@@ -452,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__
@@ -714,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:
@@ -773,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",
@@ -789,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",
@@ -904,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"
):
+1 -1
View File
@@ -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
+237 -1
View File
@@ -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"]
+386
View File
@@ -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 ``<redacted>``.
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("<redacted>", 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 ``<module>_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 ``<module_tag>_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
+132 -10
View File
@@ -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,14 @@ _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.
@@ -369,16 +429,18 @@ def _apply_auto_config(
and not disable_reasoning
and "reasoning" not in kwargs
):
if _is_ccproxy_codex(kwargs.get("base_url"), kwargs.get("api_key")):
# 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"
_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" and not disable_reasoning:
@@ -527,6 +589,20 @@ 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>' 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.setdefault("api_key", api_key)
@@ -574,8 +650,54 @@ def get_chat_model(
# 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
+336 -55
View File
@@ -284,71 +284,288 @@ def _stable_tool_call_id(message: Any, message_index: int, call_index: int) -> s
return "call_" + hashlib.sha256(seed.encode("utf-8")).hexdigest()[:24]
def _ensure_openai_tool_call_ids(messages: list[Any]) -> list[Any]:
"""Copy messages and repair missing AI/ToolMessage call identifiers."""
import copy
from collections import deque
def _tool_message_match_index(
tool_messages: list[Any],
used_indexes: set[int],
*,
call_id: str,
call_name: str,
) -> int | None:
"""Find the best unused result for one assistant tool call."""
pending_call_ids: deque[str] = deque()
normalized: list[Any] = []
def _matches(index: int, *, require_id: bool, require_name: bool) -> bool:
if index in used_indexes:
return False
message = tool_messages[index]
result_id = str(getattr(message, "tool_call_id", "") or "")
result_name = str(getattr(message, "name", "") or "")
if require_id and result_id != call_id:
return False
if not require_id and result_id:
return False
return not require_name or not result_name or result_name == call_name
for message_index, message in enumerate(messages):
message_type = getattr(message, "type", None)
if message_type == "ai":
tool_calls = list(getattr(message, "tool_calls", None) or [])
if not tool_calls:
normalized.append(message)
if call_id:
for require_name in (True, False):
for index in range(len(tool_messages)):
if _matches(index, require_id=True, require_name=require_name):
return index
for require_name in (True, False):
for index in range(len(tool_messages)):
if _matches(index, require_id=False, require_name=require_name):
return index
return None
# A result-side identifier is more authoritative than a generated fallback.
for require_name in (True, False):
for index, message in enumerate(tool_messages):
if index in used_indexes:
continue
result_id = str(getattr(message, "tool_call_id", "") or "")
result_name = str(getattr(message, "name", "") or "")
if result_id and (
not require_name or not result_name or result_name == call_name
):
return index
for require_name in (True, False):
for index in range(len(tool_messages)):
if _matches(index, require_id=False, require_name=require_name):
return index
return None
copied = copy.copy(message)
normalized_calls: list[dict[str, Any]] = []
for call_index, original_call in enumerate(tool_calls):
call = dict(original_call)
call_id = str(call.get("id") or "") or _stable_tool_call_id(
message, message_index, call_index
)
call["id"] = call_id
normalized_calls.append(call)
pending_call_ids.append(call_id)
copied.tool_calls = normalized_calls
if isinstance(copied.content, list):
call_index = 0
blocks: list[Any] = []
for original_block in copied.content:
if not isinstance(original_block, dict):
blocks.append(original_block)
continue
block = dict(original_block)
if block.get("type") in {"tool_call", "function_call"}:
if call_index < len(normalized_calls):
block["id"] = normalized_calls[call_index]["id"]
call_index += 1
blocks.append(block)
copied.content = blocks
normalized.append(copied)
def _copy_ai_message_with_tool_pairs(
message: Any,
message_index: int,
tool_messages: list[Any],
) -> tuple[Any | None, list[Any]]:
"""Return a replay-safe assistant message and its matched tool results."""
import copy
copied = copy.copy(message)
additional_kwargs = dict(getattr(message, "additional_kwargs", None) or {})
# Parsed tool_calls are canonical. Raw copies can otherwise re-introduce an
# invalid call after invalid_tool_calls has been cleared.
additional_kwargs.pop("tool_calls", None)
copied.additional_kwargs = additional_kwargs
copied.invalid_tool_calls = []
original_calls = list(getattr(message, "tool_calls", None) or [])
used_results: set[int] = set()
matched_calls: list[dict[str, Any]] = []
matched_result_indexes: list[int] = []
original_to_matched_call: dict[int, tuple[str, str]] = {}
for call_index, original_call in enumerate(original_calls):
call = dict(original_call)
call_id = str(call.get("id") or "")
call_name = str(call.get("name") or "").strip()
# A missing name is structurally unreplayable. Never infer it from
# arguments or retain its paired ToolMessage in provider history.
if not call_name:
continue
call["name"] = call_name
result_index = _tool_message_match_index(
tool_messages,
used_results,
call_id=call_id,
call_name=call_name,
)
# A historical client-side function call is only replayable together
# with its result. Incomplete calls are discarded instead of asking the
# provider to continue a broken tool turn.
if result_index is None:
continue
if not call_id:
result_id = str(
getattr(tool_messages[result_index], "tool_call_id", "") or ""
)
call_id = result_id or _stable_tool_call_id(
message, message_index, call_index
)
call["id"] = call_id
matched_calls.append(call)
matched_result_indexes.append(result_index)
original_to_matched_call[call_index] = (call_id, call_name)
used_results.add(result_index)
copied.tool_calls = matched_calls
if isinstance(copied.content, list):
original_call_index = 0
blocks: list[Any] = []
for original_block in copied.content:
if not isinstance(original_block, dict):
blocks.append(original_block)
continue
block = dict(original_block)
if block.get("type") in {"tool_call", "function_call"}:
matched_call = original_to_matched_call.get(original_call_index)
original_call_index += 1
if matched_call is None:
continue
call_id, call_name = matched_call
# LangChain content blocks use id; the Responses converter later
# maps it to call_id.
block["id"] = call_id
block["name"] = call_name
if isinstance(block.get("function"), dict):
block["function"] = {**block["function"], "name": call_name}
blocks.append(block)
copied.content = blocks
matched_results: list[Any] = []
result_to_call_id = {
result_index: matched_calls[index]["id"]
for index, result_index in enumerate(matched_result_indexes)
}
for result_index, result in enumerate(tool_messages):
call_id = result_to_call_id.get(result_index)
if call_id is None:
continue
copied_result = copy.copy(result)
copied_result.tool_call_id = call_id
matched_results.append(copied_result)
had_tool_protocol = bool(original_calls) or bool(
getattr(message, "invalid_tool_calls", None)
)
if not matched_calls and had_tool_protocol:
replayable_content = _flatten_message_content(copied.content)
if not replayable_content:
return None, matched_results
return copied, matched_results
def _sanitize_openai_tool_history(messages: list[Any]) -> list[Any]:
"""Copy history while retaining only complete, replayable tool turns."""
normalized: list[Any] = []
index = 0
while index < len(messages):
message = messages[index]
message_type = getattr(message, "type", None)
if message_type == "tool":
# A tool result without its immediately preceding assistant call is
# invalid for both Chat Completions and Responses APIs.
index += 1
continue
if message_type != "ai":
normalized.append(message)
index += 1
continue
if message_type == "tool":
tool_call_id = str(getattr(message, "tool_call_id", "") or "")
if tool_call_id:
try:
pending_call_ids.remove(tool_call_id)
except ValueError:
pass
normalized.append(message)
continue
if pending_call_ids:
copied = copy.copy(message)
copied.tool_call_id = pending_call_ids.popleft()
normalized.append(copied)
continue
normalized.append(message)
next_index = index + 1
tool_messages: list[Any] = []
while (
next_index < len(messages)
and getattr(messages[next_index], "type", None) == "tool"
):
tool_messages.append(messages[next_index])
next_index += 1
copied, matched_results = _copy_ai_message_with_tool_pairs(
message,
index,
tool_messages,
)
if copied is not None:
normalized.append(copied)
normalized.extend(matched_results)
index = next_index
return normalized
def _ensure_openai_tool_call_ids(messages: list[Any]) -> list[Any]:
"""Backward-compatible alias for replay-safe tool history normalization."""
return _sanitize_openai_tool_history(messages)
def _has_assistant_tool_protocol(messages: list[Any]) -> bool:
"""Return whether history contains assistant-side tool protocol state."""
for message in messages:
if getattr(message, "type", None) != "ai":
continue
if getattr(message, "tool_calls", None) or getattr(
message, "invalid_tool_calls", None
):
return True
additional_kwargs = getattr(message, "additional_kwargs", None) or {}
if additional_kwargs.get("tool_calls"):
return True
content = getattr(message, "content", None)
if isinstance(content, list) and any(
isinstance(block, dict)
and block.get("type") in {"tool_call", "function_call"}
for block in content
):
return True
return False
def _validate_openai_tool_history(messages: list[Any]) -> None:
"""Raise when sanitized history still contains an invalid tool protocol."""
available_call_ids: set[str] = set()
for message in messages:
message_type = getattr(message, "type", None)
if message_type == "ai":
if getattr(message, "invalid_tool_calls", None):
raise ValueError("invalid_tool_calls must not be replayed")
response_call_ids: set[str] = set()
response_calls: dict[str, str] = {}
for call in getattr(message, "tool_calls", None) or []:
call_name = str(call.get("name") or "").strip()
if not call_name:
raise ValueError("assistant tool call is missing a name")
call_id = str(call.get("id") or "").strip()
if not call_id:
raise ValueError("assistant tool call is missing an id")
if call_id in response_call_ids or call_id in available_call_ids:
raise ValueError(
"assistant tool call id is duplicated while outstanding"
)
response_call_ids.add(call_id)
available_call_ids.add(call_id)
response_calls[call_id] = call_name
content = getattr(message, "content", None)
content_call_ids: set[str] = set()
if isinstance(content, list):
for block in content:
if not isinstance(block, dict) or block.get("type") not in {
"tool_call",
"function_call",
}:
continue
block_id = str(
block.get("id") or block.get("call_id") or ""
).strip()
block_name = block.get("name") or block.get("tool_name")
function = block.get("function")
if not block_name and isinstance(function, dict):
block_name = function.get("name")
block_name = str(block_name or "").strip()
if (
not block_id
or block_id in content_call_ids
or response_calls.get(block_id) != block_name
):
raise ValueError(
"assistant content block does not match parsed tool call"
)
content_call_ids.add(block_id)
elif message_type == "tool":
call_id = str(getattr(message, "tool_call_id", "") or "")
if not call_id or call_id not in available_call_ids:
raise ValueError("tool result does not match a prior tool call")
available_call_ids.remove(call_id)
if available_call_ids:
raise ValueError("assistant tool call is missing its tool result")
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
"""Flatten list content for OpenAI-compatible APIs, preserving media.
@@ -364,7 +581,9 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
from langchain_core.messages import HumanMessage
messages = _ensure_openai_tool_call_ids(messages)
sanitize_tool_history = _has_assistant_tool_protocol(messages)
if sanitize_tool_history:
messages = _sanitize_openai_tool_history(messages)
out: list[Any] = []
pending_media: list[Any] = [] # media hoisted out of a run of tool messages
@@ -403,6 +622,8 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
msg.content = flat
out.append(msg)
_flush() # conversation may end with tool messages
if sanitize_tool_history:
_validate_openai_tool_history(out)
return out
@@ -812,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.
+17 -11
View File
@@ -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(
@@ -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()],
)
+14
View File
@@ -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",
+15 -1
View File
@@ -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
@@ -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
+34 -3
View File
@@ -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)
@@ -0,0 +1,350 @@
"""Detect deterministic tool loops and compact only provider-facing history."""
from __future__ import annotations
import json
import logging
import re
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import Any
from langchain.agents.middleware.types import (
AgentMiddleware,
ModelRequest,
ModelResponse,
)
from ..llm.errors import AgentControlError
logger = logging.getLogger(__name__)
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD = 2
DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS = 3
_TRANSIENT_PATTERNS = (
"timeout",
"timed out",
"cancelled",
"canceled",
"connection",
"rate limit",
"too many requests",
"temporarily unavailable",
"service unavailable",
"overloaded",
"bad gateway",
"gateway timeout",
"http 500",
"http 502",
"http 503",
"http 504",
)
_DETERMINISTIC_PATTERNS: tuple[tuple[str, tuple[str, ...]], ...] = (
(
"INVALID_ARGUMENTS",
("invalid argument", "validation error", "schema", "bad input"),
),
("UNKNOWN_TOOL", ("not a valid tool", "unknown tool", "tool not found")),
("UNSUPPORTED", ("not supported", "unsupported", "not implemented")),
(
"POLICY_DENIED",
("permission denied", "forbidden", "policy denied", "not allowed"),
),
)
_SAFE_CODE_RE = re.compile(r"^[A-Za-z][A-Za-z0-9_.:-]{0,95}$")
_DETERMINISTIC_CODE_MARKERS = (
"INVALID",
"VALIDATION",
"SCHEMA",
"UNKNOWN_TOOL",
"NOT_FOUND",
"UNSUPPORTED",
"NOT_IMPLEMENTED",
"POLICY",
"PERMISSION",
"FORBIDDEN",
"DENIED",
)
@dataclass(frozen=True, slots=True)
class RepetitiveToolHistoryRepair:
messages: list[Any]
blocked_tool_names: frozenset[str]
removed_rounds: int
tail_repetitions: int = 0
tail_consecutive_errors: int = 0
@dataclass(frozen=True, slots=True)
class _ToolRound:
messages: tuple[Any, ...]
signature: tuple[tuple[str, str, str], ...]
tool_names: frozenset[str]
deterministic_error: bool
def _canonical_tool_args(value: Any) -> str:
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError:
return value.strip()
try:
return json.dumps(
value,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
default=str,
)
except (TypeError, ValueError):
return repr(value)
def _deterministic_result_code(message: Any) -> str | None:
additional = getattr(message, "additional_kwargs", None)
additional = additional if isinstance(additional, Mapping) else {}
raw_code = additional.get("error_code") or additional.get("code")
status = str(getattr(message, "status", "") or "").lower()
content = str(getattr(message, "content", "") or "")
lowered = content.lower()
if any(pattern in lowered for pattern in _TRANSIENT_PATTERNS):
return None
if isinstance(raw_code, str) and _SAFE_CODE_RE.fullmatch(raw_code.strip()):
normalized = raw_code.strip().upper()
if any(
pattern.replace(" ", "_") in normalized for pattern in _TRANSIENT_PATTERNS
):
return None
if any(marker in normalized for marker in _DETERMINISTIC_CODE_MARKERS):
return normalized
return None
is_error = status == "error" or lowered.startswith("error:")
if not is_error:
return None
for code, patterns in _DETERMINISTIC_PATTERNS:
if any(pattern in lowered for pattern in patterns):
return code
return None
def _parse_tool_round(
messages: Sequence[Any], start: int
) -> tuple[_ToolRound, int] | None:
assistant = messages[start]
if getattr(assistant, "type", None) != "ai":
return None
raw_calls = list(getattr(assistant, "tool_calls", None) or [])
calls = [call for call in raw_calls if isinstance(call, Mapping)]
if not calls or len(calls) != len(raw_calls):
return None
end = start + 1
results: list[Any] = []
while end < len(messages) and getattr(messages[end], "type", None) == "tool":
results.append(messages[end])
end += 1
if not results:
return None
results_by_id = {
str(getattr(result, "tool_call_id", "") or "").strip(): result
for result in results
if str(getattr(result, "tool_call_id", "") or "").strip()
}
signature: list[tuple[str, str, str]] = []
tool_names: set[str] = set()
for index, call in enumerate(calls):
name = str(call.get("name") or "").strip()
call_id = str(call.get("id") or "").strip()
if not name or not call_id:
return None
result = results_by_id.get(call_id)
if result is None and index < len(results):
candidate = results[index]
if not str(getattr(candidate, "tool_call_id", "") or "").strip():
result = candidate
if result is None:
return None
result_code = _deterministic_result_code(result)
if result_code is None:
return _ToolRound(
messages=(assistant, *results),
signature=(),
tool_names=frozenset(),
deterministic_error=False,
), end
signature.append((name, _canonical_tool_args(call.get("args")), result_code))
tool_names.add(name)
return (
_ToolRound(
messages=(assistant, *results),
signature=tuple(signature),
tool_names=frozenset(tool_names),
deterministic_error=True,
),
end,
)
def collapse_repetitive_tool_rounds(
messages: Sequence[Any],
*,
threshold: int = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
) -> RepetitiveToolHistoryRepair:
"""Build a provider-only projection while preserving audit history.
Only the middle rounds of three-or-more identical deterministic error
groups are omitted. The first and last observations remain, and callers
must never persist this projection back to a checkpoint.
"""
if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0:
raise ValueError("repetitive tool call threshold must be non-negative")
original = list(messages)
segments: list[Any | _ToolRound] = []
index = 0
while index < len(original):
parsed = _parse_tool_round(original, index)
if parsed is None:
segments.append(original[index])
index += 1
continue
tool_round, index = parsed
segments.append(tool_round)
tail_repetitions = 0
tail_consecutive_errors = 0
if segments and isinstance(segments[-1], _ToolRound):
tail = segments[-1]
if tail.deterministic_error:
cursor = len(segments) - 1
while cursor >= 0 and isinstance(segments[cursor], _ToolRound):
current = segments[cursor]
if not current.deterministic_error:
break
tail_consecutive_errors += 1
cursor -= 1
cursor = len(segments) - 1
while cursor >= 0 and isinstance(segments[cursor], _ToolRound):
current = segments[cursor]
if (
not current.deterministic_error
or current.signature != tail.signature
):
break
tail_repetitions += 1
cursor -= 1
projected: list[Any] = []
removed_rounds = 0
index = 0
while index < len(segments):
segment = segments[index]
if not isinstance(segment, _ToolRound) or not segment.deterministic_error:
if isinstance(segment, _ToolRound):
projected.extend(segment.messages)
else:
projected.append(segment)
index += 1
continue
end = index + 1
while (
end < len(segments)
and isinstance(segments[end], _ToolRound)
and segments[end].deterministic_error
and segments[end].signature == segment.signature
):
end += 1
group = segments[index:end]
should_compact = threshold > 0 and len(group) >= threshold and len(group) > 2
if should_compact:
projected.extend(group[0].messages)
projected.extend(group[-1].messages)
removed_rounds += len(group) - 2
else:
for item in group:
projected.extend(item.messages)
index = end
blocked = (
segments[-1].tool_names
if tail_repetitions and isinstance(segments[-1], _ToolRound)
else frozenset()
)
return RepetitiveToolHistoryRepair(
messages=projected,
blocked_tool_names=blocked,
removed_rounds=removed_rounds,
tail_repetitions=tail_repetitions,
tail_consecutive_errors=tail_consecutive_errors,
)
class RepetitiveToolCallGuardMiddleware(AgentMiddleware):
"""Stop deterministic loops before another model request is made."""
name = "repetitive_tool_call_guard"
def __init__(
self,
*,
threshold: int = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
max_consecutive_errors: int = DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS,
) -> None:
super().__init__()
for name, value in {
"threshold": threshold,
"max_consecutive_errors": max_consecutive_errors,
}.items():
if not isinstance(value, int) or isinstance(value, bool) or value < 0:
raise ValueError(f"{name} must be a non-negative integer")
self.threshold = threshold
self.max_consecutive_errors = max_consecutive_errors
def _prepare_request(self, request: ModelRequest) -> ModelRequest:
repair = collapse_repetitive_tool_rounds(
request.messages,
threshold=self.threshold,
)
if self.threshold and repair.tail_repetitions >= self.threshold:
raise AgentControlError(
"MODEL_TOOL_LOOP_DETECTED",
"A deterministic repeated tool-call loop was stopped.",
status_code=422,
retryable=False,
)
if (
self.max_consecutive_errors
and repair.tail_consecutive_errors >= self.max_consecutive_errors
):
raise AgentControlError(
"MODEL_TOOL_ERROR_LIMIT",
"Too many consecutive deterministic tool errors were stopped.",
status_code=422,
retryable=False,
)
if repair.removed_rounds:
logger.info(
"Compacted deterministic tool errors for provider projection: removed_rounds=%d",
repair.removed_rounds,
)
return request.override(messages=repair.messages)
return request
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
return handler(self._prepare_request(request))
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
return await handler(self._prepare_request(request))
@@ -0,0 +1,361 @@
"""Validate completed model tool calls before they can reach ToolNode."""
from __future__ import annotations
import hashlib
import json
from collections.abc import Awaitable, Callable, Mapping, Sequence
from typing import Any
from langchain.agents.middleware.types import (
AgentMiddleware,
ExtendedModelResponse,
ModelRequest,
ModelResponse,
)
from langchain_core.messages import AIMessage
from langchain_core.tools import BaseTool
from ..llm.errors import ModelToolProtocolError, _provider_from_model
_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call"})
_MAX_DIAGNOSTIC_KEYS = 16
_MAX_DIAGNOSTIC_KEY_CHARS = 64
def _tool_name(tool: BaseTool | Mapping[str, Any] | Any) -> str | None:
if isinstance(tool, BaseTool):
return tool.name.strip() or None
if isinstance(tool, Mapping):
value = tool.get("name")
if not value and isinstance(tool.get("function"), Mapping):
value = tool["function"].get("name")
if isinstance(value, str) and value.strip():
return value.strip()
return None
value = getattr(tool, "name", None)
return value.strip() if isinstance(value, str) and value.strip() else None
def _ai_messages(response: Any) -> list[AIMessage]:
"""Extract final AI messages from every LangChain middleware response shape."""
if isinstance(response, AIMessage):
return [response]
if isinstance(response, ExtendedModelResponse):
response = response.model_response
elif not isinstance(response, ModelResponse):
nested = getattr(response, "model_response", None)
if nested is not None:
response = nested
result = getattr(response, "result", None)
if not isinstance(result, Sequence) or isinstance(result, str | bytes):
return []
return [message for message in result if isinstance(message, AIMessage)]
def _block_identity(block: Mapping[str, Any]) -> tuple[str, str]:
call_id = str(block.get("id") or block.get("call_id") or "").strip()
name = block.get("name") or block.get("tool_name")
function = block.get("function")
if not name and isinstance(function, Mapping):
name = function.get("name")
return call_id, str(name or "").strip()
def _value_digest(value: Any) -> str:
try:
encoded = json.dumps(
value,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
default=lambda item: f"<{type(item).__name__}>",
).encode("utf-8")
except (TypeError, ValueError):
encoded = f"<{type(value).__name__}:unserializable>".encode()
return "sha256:" + hashlib.sha256(encoded).hexdigest()[:16]
def _argument_diagnostic(value: Any, *, present: bool) -> dict[str, Any]:
if not present:
return {"args_present": False, "args_type": "missing"}
if isinstance(value, Mapping):
keys = sorted(str(key)[:_MAX_DIAGNOSTIC_KEY_CHARS] for key in value)
return {
"args_present": True,
"args_type": "object",
"args_key_count": len(keys),
"args_keys": keys[:_MAX_DIAGNOSTIC_KEYS],
"args_keys_truncated": len(keys) > _MAX_DIAGNOSTIC_KEYS,
"args_digest": _value_digest(value),
}
if isinstance(value, Sequence) and not isinstance(value, str | bytes):
value_type = "array"
elif isinstance(value, str):
value_type = "string"
elif value is None:
value_type = "null"
else:
value_type = type(value).__name__
return {
"args_present": True,
"args_type": value_type,
"args_digest": _value_digest(value),
}
def _summarize_call(call: Any) -> dict[str, Any]:
if not isinstance(call, Mapping):
return {"call_type": type(call).__name__}
function = call.get("function")
function = function if isinstance(function, Mapping) else {}
call_id = str(call.get("id") or call.get("call_id") or "").strip()
name = call.get("name") or call.get("tool_name") or function.get("name")
name = str(name or "").strip()
if "args" in call:
args = call.get("args")
args_present = True
elif "arguments" in call:
args = call.get("arguments")
args_present = True
elif "arguments" in function:
args = function.get("arguments")
args_present = True
else:
args = None
args_present = False
summary = {
"call_type": "object",
"name": name or "<missing>",
"id_present": bool(call_id),
**_argument_diagnostic(args, present=args_present),
}
if call_id:
summary["id_fingerprint"] = _value_digest(call_id)
return summary
def _raw_openai_call(message: AIMessage, call_index: int) -> Any | None:
additional = getattr(message, "additional_kwargs", None)
additional = additional if isinstance(additional, Mapping) else {}
raw_calls = additional.get("tool_calls")
if (
isinstance(raw_calls, Sequence)
and not isinstance(raw_calls, str | bytes)
and call_index < len(raw_calls)
):
return raw_calls[call_index]
return None
def _call_diagnostic(
message: AIMessage,
call: Any,
*,
source: str,
call_index: int,
call_count: int,
) -> dict[str, Any]:
diagnostic = {
"source": source,
"call_index": call_index,
"call_count": call_count,
**_summarize_call(call),
}
raw_call = _raw_openai_call(message, call_index)
diagnostic["raw_openai_call_available"] = raw_call is not None
if raw_call is not None:
diagnostic["raw_openai_call"] = _summarize_call(raw_call)
return diagnostic
def _route_metadata(request: ModelRequest) -> dict[str, Any]:
model = request.model
metadata = getattr(model, "metadata", None)
metadata = metadata if isinstance(metadata, Mapping) else {}
provider = metadata.get("route_provider") or _provider_from_model(model)
model_id = metadata.get("route_model")
if not model_id:
model_id = (
getattr(model, "model_name", None)
or getattr(model, "model", None)
or getattr(model, "model_id", None)
)
generation = metadata.get("route_config_generation")
try:
config_generation = int(generation) if generation is not None else None
except (TypeError, ValueError):
config_generation = None
return {
"provider": str(provider) if provider else None,
"model": str(model_id) if model_id else None,
"route_key": str(metadata.get("route_key"))
if metadata.get("route_key")
else None,
"config_generation": config_generation,
"api_mode": str(metadata.get("route_api_mode"))
if metadata.get("route_api_mode")
else None,
"endpoint": str(metadata.get("route_endpoint"))
if metadata.get("route_endpoint")
else None,
"tool_call_transport": str(metadata.get("route_tool_call_transport"))
if metadata.get("route_tool_call_transport")
else None,
}
def _raise_protocol_error(
request: ModelRequest,
reason: str,
*,
call_id: str | None = None,
call_diagnostic: dict[str, Any] | None = None,
) -> None:
raise ModelToolProtocolError(
reason,
call_id=call_id or None,
call_diagnostic=call_diagnostic,
**_route_metadata(request),
)
def _validate_message(
message: AIMessage,
request: ModelRequest,
allowed_names: frozenset[str],
) -> None:
invalid_calls = list(getattr(message, "invalid_tool_calls", None) or [])
if invalid_calls:
invalid = invalid_calls[0]
call_id = str(invalid.get("id") or "") if isinstance(invalid, Mapping) else ""
_raise_protocol_error(
request,
"invalid_final_call",
call_id=call_id,
call_diagnostic=_call_diagnostic(
message,
invalid,
source="invalid_tool_calls",
call_index=0,
call_count=len(invalid_calls),
),
)
parsed_by_id: dict[str, str] = {}
parsed_calls = list(getattr(message, "tool_calls", None) or [])
for call_index, raw_call in enumerate(parsed_calls):
diagnostic = _call_diagnostic(
message,
raw_call,
source="parsed_tool_calls",
call_index=call_index,
call_count=len(parsed_calls),
)
if not isinstance(raw_call, Mapping):
_raise_protocol_error(
request, "invalid_final_call", call_diagnostic=diagnostic
)
call_id = str(raw_call.get("id") or raw_call.get("call_id") or "").strip()
name = str(raw_call.get("name") or "").strip()
if not name:
_raise_protocol_error(
request,
"missing_name",
call_id=call_id,
call_diagnostic=diagnostic,
)
if name not in allowed_names:
_raise_protocol_error(
request,
"unknown_name",
call_id=call_id,
call_diagnostic=diagnostic,
)
if not call_id:
_raise_protocol_error(request, "missing_id", call_diagnostic=diagnostic)
if call_id in parsed_by_id:
_raise_protocol_error(
request,
"duplicate_id",
call_id=call_id,
call_diagnostic=diagnostic,
)
args = raw_call.get("args")
if not isinstance(args, Mapping):
_raise_protocol_error(
request,
"invalid_args",
call_id=call_id,
call_diagnostic=diagnostic,
)
parsed_by_id[call_id] = name
content = getattr(message, "content", None)
if not isinstance(content, list):
return
seen_block_ids: set[str] = set()
tool_blocks = [
block
for block in content
if isinstance(block, Mapping) and block.get("type") in _TOOL_BLOCK_TYPES
]
for block_index, block in enumerate(tool_blocks):
diagnostic = _call_diagnostic(
message,
block,
source="content_blocks",
call_index=block_index,
call_count=len(tool_blocks),
)
call_id, name = _block_identity(block)
if not call_id:
_raise_protocol_error(request, "missing_id", call_diagnostic=diagnostic)
if call_id in seen_block_ids:
_raise_protocol_error(
request,
"duplicate_id",
call_id=call_id,
call_diagnostic=diagnostic,
)
seen_block_ids.add(call_id)
parsed_name = parsed_by_id.get(call_id)
if parsed_name is None or (name and name != parsed_name):
_raise_protocol_error(
request,
"inconsistent_block",
call_id=call_id,
call_diagnostic=diagnostic,
)
class ToolProtocolGuardMiddleware(AgentMiddleware):
"""Fail closed on malformed final tool calls using the actual request tools."""
name = "tool_protocol_guard"
@staticmethod
def _validate(response: Any, request: ModelRequest) -> None:
allowed_names = frozenset(
name for tool in request.tools if (name := _tool_name(tool)) is not None
)
for message in _ai_messages(response):
_validate_message(message, request, allowed_names)
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
response = handler(request)
self._validate(response, request)
return response
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
response = await handler(request)
self._validate(response, request)
return response
+31 -3
View File
@@ -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. "
+32 -2
View File
@@ -7,6 +7,15 @@ All events contain a type and associated data dict.
from dataclasses import dataclass
from typing import Any
STREAM_PROTOCOL_CAPABILITIES = frozenset(
{
"task_snapshot_v1",
"complete_tool_call_v1",
"correlated_tool_call_id_v1",
"final_invalid_tool_call_v1",
}
)
@dataclass
class StreamEvent:
@@ -158,6 +167,14 @@ class StreamEventEmitter:
},
)
@staticmethod
def task_snapshot(source: str, items: list[dict[str, Any]]) -> StreamEvent:
"""Emit the complete root-agent task state without product-specific IDs."""
return StreamEvent(
"task_snapshot",
{"type": "task_snapshot", "source": source, "items": items},
)
@staticmethod
def interrupt(
interrupt_id: str,
@@ -213,6 +230,19 @@ class StreamEventEmitter:
)
@staticmethod
def error(message: str) -> StreamEvent:
def error(
message: str,
*,
code: str | None = None,
recoverable: bool | None = None,
details: dict[str, Any] | None = None,
) -> StreamEvent:
"""Error event."""
return StreamEvent("error", {"type": "error", "message": message})
data: dict[str, Any] = {"type": "error", "message": message}
if code is not None:
data["code"] = code
if recoverable is not None:
data["recoverable"] = recoverable
if details is not None:
data["details"] = details
return StreamEvent("error", data)
+276 -45
View File
@@ -16,7 +16,7 @@ from typing import Any, TypeAlias
from langchain_core._api import LangChainBetaWarning
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
from langgraph.graph import END
from langgraph.types import Command, Interrupt
from langgraph.types import Command, Interrupt, Overwrite
from ..memory.worker_activity import clear_completed_memory_activity_counts
from .emitter import StreamEventEmitter
@@ -95,9 +95,10 @@ async def _clear_interrupted_graph_state(
no output and leaves the messages channel unchanged. From the user's side the
conversation looks like it lost all history because the agent stops responding.
The fix: ``aupdate_state(config, None, as_node=END)`` clears all pending tasks
and writes a checkpoint whose ``next`` is the empty tuple, without touching
any channel values (message history is preserved).
Recovery first removes malformed/incomplete tool protocol from the messages
channel, then ``aupdate_state(config, None, as_node=END)`` clears pending
tasks and writes a checkpoint whose ``next`` is the empty tuple. Completed
tool call/result pairs and all non-tool history are preserved.
Critically, this only runs when the stuck state is *not* a legitimate
human-in-the-loop interrupt. The agent pauses via ``interrupt()`` /
@@ -114,10 +115,9 @@ async def _clear_interrupted_graph_state(
_log = logging.getLogger(__name__)
try:
snapshot = await agent.aget_state(config)
# Only act when the graph is genuinely stuck (non-empty next tuple)...
if not snapshot or not getattr(snapshot, "next", None):
if not snapshot:
return
# ...and not parked at a real human-in-the-loop interrupt.
# Never alter a real human-in-the-loop pause.
if _snapshot_has_pending_interrupt(snapshot):
_log.debug(
"Leaving interrupted graph state intact for thread %s: "
@@ -127,6 +127,13 @@ async def _clear_interrupted_graph_state(
)
return
await _repair_malformed_tool_history(agent, config, snapshot=snapshot)
# Only force END when the graph is genuinely stuck. Message repair also
# applies to failures that already left next empty.
if not getattr(snapshot, "next", None):
return
stuck_at = snapshot.next
await agent.aupdate_state(config, None, as_node=END)
_log.debug(
@@ -142,6 +149,49 @@ async def _clear_interrupted_graph_state(
)
async def _repair_malformed_tool_history(
agent: Any,
config: dict[str, Any],
*,
snapshot: Any | None = None,
) -> bool:
"""Rewrite a checkpoint's messages to a replay-safe tool history.
Only structurally invalid protocol is removed. Completed tool call/result
pairs, including repeated successes and repeated errors, are audit and
billing facts and must remain in persistent history.
"""
import logging
from ..llm.patches import _sanitize_openai_tool_history
_log = logging.getLogger(__name__)
if snapshot is None:
snapshot = await agent.aget_state(config)
if not snapshot or _snapshot_has_pending_interrupt(snapshot):
return False
values = getattr(snapshot, "values", None)
if not isinstance(values, Mapping):
return False
messages = values.get("messages")
if not isinstance(messages, list):
return False
repaired = _sanitize_openai_tool_history(messages)
if repaired == messages:
return False
await agent.aupdate_state(config, {"messages": Overwrite(repaired)})
_log.warning(
"Repaired structurally invalid tool history for thread %s: messages %d -> %d",
config.get("configurable", {}).get("thread_id", "?"),
len(messages),
len(repaired),
)
return True
@dataclass(frozen=True)
class _SubagentInfo:
path: tuple[str, ...]
@@ -217,7 +267,12 @@ class _V3EventProcessor:
tuple[tuple[str, ...], str], tuple[str, dict[str, Any]]
] = {}
self._emitted_tool_calls: set[tuple[tuple[str, ...], str]] = set()
self._pending_tool_calls: dict[
tuple[tuple[str, ...], str], tuple[str, dict[str, Any]]
] = {}
self._emitted_interrupts: set[str] = set()
self._pending_invalid_tool_calls: dict[str, tuple[str, str]] = {}
self._last_task_snapshot: tuple[tuple[str, str], ...] | None = None
self._selector = _ToolSelectionSuppressor(emitter)
@staticmethod
@@ -245,13 +300,26 @@ class _V3EventProcessor:
if method == "tools":
return self._process_tool_event(namespace, _event_data(event), subagent)
if method == "updates":
return self._process_update_event(_event_data(event))
return self._process_update_event(
_event_data(event), namespace=namespace, source="update"
)
if method == "values":
events: list[dict[str, Any]] = []
params = event.get("params") or {}
interrupts = params.get("interrupts") or ()
if interrupts:
events.extend(self._process_update_event({"__interrupt__": interrupts}))
events.extend(
self._process_update_event(
{"__interrupt__": interrupts},
namespace=namespace,
source="values",
)
)
events.extend(
self._process_update_event(
_event_data(event), namespace=namespace, source="values"
)
)
if self._process_value_message_snapshots and not namespace:
events.extend(self._process_value_messages(_event_data(event)))
return events
@@ -390,6 +458,7 @@ class _V3EventProcessor:
inp, out = _usage_counts(usage) if usage is not None else (0, 0)
if inp or out:
events.append(self.emitter.usage_stats(inp, out).data)
events.extend(self._flush_invalid_tool_calls())
return events
return []
@@ -415,14 +484,12 @@ class _V3EventProcessor:
if tool_call is None:
return events
tool_name, args, tool_call_id = tool_call
events.extend(
self._emit_tool_call_once(
namespace=namespace,
subagent=subagent,
name=tool_name,
args=args,
tool_call_id=tool_call_id,
)
self._pending_invalid_tool_calls.pop(tool_call_id, None)
self._pending_tool_calls[
(self._tool_scope(namespace, subagent), tool_call_id)
] = (
tool_name,
args,
)
return events
@@ -465,6 +532,20 @@ class _V3EventProcessor:
]
return [self.emitter.tool_call(name, args, tool_call_id).data]
def _pending_call_id(
self,
*,
scope: tuple[str, ...],
name: str,
args: dict[str, Any],
) -> str:
matches = [
call_id
for (candidate_scope, call_id), candidate in self._pending_tool_calls.items()
if candidate_scope == scope and candidate == (name, args)
]
return matches[0] if len(matches) == 1 else ""
def _process_whole_message(
self,
msg: AIMessage | AIMessageChunk,
@@ -472,6 +553,16 @@ class _V3EventProcessor:
namespace: tuple[str, ...],
) -> list[dict[str, Any]]:
events: list[dict[str, Any]] = []
for invalid in getattr(msg, "invalid_tool_calls", ()) or ():
invalid_map = _as_raw_map(invalid)
if invalid_map is None:
continue
call_id = str(
invalid_map.get("id") or invalid_map.get("tool_call_id") or ""
)
name = str(invalid_map.get("name") or invalid_map.get("tool_name") or "")
key = call_id or f"chunk_{len(self._pending_invalid_tool_calls)}"
self._pending_invalid_tool_calls[key] = (call_id, name)
additional = msg.additional_kwargs
reasoning = additional.get("reasoning_content")
emitted_reasoning = False
@@ -494,14 +585,12 @@ class _V3EventProcessor:
if tool_call is None:
continue
tool_name, args, tool_call_id = tool_call
events.extend(
self._emit_tool_call_once(
namespace=namespace,
subagent=subagent,
name=tool_name,
args=args,
tool_call_id=tool_call_id,
)
self._pending_invalid_tool_calls.pop(tool_call_id, None)
self._pending_tool_calls[
(self._tool_scope(namespace, subagent), tool_call_id)
] = (
tool_name,
args,
)
if subagent is None:
@@ -534,6 +623,9 @@ class _V3EventProcessor:
name,
args,
)
self._pending_tool_calls.pop(
(self._tool_scope(namespace, subagent), tool_call_id), None
)
events.extend(
self._emit_tool_call_once(
namespace=namespace,
@@ -580,6 +672,7 @@ class _V3EventProcessor:
content += "\n... (truncated)"
success = is_success(content)
lifecycle_key = (self._tool_scope(namespace, subagent), tool_call_id)
if subagent is not None:
events.append(
self.emitter.subagent_tool_result(
@@ -591,22 +684,78 @@ class _V3EventProcessor:
instance_id=subagent.instance_id,
).data
)
return events
events.append(
self.emitter.tool_result(
name, content, success, tool_call_id=tool_call_id
).data
)
else:
events.append(
self.emitter.tool_result(
name, content, success, tool_call_id=tool_call_id
).data
)
self._emitted_tool_calls.discard(lifecycle_key)
self._pending_tool_calls.pop(lifecycle_key, None)
return events
return []
def _process_update_event(self, data: object) -> list[dict[str, Any]]:
@staticmethod
def _normalize_task_items(value: object) -> list[dict[str, str]] | None:
if not isinstance(value, list):
return None
aliases = {
"todo": "pending",
"pending": "pending",
"active": "in_progress",
"in-progress": "in_progress",
"in_progress": "in_progress",
"done": "completed",
"completed": "completed",
}
items: list[dict[str, str]] = []
for raw in value:
raw_map = _as_raw_map(raw)
if raw_map is None:
continue
content = str(raw_map.get("content") or raw_map.get("task") or "").strip()
if not content:
continue
status = aliases.get(str(raw_map.get("status") or "pending").lower())
if status is None:
continue
items.append({"content": content, "status": status})
return items
@classmethod
def _find_task_items(cls, data: object) -> list[dict[str, str]] | None:
data_map = _as_raw_map(data)
if data_map is None:
return None
if "todos" in data_map:
return cls._normalize_task_items(data_map["todos"])
for value in data_map.values():
nested = _as_raw_map(value)
if nested is not None and "todos" in nested:
return cls._normalize_task_items(nested["todos"])
return None
def _process_update_event(
self,
data: object,
*,
namespace: tuple[str, ...] = (),
source: str = "update",
) -> list[dict[str, Any]]:
events: list[dict[str, Any]] = []
data_map = _as_raw_map(data)
if data_map is not None and "__interrupt__" in data_map:
events.extend(self._process_interrupts(data_map["__interrupt__"]))
if not namespace:
items = self._find_task_items(data)
if items is not None:
signature = tuple((item["content"], item["status"]) for item in items)
if signature != self._last_task_snapshot:
self._last_task_snapshot = signature
events.append(self.emitter.task_snapshot(source, items).data)
summarization_event = _find_summarization_event_payload(data)
if summarization_event and not self._summarization_in_progress:
signature = _summarization_event_signature(summarization_event)
@@ -622,6 +771,10 @@ class _V3EventProcessor:
events.extend(self._emit_summarization_text(summary_text))
return events
def _flush_invalid_tool_calls(self) -> list[dict[str, Any]]:
self._pending_invalid_tool_calls.clear()
return []
def _process_interrupts(self, interrupts: object) -> list[dict[str, Any]]:
events: list[dict[str, Any]] = []
if not isinstance(interrupts, list | tuple):
@@ -660,26 +813,74 @@ class _V3EventProcessor:
raw_questions = interrupt_map.get("questions")
questions = raw_questions if isinstance(raw_questions, list) else []
tc_id = str(interrupt_map.get("tool_call_id", ""))
return self._dedupe_interrupt_event(
self.emitter.ask_user_interrupt(
interrupt_id,
questions,
tc_id,
).data
events: list[dict[str, Any]] = []
candidate = self._pending_tool_calls.get(((), tc_id)) if tc_id else None
if candidate is not None:
events.extend(
self._emit_tool_call_once(
namespace=(),
subagent=None,
name=candidate[0],
args=candidate[1],
tool_call_id=tc_id,
)
)
events.extend(
self._dedupe_interrupt_event(
self.emitter.ask_user_interrupt(
interrupt_id,
questions,
tc_id,
).data
)
)
return events
raw_action_reqs = interrupt_map.get("action_requests")
action_reqs = raw_action_reqs if isinstance(raw_action_reqs, list) else []
raw_review_cfgs = interrupt_map.get("review_configs")
review_cfgs = raw_review_cfgs if isinstance(raw_review_cfgs, list) else None
if action_reqs:
return self._dedupe_interrupt_event(
self.emitter.interrupt(
interrupt_id,
action_reqs,
review_cfgs,
).data
events: list[dict[str, Any]] = []
for raw_request in action_reqs:
request_map = _as_raw_map(raw_request)
if request_map is None:
continue
call_id = str(
request_map.get("id") or request_map.get("tool_call_id") or ""
)
name = str(request_map.get("name") or request_map.get("tool_name") or "")
args_map = _as_raw_map(
request_map.get("args")
if "args" in request_map
else request_map.get("input")
)
if not call_id and name and args_map is not None:
call_id = self._pending_call_id(
scope=(),
name=name,
args=dict(args_map),
)
if call_id and name and args_map is not None:
events.extend(
self._emit_tool_call_once(
namespace=(),
subagent=None,
name=name,
args=dict(args_map),
tool_call_id=call_id,
)
)
events.extend(
self._dedupe_interrupt_event(
self.emitter.interrupt(
interrupt_id,
action_reqs,
review_cfgs,
).data
)
)
return events
return []
def _dedupe_interrupt_event(self, event: dict[str, Any]) -> list[dict[str, Any]]:
@@ -805,6 +1006,8 @@ async def stream_agent_events(
thread_id: str,
metadata: dict[str, Any] | None = None,
media: list[str] | None = None,
callbacks: list[Any] | None = None,
error_mode: str = "emit",
) -> AsyncGenerator[dict[str, Any], None]:
"""Stream events from a DeepAgents/LangGraph v3 run.
@@ -820,6 +1023,9 @@ async def stream_agent_events(
metadata: Optional metadata dict merged into the LangGraph config
(e.g. agent_name, updated_at for checkpoint persistence).
media: Optional list of local file paths for attachments.
callbacks: Optional Runnable callbacks propagated to all nested model calls.
error_mode: ``emit`` preserves the generic error event; ``raise`` lets an
embedding host produce the single terminal error envelope.
Yields:
Event dicts: thinking, text, tool_call, tool_result,
@@ -829,6 +1035,8 @@ async def stream_agent_events(
config: dict[str, Any] = {"configurable": {"thread_id": thread_id}}
if metadata:
config["metadata"] = metadata
if callbacks:
config["callbacks"] = callbacks
emitter = StreamEventEmitter()
existing_summarization_event: Mapping[str, object] | None = None
try:
@@ -957,7 +1165,30 @@ async def stream_agent_events(
yield item
except Exception as e:
_run_raised = True
yield emitter.error(str(e)).data
if error_mode == "emit":
payload = e.model_dump() if hasattr(e, "model_dump") else {}
if not isinstance(payload, Mapping):
payload = {}
code = str(payload.get("code") or "") or None
details = {
key: payload[key]
for key in (
"reason",
"provider",
"model",
"route_key",
"config_generation",
"api_mode",
"call_id",
)
if payload.get(key) is not None
}
yield emitter.error(
str(payload.get("message") or e),
code=code,
recoverable=bool(payload.get("recoverable", True)) if code else None,
details=details or None,
).data
raise
finally:
if stream is not None:
+25
View File
@@ -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,
)
+66
View File
@@ -69,3 +69,69 @@ def test_create_cli_agent_accepts_host_backend_and_memory_options(
assert calls["middleware_kwargs"]["memory_max_inline_profile_chars"] == 1000
assert calls["middleware_kwargs"]["enable_background_execution"] is False
assert calls["agent_config"] == {"recursion_limit": 321}
def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, tmp_path):
import EvoScientist.EvoScientist as agent_module
from EvoScientist.config.settings import EvoScientistConfig
calls = {}
class _Middleware:
def __init__(self, name):
self.name = name
class _Backend:
def __init__(self, **_kwargs):
pass
class _CompositeBackend:
def __init__(self, **_kwargs):
pass
class _Agent:
def with_config(self, _config):
return self
default_chain = [
_Middleware("error_normalization"),
_Middleware("configurable_model"),
_Middleware("context_editing"),
_Middleware("tool_protocol_guard"),
]
route = _Middleware("gateway_route_fallback")
monkeypatch.setattr("deepagents.backends.CompositeBackend", _CompositeBackend)
monkeypatch.setattr("deepagents.create_deep_agent", lambda **_kwargs: _Agent())
monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _Backend)
monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _Backend)
monkeypatch.setattr(agent_module, "set_active_workspace", lambda _path: None)
def fake_default_middleware(**kwargs):
calls["middleware_kwargs"] = kwargs
return list(default_chain)
monkeypatch.setattr(agent_module, "_get_default_middleware", fake_default_middleware)
def fake_load(_backend, middleware, **_kwargs):
calls["middleware"] = middleware
return {"subagents": []}
monkeypatch.setattr(agent_module, "load_mcp_and_build_kwargs", fake_load)
agent_module.create_cli_agent(
workspace_dir=str(tmp_path),
checkpointer=object(),
config=EvoScientistConfig(auto_approve=True),
chat_model=object(),
workspace_backend=object(),
main_agent_route_middleware=route,
)
assert calls["middleware_kwargs"]["enable_legacy_model_fallback"] is False
assert [middleware.name for middleware in calls["middleware"][:5]] == [
"error_normalization",
"configurable_model",
"gateway_route_fallback",
"context_editing",
"tool_protocol_guard",
]
+65 -2
View File
@@ -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
+292 -1
View File
@@ -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
+123 -6
View File
@@ -55,11 +55,6 @@ def temp_config_dir(tmp_path, monkeypatch):
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",
@@ -79,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)
@@ -106,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)
@@ -131,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
@@ -147,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."""
@@ -194,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
@@ -694,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())
@@ -754,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.
@@ -871,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
@@ -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"
+6
View File
@@ -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",
+13
View File
@@ -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"
@@ -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
+437 -2
View File
@@ -535,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
@@ -1099,6 +1293,107 @@ class TestPatchOpenAICompatContent:
assert len(set(call_ids)) == 2
assert [message.tool_call_id for message in first[1:]] == call_ids
def test_content_tool_block_is_normalized_to_parsed_call(self):
from langchain_core.messages import AIMessage, ToolMessage
from EvoScientist.llm.patches import _ensure_openai_tool_call_ids
normalized = _ensure_openai_tool_call_ids(
[
AIMessage(
content=[
{
"type": "tool_call",
"id": "wrong-id",
"name": "wrong-name",
"args": {},
}
],
tool_calls=[{"id": "call-1", "name": "execute", "args": {}}],
),
ToolMessage(content="ok", tool_call_id="call-1"),
]
)
assert normalized[0].content[0]["id"] == "call-1"
assert normalized[0].content[0]["name"] == "execute"
def test_invalid_tool_call_is_not_replayed_to_responses_api(self):
from langchain_core.messages import AIMessage, HumanMessage
from langchain_openai.chat_models.base import _construct_responses_api_input
from EvoScientist.llm.patches import _sanitize_messages
invalid = AIMessage(
content=[
{"type": "reasoning", "reasoning": "partial"},
{
"type": "tool_call",
"id": None,
"name": "execute",
"args": '{"command":',
},
],
invalid_tool_calls=[
{
"type": "invalid_tool_call",
"id": None,
"name": "execute",
"args": '{"command":',
"error": "Failed to parse tool call arguments as JSON",
}
],
)
normalized = _sanitize_messages([invalid, HumanMessage(content="retry")])
payload = _construct_responses_api_input(normalized)
assert all(item.get("type") != "function_call" for item in payload)
assert [message.type for message in normalized] == ["human"]
def test_invalid_tool_call_preserves_replayable_assistant_text(self):
from langchain_core.messages import AIMessage
from EvoScientist.llm.patches import _sanitize_messages
invalid = AIMessage(
content="I could not finish the tool request.",
invalid_tool_calls=[
{
"type": "invalid_tool_call",
"id": None,
"name": "execute",
"args": "{",
"error": "bad json",
}
],
)
normalized = _sanitize_messages([invalid])
assert len(normalized) == 1
assert normalized[0].content == "I could not finish the tool request."
assert normalized[0].invalid_tool_calls == []
def test_orphan_tool_results_and_unanswered_calls_are_removed(self):
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from EvoScientist.llm.patches import _sanitize_messages
messages = [
ToolMessage(content="orphan", tool_call_id="missing"),
AIMessage(
content="waiting",
tool_calls=[{"id": "call_unanswered", "name": "execute", "args": {}}],
),
HumanMessage(content="continue"),
]
normalized = _sanitize_messages(messages)
assert [message.type for message in normalized] == ["ai", "human"]
assert normalized[0].tool_calls == []
def test_generate_flattened(self):
from langchain_core.messages import HumanMessage
@@ -2380,6 +2675,11 @@ 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."""
@@ -2496,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"] == {
@@ -2515,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."""
@@ -2546,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
@@ -2580,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(
+113
View File
@@ -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()
+5 -1
View File
@@ -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
+177
View File
@@ -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)
+242
View File
@@ -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 "<redacted>" 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("<redacted>") == 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 "<redacted>" 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 "<redacted>" 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 "<redacted>" 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,
}
+69
View File
@@ -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."""
+125 -1
View File
@@ -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",
]
+246
View File
@@ -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:")