mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): name the connected route in OBO rejection challenges and test the initialized Agent 365 guardrail
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
06e5d9f7c5
commit
eeadf7bdc1
5 changed files with 52 additions and 5 deletions
|
|
@ -4211,6 +4211,7 @@ class MCPServerManager:
|
|||
oauth2_headers: dict[str, str] | None,
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
connected_as: str | None = None,
|
||||
) -> None:
|
||||
"""Mint an exchange-backed server's upstream credential at the transport edge.
|
||||
|
||||
|
|
@ -4247,14 +4248,16 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
if subject_token is None and caller_sign_in_for(server, user_api_key_auth) is not None:
|
||||
raise_token_exchange_challenge(server, root_path=get_request_root_path())
|
||||
raise_token_exchange_challenge(server, root_path=get_request_root_path(), connected_as=connected_as)
|
||||
return
|
||||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
spec: Final = _to_server_spec_fail_closed(resolved_server)
|
||||
if spec is None or not isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)):
|
||||
return
|
||||
if subject_token is None and isinstance(spec.config, TokenExchangeConfig):
|
||||
raise_token_exchange_challenge(resolved_server, root_path=get_request_root_path())
|
||||
raise_token_exchange_challenge(
|
||||
resolved_server, root_path=get_request_root_path(), connected_as=connected_as
|
||||
)
|
||||
match await self._cred_provider.resolve_credentials(to_subject(user_api_key_auth, subject_token), spec):
|
||||
case Ok(_):
|
||||
return
|
||||
|
|
@ -4264,6 +4267,7 @@ class MCPServerManager:
|
|||
resolved_server,
|
||||
root_path=get_request_root_path(),
|
||||
claims=err.unauthorized.claims,
|
||||
connected_as=connected_as,
|
||||
)
|
||||
raise_public(err)
|
||||
|
||||
|
|
|
|||
|
|
@ -1763,6 +1763,7 @@ if MCP_AVAILABLE:
|
|||
oauth2_headers=oauth2_headers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
raw_headers=raw_headers,
|
||||
connected_as=server_name,
|
||||
)
|
||||
|
||||
# Pass-through OAuth: when the admin has opted a server into
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from .agent_365 import Agent365Guardrail
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchanger import TokenExchanger
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
|
|
@ -15,6 +16,7 @@ def initialize_guardrail(
|
|||
guardrail: "Guardrail",
|
||||
*,
|
||||
async_handler: "AsyncHTTPHandler | None" = None,
|
||||
token_exchanger: "TokenExchanger | None" = None,
|
||||
) -> Agent365Guardrail:
|
||||
import litellm
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
|
@ -64,6 +66,7 @@ def initialize_guardrail(
|
|||
request_timeout=litellm_params.timeout if litellm_params.timeout is not None else 10.0,
|
||||
unreachable_fallback=litellm_params.unreachable_fallback,
|
||||
async_handler=async_handler,
|
||||
token_exchanger=token_exchanger,
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2930,6 +2930,31 @@ class TestMCPServerManager:
|
|||
www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or ""
|
||||
assert "resource_metadata" in www_authenticate
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_rejected_subject_challenge_names_the_connected_segment(self):
|
||||
"""A subject rejected on ``/mcp/<server_id>`` must point resource_metadata at that same
|
||||
segment, the way the sign-in preflight does, so the client's discovery fetch resolves."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject, server):
|
||||
return Error(CredError.of_unauthorized("subject token rejected by the IdP"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
server = self._token_exchange_server("te-preflight-segment")
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await manager.preflight_token_exchange(
|
||||
server=server,
|
||||
oauth2_headers={"Authorization": "Bearer rejected-subject"},
|
||||
user_api_key_auth=None,
|
||||
connected_as=server.server_id,
|
||||
)
|
||||
headers = exc_info.value.headers or {}
|
||||
www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate") or ""
|
||||
assert f"/.well-known/oauth-protected-resource/mcp/{server.server_id}" in www_authenticate, www_authenticate
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_token_exchange_maps_gateway_fault_to_public_status(self):
|
||||
"""A gateway-fault CredError (e.g. invalid_client) must surface its public status (500)
|
||||
|
|
|
|||
|
|
@ -194,7 +194,16 @@ def _make_guardrail(
|
|||
|
||||
|
||||
def _default_fallback_guardrail(handler: FakeHandler, exchanger: StubTokenExchanger | None = None) -> Agent365Guardrail:
|
||||
return _make_guardrail(handler, exchanger=exchanger)
|
||||
return Agent365Guardrail(
|
||||
guardrail_name="agent-365-guard",
|
||||
tenant_id="tenant-abc",
|
||||
client_id="client-xyz",
|
||||
client_secret="secret-123",
|
||||
async_handler=handler,
|
||||
token_exchanger=exchanger if exchanger is not None else StubTokenExchanger(_obo_ok()),
|
||||
event_hook="pre_mcp_call",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
|
||||
def _server(**overrides: Any) -> MCPServer:
|
||||
|
|
@ -320,11 +329,16 @@ class TestInitializeGuardrail:
|
|||
agent_id="yaml-agent",
|
||||
)
|
||||
handler: Final = FakeHandler([_allow_response()])
|
||||
exchanger: Final = StubTokenExchanger(_obo_ok())
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
initialize_guardrail(params, {"guardrail_name": "a365-stale"}, async_handler=handler)
|
||||
guardrail: Final = initialize_guardrail(
|
||||
params, {"guardrail_name": "a365-stale"}, async_handler=handler, token_exchanger=exchanger
|
||||
)
|
||||
assert "ignoring api_base, resource_app_id, agent_id" in caplog.text
|
||||
guardrail: Final = _make_guardrail(handler)
|
||||
await _run(guardrail, _mcp_data())
|
||||
_, server, config = exchanger.calls[0]
|
||||
assert server.resource == AGENT_365_PROD_API_BASE
|
||||
assert config.scopes == (f"{AGENT_365_PROD_RESOURCE_APP_ID}/{AGENT_365_SCOPE_NAME}",)
|
||||
evaluate_call: Final = handler.calls[0]
|
||||
assert evaluate_call.url == EVALUATE_URL
|
||||
assert evaluate_call.json["agentId"] == "my-agent-key"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue