mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(mcp): capture bounded error diagnostics without exposing credentials
This commit is contained in:
parent
1425c71c10
commit
3fc483d424
8 changed files with 341 additions and 38 deletions
|
|
@ -10,6 +10,7 @@ from contextlib import AbstractAsyncContextManager
|
|||
from datetime import timedelta
|
||||
from functools import partial
|
||||
from importlib import metadata
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Protocol, TypeAlias, TypeVar
|
||||
|
||||
import httpx
|
||||
|
|
@ -69,6 +70,7 @@ from litellm._logging import verbose_logger
|
|||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR, MCP_TOOL_LISTING_TIMEOUT
|
||||
from litellm.experimental_mcp_client.tools import list_tools_with_pagination
|
||||
from litellm.llms.custom_httpx.http_handler import get_ssl_configuration
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
MCPAuth,
|
||||
|
|
@ -607,7 +609,9 @@ class MCPClient:
|
|||
auth=effective_auth,
|
||||
verify=ssl_config,
|
||||
follow_redirects=True,
|
||||
event_hooks={"request": [guard]} if guard else {},
|
||||
event_hooks=MappingProxyType(
|
||||
{"response": [capture_upstream_error_response], "request": [guard] if guard else []}
|
||||
), # mutable-ok: httpx types require lists of hooks
|
||||
)
|
||||
|
||||
return factory
|
||||
|
|
|
|||
|
|
@ -85,12 +85,18 @@ Usage with curl::
|
|||
http://localhost:4000/mcp/atlassian_mcp
|
||||
"""
|
||||
|
||||
import re
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from itertools import islice
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from urllib.parse import parse_qsl, urlencode
|
||||
|
||||
import httpx
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
from starlette.types import Message, Send
|
||||
|
||||
from litellm.litellm_core_utils.secret_redaction import REDACTED, redact_string
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree
|
||||
|
||||
|
|
@ -321,58 +327,137 @@ class MCPDebug:
|
|||
|
||||
|
||||
_BODY_PREVIEW_CHARS: Final = 512
|
||||
_SENSITIVE_HEADER_NAMES: Final = frozenset({"authorization", "proxy-authorization", "cookie", "x-api-key"})
|
||||
_SENSITIVE_BODY_FIELD: Final = re.compile(
|
||||
r'(?P<key>"?(?:client_secret|client_assertion|refresh_token|access_token|id_token|password|code)"?\s*[=:]\s*"?)'
|
||||
r'(?P<value>[^&"\s,}]+)'
|
||||
)
|
||||
_BODY_CAPTURE_BYTES: Final = 16384
|
||||
_CAPTURE_TIMEOUT_SECONDS: Final = 1.0
|
||||
_CAPTURE_EXTENSION: Final = "litellm_mcp_error_preview"
|
||||
_SAFE_HEADER_NAMES: Final = frozenset({"content-type", "content-length", "accept"})
|
||||
_JSON_BODY: Final = TypeAdapter(JsonValue)
|
||||
_LOG_MASKER: Final = SensitiveDataMasker(visible_prefix=0, visible_suffix=0)
|
||||
|
||||
|
||||
def _mask_body_match(match: re.Match[str]) -> str:
|
||||
return f"{match.group('key')}{MCPDebug.mask_secret(match.group('value'))}"
|
||||
def _safe_text(value: str, limit: int = _BODY_PREVIEW_CHARS) -> str:
|
||||
escaped: Final = "".join(json.dumps(char)[1:-1] if ord(char) < 32 or ord(char) == 127 else char for char in value)
|
||||
return escaped if len(escaped) <= limit else f"{escaped[:limit]}...(truncated)"
|
||||
|
||||
|
||||
def _preview(raw: bytes) -> str:
|
||||
text: Final = _SENSITIVE_BODY_FIELD.sub(_mask_body_match, raw.decode("utf-8", errors="replace"))
|
||||
return (
|
||||
text
|
||||
if len(text) <= _BODY_PREVIEW_CHARS
|
||||
else f"{text[:_BODY_PREVIEW_CHARS]}...(+{len(text) - _BODY_PREVIEW_CHARS} chars)"
|
||||
)
|
||||
def safe_upstream_url(url: httpx.URL) -> str:
|
||||
return _safe_text(str(url.copy_with(username="", password="", query=None, fragment=None)))
|
||||
|
||||
|
||||
def _sensitive_field(key: str) -> bool:
|
||||
return key.lower() in ("code", "cookie", "client_assertion") or _LOG_MASKER.is_sensitive_key(key)
|
||||
|
||||
|
||||
def _redact_json(value: JsonValue, depth: int = 0) -> str:
|
||||
if depth >= 16:
|
||||
return json.dumps("(depth limit)")
|
||||
if isinstance(value, dict):
|
||||
return (
|
||||
"{"
|
||||
+ ",".join(
|
||||
json.dumps(key)
|
||||
+ ":"
|
||||
+ (json.dumps(REDACTED) if _sensitive_field(key) else _redact_json(item, depth + 1))
|
||||
for key, item in value.items()
|
||||
)
|
||||
+ "}"
|
||||
)
|
||||
if isinstance(value, list):
|
||||
return "[" + ",".join(_redact_json(item, depth + 1) for item in value) + "]"
|
||||
return json.dumps(redact_string(value) if isinstance(value, str) else value)
|
||||
|
||||
|
||||
def _preview(raw: bytes, content_type: str = "") -> str:
|
||||
if not raw:
|
||||
return "(empty)"
|
||||
if len(raw) > _BODY_CAPTURE_BYTES:
|
||||
return "(omitted: body exceeds capture limit)"
|
||||
try:
|
||||
parsed: Final = _JSON_BODY.validate_json(raw)
|
||||
except ValidationError:
|
||||
text: Final = raw.decode("utf-8", errors="replace")
|
||||
if (
|
||||
content_type.split(";", 1)[0].strip().lower() != "application/x-www-form-urlencoded"
|
||||
or "=" not in text
|
||||
or any(char in text for char in "<>\n\r")
|
||||
):
|
||||
return "(omitted: unstructured body)"
|
||||
fields: Final = parse_qsl(text, keep_blank_values=True)
|
||||
return _safe_text(
|
||||
urlencode(
|
||||
tuple((key, REDACTED if _sensitive_field(key) else redact_string(value)) for key, value in fields)
|
||||
)
|
||||
)
|
||||
if not isinstance(parsed, (dict, list)):
|
||||
return "(omitted: unstructured body)"
|
||||
return _safe_text(_redact_json(parsed))
|
||||
|
||||
|
||||
def _masked_headers(headers: httpx.Headers) -> str:
|
||||
return ", ".join(
|
||||
f"{name}={MCPDebug.mask_secret(value) if name.lower() in _SENSITIVE_HEADER_NAMES else value}"
|
||||
for name, value in headers.items()
|
||||
)
|
||||
return _safe_text(", ".join(f"{name}={value}" for name, value in headers.items() if name in _SAFE_HEADER_NAMES))
|
||||
|
||||
|
||||
def _request_body_preview(request: httpx.Request) -> str:
|
||||
try:
|
||||
return _preview(request.content) or "(empty)"
|
||||
return _preview(request.content, request.headers.get("content-type", ""))
|
||||
except httpx.RequestNotRead:
|
||||
return "(streamed, not captured)"
|
||||
|
||||
|
||||
def _response_body_preview(response: httpx.Response) -> str:
|
||||
captured: Final = response.extensions.get(_CAPTURE_EXTENSION)
|
||||
if isinstance(captured, str):
|
||||
return captured
|
||||
try:
|
||||
return _preview(response.content) or "(empty)"
|
||||
return _preview(response.content, response.headers.get("content-type", ""))
|
||||
except httpx.ResponseNotRead:
|
||||
return "(not read)"
|
||||
|
||||
|
||||
def describe_upstream_http_failure(exc: BaseException) -> str | None:
|
||||
"""One line per upstream ``httpx.Response`` in the exception tree: the request method, URL,
|
||||
masked request headers and JSON-RPC body that were sent, plus the status and body that came back.
|
||||
``None`` when the failure never reached an HTTP response (DNS, refused connection, timeout)."""
|
||||
lines: Final = tuple(
|
||||
f"{response.request.method} {response.request.url} -> HTTP {response.status_code} {response.reason_phrase}"
|
||||
f" | request headers: {_masked_headers(response.request.headers)}"
|
||||
f" | request body: {_request_body_preview(response.request)}"
|
||||
async def _read_error_prefix(chunks: AsyncIterator[bytes], remaining: int) -> bytes:
|
||||
chunk: Final = await anext(chunks, b"")
|
||||
if not chunk or len(chunk) >= remaining:
|
||||
return chunk[:remaining]
|
||||
return chunk + await _read_error_prefix(chunks, remaining - len(chunk))
|
||||
|
||||
|
||||
async def capture_upstream_error_response(response: httpx.Response) -> None:
|
||||
if not response.is_error:
|
||||
return
|
||||
try:
|
||||
prefix: Final = await asyncio.wait_for(
|
||||
_read_error_prefix(response.aiter_bytes(chunk_size=4096), _BODY_CAPTURE_BYTES + 1),
|
||||
timeout=_CAPTURE_TIMEOUT_SECONDS,
|
||||
)
|
||||
response._content = prefix # pyright: ignore[reportPrivateUsage] # rebind-ok: httpx has no public setter to retain consumed bytes for auth retries
|
||||
preview: Final = _preview(prefix, response.headers.get("content-type", ""))
|
||||
except (TimeoutError, httpx.HTTPError):
|
||||
response._content = b"" # pyright: ignore[reportPrivateUsage] # rebind-ok: httpx auth retries must survive diagnostic read failures
|
||||
response.extensions[_CAPTURE_EXTENSION] = (
|
||||
"(unavailable: error body read failed)" # rebind-ok: httpx response hooks communicate through extensions
|
||||
)
|
||||
return
|
||||
response.extensions[_CAPTURE_EXTENSION] = preview # rebind-ok: httpx response hooks communicate through extensions
|
||||
|
||||
|
||||
def describe_upstream_response(response: httpx.Response) -> str:
|
||||
try:
|
||||
request: Final = response.request
|
||||
except RuntimeError:
|
||||
return f"HTTP {response.status_code} | request unavailable"
|
||||
return (
|
||||
f"{_safe_text(request.method)} {safe_upstream_url(request.url)} -> HTTP {response.status_code}"
|
||||
f" | request headers: {_masked_headers(request.headers)}"
|
||||
f" | request body: {_request_body_preview(request)}"
|
||||
f" | response body: {_response_body_preview(response)}"
|
||||
for current in iter_exception_tree(exc)
|
||||
)
|
||||
|
||||
|
||||
def describe_upstream_http_failure(exc: BaseException) -> str | None:
|
||||
lines: Final = tuple(
|
||||
describe_upstream_response(response)
|
||||
for current in islice(iter_exception_tree(exc), 16)
|
||||
for response in (getattr(current, "response", None),)
|
||||
if isinstance(response, httpx.Response)
|
||||
)
|
||||
return "\n".join(lines) or None
|
||||
return " | ".join(lines) or None
|
||||
|
|
|
|||
|
|
@ -4266,7 +4266,7 @@ class MCPServerManager:
|
|||
raise
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"Failed to get tools from server %s: %s%s", server.name, e, _upstream_failure_suffix(e)
|
||||
"Failed to get tools from server %s: %s%s", server.name, type(e).__name__, _upstream_failure_suffix(e)
|
||||
)
|
||||
raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge)
|
||||
|
||||
|
|
@ -5012,7 +5012,9 @@ class MCPServerManager:
|
|||
verbose_logger.warning("Connection error while listing tools from %s: %s", server_name, e)
|
||||
raise MCPServerListError(ServerListFault(tag="unreachable"), server_name) from e
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Error listing tools from %s: %s%s", server_name, e, _upstream_failure_suffix(e))
|
||||
verbose_logger.warning(
|
||||
"Error listing tools from %s: %s%s", server_name, type(e).__name__, _upstream_failure_suffix(e)
|
||||
)
|
||||
raise_classified_list_failure(e, server_name)
|
||||
|
||||
_SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024
|
||||
|
|
|
|||
|
|
@ -38,7 +38,11 @@ from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, Valid
|
|||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_http_failure
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import (
|
||||
describe_upstream_http_failure,
|
||||
describe_upstream_response,
|
||||
safe_upstream_url,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
|
||||
InMemoryTokenCacheBackend,
|
||||
OAuthToken,
|
||||
|
|
@ -118,13 +122,22 @@ async def post_client_credentials_grant(
|
|||
)
|
||||
return TokenEndpointDenied(status_code=status_code, detail=f"token endpoint returned HTTP {status_code}")
|
||||
except Exception as exc: # noqa: BLE001 # any transport failure is the same outcome: unreachable
|
||||
return TokenEndpointUnreachable(detail=str(exc))
|
||||
verbose_logger.warning(
|
||||
"OAuth2 client_credentials POST %s failed: %s", safe_upstream_url(httpx.URL(url)), type(exc).__name__
|
||||
)
|
||||
return TokenEndpointUnreachable(detail=type(exc).__name__)
|
||||
try:
|
||||
body: Final = _TOKEN_BODY_ADAPTER.validate_json(response.content)
|
||||
except ValidationError:
|
||||
verbose_logger.warning("OAuth2 client_credentials invalid response: %s", describe_upstream_response(response))
|
||||
return TokenEndpointDenied(
|
||||
status_code=response.status_code, detail="token endpoint returned a non-JSON-object body"
|
||||
)
|
||||
access_token: Final = body.get("access_token")
|
||||
if not isinstance(access_token, str) or not access_token:
|
||||
verbose_logger.warning(
|
||||
"OAuth2 client_credentials response has no access token | %s", describe_upstream_response(response)
|
||||
)
|
||||
return TokenEndpointSuccess(body=body)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -900,7 +900,9 @@ if MCP_AVAILABLE:
|
|||
apply_tool_filters=apply_tool_filters,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error getting tools from %s: %s", server.name, e)
|
||||
verbose_logger.warning(
|
||||
"Error getting tools from %s: %s", server.name, classify_list_exception(e).tag
|
||||
)
|
||||
return (), classify_list_exception(e)
|
||||
return tools_result, ServerListOk(tool_count=len(tools_result))
|
||||
|
||||
|
|
|
|||
|
|
@ -478,3 +478,47 @@ async def test_bearer_auth_advertises_the_header_it_will_occupy():
|
|||
assert ClientCredentialsBearerAuth("t", refetch, ClientCredentialsConfig()).header_name == "Authorization"
|
||||
default_carrier = ClientCredentialsConfig(header_name="esb-oauth")
|
||||
assert ClientCredentialsBearerAuth("t", refetch, default_carrier).header_name == "esb-oauth"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("mode", ["denied", "invalid", "missing", "success", "timeout", "connect", "cancel"])
|
||||
async def test_token_exchange_failure_diagnostics(mode, monkeypatch, caplog):
|
||||
import asyncio
|
||||
import logging
|
||||
from litellm.llms.custom_httpx import http_handler
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import post_client_credentials_grant
|
||||
|
||||
class Poster:
|
||||
async def post(self, url, headers, data):
|
||||
request = httpx.Request("POST", url, headers=headers, data=data)
|
||||
if mode == "timeout":
|
||||
raise httpx.ReadTimeout("private-transport-message", request=request)
|
||||
if mode == "connect":
|
||||
raise httpx.ConnectError("private-transport-message", request=request)
|
||||
if mode == "cancel":
|
||||
raise asyncio.CancelledError
|
||||
response = httpx.Response(401 if mode == "denied" else 200, request=request,
|
||||
content=b"not-json-private" if mode == "invalid" else None,
|
||||
json=None if mode == "invalid" else {"error": "invalid_client", "client_secret":"first second", **({"access_token":"private-token"} if mode == "success" else {})})
|
||||
response.raise_for_status()
|
||||
return response
|
||||
|
||||
monkeypatch.setattr(http_handler, "get_async_httpx_client", lambda **kwargs: Poster())
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
if mode == "cancel":
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await post_client_credentials_grant("https://idp/token", {}, {})
|
||||
assert not caplog.text
|
||||
return
|
||||
result = await post_client_credentials_grant("https://idp/token?key=query-secret", {"client_secret":"first second"}, {"X-Custom":"header-secret"})
|
||||
for secret in ("first", "second", "query-secret", "header-secret", "private-token", "not-json-private", "private-transport-message"):
|
||||
assert secret not in caplog.text
|
||||
if mode == "success":
|
||||
assert isinstance(result, TokenEndpointSuccess) and result.body["access_token"] == "private-token"
|
||||
assert not caplog.text
|
||||
elif mode in {"timeout", "connect"}:
|
||||
assert isinstance(result, TokenEndpointUnreachable)
|
||||
assert "POST https://idp/token failed" in caplog.text
|
||||
else:
|
||||
assert "POST https://idp/token -> HTTP" in caplog.text
|
||||
assert {"denied":"denied", "invalid":"invalid response", "missing":"no access token"}[mode] in caplog.text
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import asyncio
|
|||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import (
|
||||
MCP_DEBUG_REQUEST_HEADER,
|
||||
|
|
@ -280,7 +281,7 @@ class TestDescribeUpstreamHttpFailure:
|
|||
request = httpx.Request(
|
||||
"POST",
|
||||
"https://upstream.example/apis/mcp",
|
||||
headers={"Authorization": "Bearer secret-token-abcdef0123456789", "Content-Type": "application/json"},
|
||||
headers={"Authorization": "Bearer secret-token-abcdef0123456789", "Content-Type": "application/json" if body.startswith(b"{") else "application/x-www-form-urlencoded"},
|
||||
content=body,
|
||||
)
|
||||
response = (
|
||||
|
|
@ -327,3 +328,139 @@ class TestDescribeUpstreamHttpFailure:
|
|||
|
||||
def test_returns_none_without_http_response(self):
|
||||
assert describe_upstream_http_failure(ConnectionError("refused")) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("body", [
|
||||
b'{"password":"first second","token":"demo-secret"}',
|
||||
b'{"nested":[{"access_token":"first,second"}]}',
|
||||
b'client%5Fsecret=first+second&token=demo-secret',
|
||||
])
|
||||
def test_failure_log_fully_redacts_structured_secrets(body):
|
||||
request = httpx.Request("POST", "https://upstream/mcp?credential=query-secret",
|
||||
headers={"X-Custom-Credential": "custom-secret"}, content=body)
|
||||
response = httpx.Response(500, request=request, content=body)
|
||||
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
|
||||
assert detail is not None
|
||||
for secret in ("first", "second", "demo-secret", "custom-secret", "query-secret"):
|
||||
assert secret not in detail
|
||||
|
||||
|
||||
def test_failure_log_omits_unstructured_body():
|
||||
request = httpx.Request("POST", "https://upstream/mcp", content=b"arbitrary-secret")
|
||||
response = httpx.Response(500, request=request, content=b"<html>arbitrary-secret</html>")
|
||||
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
|
||||
assert detail is not None
|
||||
assert "arbitrary-secret" not in detail
|
||||
assert "omitted" in detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("mode", ["error", "large", "timeout", "read_failure", "success", "cancel"])
|
||||
async def test_error_capture_is_bounded_and_preserves_success_and_cancellation(mode):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
|
||||
|
||||
class Stream(httpx.AsyncByteStream):
|
||||
def __init__(self):
|
||||
self.reads = 0
|
||||
|
||||
async def __aiter__(self):
|
||||
self.reads += 1
|
||||
if mode == "timeout":
|
||||
await asyncio.sleep(10)
|
||||
if mode == "read_failure":
|
||||
raise httpx.ReadError("private-read-error")
|
||||
if mode == "cancel":
|
||||
raise asyncio.CancelledError
|
||||
yield b'{"error":"missing_scope","password":"first second"}' if mode != "large" else b"x" * 20000
|
||||
|
||||
stream = Stream()
|
||||
request = httpx.Request("POST", "https://upstream/mcp")
|
||||
response = httpx.Response(200 if mode == "success" else 500, request=request, stream=stream)
|
||||
if mode == "cancel":
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await capture_upstream_error_response(response)
|
||||
return
|
||||
await capture_upstream_error_response(response)
|
||||
if mode == "success":
|
||||
assert stream.reads == 0
|
||||
assert await response.aread() == b'{"error":"missing_scope","password":"first second"}'
|
||||
return
|
||||
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
|
||||
assert detail is not None
|
||||
assert "first" not in detail and "second" not in detail and "private-read-error" not in detail
|
||||
expected = {"error": "missing_scope", "large": "capture limit", "timeout": "read failed", "read_failure": "read failed"}
|
||||
assert expected[mode] in detail
|
||||
if mode == "error":
|
||||
assert await response.aread() == b'{"error":"missing_scope","password":"first second"}'
|
||||
|
||||
|
||||
@pytest.mark.parametrize("body", [b"", b'"scalar"', b'{"hint":"line1\\nline2"}', b'{"hint":"' + b'x' * 600 + b'"}'])
|
||||
def test_failure_preview_handles_empty_scalar_control_and_long_bodies(body):
|
||||
request = httpx.Request("POST", "https://user:secret@upstream/mcp?key=private#private", content=body)
|
||||
response = httpx.Response(500, request=request, content=body)
|
||||
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
|
||||
assert detail is not None
|
||||
assert "private" not in detail and "user:secret" not in detail and "\n" not in detail
|
||||
if not body:
|
||||
assert "(empty)" in detail
|
||||
elif body.startswith(b'"'):
|
||||
assert "omitted" in detail
|
||||
elif len(body) > 512:
|
||||
assert "truncated" in detail and len(detail) < 1300
|
||||
else:
|
||||
assert "line1\\nline2" in detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_error_capture_preserves_httpx_auth_retry():
|
||||
from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response
|
||||
|
||||
class RetryAuth(httpx.Auth):
|
||||
def auth_flow(self, request):
|
||||
response = yield request
|
||||
if response.status_code == 401:
|
||||
request.headers["Authorization"] = "Bearer refreshed"
|
||||
yield request
|
||||
|
||||
def upstream(request):
|
||||
if request.headers.get("Authorization"):
|
||||
return httpx.Response(200, json={"ok": True})
|
||||
return httpx.Response(401, json={"error":"expired_token"})
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(upstream), auth=RetryAuth(),
|
||||
event_hooks={"response":[capture_upstream_error_response]}) as client:
|
||||
response = await client.get("https://upstream/mcp")
|
||||
assert response.status_code == 200 and response.json() == {"ok":True}
|
||||
assert response.history[0].json() == {"error":"expired_token"}
|
||||
|
||||
|
||||
def test_failure_diagnostics_without_request_and_with_streamed_request():
|
||||
response = httpx.Response(503)
|
||||
exc = httpx.HTTPStatusError("failed", request=httpx.Request("GET", "https://upstream"), response=response)
|
||||
assert describe_upstream_http_failure(exc) == "HTTP 503 | request unavailable"
|
||||
request = httpx.Request("POST", "https://upstream", content=iter((b"private-body",)))
|
||||
response = httpx.Response(503, request=request)
|
||||
described = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response))
|
||||
assert described is not None and "streamed, not captured" in described and "private-body" not in described
|
||||
|
||||
|
||||
|
||||
def test_deep_error_body_is_bounded_without_exposing_nested_values():
|
||||
import json
|
||||
body = b'{"nested":' * 18 + b'{"password":"hidden-value"}' + b'}' * 18
|
||||
request = httpx.Request("POST", "https://upstream/mcp", content=body)
|
||||
response = httpx.Response(500, request=request, content=body)
|
||||
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
|
||||
assert detail is not None and "depth limit" in detail and "hidden-value" not in detail
|
||||
assert json.loads(detail.split("response body: ")[1])["nested"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("body", [b'client%5Fsecret=first+second&client_id=visible', b'client_secret=first%26second&client_id=visible'])
|
||||
def test_encoded_form_credentials_are_decoded_before_redaction(body):
|
||||
request = httpx.Request("POST", "https://upstream/token", content=body,
|
||||
headers={"Content-Type":"application/x-www-form-urlencoded"})
|
||||
response = httpx.Response(400, request=request, content=body,
|
||||
headers={"Content-Type":"application/x-www-form-urlencoded"})
|
||||
detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response))
|
||||
assert detail is not None and "client_id=visible" in detail
|
||||
assert "first" not in detail and "second" not in detail
|
||||
|
|
|
|||
|
|
@ -459,3 +459,19 @@ async def test_fetch_tools_logs_upstream_request_details_on_500(caplog):
|
|||
assert "POST https://upstream/apis/mcp -> HTTP 500" in caplog.text
|
||||
assert '"method":"initialize"' in caplog.text
|
||||
assert "upstream-token-0123456789" not in caplog.text
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_creation_failure_logs_sanitized_exchange(monkeypatch, caplog):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(server_id="sample", name="sample", url="https://upstream/mcp", transport=MCPTransport.http, auth_type=MCPAuth.none)
|
||||
request = httpx.Request("POST", "https://upstream/mcp?credential=query-secret")
|
||||
response = httpx.Response(500, request=request, json={"error":"missing_scope"})
|
||||
error = httpx.HTTPStatusError("query-secret", request=request, response=response)
|
||||
monkeypatch.setattr(manager, "_create_mcp_client", AsyncMock(side_effect=error))
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
with pytest.raises(MCPServerListError):
|
||||
await manager._get_tools_from_server(server)
|
||||
assert "POST https://upstream/mcp -> HTTP 500" in caplog.text
|
||||
assert "missing_scope" in caplog.text and "query-secret" not in caplog.text
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue