fix(mcp): let a salt-key-orphaned OAuth credential be replaced by re-authorization (#37672)

store_user_oauth_credential refused to overwrite any existing row that did not
decode as an OAuth2 payload, which conflated two states: a live BYOK secret that
reads back as plaintext, and ciphertext written under a LITELLM_SALT_KEY the proxy
no longer holds. The second is unrecoverable by any caller, so refusing preserved
nothing and instead wedged the user out of the OAuth flow permanently, since
re-authorizing is their only recovery.

The guard now raises only when the existing value is genuinely readable. An
undecryptable row is logged and replaced by the newly authorized token.

Both read paths were equally silent: get_user_oauth_credential and
list_user_oauth_credentials (which backs the bulk prefetch) each dropped an
undecryptable row indistinguishably from "user never authorized", so an operator
saw an upstream 401 and no hint that a credential had failed to decrypt. Both now
warn with the user and server ids, never the stored value.
This commit is contained in:
Yassin Kortam 2026-08-20 14:11:17 -07:00 • committed by GitHub
parent fc3b160fb5
commit f3639a6fb3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 207 additions and 14 deletions

View file

@ -1224,14 +1224,28 @@ def _decode_user_credential(stored: str) -> str | None:
return None
def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None:
"""Return the OAuth2 payload dict if ``stored`` holds one, else ``None``.
def _warn_undecryptable_credential(user_id: str, server_id: str) -> None:
"""Log the one credential state that otherwise reads as "user never authorized"."""
verbose_proxy_logger.warning(
"MCP user credential for user=%s server=%s could not be decrypted (likely written under a "
"previous LITELLM_SALT_KEY); the user is treated as not connected and must re-authorize.",
user_id,
server_id,
)
def _parse_oauth_payload(decoded: str | None) -> OAuthCredentialPayload | None:
"""Return the OAuth2 payload dict if ``decoded`` holds one, else ``None``.
A row is considered an OAuth2 credential iff its decoded value parses as
a JSON object with ``"type": "oauth2"``. Plain BYOK credentials (which
share the same column) decode to a non-JSON string and return ``None``.
Callers that need to tell an unreadable row from a readable non-OAuth2 one
pass the result of :func:`_decode_user_credential` so a single decode
answers both questions: ``None`` there means the value can be neither
decrypted nor base64-decoded, so no caller can ever recover it.
"""
decoded: Final = _decode_user_credential(stored)
if decoded is None:
return None
parsed: OAuthCredentialPayload | None
@ -1244,6 +1258,11 @@ def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None:
return None
def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None:
"""Return the OAuth2 payload dict held in ``stored``, else ``None``."""
return _parse_oauth_payload(_decode_user_credential(stored))
async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, new_master_key: str):
"""Re-encrypt every ``LiteLLM_MCPUserCredentials`` row with ``new_master_key``.
@ -1415,15 +1434,25 @@ async def store_user_oauth_credential(
# (e.g. during token refresh), saving an extra DB round-trip.
if not skip_byok_guard:
existing: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
if existing is not None and _decode_oauth_payload(existing.credential_b64) is None:
# Existing row is either a BYOK secret or an OAuth2 row that no
# longer decrypts (e.g. after a salt-key rotation). In either
# case, refuse to overwrite — the caller would clobber data
# that may still be recoverable.
raise ValueError(
f"Existing credential for user {user_id} and server "
f"{server_id} could not be verified as an OAuth2 token. "
f"Refusing to overwrite."
decoded: Final = _decode_user_credential(existing.credential_b64) if existing is not None else None
if existing is not None and _parse_oauth_payload(decoded) is None:
# Refuse only while the row still holds readable content, which is a live BYOK
# secret that overwriting would destroy. A row that does not decode was written
# under a different LITELLM_SALT_KEY, and one that decodes to nothing holds no
# secret at all; refusing either preserves nothing and instead wedges the user
# out of the OAuth flow for good, since re-authorizing is their only recovery.
if decoded:
raise ValueError(
f"Existing credential for user {user_id} and server "
f"{server_id} could not be verified as an OAuth2 token. "
f"Refusing to overwrite."
)
verbose_proxy_logger.warning(
"store_user_oauth_credential: existing credential for user=%s server=%s could not be "
"decrypted (likely written under a previous LITELLM_SALT_KEY); replacing it with the "
"newly authorized OAuth2 token.",
user_id,
server_id,
)
encoded: Final = encrypt_value_helper(json.dumps(payload))
@ -1461,7 +1490,10 @@ async def get_user_oauth_credential(
row: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id)
if row is None:
return None
return _decode_oauth_payload(row.credential_b64)
decoded: Final = _decode_user_credential(row.credential_b64)
if decoded is None:
_warn_undecryptable_credential(user_id, server_id)
return _parse_oauth_payload(decoded)
async def list_user_oauth_credentials(
@ -1473,7 +1505,10 @@ async def list_user_oauth_credentials(
rows: Final = await _db_find_user_credential_rows(prisma_client, {"user_id": user_id})
results: Final[list[OAuthCredentialPayload]] = []
for row in rows:
payload = _decode_oauth_payload(row.credential_b64)
decoded = _decode_user_credential(row.credential_b64)
if decoded is None:
_warn_undecryptable_credential(user_id, row.server_id)
payload = _parse_oauth_payload(decoded)
if payload is None:
continue
payload["server_id"] = row.server_id

View file

@ -535,6 +535,164 @@ async def test_byok_guard_allows_overwriting_existing_oauth():
assert _stored_value(prisma) != oauth_row.credential_b64
# ── Recovery from an unplanned LITELLM_SALT_KEY change ────────────────────────
PREVIOUS_SALT_KEY = "the-salt-key-this-deployment-used-before-9999"
def _row_written_under_previous_salt_key(monkeypatch, payload: str):
"""A row encrypted under a salt key the proxy no longer holds.
Asserts the fixture really is undecryptable under the current key, so a test
built on it cannot pass by accident.
"""
monkeypatch.setenv("LITELLM_SALT_KEY", PREVIOUS_SALT_KEY)
encrypted = encrypt_value_helper(payload)
monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY)
assert _decode_user_credential(encrypted) is None, "fixture must not decrypt under the current salt key"
row = MagicMock()
row.credential_b64 = encrypted
row.user_id = "alice"
row.server_id = "srv-1"
return row
@pytest.mark.asyncio
async def test_reauthorization_replaces_row_written_under_previous_salt_key(monkeypatch):
# The wedged user: their row cannot be decrypted, so refusing preserves nothing.
old_payload = json.dumps({"type": "oauth2", "access_token": "tok-written-before-rotation"})
prisma = _make_prisma_with_existing(row=_row_written_under_previous_salt_key(monkeypatch, old_payload))
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok-after-reauthorization")
# The replacement must decrypt under the CURRENT key and be the newly authorized token.
replacement = MagicMock()
replacement.credential_b64 = _stored_value(prisma)
replacement.server_id = "srv-1"
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=replacement)
stored = await get_user_oauth_credential(prisma, "alice", "srv-1")
assert stored is not None
assert stored["access_token"] == "tok-after-reauthorization"
@pytest.mark.asyncio
async def test_readable_byok_is_still_refused_after_a_salt_key_change(monkeypatch):
# A legacy plain-base64 BYOK secret stays readable across a salt-key change, so
# the recovery path must not use it as an excuse to clobber a live credential.
monkeypatch.setenv("LITELLM_SALT_KEY", "a-completely-different-salt-key-4321")
prisma = _make_prisma_with_existing(row=_legacy_row("sk-live-byok-secret"))
with pytest.raises(ValueError, match="could not be verified as an OAuth2"):
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok")
prisma.db.litellm_mcpusercredentials.upsert.assert_not_awaited()
@pytest.mark.asyncio
async def test_recovery_warns_with_identifiers_and_never_logs_credentials(monkeypatch, caplog):
import logging
old_payload = json.dumps({"type": "oauth2", "access_token": "tok-written-before-rotation"})
row = _row_written_under_previous_salt_key(monkeypatch, old_payload)
prisma = _make_prisma_with_existing(row=row)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok-after-reauthorization")
messages = [rec.getMessage() for rec in caplog.records]
matching = [m for m in messages if "could not be decrypted" in m and "replacing it" in m]
assert len(matching) == 1, f"expected one recovery warning, got {messages}"
assert "user=alice" in matching[0] and "server=srv-1" in matching[0]
for secret in ("tok-after-reauthorization", "tok-written-before-rotation", row.credential_b64):
assert secret not in matching[0]
@pytest.mark.asyncio
async def test_get_user_oauth_credential_warns_when_row_cannot_be_decrypted(monkeypatch, caplog):
import logging
old_payload = json.dumps({"type": "oauth2", "access_token": "tok-written-before-rotation"})
prisma = _make_prisma_with_existing(row=_row_written_under_previous_salt_key(monkeypatch, old_payload))
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
assert await get_user_oauth_credential(prisma, "alice", "srv-1") is None
matching = [rec.getMessage() for rec in caplog.records if "could not be decrypted" in rec.getMessage()]
assert len(matching) == 1, f"expected one read-path warning, got {[r.getMessage() for r in caplog.records]}"
assert "user=alice" in matching[0] and "server=srv-1" in matching[0]
@pytest.mark.asyncio
async def test_list_user_oauth_credentials_warns_per_row_when_rows_cannot_be_decrypted(monkeypatch, caplog):
# The bulk prefetch is the other read path, and it is by definition the multi-server case:
# a warning naming the wrong server sends the operator to the wrong place. Two wedged rows
# plus one healthy one, so a warning built from a constant or from the first row is caught.
import logging
old_payload = json.dumps({"type": "oauth2", "access_token": "tok-written-before-rotation"})
wedged_one = _row_written_under_previous_salt_key(monkeypatch, old_payload)
wedged_two = _row_written_under_previous_salt_key(monkeypatch, old_payload)
wedged_two.server_id = "srv-2"
prisma = _make_prisma_with_existing(row=None)
await store_user_oauth_credential(prisma, "alice", "srv-3", "tok-healthy")
healthy = MagicMock()
healthy.credential_b64 = _stored_value(prisma)
healthy.server_id = "srv-3"
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[wedged_one, healthy, wedged_two])
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
result = await list_user_oauth_credentials(prisma, "alice")
assert [cred["server_id"] for cred in result] == ["srv-3"]
matching = [rec.getMessage() for rec in caplog.records if "could not be decrypted" in rec.getMessage()]
assert len(matching) == 2, f"expected one warning per wedged row, got {matching}"
assert all("user=alice" in message for message in matching)
assert {"srv-1", "srv-2"} == {message.split("server=")[1].split(" ")[0] for message in matching}
@pytest.mark.asyncio
async def test_skip_byok_guard_does_not_read_the_existing_row(monkeypatch):
# The refresh paths pass skip_byok_guard=True precisely to save a DB round-trip on the
# hottest MCP path, so the flag has to actually suppress the lookup, not just the raise.
prisma = _make_prisma_with_existing(row=_legacy_row("plain-byok-key"))
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok", skip_byok_guard=True)
prisma.db.litellm_mcpusercredentials.find_unique.assert_not_awaited()
prisma.db.litellm_mcpusercredentials.upsert.assert_awaited_once()
@pytest.mark.asyncio
async def test_blank_credential_row_is_replaced_rather_than_refused():
# A blank value decodes to "" rather than None, so it is not a decryption failure, but it
# holds no secret either. Pinned deliberately: the guard exists to protect readable
# content, and refusing here would wedge the user while preserving nothing.
blank = MagicMock()
blank.credential_b64 = ""
blank.user_id = "alice"
blank.server_id = "srv-1"
assert _decode_user_credential(blank.credential_b64) == "", "fixture must decode to empty, not None"
prisma = _make_prisma_with_existing(row=blank)
await store_user_oauth_credential(prisma, "alice", "srv-1", "tok-after-reauthorization")
prisma.db.litellm_mcpusercredentials.upsert.assert_awaited_once()
@pytest.mark.asyncio
async def test_readable_byok_row_does_not_warn_on_the_read_path(caplog):
# A BYOK row is not a decryption failure; warning on it would train operators to ignore the log.
import logging
prisma = _make_prisma_with_existing(row=_legacy_row("sk-live-byok-secret"))
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
assert await get_user_oauth_credential(prisma, "alice", "srv-1") is None
assert [rec.getMessage() for rec in caplog.records if "could not be decrypted" in rec.getMessage()] == []
# ── list_user_oauth_credentials ───────────────────────────────────────────────