mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/competent-williams-410458
This commit is contained in:
commit
b76073e4ee
2 changed files with 64 additions and 1 deletions
|
|
@ -581,8 +581,12 @@ async def exchange_token_with_server(
|
|||
if mcp_server.token_url is None:
|
||||
raise HTTPException(status_code=400, detail="MCP server token url is not set")
|
||||
|
||||
# The id and secret must come from the same source. When the server-side client_id wins,
|
||||
# falling back to the caller's secret pairs the persisted client with a foreign secret; the
|
||||
# register short-circuit hands clients a placeholder secret ("dummy"), so a re-auth against a
|
||||
# persisted public PKCE client (no stored secret) would send that placeholder and the IdP 401s.
|
||||
resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id
|
||||
resolved_client_secret = mcp_server.client_secret if mcp_server.client_secret else client_secret
|
||||
resolved_client_secret = mcp_server.client_secret if mcp_server.client_id else client_secret
|
||||
try:
|
||||
client_auth = build_token_endpoint_client_auth(
|
||||
auth_method=mcp_server.token_endpoint_auth_method,
|
||||
|
|
|
|||
|
|
@ -4156,3 +4156,62 @@ async def test_store_per_user_token_server_side_skips_invalidate_when_db_write_f
|
|||
|
||||
invalidate_mock.assert_not_awaited()
|
||||
cache_set_mock.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_exchange_pairs_client_secret_with_server_client_id():
|
||||
"""Re-auth regression: the register short-circuit hands the browser a placeholder
|
||||
``client_secret: "dummy"``, which the browser echoes back to /token. The server-side
|
||||
persisted client_id wins the resolution, so the secret must come from the same (server)
|
||||
source; pairing the persisted public PKCE client (no stored secret) with the caller's
|
||||
placeholder makes the IdP reject the exchange with 401 on every re-auth."""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
server = MCPServer(
|
||||
server_id="srv-1",
|
||||
name="srv-1",
|
||||
server_name="srv-1",
|
||||
alias="srv-1",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="persisted-client",
|
||||
client_secret=None,
|
||||
authorization_url="https://provider.example/oauth/authorize",
|
||||
token_url="https://provider.example/oauth/token",
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://litellm.example.com/"
|
||||
mock_request.headers = {}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
mock_response.json.return_value = {"access_token": "at", "token_type": "Bearer"}
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
|
||||
return_value=mock_async_client,
|
||||
):
|
||||
await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
grant_type="authorization_code",
|
||||
code="auth-code",
|
||||
redirect_uri="https://litellm.example.com/ui/mcp/oauth/callback",
|
||||
client_id="srv-1",
|
||||
client_secret="dummy",
|
||||
code_verifier="verifier",
|
||||
)
|
||||
|
||||
sent = mock_async_client.post.call_args.kwargs["data"]
|
||||
assert sent["client_id"] == "persisted-client"
|
||||
assert "client_secret" not in sent
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue