Feat/ai4sci tool protocol #1
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
@@ -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()],
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user