Feat/ai4sci tool protocol #1

Open
ouyangbo wants to merge 10 commits from feat/ai4sci-tool-protocol into Ai4Sci
12 changed files with 1369 additions and 19 deletions
Showing only changes of commit 88ac9f5ba1 - Show all commits
+12
View File
@@ -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(),
+294
View File
@@ -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 ``<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
+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()],
)
+2
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,
@@ -44,6 +45,7 @@ __all__ = [
"Choice",
"ConfigurableModelMiddleware",
"ContextOverflowMapperMiddleware",
"ErrorNormalizationMiddleware",
"EvoMemoryLifecycleMiddleware",
"EvoMemoryMiddleware",
"ModelFallbackMiddleware",
@@ -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
+29 -3
View File
@@ -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)
+21 -3
View File
@@ -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
@@ -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"
+112
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
@@ -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()
+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
+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,
}