fix(mcp): warn on undecryptable per-user env vars and align status with global fallbacks

This commit is contained in:
mateo-berri 2026-06-04 19:10:15 +00:00
parent ae026d2a4d
commit 16f1d3ecb8
No known key found for this signature in database
4 changed files with 64 additions and 4 deletions

View file

@ -1208,6 +1208,12 @@ def _decode_user_env_vars(stored: str) -> Dict[str, str]:
return_original_value=False,
)
if decrypted is None:
if stored:
verbose_proxy_logger.warning(
"MCP per-user env vars failed to decrypt (LITELLM_SALT_KEY "
"changed?); treating as unset so the user is prompted to "
"re-enter them rather than silently forwarding ciphertext"
)
return {}
try:
parsed = json.loads(decrypted)

View file

@ -2214,10 +2214,14 @@ if MCP_AVAILABLE:
each value ``is_set`` and never echoes the decrypted secret back, so a
leaked token can't be used to exfiltrate the raw upstream credential.
"""
_, user_specs = parse_admin_env_vars(getattr(server, "env_vars", None))
global_values, user_specs = parse_admin_env_vars(
getattr(server, "env_vars", None)
)
# Limit "required" to vars that are actually referenced by static_headers.
# If an admin defined a per-user var but never used it, it's not blocking.
# A var only blocks when it's referenced by static_headers and has no
# admin global fallback, mirroring _resolve_static_headers_with_env_vars
# (globals win the merge) so the status endpoint never asks the user for
# credentials a tool call wouldn't actually require.
static_headers = getattr(server, "static_headers", None) or {}
if isinstance(static_headers, str):
try:
@ -2226,7 +2230,9 @@ if MCP_AVAILABLE:
static_headers = {}
referenced = collect_env_var_references(strings=static_headers.values())
user_var_names = {spec["name"] for spec in user_specs}
blocking = referenced & user_var_names
blocking = {
name for name in (referenced & user_var_names) if name not in global_values
}
required: List[MCPUserEnvVarSpec] = []
missing_count = 0

View file

@ -945,6 +945,34 @@ async def test_get_user_env_vars_returns_empty_for_missing_row():
assert await get_user_env_vars(prisma, "alice", "srv-1") == {}
@pytest.mark.asyncio
async def test_decode_user_env_vars_warns_when_undecryptable(
env_vars_salt_key, monkeypatch
):
"""A stored blob encrypted under a previous salt key must surface a warning
(not just a debug line) and decode to ``{}`` so a rotated ``LITELLM_SALT_KEY``
is diagnosable instead of silently sending the user a misleading "set up your
credentials" 412 for values they already stored."""
from unittest.mock import MagicMock
import litellm.proxy._experimental.mcp_server.db as mcp_db
from litellm.proxy._experimental.mcp_server.db import (
_decode_user_env_vars,
store_user_env_vars,
)
prisma = _mock_env_vars_prisma()
await store_user_env_vars(prisma, "alice", "srv-1", {"CORP_PASSWORD": "s3cret"})
blob = _captured_values_blob(prisma)
monkeypatch.setenv("LITELLM_SALT_KEY", "a-totally-different-salt-key-0000")
logger = MagicMock()
monkeypatch.setattr(mcp_db, "verbose_proxy_logger", logger)
assert _decode_user_env_vars(blob) == {}
logger.warning.assert_called_once()
@pytest.mark.asyncio
async def test_get_user_env_vars_bulk_distributes_results(env_vars_salt_key):
from unittest.mock import AsyncMock, MagicMock

View file

@ -3838,6 +3838,26 @@ class TestComputeUserEnvVarStatus:
assert status.required == []
assert status.setup_url is None
def test_dual_scope_var_with_global_fallback_is_not_required(self):
# SHARED_TOKEN is declared both global and user. The global value covers
# the reference (globals win in _resolve_static_headers_with_env_vars),
# so the tool-call path never raises a 412 for it. The status endpoint
# must agree and not report it as required/missing, otherwise it asks the
# user for a credential the request would never actually need.
server = _make_env_var_server(
env_vars=[
{"name": "SHARED_TOKEN", "value": "global-secret", "scope": "global"},
{"name": "SHARED_TOKEN", "value": "", "scope": "user"},
],
static_headers={"Authorization": "Bearer ${SHARED_TOKEN}"},
)
status = mgmt_endpoints._compute_user_env_var_status(
server=server, stored_values={}
)
assert status.required == []
assert status.missing_count == 0
assert status.setup_url is None
class TestGetMCPUserEnvVars:
@pytest.mark.asyncio