mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(mcp): consolidate exception-tree walkers into one shared faults traversal
This commit is contained in:
parent
b200d664ee
commit
40f02b2eb5
9 changed files with 164 additions and 49 deletions
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
35
litellm/proxy/_experimental/mcp_server/faults/traversal.py
Normal file
35
litellm/proxy/_experimental/mcp_server/faults/traversal.py
Normal 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__)
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 23409
|
||||
"limit": 23408
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27511
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue