mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(mcp): keep user token in authorization_code tools preview
After to_server_spec maps oauth2 onto the v2 resolver, the interactive tools preview for an unsaved authorization_code server read the per-user token store, found nothing, and fail-closed with a 401, so the create/test tab could no longer list tools The preview now routes the just-authorized token (forwarded in oauth2_headers) through mcp_auth_header, so _create_mcp_client takes the per-request-override v1 path and uses it directly, matching v1's preview. Gated to the v2-mapped oauth2 case; M2M, delegate/passthrough, and token-exchange keep their existing preview path Adds tests: interactive oauth routes the forwarded token to mcp_auth_header, M2M and token-exchange do not
This commit is contained in:
parent
ea11a97a1e
commit
f028f32c9e
2 changed files with 179 additions and 9 deletions
|
|
@ -1039,6 +1039,32 @@ if MCP_AVAILABLE:
|
|||
static_headers=request.static_headers,
|
||||
)
|
||||
|
||||
# Interactive authorization_code tools preview: the UI forwards the just-authorized token
|
||||
# in oauth2_headers before any per-user credential is persisted. Route it through
|
||||
# mcp_auth_header so _create_mcp_client's per-request-override deferral takes the v1 path
|
||||
# (use the forwarded token directly - no resolver, no fail-closed challenge), exactly as
|
||||
# v1's preview did. Gated to the v2-mapped oauth2 case (to_server_spec non-None ==
|
||||
# authorization_code): M2M (client_credentials) and delegate/passthrough return None from
|
||||
# to_server_spec, and token-exchange is a different auth_type, so all three keep their
|
||||
# existing v1 preview behavior untouched.
|
||||
if (
|
||||
mcp_auth_header is None
|
||||
and server_model.auth_type == MCPAuth.oauth2
|
||||
and oauth2_headers
|
||||
):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415
|
||||
to_server_spec,
|
||||
)
|
||||
|
||||
forwarded_authorization = oauth2_headers.get("Authorization")
|
||||
if forwarded_authorization and to_server_spec(server_model) is not None:
|
||||
# The oauth2 client re-adds the "Bearer " scheme, so pass the bare token.
|
||||
mcp_auth_header = (
|
||||
forwarded_authorization[7:]
|
||||
if forwarded_authorization[:7].lower() == "bearer "
|
||||
else forwarded_authorization
|
||||
)
|
||||
|
||||
client = await global_mcp_server_manager._create_mcp_client(
|
||||
server=server_model,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
|
|
|
|||
|
|
@ -270,6 +270,154 @@ class TestExecuteWithMcpClient:
|
|||
or "Authorization" not in captured["extra_headers"]
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_interactive_oauth_routes_forwarded_token_to_mcp_auth_header(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""Interactive authorization_code preview (oauth2, no client credentials): the forwarded
|
||||
just-authorized token must be routed to mcp_auth_header so _create_mcp_client takes the v1
|
||||
per-request-override path and uses it directly, instead of the v2 resolver fail-closing on
|
||||
the not-yet-persisted credential. The "Bearer " scheme is stripped since the oauth2 client
|
||||
re-adds it."""
|
||||
captured: dict = {}
|
||||
|
||||
def fake_build_stdio_env(server, raw_headers):
|
||||
return None
|
||||
|
||||
async def fake_create_client(*args, **kwargs):
|
||||
captured["mcp_auth_header"] = kwargs.get("mcp_auth_header")
|
||||
return object()
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_build_stdio_env",
|
||||
fake_build_stdio_env,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_create_mcp_client",
|
||||
fake_create_client,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
async def ok_operation(client):
|
||||
return {"status": "ok"}
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="linear",
|
||||
url="https://mcp.linear.app/mcp",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
authorization_url="https://mcp.linear.app/authorize",
|
||||
)
|
||||
|
||||
result = await rest_endpoints._execute_with_mcp_client(
|
||||
payload,
|
||||
ok_operation,
|
||||
oauth2_headers={"Authorization": "Bearer forwarded-user-token"},
|
||||
)
|
||||
|
||||
assert result["status"] == "ok"
|
||||
assert captured["mcp_auth_header"] == "forwarded-user-token"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_m2m_does_not_route_forwarded_token_to_mcp_auth_header(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""M2M (client_credentials) must NOT route the forwarded header to mcp_auth_header - it stays
|
||||
on the auto-fetch path. to_server_spec returns None for M2M, so the interactive shortcut is
|
||||
skipped and mcp_auth_header remains None (the incoming header is dropped, as before)."""
|
||||
captured: dict = {}
|
||||
|
||||
def fake_build_stdio_env(server, raw_headers):
|
||||
return None
|
||||
|
||||
async def fake_create_client(*args, **kwargs):
|
||||
captured["mcp_auth_header"] = kwargs.get("mcp_auth_header")
|
||||
return object()
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_build_stdio_env",
|
||||
fake_build_stdio_env,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_create_mcp_client",
|
||||
fake_create_client,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
async def ok_operation(client):
|
||||
return {"status": "ok"}
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="m2m-server",
|
||||
url="https://example.com",
|
||||
auth_type=MCPAuth.oauth2,
|
||||
token_url="https://auth.example.com/token",
|
||||
credentials={"client_id": "my-id", "client_secret": "my-secret"},
|
||||
)
|
||||
|
||||
result = await rest_endpoints._execute_with_mcp_client(
|
||||
payload,
|
||||
ok_operation,
|
||||
oauth2_headers={"Authorization": "Bearer sk-litellm-api-key"},
|
||||
)
|
||||
|
||||
assert result["status"] == "ok"
|
||||
assert captured["mcp_auth_header"] is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_does_not_route_forwarded_token_to_mcp_auth_header(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""OBO / token-exchange (auth_type oauth2_token_exchange, not oauth2) must NOT route the
|
||||
forwarded header to mcp_auth_header - resolve_mcp_auth performs the exchange on the v1 path.
|
||||
The guard's auth_type == oauth2 check excludes it, so mcp_auth_header stays None (setting it
|
||||
would short-circuit the exchange, since resolve_mcp_auth returns mcp_auth_header first)."""
|
||||
captured: dict = {}
|
||||
|
||||
def fake_build_stdio_env(server, raw_headers):
|
||||
return None
|
||||
|
||||
async def fake_create_client(*args, **kwargs):
|
||||
captured["mcp_auth_header"] = kwargs.get("mcp_auth_header")
|
||||
return object()
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_build_stdio_env",
|
||||
fake_build_stdio_env,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_create_mcp_client",
|
||||
fake_create_client,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
async def ok_operation(client):
|
||||
return {"status": "ok"}
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="obo-server",
|
||||
url="https://example.com",
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
token_url="https://auth.example.com/token",
|
||||
)
|
||||
|
||||
result = await rest_endpoints._execute_with_mcp_client(
|
||||
payload,
|
||||
ok_operation,
|
||||
oauth2_headers={"Authorization": "Bearer subject-jwt"},
|
||||
)
|
||||
|
||||
assert result["status"] == "ok"
|
||||
assert captured["mcp_auth_header"] is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_catches_exception_group(self, monkeypatch):
|
||||
"""MCP SDK's anyio TaskGroup raises BaseExceptionGroup which does not
|
||||
|
|
@ -1660,9 +1808,9 @@ class TestPreviewOpenAPITools:
|
|||
names = [t["name"] for t in result["tools"]]
|
||||
anthropic_re = re.compile(r"^[a-zA-Z0-9_-]{1,128}$")
|
||||
for name in names:
|
||||
assert anthropic_re.match(
|
||||
name
|
||||
), f"preview tool name {name!r} violates ^[a-zA-Z0-9_-]+$"
|
||||
assert anthropic_re.match(name), (
|
||||
f"preview tool name {name!r} violates ^[a-zA-Z0-9_-]+$"
|
||||
)
|
||||
assert "actions_download-job-logs-for-workflow-run" in names
|
||||
assert "pulls_list-files" in names
|
||||
|
||||
|
|
@ -1719,9 +1867,7 @@ class TestPreviewOpenAPITools:
|
|||
|
||||
registered_summary_to_name: dict = {}
|
||||
|
||||
def fake_create_tool_function(
|
||||
path, method, operation, base_url
|
||||
): # noqa: ANN001
|
||||
def fake_create_tool_function(path, method, operation, base_url): # noqa: ANN001
|
||||
def _f():
|
||||
return None
|
||||
|
||||
|
|
@ -1734,9 +1880,7 @@ class TestPreviewOpenAPITools:
|
|||
)
|
||||
|
||||
class _StubRegistry:
|
||||
def register_tool(
|
||||
self, name, description, input_schema, handler
|
||||
): # noqa: ANN001
|
||||
def register_tool(self, name, description, input_schema, handler): # noqa: ANN001
|
||||
registered_summary_to_name[description] = name
|
||||
|
||||
monkeypatch.setattr(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue