mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(mcp): keep issuer provenance consistent when it changes or is discovered
Two lifecycle gaps let the issuer trust anchor drift out of sync with the endpoints it governs. Changing or clearing a previously pinned issuer left the authorization_url and token_url that were resolved under the old issuer in the row, so clearing the anchor could revive stale, possibly untrusted endpoints instead of re-discovering. And a build that discovered an issuer trust-on-first-use persisted it to the row while the returned in-memory server kept the issuer unset, so the registry and the row disagreed and the per-user OAuth token identity, which includes the issuer, differed between that build and the next rebuild and forced a spurious re-auth. update_mcp_server now treats a change to a previously pinned issuer the same as a url or auth_type change and clears the auth-flow-scoped endpoint fields that were resolved under it. The trigger fires only when an issuer was already pinned and is now changed or cleared, so establishing one for the first time, including the trust-on-first-use discovery write-back, does not wipe the fields it just resolved. Both build paths, build_mcp_server_from_table and load_servers_from_config, now construct the server with effective_issuer = manual_issuer or the discovered issuer, skipping an origin-fallback guess exactly as the persistence does, so the in-memory object always reflects what the row will hold. Regression tests pin each case: clearing and re-pointing a pinned issuer clear the stale endpoints, a first-time establish preserves the discovered fields, and a build reflects the discovered issuer while an origin-fallback guess is not reflected
This commit is contained in:
parent
032a2f2d76
commit
ad73f3a7a2
4 changed files with 176 additions and 4 deletions
|
|
@ -61,6 +61,13 @@ _AUTH_FLOW_SCOPED_FIELDS: frozenset = frozenset(
|
|||
}
|
||||
)
|
||||
|
||||
|
||||
def _blank_to_none(value: Optional[str]) -> Optional[str]:
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
return value.strip() or None
|
||||
|
||||
|
||||
# Token-exchange settings with dedicated columns that also exist on
|
||||
# ``MCPCredentials`` as a legacy shape (rows and REST callers that predate the
|
||||
# columns). Every write lifts blob values into the columns and strips them from
|
||||
|
|
@ -705,7 +712,8 @@ async def update_mcp_server(
|
|||
# legacy blob copies below, so the existing row is needed for those updates.
|
||||
explicit_te_write = bool(_TOKEN_EXCHANGE_COLUMN_FIELDS & data_dict.keys())
|
||||
url_provided = "url" in data_dict and data_dict["url"] is not None
|
||||
if data.auth_type or has_credentials or explicit_te_write or url_provided:
|
||||
issuer_provided = "issuer" in data_dict
|
||||
if data.auth_type or has_credentials or explicit_te_write or url_provided or issuer_provided:
|
||||
existing = await MCPServerRepository(prisma_client).table.find_unique(where={"server_id": data.server_id})
|
||||
|
||||
auth_type_changed = bool(
|
||||
|
|
@ -716,12 +724,16 @@ async def update_mcp_server(
|
|||
# A url change re-points the server at a potentially different upstream, so any discovered or
|
||||
# trust-on-first-use OAuth endpoints/issuer belong to the old upstream and must re-discover.
|
||||
url_changed = bool(url_provided and existing and existing.url != data_dict["url"])
|
||||
old_issuer = _blank_to_none(getattr(existing, "issuer", None)) if existing else None
|
||||
issuer_changed = bool(
|
||||
issuer_provided and old_issuer is not None and _blank_to_none(data_dict.get("issuer")) != old_issuer
|
||||
)
|
||||
|
||||
# Clear stale credentials when auth_type changes but no new credentials provided
|
||||
if auth_type_changed and "credentials" not in data_dict:
|
||||
data_dict["credentials"] = None
|
||||
|
||||
if auth_type_changed or url_changed:
|
||||
if auth_type_changed or url_changed or issuer_changed:
|
||||
# Clear each auth-flow-scoped field that the caller either omitted (partial update) or
|
||||
# resubmitted unchanged. The edit form re-sends every field, so a stale issuer/endpoint
|
||||
# belonging to the old upstream would otherwise survive a url/auth_type change and win in the
|
||||
|
|
|
|||
|
|
@ -1236,6 +1236,12 @@ class MCPServerManager:
|
|||
resolved_registration_url = manual_registration_url or (
|
||||
gated_oauth_metadata.registration_url if gated_oauth_metadata else None
|
||||
)
|
||||
discovered_issuer = (
|
||||
gated_oauth_metadata.discovered_issuer
|
||||
if gated_oauth_metadata and not gated_oauth_metadata.from_origin_fallback
|
||||
else None
|
||||
)
|
||||
effective_issuer = manual_issuer or discovered_issuer
|
||||
|
||||
config_oauth2_flow = server_config.get("oauth2_flow", None)
|
||||
if auth_type == MCPAuth.oauth2 and config_oauth2_flow not in (
|
||||
|
|
@ -1284,7 +1290,7 @@ class MCPServerManager:
|
|||
client_secret=server_config.get("client_secret", None),
|
||||
oauth2_flow=self._explicit_oauth2_flow(config_oauth2_flow),
|
||||
scopes=resolved_scopes,
|
||||
issuer=manual_issuer,
|
||||
issuer=effective_issuer,
|
||||
authorization_url=resolved_authorization_url,
|
||||
token_url=resolved_token_url,
|
||||
registration_url=resolved_registration_url,
|
||||
|
|
@ -1700,6 +1706,12 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
resolved_scopes = scopes or (gated_oauth_metadata.scopes if gated_oauth_metadata else None)
|
||||
discovered_issuer = (
|
||||
gated_oauth_metadata.discovered_issuer
|
||||
if gated_oauth_metadata and not gated_oauth_metadata.from_origin_fallback
|
||||
else None
|
||||
)
|
||||
effective_issuer = manual_issuer or discovered_issuer
|
||||
|
||||
new_server = MCPServer(
|
||||
server_id=mcp_server.server_id,
|
||||
|
|
@ -1719,7 +1731,7 @@ class MCPServerManager:
|
|||
client_secret=client_secret_value or getattr(mcp_server, "client_secret", None),
|
||||
oauth2_flow=self._explicit_oauth2_flow(getattr(mcp_server, "oauth2_flow", None)),
|
||||
scopes=resolved_scopes,
|
||||
issuer=manual_issuer,
|
||||
issuer=effective_issuer,
|
||||
authorization_url=manual_authorization_url or getattr(gated_oauth_metadata, "authorization_url", None),
|
||||
token_url=manual_token_url or getattr(gated_oauth_metadata, "token_url", None),
|
||||
registration_url=manual_registration_url or getattr(gated_oauth_metadata, "registration_url", None),
|
||||
|
|
|
|||
|
|
@ -270,6 +270,94 @@ async def test_url_change_clears_stale_oauth_fields_even_when_resubmitted_unchan
|
|||
assert data_dict["authorization_url"] == "https://new-idp.example.com/authorize"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_clearing_pinned_issuer_clears_stale_oauth_endpoints():
|
||||
"""Clearing a previously pinned issuer must not revive the endpoints resolved under it. Under an
|
||||
issuer anchor the endpoints come solely from the issuer document and are not persisted, but a row
|
||||
that was resource-rooted before the pin can still hold stale authorization_url/token_url; clearing
|
||||
the anchor without clearing those would let them win the resolution merge and be posted to without
|
||||
fresh discovery (RFC 8414 §3.3 provenance)."""
|
||||
mock_prisma = _mock_prisma()
|
||||
existing = MagicMock()
|
||||
existing.auth_type = "oauth2"
|
||||
existing.url = "https://same.example.com/mcp"
|
||||
existing.credentials = None
|
||||
existing.issuer = "https://pinned-idp.example.com"
|
||||
existing.token_url = "https://pinned-idp.example.com/token"
|
||||
existing.authorization_url = "https://pinned-idp.example.com/authorize"
|
||||
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
|
||||
|
||||
data = UpdateMCPServerRequest(
|
||||
server_id="my-test-server",
|
||||
issuer="", # admin clears the anchor; url and auth_type unchanged
|
||||
token_url="https://pinned-idp.example.com/token",
|
||||
authorization_url="https://pinned-idp.example.com/authorize",
|
||||
)
|
||||
await update_mcp_server(mock_prisma, data, "test-user")
|
||||
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
|
||||
|
||||
assert data_dict["token_url"] is None
|
||||
assert data_dict["authorization_url"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_repointing_pinned_issuer_clears_stale_endpoints_keeps_new_issuer():
|
||||
"""Re-pointing the issuer to a different authorization server invalidates the old issuer's
|
||||
endpoints while keeping the new issuer the admin submitted."""
|
||||
mock_prisma = _mock_prisma()
|
||||
existing = MagicMock()
|
||||
existing.auth_type = "oauth2"
|
||||
existing.url = "https://same.example.com/mcp"
|
||||
existing.credentials = None
|
||||
existing.issuer = "https://old-idp.example.com"
|
||||
existing.token_url = "https://old-idp.example.com/token"
|
||||
existing.authorization_url = "https://old-idp.example.com/authorize"
|
||||
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
|
||||
|
||||
data = UpdateMCPServerRequest(
|
||||
server_id="my-test-server",
|
||||
issuer="https://new-idp.example.com",
|
||||
token_url="https://old-idp.example.com/token", # resubmitted stale -> must clear
|
||||
authorization_url="https://old-idp.example.com/authorize", # resubmitted stale -> must clear
|
||||
)
|
||||
await update_mcp_server(mock_prisma, data, "test-user")
|
||||
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
|
||||
|
||||
assert data_dict["issuer"] == "https://new-idp.example.com"
|
||||
assert data_dict["token_url"] is None
|
||||
assert data_dict["authorization_url"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_establishing_issuer_first_time_preserves_discovered_fields():
|
||||
"""Establishing an issuer for the first time (None -> X), which is exactly what the trust-on-first-use
|
||||
discovery write-back does, must NOT clear the endpoints or oauth2_flow it discovered in the same
|
||||
write. Only an issuer that was already pinned and is now changed or cleared invalidates its
|
||||
endpoints, so the discovery persist cannot wipe the fields it just resolved."""
|
||||
mock_prisma = _mock_prisma()
|
||||
existing = MagicMock()
|
||||
existing.auth_type = "oauth2"
|
||||
existing.url = "https://same.example.com/mcp"
|
||||
existing.credentials = None
|
||||
existing.issuer = None
|
||||
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing)
|
||||
|
||||
data = UpdateMCPServerRequest(
|
||||
server_id="my-test-server",
|
||||
issuer="https://discovered-idp.example.com",
|
||||
authorization_url="https://discovered-idp.example.com/authorize",
|
||||
token_url="https://discovered-idp.example.com/token",
|
||||
oauth2_flow="authorization_code",
|
||||
)
|
||||
await update_mcp_server(mock_prisma, data, "mcp_oauth_discovery")
|
||||
data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"]
|
||||
|
||||
assert data_dict["issuer"] == "https://discovered-idp.example.com"
|
||||
assert data_dict["authorization_url"] == "https://discovered-idp.example.com/authorize"
|
||||
assert data_dict["token_url"] == "https://discovered-idp.example.com/token"
|
||||
assert data_dict.get("oauth2_flow") == "authorization_code"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unchanged_url_does_not_clear_discovered_oauth_fields():
|
||||
"""A partial update that resends the same url (or omits it) must not clear the discovered OAuth
|
||||
|
|
|
|||
|
|
@ -1166,6 +1166,66 @@ class TestMCPServerManager:
|
|||
assert built.token_url == "https://idp.example.com/token"
|
||||
assert built.scopes == ["read", "admin"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_from_table_reflects_discovered_issuer_trust_on_first_use(self):
|
||||
"""An unpinned server resolves endpoints resource-rooted on first discovery and records the
|
||||
discovered issuer trust-on-first-use. The returned in-memory server must carry that discovered
|
||||
issuer so the registry matches what gets persisted to the row; otherwise the OAuth token
|
||||
identity (which includes issuer) differs between this build and the next rebuild, forcing a
|
||||
spurious re-auth. Endpoints and issuer come from the same authorization-server document, so
|
||||
they are consistent."""
|
||||
manager = MCPServerManager()
|
||||
row = LiteLLM_MCPServerTable(
|
||||
server_id="tofu-issuer-1",
|
||||
alias="tofu_issuer",
|
||||
description="unpinned, discovers its issuer",
|
||||
url="https://up.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
metadata = MCPOAuthMetadata(
|
||||
authorization_url="https://idp.example.com/authorize",
|
||||
token_url="https://idp.example.com/token",
|
||||
scopes=["read"],
|
||||
discovered_issuer="https://idp.example.com",
|
||||
)
|
||||
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
|
||||
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
|
||||
|
||||
assert built.issuer == "https://idp.example.com"
|
||||
assert built.authorization_url == "https://idp.example.com/authorize"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_from_table_origin_fallback_issuer_is_not_reflected(self):
|
||||
"""An origin-fallback discovery is a guess that is deliberately never persisted, so the built
|
||||
server must not claim an issuer the row will not hold; otherwise in-memory and DB would
|
||||
disagree in the opposite direction."""
|
||||
manager = MCPServerManager()
|
||||
row = LiteLLM_MCPServerTable(
|
||||
server_id="origin-fallback-1",
|
||||
alias="origin_fallback",
|
||||
description="unpinned, origin-fallback discovery",
|
||||
url="https://up.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
metadata = MCPOAuthMetadata(
|
||||
authorization_url="https://up.example.com/authorize",
|
||||
token_url="https://up.example.com/token",
|
||||
discovered_issuer="https://up.example.com",
|
||||
from_origin_fallback=True,
|
||||
)
|
||||
with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)):
|
||||
built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False)
|
||||
|
||||
assert built.issuer is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_from_table_whitespace_authorization_url_is_not_a_pin(self):
|
||||
"""A whitespace-only authorization_url on the row must not be kept for redirects while the
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue