mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): spare BYOK rows when purging stale OAuth tokens and invalidate caches on server delete
LiteLLM_MCPUserCredentials stores BYOK API keys in the same column as per-user OAuth tokens, so the purge on a mint-relevant config change now deletes only rows whose payload decodes as an OAuth2 credential, each by its (user_id, server_id) pair, instead of every row for the server. An api_key server whose url changes purges nothing. delete_mcp_server now also invalidates each enumerated user's cached token so a re-created server reusing the id cannot serve tokens minted for the deleted one, and both cache drops are best-effort
This commit is contained in:
parent
e720b5e25a
commit
aa351311c0
6 changed files with 222 additions and 30 deletions
|
|
@ -558,7 +558,11 @@ async def delete_mcp_server_from_virtualkey():
|
|||
pass
|
||||
|
||||
|
||||
async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Optional[LiteLLM_MCPServerTable]:
|
||||
async def delete_mcp_server(
|
||||
prisma_client: PrismaClient,
|
||||
server_id: str,
|
||||
invalidate_token_cache: Optional[Callable[[str, str], Awaitable[None]]] = None,
|
||||
) -> Optional[LiteLLM_MCPServerTable]:
|
||||
"""
|
||||
Delete the mcp server from the db by server_id
|
||||
|
||||
|
|
@ -569,6 +573,12 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
caller-visible error. Each table is cleaned independently so a failure on one
|
||||
still attempts the other.
|
||||
|
||||
Each enumerated credential row's user also gets their cached per-user token
|
||||
invalidated (legacy cache + v2 store, via invalidate_token_cache, defaulting
|
||||
to the manager's shared invalidation): the caches are keyed by
|
||||
(user_id, server_id), so without this a re-created server reusing the same
|
||||
server_id would serve tokens minted for the deleted server until TTL.
|
||||
|
||||
Returns the deleted mcp server record if it exists, otherwise None
|
||||
"""
|
||||
deleted_server = await MCPServerRepository(prisma_client).table.delete(
|
||||
|
|
@ -577,6 +587,18 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
},
|
||||
)
|
||||
if deleted_server is not None:
|
||||
credential_user_ids: List[str] = []
|
||||
try:
|
||||
credential_rows = await prisma_client.db.litellm_mcpusercredentials.find_many(
|
||||
where={"server_id": server_id}
|
||||
)
|
||||
credential_user_ids = [row.user_id for row in credential_rows]
|
||||
except Exception as e: # noqa: BLE001 - enumeration is best-effort; cached tokens expire by TTL
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP server %s deleted but per-user credential enumeration failed; cached tokens expire by TTL: %s",
|
||||
server_id,
|
||||
e,
|
||||
)
|
||||
for model, label in (
|
||||
(prisma_client.db.litellm_mcpusercredentials, "credential"),
|
||||
(prisma_client.db.litellm_mcpuserenvvars, "env var"),
|
||||
|
|
@ -591,6 +613,15 @@ async def delete_mcp_server(prisma_client: PrismaClient, server_id: str) -> Opti
|
|||
label,
|
||||
e,
|
||||
)
|
||||
if credential_user_ids:
|
||||
if invalidate_token_cache is None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
|
||||
for user_id in credential_user_ids:
|
||||
await invalidate_token_cache(user_id, server_id)
|
||||
return deleted_server
|
||||
|
||||
|
||||
|
|
@ -1123,21 +1154,30 @@ async def purge_user_oauth_credentials_for_server(
|
|||
server_id: str,
|
||||
invalidate_token_cache: Optional[Callable[[str, str], Awaitable[None]]] = None,
|
||||
) -> int:
|
||||
"""Delete every stored per-user OAuth credential for a server and invalidate each user's cached
|
||||
"""Delete every stored per-user OAuth token for a server and invalidate each user's cached
|
||||
token everywhere it can be served from (the legacy per-user token cache and the v2 per-user OAuth
|
||||
token store), so no user keeps a token minted for a superseded configuration. Called when a server
|
||||
update changes a mint-relevant field (see mcp_oauth_token_identity). Returns the number of rows
|
||||
removed. A row inserted between the find and the delete is removed from the DB but cannot be
|
||||
evicted from the caches (its user_id was never seen); that case is detected, logged, and bounded
|
||||
by the cache TTL.
|
||||
removed.
|
||||
|
||||
LiteLLM_MCPUserCredentials also stores BYOK API keys in the same column; only rows whose payload
|
||||
decodes as an OAuth2 credential (see _decode_oauth_payload) are deleted, because a config change
|
||||
only invalidates minted tokens, never a user's own stored key. Rows are therefore deleted per
|
||||
(user_id, server_id) pair rather than by a blanket server_id filter. An OAuth row inserted while
|
||||
the purge runs for a user not yet enumerated survives; a re-auth completing in the window for an
|
||||
already-enumerated user is deleted along with the stale row (the pair delete cannot tell them
|
||||
apart), which costs that user one extra re-auth and nothing else.
|
||||
|
||||
invalidate_token_cache is injectable for tests; it defaults to the manager's shared
|
||||
invalidate_user_oauth_token_cache, the single invalidation point for per-user tokens."""
|
||||
repo = MCPUserCredentialsRepository(prisma_client)
|
||||
rows = await repo.table.find_many(where={"server_id": server_id})
|
||||
if not rows:
|
||||
oauth_rows = [row for row in rows if _decode_oauth_payload(row.credential_b64) is not None]
|
||||
if not oauth_rows:
|
||||
return 0
|
||||
deleted_count = await repo.table.delete_many(where={"server_id": server_id})
|
||||
deleted_count = sum(
|
||||
[await repo.table.delete_many(where={"user_id": row.user_id, "server_id": server_id}) for row in oauth_rows]
|
||||
)
|
||||
if invalidate_token_cache is None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
|
|
@ -1145,15 +1185,15 @@ async def purge_user_oauth_credentials_for_server(
|
|||
|
||||
invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache
|
||||
|
||||
for row in rows:
|
||||
for row in oauth_rows:
|
||||
await invalidate_token_cache(row.user_id, server_id)
|
||||
if deleted_count != len(rows):
|
||||
if deleted_count != len(oauth_rows):
|
||||
verbose_proxy_logger.warning(
|
||||
"MCP server %s: purge removed %d credential row(s) but %d were enumerated; "
|
||||
"row(s) raced in during the purge and their cached tokens will expire by TTL",
|
||||
"MCP server %s: purge removed %d OAuth credential row(s) but %d were enumerated; "
|
||||
"row(s) were deleted concurrently during the purge",
|
||||
server_id,
|
||||
deleted_count,
|
||||
len(rows),
|
||||
len(oauth_rows),
|
||||
)
|
||||
return deleted_count
|
||||
|
||||
|
|
|
|||
|
|
@ -4073,7 +4073,12 @@ class MCPServerManager:
|
|||
verbose_logger.warning(
|
||||
"Failed to invalidate cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
|
||||
)
|
||||
await self._per_user_token_cache.delete(user_id, server_id)
|
||||
try:
|
||||
await self._per_user_token_cache.delete(user_id, server_id)
|
||||
except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop
|
||||
verbose_logger.warning(
|
||||
"Failed to drop legacy cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc
|
||||
)
|
||||
|
||||
async def _resolve_oauth2_headers_for_tool_call(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -148,19 +148,30 @@ def test_mcp_oauth_token_identity_detects_change_under_encryption():
|
|||
assert mcp_oauth_token_identity(unchanged) != mcp_oauth_token_identity(changed)
|
||||
|
||||
|
||||
def _oauth_row(user_id: str, server_id: str = "srv-1"):
|
||||
"""A stored per-user OAuth token row (payload tagged type=oauth2, legacy plain-base64 encoding)."""
|
||||
row = _legacy_row(json.dumps({"type": "oauth2", "access_token": "tok-" + user_id}))
|
||||
row.user_id = user_id
|
||||
row.server_id = server_id
|
||||
return row
|
||||
|
||||
|
||||
def _byok_row(user_id: str, server_id: str = "srv-1"):
|
||||
"""A stored BYOK API key row: the same column, but the payload is a plain string, not OAuth JSON."""
|
||||
row = _legacy_row("sk-byok-" + user_id)
|
||||
row.user_id = user_id
|
||||
row.server_id = server_id
|
||||
return row
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_invalidates_every_store():
|
||||
"""The purge must route each (user, server) through the injected invalidator (defaulting to the
|
||||
manager's shared invalidation, the single point covering both the legacy per-user token cache and
|
||||
the v2 per-user OAuth token store); evicting only one cache lets the other keep serving a token
|
||||
minted for the old config."""
|
||||
async def test_purge_user_oauth_credentials_for_server_invalidates_each_user():
|
||||
"""The purge must route each (user, server) row through the invalidator exactly once."""
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
r1 = MagicMock(user_id="alice", server_id="srv-1")
|
||||
r2 = MagicMock(user_id="bob", server_id="srv-1")
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[r1, r2])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice"), _oauth_row("bob")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=1)
|
||||
|
||||
invalidations = []
|
||||
|
||||
|
|
@ -170,29 +181,131 @@ async def test_purge_user_oauth_credentials_for_server_invalidates_every_store()
|
|||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1", invalidate_token_cache=record_invalidation)
|
||||
|
||||
assert purged == 2
|
||||
prisma.db.litellm_mcpusercredentials.delete_many.assert_awaited_once()
|
||||
assert prisma.db.litellm_mcpusercredentials.delete_many.await_count == 2
|
||||
assert set(invalidations) == {("alice", "srv-1"), ("bob", "srv-1")}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_spares_byok_rows():
|
||||
"""Regression: the purge used to delete_many on server_id alone, wiping BYOK API keys that share
|
||||
the LiteLLM_MCPUserCredentials table. Only rows holding an OAuth2 payload may be deleted, each by
|
||||
its (user_id, server_id) pair, and only their users' token caches invalidated."""
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_byok_row("carol"), _oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=1)
|
||||
|
||||
invalidations = []
|
||||
|
||||
async def record_invalidation(user_id: str, server_id: str) -> None:
|
||||
invalidations.append((user_id, server_id))
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1", invalidate_token_cache=record_invalidation)
|
||||
|
||||
assert purged == 1
|
||||
prisma.db.litellm_mcpusercredentials.delete_many.assert_awaited_once_with(
|
||||
where={"user_id": "alice", "server_id": "srv-1"}
|
||||
)
|
||||
assert invalidations == [("alice", "srv-1")]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_all_byok_is_noop():
|
||||
"""An api_key (BYOK-only) server whose identity tuple changes (e.g. its url) must purge nothing."""
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_byok_row("carol"), _byok_row("dave")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1")
|
||||
|
||||
assert purged == 0
|
||||
prisma.db.litellm_mcpusercredentials.delete_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_defaults_to_manager_invalidator(monkeypatch):
|
||||
"""When no invalidator is injected, the purge must resolve to the manager's shared
|
||||
invalidate_user_oauth_token_cache, the single point covering both the legacy per-user token cache
|
||||
and the v2 per-user OAuth token store; a wrong or no-op default silently leaves every cache
|
||||
serving tokens minted for the superseded config."""
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=1)
|
||||
|
||||
shared_invalidator = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
mcp_server_manager.global_mcp_server_manager,
|
||||
"invalidate_user_oauth_token_cache",
|
||||
shared_invalidator,
|
||||
)
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1")
|
||||
|
||||
assert purged == 1
|
||||
shared_invalidator.assert_awaited_once_with("alice", "srv-1")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_logs_raced_rows(monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server import db as db_module
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(
|
||||
return_value=[MagicMock(user_id="alice", server_id="srv-1")]
|
||||
)
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=0)
|
||||
warning = MagicMock()
|
||||
monkeypatch.setattr(db_module.verbose_proxy_logger, "warning", warning)
|
||||
|
||||
purged = await purge_user_oauth_credentials_for_server(prisma, "srv-1", invalidate_token_cache=AsyncMock())
|
||||
|
||||
assert purged == 2
|
||||
assert purged == 0
|
||||
warning.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_mcp_server_invalidates_cached_tokens_for_enumerated_users():
|
||||
"""Deleting a server must invalidate each enumerated user's cached per-user token: the caches are
|
||||
keyed by (user_id, server_id), so a re-created server reusing the same server_id would otherwise
|
||||
serve tokens minted for the deleted server until TTL."""
|
||||
from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=MagicMock(server_id="srv-1"))
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice"), _byok_row("bob")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock(return_value=0)
|
||||
|
||||
invalidations = []
|
||||
|
||||
async def record_invalidation(user_id: str, server_id: str) -> None:
|
||||
invalidations.append((user_id, server_id))
|
||||
|
||||
deleted = await delete_mcp_server(prisma, "srv-1", invalidate_token_cache=record_invalidation)
|
||||
|
||||
assert deleted is not None
|
||||
assert set(invalidations) == {("alice", "srv-1"), ("bob", "srv-1")}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_mcp_server_returns_none_without_cleanup_when_server_missing():
|
||||
from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=None)
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock()
|
||||
|
||||
deleted = await delete_mcp_server(prisma, "srv-1", invalidate_token_cache=AsyncMock())
|
||||
|
||||
assert deleted is None
|
||||
prisma.db.litellm_mcpusercredentials.find_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_purge_user_oauth_credentials_for_server_noop_when_empty():
|
||||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
|
|
|||
|
|
@ -6869,6 +6869,7 @@ async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow():
|
|||
(None, None),
|
||||
("", None),
|
||||
("not a url", None),
|
||||
("http://[::1", None),
|
||||
],
|
||||
)
|
||||
def test_redact_mcp_resource_url_strips_credentials(url, expected):
|
||||
|
|
|
|||
|
|
@ -3376,6 +3376,25 @@ class TestMCPServerManager:
|
|||
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
|
||||
assert legacy_cache.deletes == [("alice", "srv-1")]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalidate_user_oauth_token_cache_swallows_legacy_cache_errors(self):
|
||||
"""The legacy cache drop is best-effort like the v2 drop: a failure must be logged, never
|
||||
raised into the credential write that triggered the invalidation."""
|
||||
|
||||
class _Store:
|
||||
async def fetch(self, user_id: str, server_id: str):
|
||||
return None
|
||||
|
||||
async def invalidate(self, user_id: str, server_id: str) -> None:
|
||||
return None
|
||||
|
||||
class _RaisingLegacyCache:
|
||||
async def delete(self, user_id: str, server_id: str) -> None:
|
||||
raise RuntimeError("redis down")
|
||||
|
||||
manager = MCPServerManager(per_user_oauth_token_store=_Store(), per_user_token_cache=_RaisingLegacyCache())
|
||||
await manager.invalidate_user_oauth_token_cache("alice", "srv-1")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_oauth2_headers_no_user_id(self):
|
||||
"""Skip lookup entirely when user_api_key_auth has no user_id."""
|
||||
|
|
|
|||
|
|
@ -5136,7 +5136,7 @@ def test_stamp_oauth2_flow_ignores_non_oauth2():
|
|||
assert payload.oauth2_flow is None
|
||||
|
||||
|
||||
async def _run_edit(old_record, updated_record):
|
||||
async def _run_edit(old_record, updated_record, purge_mock=None):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import edit_mcp_server
|
||||
|
||||
server_id = updated_record.server_id
|
||||
|
|
@ -5163,7 +5163,7 @@ async def _run_edit(old_record, updated_record):
|
|||
patch("litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager") as mock_manager,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.purge_user_oauth_credentials_for_server",
|
||||
AsyncMock(return_value=1),
|
||||
purge_mock if purge_mock is not None else AsyncMock(return_value=1),
|
||||
) as mock_purge,
|
||||
):
|
||||
mock_manager.update_server = AsyncMock()
|
||||
|
|
@ -5199,6 +5199,20 @@ async def test_edit_mcp_server_skips_purge_when_identity_unchanged():
|
|||
mock_purge.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_edit_mcp_server_purge_failure_does_not_fail_the_edit():
|
||||
"""The purge is best-effort: a purge exception after a successful update must be swallowed and
|
||||
logged, never turned into an error response for an edit whose primary job already succeeded."""
|
||||
server_id = str(uuid.uuid4())
|
||||
old = generate_mock_mcp_server_db_record(server_id=server_id, url="https://old.example.com/mcp")
|
||||
updated = generate_mock_mcp_server_db_record(server_id=server_id, url="https://new.example.com/mcp")
|
||||
|
||||
result, mock_purge = await _run_edit(old, updated, purge_mock=AsyncMock(side_effect=RuntimeError("db down")))
|
||||
|
||||
assert result.server_id == server_id
|
||||
mock_purge.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_edit_mcp_server_snapshot_failure_skips_purge_but_edit_succeeds():
|
||||
"""The pre-update snapshot read is advisory (it only feeds the purge decision); a read failure
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue