From 88ac9f5ba1eeb5a89587f9ee551925952404f37a Mon Sep 17 00:00:00 2001 From: jfilipiuk Date: Mon, 13 Jul 2026 15:17:56 +0200 Subject: [PATCH] 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 --- EvoScientist/EvoScientist.py | 12 + EvoScientist/llm/errors.py | 294 +++++++++++++ EvoScientist/memory/agents/memory_worker.py | 28 +- .../memory/agents/observation_linker.py | 6 +- EvoScientist/middleware/__init__.py | 2 + .../middleware/error_normalization.py | 230 ++++++++++ EvoScientist/middleware/model_fallback.py | 32 +- EvoScientist/middleware/tool_selector.py | 24 +- tests/test_error_normalization_middleware.py | 400 ++++++++++++++++++ tests/test_model_fallback.py | 112 +++++ tests/test_observation_memory.py | 6 +- tests/test_serde_default_rich_exception.py | 242 +++++++++++ 12 files changed, 1369 insertions(+), 19 deletions(-) create mode 100644 EvoScientist/llm/errors.py create mode 100644 EvoScientist/middleware/error_normalization.py create mode 100644 tests/test_error_normalization_middleware.py create mode 100644 tests/test_serde_default_rich_exception.py diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index 326e1e1..7239472 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -305,6 +305,7 @@ def _inject_subagent_middleware( """ from .middleware import ( ContextOverflowMapperMiddleware, + ErrorNormalizationMiddleware, ToolErrorHandlerMiddleware, create_context_editing_middleware, create_memory_lifecycle_middleware, @@ -333,6 +334,11 @@ 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(), # 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``). @@ -667,6 +673,7 @@ def _get_default_middleware( from .middleware import ( ConfigurableModelMiddleware, ContextOverflowMapperMiddleware, + ErrorNormalizationMiddleware, ModelFallbackMiddleware, ToolErrorHandlerMiddleware, create_code_interpreter_middleware, @@ -729,6 +736,11 @@ def _get_default_middleware( tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider) 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(), diff --git a/EvoScientist/llm/errors.py b/EvoScientist/llm/errors.py new file mode 100644 index 0000000..1140f12 --- /dev/null +++ b/EvoScientist/llm/errors.py @@ -0,0 +1,294 @@ +"""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 ProviderStreamError(Exception): + """Envelope-shaped wrapper for a provider SDK exception raised + inside a chat model call. + + Attributes mirror the SSE envelope one-for-one: + + - ``provider`` — concrete provider tag (``openai`` / ``anthropic`` + / ``deepseek`` / ``openrouter`` / ``openai_compat`` / …) + - ``class_qualname`` — fully qualified name of the underlying + exception's class (e.g. ``openrouter.errors.…``) + - ``message`` — API-key-redacted ``str(exc)`` + - ``status_code`` — HTTP status if the SDK exposed one + - ``code`` — provider error code (``insufficient_quota``, …) + - ``err_type`` — provider error type label (openai's ``.type``) + - ``request_id`` — SDK-provided correlation id + + The underlying exception is available via ``__cause__`` (set by + ``raise ProviderStreamError(...) from exc`` in the middleware). + """ + + def __init__( + self, + provider: str, + class_qualname: str, + message: str, + *, + status_code: int | None = None, + code: str | None = None, + err_type: str | None = None, + request_id: str | None = None, + ) -> None: + super().__init__(message) + self.provider = provider + self.class_qualname = class_qualname + self.message = message + self.status_code = status_code + self.code = code + self.err_type = err_type + self.request_id = request_id + + def as_envelope(self) -> dict[str, Any]: + """Return the SSE envelope dict — the shape the WebUI consumes.""" + payload: dict[str, Any] = { + "error": self.class_qualname.rsplit(".", 1)[-1], + "class": self.class_qualname, + "message": self.message, + "provider": self.provider, + } + if self.status_code is not None: + payload["status_code"] = self.status_code + if self.code is not None: + payload["code"] = self.code + if self.err_type is not None: + payload["type"] = self.err_type + if self.request_id: + payload["request_id"] = self.request_id + return payload + + def model_dump(self) -> dict[str, Any]: + """Serialization hook consumed by ``langgraph_api.serde.default``. + + Upstream's dispatch checks ``hasattr(obj, 'model_dump')`` BEFORE + the ``isinstance(obj, BaseException)`` branch, so exposing this + method lets upstream emit our envelope with no monkey-patch on + its ``default`` callable. The name matches Pydantic's + convention deliberately — it's the hook upstream is looking + for. + """ + return self.as_envelope() + + +# --------------------------------------------------------------------------- +# API-key redaction — env-driven, prefix-only +# --------------------------------------------------------------------------- +# +# Redaction is built from credentials actually deployed via env vars, +# not from generic key shapes. Rationale: (a) zero false positives — +# we only scrub strings we know are secrets, (b) defense-in-depth — +# the compiled regex holds only the first 8 chars of each key, so a +# leak of the regex object itself (traceback locals, process dump) +# can't expose the secret. Suffix-greedy match consumes the rest of +# the key shape at runtime. The table is rebuilt on every +# ``_redact_api_keys`` call so credentials loaded after import +# (typically ``load_dotenv`` in a main entry point) still get +# scrubbed. ``re.compile`` caches by source string internally, so an +# unchanged env costs a dict lookup. + +_API_KEY_ENV_SUFFIXES = ("_API_KEY", "_TOKEN", "_SECRET") +_API_KEY_MIN_LEN = 12 +_API_KEY_PREFIX_LEN = 8 + + +def _build_env_key_redaction_re() -> re.Pattern[str] | None: + prefixes: list[str] = [] + for k, v in os.environ.items(): + if not k.endswith(_API_KEY_ENV_SUFFIXES): + continue + if not isinstance(v, str) or len(v) < _API_KEY_MIN_LEN: + continue + prefixes.append(re.escape(v[:_API_KEY_PREFIX_LEN])) + if not prefixes: + return None + alternation = "|".join(f"{p}[A-Za-z0-9_+/=.-]*" for p in prefixes) + return re.compile(alternation) + + +def _redact_api_keys(message: str) -> str: + """Replace any deployed key prefix in *message* with ````. + + Defensive; provider error messages occasionally echo the + authorization header back. Rebuilt per call so credentials loaded + after import (typical ``load_dotenv`` pattern) are still redacted. + """ + pattern = _build_env_key_redaction_re() + if pattern is None: + return message + return pattern.sub("", message) + + +# --------------------------------------------------------------------------- +# Provider inference from ModelRequest.model +# --------------------------------------------------------------------------- +# +# Host → concrete provider. Hand-maintained snapshot mirroring the +# routed-provider tables in ``llm/models.py`` +# (``_OPENAI_ROUTED_PROVIDERS`` + ``_ANTHROPIC_ROUTED_PROVIDERS``). +# Kept here rather than imported from ``models.py`` to keep the +# import surface of ``errors.py`` minimal — importing ``models.py`` +# would pull in every langchain chat-model client at first +# middleware access. Consumed by ``_lookup_host_or_compat``; unknown +# hosts fall back to ``_compat`` so the WebUI knows +# "openai/anthropic SDK, but not native" instead of getting a +# misleading concrete tag. Update when a new routed provider is +# added to ``models.py``. +# +# Related sibling: ``_PROVIDER_EXC_MODULE_PREFIXES`` in +# ``middleware/error_normalization.py`` — the exception-side +# provider allow-list. Adding a whole new provider SDK (not just a +# new base_url routed through an existing one) means updating that +# list too. + +_HOST_TO_PROVIDER: dict[str, str] = { + "api.openai.com": "openai", + "api.anthropic.com": "anthropic", + "api.deepseek.com": "deepseek", + "api.moonshot.cn": "moonshot", + "api.siliconflow.cn": "siliconflow", + "open.bigmodel.cn": "zhipu", # zhipu + zhipu-code share this host + "ark.cn-beijing.volces.com": "volcengine", + "dashscope.aliyuncs.com": "dashscope", + "coding.dashscope.aliyuncs.com": "dashscope", + "api.minimaxi.com": "minimax", + "api.kimi.com": "kimi", # kimi-coding shares this host + "openrouter.ai": "openrouter", +} + + +def _provider_from_model(model: Any) -> str | None: + """Derive the concrete provider tag from a chat model instance. + + Class-based dispatch for unambiguous providers (``ChatOpenRouter``, + ``ChatGoogleGenerativeAI``); ``openai_api_base`` / + ``anthropic_api_url`` looked up in ``_HOST_TO_PROVIDER`` for + openai/anthropic-shape clients (native + routed). Returns ``None`` + when the model isn't from a recognized provider SDK — the caller + (``ErrorNormalizationMiddleware``) then passes the exception + through unchanged. + """ + cls_module = type(model).__module__ or "" + if cls_module.startswith("langchain_openrouter"): + return "openrouter" + if cls_module.startswith("langchain_google_genai"): + return "google_genai" + if cls_module.startswith("langchain_openai"): + return _lookup_host_or_compat( + getattr(model, "openai_api_base", None), module_tag="openai" + ) + if cls_module.startswith("langchain_anthropic"): + return _lookup_host_or_compat( + getattr(model, "anthropic_api_url", None), module_tag="anthropic" + ) + return None + + +def _lookup_host_or_compat(base_url: str | None, module_tag: str) -> str: + """Extract host from *base_url* and look up in ``_HOST_TO_PROVIDER``. + + Falls back to *module_tag* when no ``base_url`` is set (native SDK + default endpoint) or ``_compat`` for an unrecognized + host — the honest "openai SDK shape but unknown upstream" tag. + """ + if not base_url: + return module_tag + try: + from urllib.parse import urlparse + + host = urlparse(base_url).hostname + except Exception: + host = None + if not host: + return module_tag + return _HOST_TO_PROVIDER.get(host.lower(), f"{module_tag}_compat") + + +# --------------------------------------------------------------------------- +# SDK-field extractors — populate the envelope's optional fields +# --------------------------------------------------------------------------- + + +def _extract_status_code(exc: BaseException) -> int | None: + """Best-effort HTTP status code from a provider SDK exception. + + Order matters: openai/anthropic store it on ``.status_code``; + httpx-wrappers expose it via ``.response.status_code``; + ``google.genai.errors.APIError`` (unusually) stores it as an + integer ``.code`` — type-disambiguated from openai/anthropic's + string ``.code`` (provider error code, surfaced separately). + """ + status_code = getattr(exc, "status_code", None) + if isinstance(status_code, int): + return status_code + response = getattr(exc, "response", None) + if response is not None: + rsc = getattr(response, "status_code", None) + if isinstance(rsc, int): + return rsc + code = getattr(exc, "code", None) + if isinstance(code, int): + return code + return None + + +def _extract_provider_code(exc: BaseException) -> str | None: + """Provider error code (e.g. ``insufficient_quota``, + ``invalid_api_key``). Distinct from HTTP status; higher signal for + a WebUI toast than the integer alone. + """ + code = getattr(exc, "code", None) + if isinstance(code, str) and code: + return code + return None + + +def _extract_error_type(exc: BaseException) -> str | None: + """Provider error type label. + + - openai exposes this as ``.type`` (``rate_limit_error`` etc.) + - ``google.genai.errors.APIError`` stores a string label at + ``.status`` (``"NOT_FOUND"``, ``"RESOURCE_EXHAUSTED"``, …) — a + good fit for the same field. + + ``.type`` takes precedence when both are set. + """ + err_type = getattr(exc, "type", None) + if isinstance(err_type, str) and err_type: + return err_type + status = getattr(exc, "status", None) + if isinstance(status, str) and status: + return status + return None diff --git a/EvoScientist/memory/agents/memory_worker.py b/EvoScientist/memory/agents/memory_worker.py index 9fa0e19..5009d72 100644 --- a/EvoScientist/memory/agents/memory_worker.py +++ b/EvoScientist/memory/agents/memory_worker.py @@ -428,6 +428,7 @@ def _memory_worker_middleware( enable_observation_memory: bool = True, ): """Build middleware for memory workers, excluding task execution tools.""" + from ...middleware.error_normalization import ErrorNormalizationMiddleware from ...middleware.memory import create_memory_middleware memory_controls = MemoryControls( @@ -439,18 +440,23 @@ def _memory_worker_middleware( enable_observation_tool = memory_controls.observation_tool_enabled( _memory_worker_observation_target(source_type) ) - return memory_agent_middleware( - create_memory_middleware( - str(memory_dir), - workspace_dir=workspace_dir, - source_type=source_type, - source_agent=_memory_worker_agent_name(source_type), - enable_profile_memory=enable_profile_memory, - enable_observation_memory=enable_observation_memory, - enable_observation_tool=enable_observation_tool, + return [ + # Outermost — normalize provider-SDK exceptions from the + # auxiliary model call before any inner middleware sees them. + ErrorNormalizationMiddleware(), + *memory_agent_middleware( + create_memory_middleware( + str(memory_dir), + workspace_dir=workspace_dir, + source_type=source_type, + source_agent=_memory_worker_agent_name(source_type), + enable_profile_memory=enable_profile_memory, + enable_observation_memory=enable_observation_memory, + enable_observation_tool=enable_observation_tool, + ), + excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS, ), - excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS, - ) + ] def _build_memory_worker_agent( diff --git a/EvoScientist/memory/agents/observation_linker.py b/EvoScientist/memory/agents/observation_linker.py index fb170c0..044158b 100644 --- a/EvoScientist/memory/agents/observation_linker.py +++ b/EvoScientist/memory/agents/observation_linker.py @@ -71,6 +71,8 @@ def build_observation_linker_graph( workspace_dir: str | Path | None = None, ) -> CompiledStateGraph: """Build the registered LangGraph observation linker.""" + from ...middleware.error_normalization import ErrorNormalizationMiddleware + agent_paths = resolve_memory_agent_paths( memory_dir=memory_dir, workspace_dir=workspace_dir, @@ -85,5 +87,7 @@ def build_observation_linker_graph( tools=tools, memory_dir=agent_paths.memory_dir, workspace_dir=agent_paths.workspace_dir, - middleware=memory_agent_middleware(), + # Outermost — normalize provider-SDK exceptions from the + # auxiliary model call before any inner middleware sees them. + middleware=[ErrorNormalizationMiddleware(), *memory_agent_middleware()], ) diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py index cd28de6..bb8afcb 100644 --- a/EvoScientist/middleware/__init__.py +++ b/EvoScientist/middleware/__init__.py @@ -18,6 +18,7 @@ from .context_editing import ( create_context_editing_middleware, ) from .context_overflow import ContextOverflowMapperMiddleware +from .error_normalization import ErrorNormalizationMiddleware from .memory import ( EvoMemoryMiddleware, create_memory_middleware, @@ -44,6 +45,7 @@ __all__ = [ "Choice", "ConfigurableModelMiddleware", "ContextOverflowMapperMiddleware", + "ErrorNormalizationMiddleware", "EvoMemoryLifecycleMiddleware", "EvoMemoryMiddleware", "ModelFallbackMiddleware", diff --git a/EvoScientist/middleware/error_normalization.py b/EvoScientist/middleware/error_normalization.py new file mode 100644 index 0000000..3ec816d --- /dev/null +++ b/EvoScientist/middleware/error_normalization.py @@ -0,0 +1,230 @@ +"""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. Every exception raised inside a call to a recognized +provider model gets wrapped — SDK exceptions, httpx errors, +langchain-wrapper failures, and even builtins like ``RuntimeError``. +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. + - Models we don't recognize as a provider SDK. + """ + from langchain_core.exceptions import ContextOverflowError + + from ..llm.errors import ( + 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 + + # LangGraph control-flow / structural signals must propagate + # untouched, regardless of which caller invoked us. + if _should_pass_through(exc): + return None + + # SummarizationMiddleware sits outside our stack and catches this + # by exact type to trigger reactive history compression + retry. + if isinstance(exc, ContextOverflowError): + return None + + provider = _provider_from_model(getattr(request, "model", None)) + if provider is None: + return None + cls = type(exc) + mod = cls.__module__ or "" + class_qualname = f"{mod}.{cls.__qualname__}" if mod else cls.__qualname__ + + request_id_attr = getattr(exc, "request_id", None) + request_id = ( + request_id_attr + if isinstance(request_id_attr, str) and request_id_attr + else None + ) + + return ProviderStreamError( + provider=provider, + class_qualname=class_qualname, + message=_redact_api_keys(str(exc)), + status_code=_extract_status_code(exc), + code=_extract_provider_code(exc), + err_type=_extract_error_type(exc), + request_id=request_id, + ) + + +class ErrorNormalizationMiddleware(AgentMiddleware): + """Wrap the model call in try/except and normalize provider SDK + exceptions into a non-dataclass envelope wrapper. + + Place this middleware **outermost** in the chain (first in the + middleware list) so it catches exceptions raised by inner + middlewares as well as the model handler itself. + """ + + name = "error_normalization" + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + try: + return handler(request) + except Exception as exc: + normalized = _normalize(request, exc) + if normalized is None: + raise + raise normalized from exc + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + try: + return await handler(request) + except Exception as exc: + normalized = _normalize(request, exc) + if normalized is None: + raise + raise normalized from exc diff --git a/EvoScientist/middleware/model_fallback.py b/EvoScientist/middleware/model_fallback.py index 0dee6ff..29a0de7 100644 --- a/EvoScientist/middleware/model_fallback.py +++ b/EvoScientist/middleware/model_fallback.py @@ -263,7 +263,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 +296,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 +312,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 +356,7 @@ def _guard_and_fallback( f"Model error ({reason}) -- not eligible for fallback, re-raising", style="red", ) - raise primary_exc + _raise_normalized(request, primary_exc) return _try_fallbacks(request, invoke, primary_exc) diff --git a/EvoScientist/middleware/tool_selector.py b/EvoScientist/middleware/tool_selector.py index 3e28432..b498d0f 100644 --- a/EvoScientist/middleware/tool_selector.py +++ b/EvoScientist/middleware/tool_selector.py @@ -132,10 +132,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 +182,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 diff --git a/tests/test_error_normalization_middleware.py b/tests/test_error_normalization_middleware.py new file mode 100644 index 0000000..62649bf --- /dev/null +++ b/tests/test_error_normalization_middleware.py @@ -0,0 +1,400 @@ +"""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 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 + + +# --------------------------------------------------------------------------- +# _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_wraps_any_exception_from_recognized_model(self): + """Any exception raised inside a call to a provider-recognized + model gets wrapped — including builtins like ``RuntimeError``. + Rationale: at the middleware boundary we can tell the model is + a provider, but not the exception's origin (SDK vs + langchain-wrapper vs httpx vs our code). Wrapping uniformly + gives the WebUI a consistent envelope; upstream's + ``RuntimeError``-whitelist would emit ``{"error": + "RuntimeError", "message": str(exc)}`` which isn't more + useful. + """ + raised = RuntimeError("internal glitch") + + async def handler(_req): + raise raised + + req = _request(_openai_model()) + mw = ErrorNormalizationMiddleware() + with pytest.raises(ProviderStreamError) as excinfo: + self._run_awrap(mw, req, handler) + assert excinfo.value.provider == "openai" + assert excinfo.value.__cause__ is raised + assert excinfo.value.class_qualname == "builtins.RuntimeError" + + def test_sync_wrap_normalizes_provider_exception(self): + raised = _make_exc(message="boom") + + def handler(_req): + raise raised + + req = _request(_openrouter_model()) + mw = ErrorNormalizationMiddleware() + with pytest.raises(ProviderStreamError) as excinfo: + mw.wrap_model_call(request=req, handler=handler) + assert excinfo.value.provider == "openrouter" + + def test_success_path_returns_handler_result(self): + async def handler(_req): + return "ok" + + req = _request(_openrouter_model()) + mw = ErrorNormalizationMiddleware() + assert self._run_awrap(mw, req, handler) == "ok" diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py index fa9e346..30f1a03 100644 --- a/tests/test_model_fallback.py +++ b/tests/test_model_fallback.py @@ -6,6 +6,7 @@ fallback chain behaviour via _try_fallbacks / _guard_and_fallback. from __future__ import annotations +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -223,6 +224,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 +326,34 @@ class TestGuardAndFallback: invoke.assert_not_awaited() + async def test_context_overflow_with_provider_model_passes_through_unwrapped(self): + """Regression: a ``ContextOverflowError`` entering + ``_guard_and_fallback`` under a recognized-provider model must + come out unwrapped. Otherwise ``_raise_normalized`` → + ``_normalize`` would wrap it as a ``ProviderStreamError`` and + deepagents' ``SummarizationMiddleware`` (which sits outside + the user middleware stack and catches by exact type) would + stop compressing history and retrying. + """ + add_fallback("fb", "prov") + req = _fake_request() + # Recognized provider — without the gate in ``_normalize`` this + # would wrap. With the gate, the raw type propagates. + cls = type( + "ChatOpenAI", (), {"__module__": "langchain_openai.chat_models.base"} + ) + model = cls() + model.openai_api_base = None + req.model = model + invoke = AsyncMock() + + raised = ContextOverflowError("context length exceeded") + with pytest.raises(ContextOverflowError) as exc_info: + await _guard_and_fallback(raised, req, invoke) + + assert exc_info.value is raised + invoke.assert_not_awaited() + async def test_malformed_400_raises_immediately(self): add_fallback("fb", "prov") req = _fake_request() diff --git a/tests/test_observation_memory.py b/tests/test_observation_memory.py index 5989034..3db0d79 100644 --- a/tests/test_observation_memory.py +++ b/tests/test_observation_memory.py @@ -2294,7 +2294,11 @@ def test_memory_worker_observation_writer_modes( observation_writer=observation_writer, ) - assert type(middleware[0]).__name__ == "ToolErrorHandlerMiddleware" + # ErrorNormalizationMiddleware wraps outermost so provider-SDK + # exceptions from the auxiliary model call get normalized before + # the tool-error handler sees them. + assert type(middleware[0]).__name__ == "ErrorNormalizationMiddleware" + assert type(middleware[1]).__name__ == "ToolErrorHandlerMiddleware" assert _memory_tool_names(middleware) == expected_tools diff --git a/tests/test_serde_default_rich_exception.py b/tests/test_serde_default_rich_exception.py new file mode 100644 index 0000000..aa2fdbb --- /dev/null +++ b/tests/test_serde_default_rich_exception.py @@ -0,0 +1,242 @@ +"""Regression tests for the helpers ``ErrorNormalizationMiddleware`` +uses to build the SSE error envelope. + +- ``_redact_api_keys`` + ``_build_env_key_redaction_re`` — scrubs + deployed credentials that the SDK might echo back. +- ``_extract_status_code`` / ``_extract_provider_code`` / + ``_extract_error_type`` — read SDK-specific fields off the raised + exception. + +Middleware wire behavior + ``_provider_from_model`` live in +``test_error_normalization_middleware.py``. One end-to-end orjson test +at the bottom guards that a ``ProviderStreamError`` survives +langgraph_api's UNPATCHED ``serde.default`` under +``OPT_SERIALIZE_DATACLASS`` — the whole reason the wrapper exists. +""" + +from __future__ import annotations + +import os + +import langgraph_api.serde as _serde_mod + +from EvoScientist.llm.errors import ( + _API_KEY_ENV_SUFFIXES, + _build_env_key_redaction_re, + _extract_error_type, + _extract_provider_code, + _extract_status_code, + _redact_api_keys, +) + +# --------------------------------------------------------------------------- +# Redaction +# --------------------------------------------------------------------------- + + +def test_env_deployed_key_redacted_in_message(monkeypatch): + """A credential exported via env var is scrubbed by + ``_redact_api_keys``. The redaction table is rebuilt per call — + ``monkeypatch.setenv`` alone is enough, no attribute reassignment. + """ + key = "sk-proj-aBcDeFgHiJkLmNoPqRsTuVwXyZ1234567890" + monkeypatch.setenv("OPENAI_API_KEY", key) + + msg = ( + f"Invalid API key: {key}. Get a new one at https://platform.openai.com/api-keys" + ) + redacted = _redact_api_keys(msg) + assert key not in redacted + assert "" in redacted + assert "Invalid API key" in redacted + assert "platform.openai.com" in redacted + + +def test_multiple_env_keys_redacted_independently(monkeypatch): + """Each ``*_API_KEY`` / ``*_TOKEN`` / ``*_SECRET`` env var + contributes its own prefix to the alternation. + """ + k1 = "sk-or-aBcDeFg012345678901234" + k2 = "AIzaABCDEFGHIJ0123456789" + k3 = "ghp_p4t70k3n0123456789abcdef" + monkeypatch.setenv("OPENROUTER_API_KEY", k1) + monkeypatch.setenv("GOOGLE_API_KEY", k2) + monkeypatch.setenv("GITHUB_TOKEN", k3) + + msg = _redact_api_keys(f"Failures: {k1}, {k2}, {k3}") + assert k1 not in msg + assert k2 not in msg + assert k3 not in msg + assert msg.count("") == 3 + + +def test_base64_suffix_secret_fully_redacted(monkeypatch): + """A base64-style secret (``/`` ``+`` ``=``) must redact end-to-end, + not leak its tail past the first padding char. + """ + key = "AbCdEfGh/secret+tail==" + monkeypatch.setenv("SOME_SECRET", key) + + msg = _redact_api_keys(f"auth failed with token={key} on retry") + assert "secret" not in msg + assert "tail" not in msg + assert "" in msg + assert "auth failed" in msg + assert "on retry" in msg + + +def test_unknown_shape_not_redacted_without_env(monkeypatch): + """Env-only redaction: a key-shaped string not deployed via env is + left alone. Tradeoff — we only scrub what we know is a secret. + """ + for k in list(os.environ): + if k.endswith(_API_KEY_ENV_SUFFIXES): + monkeypatch.delenv(k, raising=False) + + msg = _redact_api_keys("Unknown key seen: sk-or-aBcDeFg012345678901234") + assert "sk-or-aBcDeFg012345678901234" in msg + assert "" not in msg + + +def test_env_key_loaded_after_first_call_is_redacted(monkeypatch): + """The pattern rebuilds every call so keys loaded after + ``patches.py`` imports (typical ``load_dotenv`` sequence) are + still scrubbed on the next call. + """ + for k in list(os.environ): + if k.endswith(_API_KEY_ENV_SUFFIXES): + monkeypatch.delenv(k, raising=False) + + key = "sk-proj-loaded_after_import_1234567890abcdef" + # Pass 1: env empty — key leaks. + assert key in _redact_api_keys(f"leak: {key}") + + # Pass 2: after simulated load_dotenv. + monkeypatch.setenv("OPENAI_API_KEY", key) + redacted = _redact_api_keys(f"leak: {key}") + assert key not in redacted + assert "" in redacted + + +def test_redaction_regex_holds_only_prefix(monkeypatch): + """Defense-in-depth: the compiled regex must not embed the full key. + A process-memory leak (traceback locals, debugger) exposes at most + the first 8 chars — not the secret. + """ + key = "sk-proj-aBcDeFgHiJkLmNoPqRsTuVwXyZ1234567890_secret_suffix" + monkeypatch.setenv("OPENAI_API_KEY", key) + pattern = _build_env_key_redaction_re() + assert pattern is not None + assert key not in pattern.pattern + assert "aBcDeFgHiJkLmNoPqRs" not in pattern.pattern + # Sanity: still matches the full key at runtime via prefix + suffix + # greedy. + m = pattern.search(f"err: {key}") + assert m is not None + assert m.group(0) == key + + +# --------------------------------------------------------------------------- +# Field extractors +# --------------------------------------------------------------------------- + + +def _fake_exc(**attrs): + return type("APIError", (Exception,), attrs)("boom") + + +def test_status_code_read_from_direct_attribute(): + """openai / anthropic ``APIStatusError`` carries integer + ``.status_code`` — the primary path. + """ + assert _extract_status_code(_fake_exc(status_code=429)) == 429 + + +def test_status_code_read_via_response_attribute(): + """Wrappers that don't promote status to top level expose it via + ``.response.status_code`` (httpx pattern). + """ + + class FakeResponse: + status_code = 504 + + assert _extract_status_code(_fake_exc(response=FakeResponse())) == 504 + + +def test_status_code_read_via_integer_code_attribute(): + """``google.genai.errors.APIError`` stores HTTP status as integer + ``.code`` — type-disambiguated from openai/anthropic's string + ``.code`` (provider error code). + """ + assert _extract_status_code(_fake_exc(code=400)) == 400 + + +def test_provider_code_read_from_string_code_attribute(): + """Provider error code (``insufficient_quota`` etc.) is a string + ``.code`` — higher signal than the integer HTTP status alone. + """ + assert ( + _extract_provider_code(_fake_exc(code="insufficient_quota")) + == "insufficient_quota" + ) + + +def test_provider_code_ignores_integer_code(): + """An integer ``.code`` is HTTP status (see above); must not bleed + into the provider-code path. + """ + assert _extract_provider_code(_fake_exc(code=429)) is None + + +def test_error_type_read_from_type_attribute(): + """openai exposes a ``.type`` label (``rate_limit_error``).""" + assert _extract_error_type(_fake_exc(type="rate_limit_error")) == "rate_limit_error" + + +def test_extractors_return_none_when_attributes_absent(): + """A bare exception with no SDK-shape attributes — every extractor + returns None so the envelope drops the optional fields. + """ + exc = _fake_exc() + assert _extract_status_code(exc) is None + assert _extract_provider_code(exc) is None + assert _extract_error_type(exc) is None + + +# --------------------------------------------------------------------------- +# End-to-end: ProviderStreamError survives orjson under +# OPT_SERIALIZE_DATACLASS via upstream's UNPATCHED serde.default. +# --------------------------------------------------------------------------- + + +def test_provider_stream_error_survives_orjson_dataclass_option(): + """Guard: ``ProviderStreamError`` — a plain Exception subclass with + a ``model_dump()`` hook — must emerge as the envelope on the wire + even under ``OPT_SERIALIZE_DATACLASS``, using ONLY upstream's + stock ``serde.default``. Proof that we no longer need to patch + the serde module. + """ + import orjson + + from EvoScientist.llm.errors import ProviderStreamError + + err = ProviderStreamError( + provider="openrouter", + class_qualname="openrouter.errors.foo.UnauthorizedResponseError", + message="User not found.", + status_code=401, + ) + + wire = orjson.dumps( + err, + default=_serde_mod.default, # upstream, unpatched + option=orjson.OPT_SERIALIZE_DATACLASS, + ) + decoded = orjson.loads(wire) + assert decoded == { + "error": "UnauthorizedResponseError", + "class": "openrouter.errors.foo.UnauthorizedResponseError", + "message": "User not found.", + "provider": "openrouter", + "status_code": 401, + }