refactor(mcp): consolidate exception-tree walkers into one shared faults traversal

This commit is contained in:
Tin Chi Lo 2026-07-14 00:21:22 -07:00
parent b200d664ee
commit 40f02b2eb5
9 changed files with 164 additions and 49 deletions

View file

@ -15,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.faults.render_oauth import (
dcr_fault_detail,
render_token_fault,
)
from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree
from litellm.proxy._experimental.mcp_server.faults.types import (
CallerRejected,
CredentialSource,
@ -34,5 +35,6 @@ __all__ = [
"classify_upstream_dcr_rejection",
"classify_upstream_token_rejection",
"dcr_fault_detail",
"iter_exception_tree",
"render_token_fault",
]

View file

@ -0,0 +1,35 @@
"""Shared exception-tree traversal for fault classification.
Failures cross the MCP SDK's anyio task groups wrapped in ``ExceptionGroup``s and chained through
``raise ... from`` causes, so every classifier that needs an exception buried in the tree (an
upstream ``httpx.Response``, a context-window overflow) has to walk the same shapes. One traversal
with one deliberate order keeps blame assignment consistent across classifiers: explicit links are
searched before incidental ones, so an exception raised while handling the real failure can never
shadow the failure itself.
"""
from __future__ import annotations
from collections.abc import Iterator
def iter_exception_tree(exc: BaseException) -> Iterator[BaseException]:
"""Yield ``exc`` and every exception reachable from it, explicit links first: each node's
``raise ... from`` cause subtree, then ``ExceptionGroup`` members in raise order, then the
incidental ``__context__`` chain last. Cycle-safe via identity tracking, and iterative so a
deep chain cannot overflow the interpreter stack."""
seen: set[int] = set()
stack = [exc]
while stack:
current = stack.pop()
if id(current) in seen:
continue
seen.add(id(current))
yield current
if current.__context__ is not None:
stack.append(current.__context__)
exceptions = getattr(current, "exceptions", None)
if isinstance(exceptions, tuple):
stack.extend(reversed(exceptions))
if current.__cause__ is not None:
stack.append(current.__cause__)

View file

@ -51,6 +51,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
)
@ -440,44 +441,17 @@ def _extract_upstream_auth_failure(
upstream MCP server.
The MCP SDK wraps transport errors in anyio ``ExceptionGroup`` objects and
may chain through ``__cause__`` / ``__context__``. We inspect all of those
layers for an ``httpx.Response``-bearing exception (typically
``httpx.HTTPStatusError``) and extract the status code and any upstream
``WWW-Authenticate`` header.
may chain through ``__cause__`` / ``__context__``; ``iter_exception_tree``
visits all of those layers, explicit links first. The first exception
bearing a real ``httpx.Response`` with a 401/403 wins, and its status code
and upstream ``WWW-Authenticate`` header are extracted.
Returns ``(status_code, www_authenticate)`` on match, else ``None``.
"""
seen: set[int] = set()
stack: list[BaseException] = [exc]
while stack:
current = stack.pop()
if id(current) in seen:
continue
seen.add(id(current))
for current in iter_exception_tree(exc):
response = getattr(current, "response", None)
if response is not None:
status_code = getattr(response, "status_code", None)
if isinstance(status_code, int) and status_code in (401, 403):
www_authenticate: Optional[str] = None
headers = getattr(response, "headers", None)
if headers is not None:
try:
www_authenticate = headers.get("www-authenticate")
except Exception:
www_authenticate = None
return status_code, www_authenticate
# anyio / PEP 654 ExceptionGroup
sub_exceptions = getattr(current, "exceptions", None)
if sub_exceptions:
stack.extend(sub_exceptions)
if current.__cause__ is not None:
stack.append(current.__cause__)
if current.__context__ is not None and current.__context__ is not current.__cause__:
stack.append(current.__context__)
if isinstance(response, httpx.Response) and response.status_code in (401, 403):
return response.status_code, response.headers.get("www-authenticate")
return None

View file

@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional
from litellm._logging import verbose_logger
from litellm.exceptions import ContextWindowExceededError
from litellm.litellm_core_utils.exception_mapping_utils import ExceptionCheckers
from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree
from litellm.proxy._experimental.mcp_server.utils import MCP_TOOL_PREFIX_SEPARATOR
if TYPE_CHECKING:
@ -33,18 +34,15 @@ class SemanticToolFilterContextWindowError(Exception):
)
def _is_context_window_error(error: Optional[BaseException], max_depth: int = 5) -> bool:
"""Detect a context-window overflow anywhere in an exception's cause chain."""
current = error
for _ in range(max_depth):
if current is None:
return False
if isinstance(current, ContextWindowExceededError):
return True
if ExceptionCheckers.is_error_str_context_window_exceeded(str(current)):
return True
current = current.__cause__ or current.__context__
return False
def _is_context_window_error(error: Optional[BaseException]) -> bool:
"""Detect a context-window overflow anywhere in an exception's tree."""
if error is None:
return False
return any(
isinstance(current, ContextWindowExceededError)
or ExceptionCheckers.is_error_str_context_window_exceeded(str(current))
for current in iter_exception_tree(error)
)
class SemanticMCPToolFilter:

View file

@ -60,7 +60,7 @@
"limit": 4
},
"BLE001": {
"limit": 2903
"limit": 2902
},
"C401": {
"limit": 11
@ -363,6 +363,6 @@
"limit": 105
},
"UP045": {
"limit": 18462
"limit": 18461
}
}

View file

@ -0,0 +1,48 @@
"""Traversal contract for the shared exception-tree walk: the root is yielded first, explicit
links win (the ``raise ... from`` cause subtree, then ExceptionGroup members in raise order,
then the incidental ``__context__`` chain last), and adversarial shapes terminate."""
from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree
def test_yields_the_root_itself_first():
exc = ValueError("root")
assert list(iter_exception_tree(exc)) == [exc]
def test_cause_subtree_is_exhausted_before_context():
deep = KeyError("deep")
cause = RuntimeError("cause")
cause.__cause__ = deep
context = OSError("context")
root = ValueError("root")
root.__cause__ = cause
root.__context__ = context
assert list(iter_exception_tree(root)) == [root, cause, deep, context]
def test_group_members_yield_in_raise_order_between_cause_and_context():
first = KeyError("first")
second = IndexError("second")
group = BaseExceptionGroup("group", [first, second])
cause = RuntimeError("cause")
context = OSError("context")
group.__cause__ = cause
group.__context__ = context
assert list(iter_exception_tree(group)) == [group, cause, first, second, context]
def test_terminates_on_a_cause_cycle():
a = ValueError("a")
b = RuntimeError("b")
a.__cause__ = b
b.__cause__ = a
assert list(iter_exception_tree(a)) == [a, b]
def test_node_reachable_as_both_cause_and_context_yields_once():
inner = KeyError("inner")
root = ValueError("root")
root.__cause__ = inner
root.__context__ = inner
assert list(iter_exception_tree(root)) == [root, inner]

View file

@ -50,6 +50,36 @@ def test_extract_upstream_auth_failure_returns_none_for_non_auth():
assert _extract_upstream_auth_failure(RuntimeError("boom")) is None
def _auth_status_error(status_code: int, www_authenticate: str) -> httpx.HTTPStatusError:
response = httpx.Response(
status_code=status_code,
headers={"www-authenticate": www_authenticate},
request=httpx.Request("GET", "https://upstream/mcp"),
)
return httpx.HTTPStatusError(str(status_code), request=response.request, response=response)
def test_extract_upstream_auth_failure_finds_401_behind_cause_chain():
wrapper = RuntimeError("wrapped")
wrapper.__cause__ = _auth_status_error(401, "Bearer")
assert _extract_upstream_auth_failure(wrapper) == (401, "Bearer")
def test_extract_upstream_auth_failure_finds_401_behind_context_chain():
wrapper = RuntimeError("wrapped")
wrapper.__context__ = _auth_status_error(401, "Bearer")
assert _extract_upstream_auth_failure(wrapper) == (401, "Bearer")
def test_extract_upstream_auth_failure_prefers_causal_chain_over_context():
"""A 403 raised incidentally while handling the real 401 (surviving only as ``__context__``)
must not shadow the 401 on the explicit ``raise ... from`` chain."""
wrapper = RuntimeError("wrapped")
wrapper.__cause__ = _auth_status_error(401, "Bearer realm=real")
wrapper.__context__ = _auth_status_error(403, "Bearer realm=incidental")
assert _extract_upstream_auth_failure(wrapper) == (401, "Bearer realm=real")
@pytest.mark.asyncio
async def test_fetch_tools_from_passthrough_raises_on_upstream_401():
manager = MCPServerManager()

View file

@ -1664,3 +1664,31 @@ def test_is_context_window_error_detection_variants():
assert _is_context_window_error(ValueError("Invalid 'input[0]': maximum input length is 8192 tokens."))
assert not _is_context_window_error(ValueError("A generic API error occurred."))
assert not _is_context_window_error(None)
def test_is_context_window_error_sees_through_trees_the_chain_walk_missed():
"""Overflow shapes the old single-path depth-5 chain walk could not reach: hidden in
``__context__`` behind a non-matching ``__cause__``, buried inside an anyio-style
``ExceptionGroup``, and chained deeper than five links."""
import litellm
from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
_is_context_window_error,
)
def _cwe() -> litellm.ContextWindowExceededError:
return litellm.ContextWindowExceededError(message="overflow", model="m", llm_provider="openai")
shadowed = ValueError("wrapper")
shadowed.__cause__ = TypeError("unrelated failure")
shadowed.__context__ = _cwe()
assert _is_context_window_error(shadowed)
grouped = BaseExceptionGroup("task group", [RuntimeError("sibling"), _cwe()])
assert _is_context_window_error(grouped)
deep: BaseException = _cwe()
for depth in range(6):
wrapper = ValueError(f"layer {depth}")
wrapper.__cause__ = deep
deep = wrapper
assert _is_context_window_error(deep)

View file

@ -1,6 +1,6 @@
{
"LIT001": {
"limit": 23409
"limit": 23408
},
"LIT002": {
"limit": 27511