This commit is contained in:
devin-ai-integration[bot] 2026-10-05 18:05:58 +00:00 • committed by GitHub
commit 2747214749
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 315 additions and 214 deletions

View file

@ -62,23 +62,27 @@ jobs:
# a content-integrity hash, not a password or secret hash. SHA-256 is mandated
# by Oracle for this header; see
# https://docs.oracle.com/en-us/iaas/Content/API/Concepts/signingrequests.htm
# The `usedforsecurity=False` flag on the hashlib.sha256 call already declares
# non-security intent, but CodeQL's taint flow still re-fires when callers
# further up the stack are modified. The suppression is scoped to this one
# file/rule pair via SARIF post-filtering so every other callsite of
# The sha256 call carries no usedforsecurity=False hint (kept out for FIPS review), so
# CodeQL's taint flow fires whenever callers further up the stack are modified. The
# suppression is scoped to this one file/rule pair via SARIF post-filtering so every other callsite of
# py/weak-sensitive-data-hashing in the repository continues to be analyzed.
# The same query fires on the HIBP k-anonymity lookup in
# litellm/proxy/auth/password_policy.py, where the password's SHA-1 is only
# a lookup key into the haveibeenpwned range API (the protocol mandates
# SHA-1) and the digest itself never leaves the proxy beyond its first 5
# characters.
- name: Filter SARIF (OCI sha256, HIBP sha1)
# It also fires on the MCP advisory-lock derivations in
# litellm/proxy/_experimental/mcp_server/db.py, which hash identifier
# strings (user_id, server_id, server names) into Postgres advisory lock
# keys, so neither a password nor a secret is hashed.
- name: Filter SARIF (OCI sha256, HIBP sha1, MCP lock ids)
if: matrix.language == 'python'
uses: advanced-security/filter-sarif@2da736ff05ef065cb2894ac6892e47b5eac2c3c0 # v1.1
with:
patterns: |
-litellm/llms/oci/common_utils.py:py/weak-sensitive-data-hashing
-litellm/proxy/auth/password_policy.py:py/weak-sensitive-data-hashing
-litellm/proxy/_experimental/mcp_server/db.py:py/weak-sensitive-data-hashing
input: sarif-results/python.sarif
output: sarif-results/python.sarif

View file

@ -96,15 +96,7 @@ class OCIRequestWrapper:
def sha256_base64(data: bytes) -> str:
# SHA-256 is used here to compute the x-content-sha256 header required by the
# OCI HTTP signing specification (RSA-SHA256 request signing), not for password
# or secret hashing. This is the correct and mandated algorithm for this purpose.
# See: https://docs.oracle.com/en-us/iaas/Content/API/Concepts/signingrequests.htm
#
# ``usedforsecurity=False`` declares non-security intent to static analyzers
# (CodeQL ``py/weak-sensitive-data-hashing``) — without it the request body
# gets flagged as "password-like data" via taint tracking.
digest: Final = hashlib.sha256(data, usedforsecurity=False).digest() # noqa: S324
digest: Final = hashlib.sha256(data).digest() # noqa: S324 # OCI request-signing content digest
return base64.b64encode(digest).decode()

View file

@ -668,7 +668,7 @@ def _mcp_identifier_lock_keys(*identifiers: str | None) -> tuple[int, ...]:
so concurrent requests for the same pair always lock in the same order."""
return tuple(
int.from_bytes(
hashlib.blake2b(f"mcp_identifier:{normalized}".encode(), digest_size=8).digest(),
hashlib.sha256(f"mcp_identifier:{normalized}".encode()).digest()[:8],
"big",
signed=True,
)
@ -2438,7 +2438,7 @@ async def merge_user_env_vars(
"""
allowed: Final = set(allowed_names)
lock_key: Final = int.from_bytes(
hashlib.blake2b(f"{user_id}:{server_id}".encode(), digest_size=8).digest(),
hashlib.sha256(f"{user_id}:{server_id}".encode()).digest()[:8],
"big",
signed=True,
)

View file

@ -204,6 +204,8 @@ class SAMLAuthHandler:
"wantAssertionsSigned": SAMLAuthHandler._bool_env("SAML_WANT_ASSERTIONS_SIGNED", True),
"wantMessagesSigned": SAMLAuthHandler._bool_env("SAML_WANT_MESSAGES_SIGNED", False),
"authnRequestsSigned": SAMLAuthHandler._bool_env("SAML_AUTHN_REQUESTS_SIGNED", False),
"signatureAlgorithm": "http://www.w3.org/2001/04/xmldsig-more#rsa-sha256",
"digestAlgorithm": "http://www.w3.org/2001/04/xmlenc#sha256",
"wantNameId": True,
"requestedAuthnContext": False,
"rejectUnsolicitedResponsesWithInResponseTo": False,

View file

@ -1,3 +1,4 @@
import hashlib
import os
import uuid
from collections.abc import Generator
@ -21,6 +22,23 @@ def read_rows(
return ROWS.validate_python(connection.execute(query, parameters).fetchall())
def advisory_lock_key(*parts: str) -> int:
return int.from_bytes(hashlib.sha256(":".join(parts).encode()).digest()[:8], "big", signed=True)
def legacy_advisory_lock_key(text: str) -> int:
return int.from_bytes(hashlib.blake2b(text.encode(), digest_size=8).digest(), "big", signed=True)
def advisory_waiters(lock_key: int) -> list[dict[str, JsonValue]]:
unsigned: Final = lock_key & 0xFFFFFFFFFFFFFFFF
return read_rows(
"SELECT pid FROM pg_locks WHERE locktype = 'advisory' AND objsubid = 1 "
"AND classid = %s::oid AND objid = %s::oid AND NOT granted",
(str(unsigned >> 32), str(unsigned & 0xFFFFFFFF)),
)
def write_rows(query: LiteralString, parameters: tuple[str, ...], *, database_url: str | None = None) -> None:
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
connection.execute(query, parameters)

View file

@ -1,11 +1,15 @@
import concurrent.futures
import itertools
import os
import uuid
from pathlib import Path
from typing import Final
import psycopg
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.database import advisory_lock_key, advisory_waiters, legacy_advisory_lock_key
from integration._support.mcp import (
McpCaller,
McpPeer,
@ -348,6 +352,46 @@ def test_config_declared_server_behaves_like_database_server_but_is_read_only(ga
assert len(tool_calls(declared_peer.drain())) == 1 and tool_calls(database_peer.drain()) == ()
def test_create_waits_on_the_sha256_advisory_lock_for_the_identifier(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
name: Final = "mgmt" + uuid.uuid4().hex[:8]
lock_key: Final = advisory_lock_key(f"mcp_identifier:{name.lower()}")
with (
concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool,
psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as holder,
):
holder.execute("SELECT pg_advisory_lock(%s::bigint)", (lock_key,))
pending: Final = pool.submit(
gateway.request,
"POST",
"/v1/mcp/server",
{"server_name": name, "alias": name, **peer.registration()},
)
eventually(lambda: advisory_waiters(lock_key), lambda rows: len(rows) == 1, seconds=20)
assert not pending.done()
holder.execute("SELECT pg_advisory_unlock(%s::bigint)", (lock_key,))
response: Final = pending.result(timeout=30)
assert response.status_code == 201, response.text
identity: Final = str(response.json()["server_id"])
scenario.cleanups.callback(forget_mcp, gateway, identity)
assert _servers(gateway)[identity]["server_name"] == name
def test_create_does_not_wait_on_the_legacy_blake2b_lock_id(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
name: Final = "mgmt" + uuid.uuid4().hex[:8]
lock_key: Final = legacy_advisory_lock_key(f"mcp_identifier:{name.lower()}")
with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as holder:
holder.execute("SELECT pg_advisory_lock(%s::bigint)", (lock_key,))
response: Final = gateway.request(
"POST", "/v1/mcp/server", {"server_name": name, "alias": name, **peer.registration()}
)
assert response.status_code == 201, response.text
identity: Final = str(response.json()["server_id"])
scenario.cleanups.callback(forget_mcp, gateway, identity)
assert _servers(gateway)[identity]["server_name"] == name
def test_team_granted_database_server_detail_is_available_to_team_key(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
alias: Final = "lit3974_team_" + uuid.uuid4().hex[:8]
@ -396,11 +440,17 @@ def test_ui_session_lists_and_fetches_team_granted_config_server(
def test_modern_sse_registration_rejected_without_saving(gateway: Gateway, explicit_transport: bool) -> None:
identity: Final = str(uuid.uuid4())
with mcp_peer() as peer:
response: Final = gateway.request("POST", "/v1/mcp/server", {
"server_id": identity, "server_name": "invalid" + uuid.uuid4().hex[:8],
"url": peer.url, "mcp_info": {"protocol_version": "2026-07-28"},
**({"transport": "sse"} if explicit_transport else {}),
})
response: Final = gateway.request(
"POST",
"/v1/mcp/server",
{
"server_id": identity,
"server_name": "invalid" + uuid.uuid4().hex[:8],
"url": peer.url,
"mcp_info": {"protocol_version": "2026-07-28"},
**({"transport": "sse"} if explicit_transport else {}),
},
)
try:
assert response.status_code == 422, response.text
assert "Modern MCP requires HTTP or stdio" in response.text
@ -416,31 +466,67 @@ def test_protocol_transport_updates_validate_effective_configuration(gateway: Ga
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_mcp(scenario, peer, "protocol" + uuid.uuid4().hex[:8])
modern: Final = gateway.request("PUT", "/v1/mcp/server", {
"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"},
})
modern: Final = gateway.request(
"PUT",
"/v1/mcp/server",
{
"server_id": identity,
"mcp_info": {"protocol_version": "2026-07-28"},
},
)
assert modern.status_code == 202, modern.text
assert modern.json()["transport"] == "http"
changed: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "renamed"})
assert changed.status_code == 202, changed.text
assert changed.json()["mcp_info"]["protocol_version"] == "2026-07-28"
snapshot: Final = read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))
snapshot: Final = read_rows(
'SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)
)
peer.drain()
rejected: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "sse", "url": peer.url})
rejected: Final = gateway.request(
"PUT", "/v1/mcp/server", {"server_id": identity, "transport": "sse", "url": peer.url}
)
assert rejected.status_code == 400, rejected.text
assert "Modern MCP requires HTTP or stdio" in rejected.text
assert read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)) == snapshot
assert (
read_rows(
'SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s',
(identity,),
)
== snapshot
)
assert tool_calls(peer.drain()) == ()
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200
legacy: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "sse", "url": peer.url, "mcp_info": {}})
legacy: Final = gateway.request(
"PUT", "/v1/mcp/server", {"server_id": identity, "transport": "sse", "url": peer.url, "mcp_info": {}}
)
assert legacy.status_code == 202, legacy.text
legacy_snapshot: Final = read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))
rejected_protocol: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"}})
legacy_snapshot: Final = read_rows(
'SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)
)
rejected_protocol: Final = gateway.request(
"PUT", "/v1/mcp/server", {"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"}}
)
assert rejected_protocol.status_code == 400, rejected_protocol.text
assert read_rows('SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)) == legacy_snapshot
repaired: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "transport": "http", "url": peer.url, "mcp_info": {"protocol_version": "2026-07-28"}})
assert (
read_rows(
'SELECT transport, mcp_info, updated_at::text FROM "LiteLLM_MCPServerTable" WHERE server_id=%s',
(identity,),
)
== legacy_snapshot
)
repaired: Final = gateway.request(
"PUT",
"/v1/mcp/server",
{
"server_id": identity,
"transport": "http",
"url": peer.url,
"mcp_info": {"protocol_version": "2026-07-28"},
},
)
assert repaired.status_code == 202, repaired.text
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200
@ -460,27 +546,53 @@ def test_concurrent_protocol_transport_edits_cannot_save_incompatible_configurat
with ThreadPoolExecutor(max_workers=2) as pool:
# Hold the row so both workers read the old configuration before either can write.
with psycopg.connect(os.environ["DATABASE_URL"]) as blocker:
blocker.execute('SELECT server_id FROM "LiteLLM_MCPServerTable" WHERE server_id=%s FOR UPDATE', (identity,))
protocol = pool.submit(gateway.request, "PUT", "/v1/mcp/server", {
"server_id": identity, "mcp_info": {"protocol_version": "2026-07-28"},
**({"alias": "renamed" + uuid.uuid4().hex[:8]} if rename else {}),
})
transport = pool.submit(peer.request, "PUT", "/v1/mcp/server", {
"server_id": identity, "transport": "sse", "url": upstream.url,
})
blocker.execute(
'SELECT server_id FROM "LiteLLM_MCPServerTable" WHERE server_id=%s FOR UPDATE', (identity,)
)
protocol = pool.submit(
gateway.request,
"PUT",
"/v1/mcp/server",
{
"server_id": identity,
"mcp_info": {"protocol_version": "2026-07-28"},
**({"alias": "renamed" + uuid.uuid4().hex[:8]} if rename else {}),
},
)
transport = pool.submit(
peer.request,
"PUT",
"/v1/mcp/server",
{
"server_id": identity,
"transport": "sse",
"url": upstream.url,
},
)
eventually(
lambda: read_rows("SELECT pid FROM pg_stat_activity WHERE datname=current_database() AND wait_event_type='Lock' AND query LIKE %s", ('%LiteLLM_MCPServerTable%',)),
lambda: read_rows(
"SELECT pid FROM pg_stat_activity WHERE datname=current_database() AND wait_event_type='Lock' AND query LIKE %s",
("%LiteLLM_MCPServerTable%",),
),
lambda rows: len(rows) >= 2,
seconds=3,
)
responses: Final = [protocol.result(), transport.result()]
assert sorted(response.status_code for response in responses) == [202, 400], [r.text for r in responses]
saved: Final = read_rows('SELECT transport, mcp_info FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,))[0]
saved: Final = read_rows(
'SELECT transport, mcp_info FROM "LiteLLM_MCPServerTable" WHERE server_id=%s', (identity,)
)[0]
assert saved["transport"] == "http" or saved["mcp_info"] != {"protocol_version": "2026-07-28"}
repaired: Final = gateway.request("PUT", "/v1/mcp/server", {
"server_id": identity, "transport": "http", "url": upstream.url,
"mcp_info": {"protocol_version": "2026-07-28"},
})
repaired: Final = gateway.request(
"PUT",
"/v1/mcp/server",
{
"server_id": identity,
"transport": "http",
"url": upstream.url,
"mcp_info": {"protocol_version": "2026-07-28"},
},
)
assert repaired.status_code == 202, repaired.text
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
assert call_tool(gateway, key, identity, "add", ADD).status_code == 200

View file

@ -1,3 +1,4 @@
import os
import signal
import uuid
from collections.abc import Mapping
@ -8,9 +9,10 @@ from pathlib import Path
from typing import Final
import httpx
import psycopg
import pytest
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.database import advisory_lock_key, advisory_waiters, legacy_advisory_lock_key, read_rows
from integration._support.mcp import McpPeer, call_tool, mcp_peer, register_mcp, tool_names
from integration._support.process import owned_proxy_process
from pydantic import JsonValue, TypeAdapter
@ -357,3 +359,37 @@ def test_concurrent_stores_of_different_variables_do_not_lose_an_update(gateway:
with ThreadPoolExecutor(max_workers=2) as pool:
for _ in range(5):
race_once(pool)
def test_store_waits_on_the_sha256_advisory_lock_for_the_user_and_server(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_user_var_server(scenario, peer, TOKEN)
user: Final = scenario.user()
key: Final = scenario.key(user_id=user, object_permission=grants(identity))
wait_for_tools(gateway, key, identity)
lock_key: Final = advisory_lock_key(user, identity)
with (
ThreadPoolExecutor(max_workers=1) as pool,
psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as holder,
):
holder.execute("SELECT pg_advisory_lock(%s::bigint)", (lock_key,))
pending: Final = pool.submit(store, gateway, key, identity, {TOKEN: "held-value"})
eventually(lambda: advisory_waiters(lock_key), lambda rows: len(rows) == 1, seconds=20)
assert not pending.done()
holder.execute("SELECT pg_advisory_unlock(%s::bigint)", (lock_key,))
response: Final = pending.result(timeout=30)
assert response.status_code == 200, response.text
assert set_names(env_status(gateway, key, identity)) == {TOKEN: True}
def test_store_does_not_wait_on_the_legacy_blake2b_lock_id(gateway: Gateway) -> None:
with mcp_peer() as peer, gateway.scenario() as scenario:
identity: Final = register_user_var_server(scenario, peer, TOKEN)
user: Final = scenario.user()
key: Final = scenario.key(user_id=user, object_permission=grants(identity))
lock_key: Final = legacy_advisory_lock_key(f"{user}:{identity}")
with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as holder:
holder.execute("SELECT pg_advisory_lock(%s::bigint)", (lock_key,))
response: Final = store(gateway, key, identity, {TOKEN: "free-value"})
assert response.status_code == 200, response.text
assert set_names(env_status(gateway, key, identity)) == {TOKEN: True}

View file

@ -6,14 +6,13 @@ connection. The DB-backed per-user flow is exercised in higher-level
tests in tests/mcp_tests.
"""
import hashlib
from typing import Final
from unittest.mock import AsyncMock
import pytest
from respx import MockRouter
from litellm.types.mcp_server.mcp_server_manager import MCPServer
# Look up these names lazily on every access. Tests in this directory call
# ``importlib.reload`` on the utils module to exercise registration logic,
# which replaces ``MCPMissingUserEnvVarsError`` with a freshly-constructed
@ -22,6 +21,8 @@ from litellm.types.mcp_server.mcp_server_manager import MCPServer
# stop matching the new class. Accessing the attribute through the module
# always picks up the current version.
import litellm.proxy._experimental.mcp_server.utils as _mcp_utils
from litellm.proxy._experimental.mcp_server.db import _mcp_identifier_lock_keys, merge_user_env_vars
from litellm.types.mcp_server.mcp_server_manager import MCPServer
def _u(name: str):
@ -78,16 +79,12 @@ def test_find_env_var_references():
def test_collect_env_var_references():
refs = _u("collect_env_var_references")(
strings=["${A}", "static", "${B}-${C}", None]
)
refs = _u("collect_env_var_references")(strings=["${A}", "static", "${B}-${C}", None])
assert refs == {"A", "B", "C"}
def test_interpolate_env_vars_replaces_known_and_leaves_unknown():
assert _u("interpolate_env_vars")(
"${A}://${B}/${C}", {"A": "https", "B": "host"}
) == ("https://host/${C}")
assert _u("interpolate_env_vars")("${A}://${B}/${C}", {"A": "https", "B": "host"}) == ("https://host/${C}")
def test_interpolate_headers_returns_independent_copy():
@ -188,9 +185,7 @@ def mock_server():
@pytest.mark.asyncio
async def test_resolve_static_headers_interpolates_globals_and_user(
mock_server, monkeypatch
):
async def test_resolve_static_headers_interpolates_globals_and_user(mock_server, monkeypatch):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
@ -203,9 +198,7 @@ async def test_resolve_static_headers_interpolates_globals_and_user(
monkeypatch.setattr(manager, "_load_user_env_vars", fake_load_user_env_vars)
headers = await manager._resolve_static_headers_with_env_vars(
mock_server, user_api_key_auth=object()
)
headers = await manager._resolve_static_headers_with_env_vars(mock_server, user_api_key_auth=object())
assert headers == {
"X-DB-URL": "postgres://alice:s3cret@db.local/db",
"X-Other": "literal",
@ -213,36 +206,28 @@ async def test_resolve_static_headers_interpolates_globals_and_user(
@pytest.mark.asyncio
async def test_resolve_static_headers_raises_when_user_vars_missing(
mock_server, monkeypatch
):
async def test_resolve_static_headers_raises_when_user_vars_missing(mock_server, monkeypatch):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
manager = MCPServerManager()
async def fake_load_user_env_vars(
server, user_api_key_auth, *, force_refresh=False
):
async def fake_load_user_env_vars(server, user_api_key_auth, *, force_refresh=False):
# User has only filled in one of the two required vars
return {"CORP_USERNAME": "alice"}
monkeypatch.setattr(manager, "_load_user_env_vars", fake_load_user_env_vars)
with pytest.raises(_u("MCPMissingUserEnvVarsError")) as exc:
await manager._resolve_static_headers_with_env_vars(
mock_server, user_api_key_auth=object()
)
await manager._resolve_static_headers_with_env_vars(mock_server, user_api_key_auth=object())
assert exc.value.missing == ["CORP_PASSWORD"]
assert exc.value.server_id == "srv-1"
assert "fill_env_vars=srv-1" in exc.value.setup_url
@pytest.mark.asyncio
async def test_resolve_static_headers_rechecks_db_before_raising_412(
mock_server, monkeypatch
):
async def test_resolve_static_headers_rechecks_db_before_raising_412(mock_server, monkeypatch):
"""A stale cached negative must not produce a 412 on the tool-call path.
Cache invalidation is process-local, so a user who stored values on another
@ -258,9 +243,7 @@ async def test_resolve_static_headers_rechecks_db_before_raising_412(
calls = []
async def fake_load_user_env_vars(
server, user_api_key_auth, *, force_refresh=False
):
async def fake_load_user_env_vars(server, user_api_key_auth, *, force_refresh=False):
calls.append(force_refresh)
if force_refresh:
# Fresh DB read sees the values the user stored on another worker.
@ -270,9 +253,7 @@ async def test_resolve_static_headers_rechecks_db_before_raising_412(
monkeypatch.setattr(manager, "_load_user_env_vars", fake_load_user_env_vars)
headers = await manager._resolve_static_headers_with_env_vars(
mock_server, user_api_key_auth=object()
)
headers = await manager._resolve_static_headers_with_env_vars(mock_server, user_api_key_auth=object())
assert headers == {
"X-DB-URL": "postgres://alice:s3cret@db.local/db",
"X-Other": "literal",
@ -282,9 +263,7 @@ async def test_resolve_static_headers_rechecks_db_before_raising_412(
@pytest.mark.asyncio
async def test_resolve_static_headers_missing_is_non_blocking_for_listing(
mock_server, monkeypatch
):
async def test_resolve_static_headers_missing_is_non_blocking_for_listing(mock_server, monkeypatch):
"""With raise_on_missing=False (the tool-list path), missing per-user vars
must NOT raise. Available vars interpolate; unfilled ${NAME} refs are left
untouched so the server's tools still appear in the listing."""
@ -312,9 +291,7 @@ async def test_resolve_static_headers_missing_is_non_blocking_for_listing(
@pytest.mark.asyncio
async def test_resolve_static_headers_propagates_db_error_on_tool_call(
mock_server, monkeypatch
):
async def test_resolve_static_headers_propagates_db_error_on_tool_call(mock_server, monkeypatch):
"""A DB failure on the tool-call path must surface as a real error, not be
masked as a "missing credentials" MCPMissingUserEnvVarsError (412)."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -329,15 +306,11 @@ async def test_resolve_static_headers_propagates_db_error_on_tool_call(
monkeypatch.setattr(manager, "_load_user_env_vars", boom)
with pytest.raises(RuntimeError, match="db down"):
await manager._resolve_static_headers_with_env_vars(
mock_server, user_api_key_auth=object()
)
await manager._resolve_static_headers_with_env_vars(mock_server, user_api_key_auth=object())
@pytest.mark.asyncio
async def test_resolve_static_headers_swallows_db_error_on_listing(
mock_server, monkeypatch
):
async def test_resolve_static_headers_swallows_db_error_on_listing(mock_server, monkeypatch):
"""On the listing path a DB failure is non-blocking: globals interpolate
and unfilled per-user ${NAME} refs are left untouched."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
@ -480,17 +453,13 @@ async def test_resolve_static_headers_dual_scope_var_uses_global_without_412(
load_calls = []
async def fake_load_user_env_vars(
server, user_api_key_auth, *, force_refresh=False
):
async def fake_load_user_env_vars(server, user_api_key_auth, *, force_refresh=False):
load_calls.append(force_refresh)
return {}
monkeypatch.setattr(manager, "_load_user_env_vars", fake_load_user_env_vars)
headers = await manager._resolve_static_headers_with_env_vars(
server, user_api_key_auth=object()
)
headers = await manager._resolve_static_headers_with_env_vars(server, user_api_key_auth=object())
assert headers == {"Authorization": "Bearer global-secret"}
# The global fully covers the reference, so no per-user lookup is needed.
assert load_calls == []
@ -522,17 +491,13 @@ async def test_resolve_static_headers_empty_global_does_not_cover_user_var(
],
)
async def fake_load_user_env_vars(
server, user_api_key_auth, *, force_refresh=False
):
async def fake_load_user_env_vars(server, user_api_key_auth, *, force_refresh=False):
return {}
monkeypatch.setattr(manager, "_load_user_env_vars", fake_load_user_env_vars)
with pytest.raises(_u("MCPMissingUserEnvVarsError")) as exc:
await manager._resolve_static_headers_with_env_vars(
server, user_api_key_auth=object()
)
await manager._resolve_static_headers_with_env_vars(server, user_api_key_auth=object())
assert exc.value.missing == ["SHARED_TOKEN"]
@ -561,16 +526,12 @@ async def test_resolve_static_headers_user_value_wins_over_empty_global(
],
)
async def fake_load_user_env_vars(
server, user_api_key_auth, *, force_refresh=False
):
async def fake_load_user_env_vars(server, user_api_key_auth, *, force_refresh=False):
return {"SHARED_TOKEN": "user-secret"}
monkeypatch.setattr(manager, "_load_user_env_vars", fake_load_user_env_vars)
headers = await manager._resolve_static_headers_with_env_vars(
server, user_api_key_auth=object()
)
headers = await manager._resolve_static_headers_with_env_vars(server, user_api_key_auth=object())
assert headers == {"Authorization": "Bearer user-secret"}
@ -656,9 +617,7 @@ async def test_load_user_env_vars_returns_empty_without_user():
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager = MCPServerManager()
server = MCPServer(
server_id="s", name="s", transport="http", url="https://example.com"
)
server = MCPServer(server_id="s", name="s", transport="http", url="https://example.com")
assert await manager._load_user_env_vars(server, None) == {}
@ -673,9 +632,7 @@ async def test_load_user_env_vars_returns_empty_without_user_id():
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager = MCPServerManager()
server = MCPServer(
server_id="s", name="s", transport="http", url="https://example.com"
)
server = MCPServer(server_id="s", name="s", transport="http", url="https://example.com")
fake_auth = MagicMock()
fake_auth.user_id = None
assert await manager._load_user_env_vars(server, fake_auth) == {}
@ -695,9 +652,7 @@ async def test_load_user_env_vars_raises_when_db_unavailable(monkeypatch):
from litellm.types.mcp_server.mcp_server_manager import MCPServer
manager = MCPServerManager()
server = MCPServer(
server_id="s", name="s", transport="http", url="https://example.com"
)
server = MCPServer(server_id="s", name="s", transport="http", url="https://example.com")
fake_auth = MagicMock()
fake_auth.user_id = "alice"
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
@ -706,9 +661,7 @@ async def test_load_user_env_vars_raises_when_db_unavailable(monkeypatch):
@pytest.mark.asyncio
async def test_resolve_static_headers_db_unavailable_is_not_missing_412(
mock_server, monkeypatch
):
async def test_resolve_static_headers_db_unavailable_is_not_missing_412(mock_server, monkeypatch):
"""On the tool-call path, an unavailable DB must surface as a real error
rather than a misleading MCPMissingUserEnvVarsError (412). This guards the
regression where ``_load_user_env_vars`` returned ``{}`` when prisma_client
@ -725,9 +678,7 @@ async def test_resolve_static_headers_db_unavailable_is_not_missing_412(
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
with pytest.raises(RuntimeError, match="database connection"):
await manager._resolve_static_headers_with_env_vars(
mock_server, user_api_key_auth=fake_auth
)
await manager._resolve_static_headers_with_env_vars(mock_server, user_api_key_auth=fake_auth)
@pytest.mark.asyncio
@ -751,9 +702,7 @@ async def test_load_user_env_vars_caches_within_ttl(env_vars_salt_key, monkeypat
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
manager = MCPServerManager()
server = MCPServer(
server_id="srv-1", name="s", transport="http", url="https://example.com"
)
server = MCPServer(server_id="srv-1", name="s", transport="http", url="https://example.com")
fake_auth = MagicMock()
fake_auth.user_id = "alice"
@ -766,9 +715,7 @@ async def test_load_user_env_vars_caches_within_ttl(env_vars_salt_key, monkeypat
@pytest.mark.asyncio
async def test_load_user_env_vars_force_refresh_bypasses_cache(
env_vars_salt_key, monkeypatch
):
async def test_load_user_env_vars_force_refresh_bypasses_cache(env_vars_salt_key, monkeypatch):
"""force_refresh re-reads from the DB even with a fresh cached entry, so a
process-local stale value cannot mask credentials stored on another worker."""
from unittest.mock import AsyncMock, MagicMock
@ -787,33 +734,25 @@ async def test_load_user_env_vars_force_refresh_bypasses_cache(
new_row.values_b64 = _encrypted_user_env_blob({"TOKEN": "new"})
prisma = _mock_env_vars_prisma()
prisma.db.litellm_mcpuserenvvars.find_unique = AsyncMock(
side_effect=[old_row, new_row]
)
prisma.db.litellm_mcpuserenvvars.find_unique = AsyncMock(side_effect=[old_row, new_row])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
manager = MCPServerManager()
server = MCPServer(
server_id="srv-1", name="s", transport="http", url="https://example.com"
)
server = MCPServer(server_id="srv-1", name="s", transport="http", url="https://example.com")
fake_auth = MagicMock()
fake_auth.user_id = "alice"
assert await manager._load_user_env_vars(server, fake_auth) == {"TOKEN": "old"}
# A normal load is served from cache (still "old"); force_refresh re-reads.
assert await manager._load_user_env_vars(server, fake_auth) == {"TOKEN": "old"}
assert await manager._load_user_env_vars(server, fake_auth, force_refresh=True) == {
"TOKEN": "new"
}
assert await manager._load_user_env_vars(server, fake_auth, force_refresh=True) == {"TOKEN": "new"}
assert prisma.db.litellm_mcpuserenvvars.find_unique.await_count == 2
mgr_mod._user_env_vars_cache.clear()
@pytest.mark.asyncio
async def test_load_user_env_vars_invalidation_forces_refetch(
env_vars_salt_key, monkeypatch
):
async def test_load_user_env_vars_invalidation_forces_refetch(env_vars_salt_key, monkeypatch):
"""After invalidation (store/clear) the next load reads fresh from the DB
instead of serving the stale cached value."""
from unittest.mock import AsyncMock, MagicMock
@ -833,15 +772,11 @@ async def test_load_user_env_vars_invalidation_forces_refetch(
new_row.values_b64 = _encrypted_user_env_blob({"TOKEN": "new"})
prisma = _mock_env_vars_prisma()
prisma.db.litellm_mcpuserenvvars.find_unique = AsyncMock(
side_effect=[old_row, new_row]
)
prisma.db.litellm_mcpuserenvvars.find_unique = AsyncMock(side_effect=[old_row, new_row])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
manager = MCPServerManager()
server = MCPServer(
server_id="srv-1", name="s", transport="http", url="https://example.com"
)
server = MCPServer(server_id="srv-1", name="s", transport="http", url="https://example.com")
fake_auth = MagicMock()
fake_auth.user_id = "alice"
@ -980,9 +915,7 @@ async def test_merge_user_env_vars_does_not_persist_plaintext(env_vars_salt_key)
prisma = _transactional_env_vars_prisma()
values = {"CORP_USERNAME": "alice", "CORP_PASSWORD": "s3cret"}
await merge_user_env_vars(
prisma, "alice", "srv-1", values, allowed_names=values.keys()
)
await merge_user_env_vars(prisma, "alice", "srv-1", values, allowed_names=values.keys())
row = await prisma.db.litellm_mcpuserenvvars.find_unique(
where={"user_id_server_id": {"user_id": "alice", "server_id": "srv-1"}}
@ -1045,9 +978,7 @@ async def test_get_user_env_vars_treats_a_blob_that_is_not_a_json_object_as_unse
@pytest.mark.asyncio
async def test_decode_user_env_vars_warns_when_undecryptable(
env_vars_salt_key, monkeypatch
):
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
@ -1196,9 +1127,7 @@ async def test_merge_user_env_vars_acquires_lock_without_deserializing_void(
{
"user_facing_error": {
"error_code": "P2010",
"meta": {
"message": "Failed to deserialize column of type 'void'."
},
"meta": {"message": "Failed to deserialize column of type 'void'."},
}
}
)
@ -1217,9 +1146,7 @@ async def test_merge_user_env_vars_acquires_lock_without_deserializing_void(
prisma.db.tx = MagicMock(return_value=tx)
values = {"CORP_TOKEN": "t0ken"}
merged = await merge_user_env_vars(
prisma, "alice", "srv-1", values, allowed_names=values.keys()
)
merged = await merge_user_env_vars(prisma, "alice", "srv-1", values, allowed_names=values.keys())
assert merged == values
assert tx.stored is not None
@ -1272,9 +1199,7 @@ async def test_delete_mcp_server_succeeds_when_orphan_cleanup_fails():
deleted = object()
prisma = _mock_env_vars_prisma()
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=deleted)
prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock(
side_effect=Exception("connection pool exhausted")
)
prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock(side_effect=Exception("connection pool exhausted"))
result = await delete_mcp_server(prisma, "srv-1")
@ -1329,9 +1254,7 @@ async def test_delete_mcp_server_credential_cleanup_failure_still_cleans_env_var
deleted = object()
prisma = _mock_env_vars_prisma()
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=deleted)
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(
side_effect=Exception("connection pool exhausted")
)
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(side_effect=Exception("connection pool exhausted"))
result = await delete_mcp_server(prisma, "srv-1")
@ -1449,9 +1372,7 @@ async def test_build_mcp_server_from_table_decrypts_global_env_vars(env_vars_sal
)
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPEnvVar
req = _global_env_var_server_request(
[MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")]
)
req = _global_env_var_server_request([MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")])
prepared = _prepare_mcp_server_data(req)
table = LiteLLM_MCPServerTable(
@ -1490,9 +1411,7 @@ async def test_add_server_does_not_double_decrypt_global_env_vars(env_vars_salt_
)
from litellm.proxy._types import LiteLLM_MCPServerTable, MCPEnvVar
req = _global_env_var_server_request(
[MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")]
)
req = _global_env_var_server_request([MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")])
env_vars = json.loads(_prepare_mcp_server_data(req)["env_vars"])
# Mirror what create_mcp_server / get_mcp_server return to add_server.
decrypt_global_env_var_values(env_vars)
@ -1540,9 +1459,7 @@ async def test_create_mcp_server_decrypts_env_vars_when_prisma_returns_json_stri
UpdateMCPServerRequest,
)
req = _global_env_var_server_request(
[MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")]
)
req = _global_env_var_server_request([MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")])
encrypted_env_vars_str = _prepare_mcp_server_data(req)["env_vars"]
assert "s3cr3t-p@ss" not in encrypted_env_vars_str
@ -1551,17 +1468,27 @@ async def test_create_mcp_server_decrypts_env_vars_when_prisma_returns_json_stri
from prisma.models import LiteLLM_MCPServerTable
return LiteLLM_MCPServerTable.model_validate({
"server_id": "srv-returned", "transport": "http", "mcp_access_groups": [], "allowed_tools": [],
"extra_headers": [], "args": [], "allow_all_keys": False, "available_on_public_internet": True,
"delegate_auth_to_upstream": False, "oauth_passthrough": False, "per_server_oauth_discovery": False,
"is_byok": False, "byok_description": [], "env_vars": json.dumps(encrypted_env_vars_str),
})
return LiteLLM_MCPServerTable.model_validate(
{
"server_id": "srv-returned",
"transport": "http",
"mcp_access_groups": [],
"allowed_tools": [],
"extra_headers": [],
"args": [],
"allow_all_keys": False,
"available_on_public_internet": True,
"delegate_auth_to_upstream": False,
"oauth_passthrough": False,
"per_server_oauth_discovery": False,
"is_byok": False,
"byok_description": [],
"env_vars": json.dumps(encrypted_env_vars_str),
}
)
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.create = AsyncMock(
return_value=_prisma_row_with_json_string_env_vars()
)
mock_prisma.db.litellm_mcpservertable.create = AsyncMock(return_value=_prisma_row_with_json_string_env_vars())
created = await create_mcp_server(
mock_prisma,
@ -1578,9 +1505,7 @@ async def test_create_mcp_server_decrypts_env_vars_when_prisma_returns_json_stri
assert created.env == {}
mock_prisma_upd = MagicMock()
mock_prisma_upd.db.litellm_mcpservertable.update = AsyncMock(
return_value=_prisma_row_with_json_string_env_vars()
)
mock_prisma_upd.db.litellm_mcpservertable.update = AsyncMock(return_value=_prisma_row_with_json_string_env_vars())
updated = await update_mcp_server(
mock_prisma_upd,
UpdateMCPServerRequest(server_id="srv-update"),
@ -1604,15 +1529,11 @@ def test_reencrypt_global_env_var_values_handles_json_string(env_vars_salt_key):
)
from litellm.proxy._types import MCPEnvVar
req = _global_env_var_server_request(
[MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")]
)
req = _global_env_var_server_request([MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")])
encrypted_env_vars_str = _prepare_mcp_server_data(req)["env_vars"]
original_ciphertext = json.loads(encrypted_env_vars_str)[0]["value"]
rebuilt = _reencrypt_global_env_var_values(
encrypted_env_vars_str, new_encryption_key="rotated-master-key-0000"
)
rebuilt = _reencrypt_global_env_var_values(encrypted_env_vars_str, new_encryption_key="rotated-master-key-0000")
assert rebuilt is not None
assert rebuilt[0]["name"] == "DB_PASSWORD"
@ -1621,9 +1542,7 @@ def test_reencrypt_global_env_var_values_handles_json_string(env_vars_salt_key):
@pytest.mark.asyncio
async def test_rotate_mcp_user_env_vars_logs_rotated_and_skipped_counts(
env_vars_salt_key, monkeypatch
):
async def test_rotate_mcp_user_env_vars_logs_rotated_and_skipped_counts(env_vars_salt_key, monkeypatch):
"""Master-key rotation is a rare, high-stakes batch op, so it emits one
summary line. The counts must track real work: a decryptable row is
re-encrypted and counted as rotated, while a row that no longer decrypts is
@ -1647,17 +1566,13 @@ async def test_rotate_mcp_user_env_vars_logs_rotated_and_skipped_counts(
# Encrypted under an unrelated key, so it won't decrypt under the active salt
# key and must be skipped rather than re-encrypted.
undecryptable = encrypt_value_helper(
json.dumps({"X": "y"}), new_encryption_key="unrelated-key-9999"
)
undecryptable = encrypt_value_helper(json.dumps({"X": "y"}), new_encryption_key="unrelated-key-9999")
good_one = _row("alice", "srv-1", _encrypted_user_env_blob({"GH_TOKEN": "tok-1"}))
good_two = _row("bob", "srv-2", _encrypted_user_env_blob({"GH_TOKEN": "tok-2"}))
bad = _row("carol", "srv-3", undecryptable)
prisma = MagicMock()
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(
return_value=[good_one, good_two, bad]
)
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[good_one, good_two, bad])
prisma.db.litellm_mcpuserenvvars.update = AsyncMock()
logger = MagicMock()
@ -1667,10 +1582,7 @@ async def test_rotate_mcp_user_env_vars_logs_rotated_and_skipped_counts(
update = prisma.db.litellm_mcpuserenvvars.update
assert update.await_count == 2
updated_servers = {
call.kwargs["where"]["user_id_server_id"]["server_id"]
for call in update.call_args_list
}
updated_servers = {call.kwargs["where"]["user_id_server_id"]["server_id"] for call in update.call_args_list}
assert updated_servers == {"srv-1", "srv-2"} # srv-3 was skipped, not rotated
for call in update.call_args_list:
assert call.kwargs["data"]["values_b64"] not in (
@ -1684,9 +1596,7 @@ async def test_rotate_mcp_user_env_vars_logs_rotated_and_skipped_counts(
assert info_args[2] == 1 # skipped
def test_decrypt_global_env_var_drops_undecryptable_value(
env_vars_salt_key, monkeypatch
):
def test_decrypt_global_env_var_drops_undecryptable_value(env_vars_salt_key, monkeypatch):
"""A global value encrypted under a previous salt key must be dropped (not
forwarded as ciphertext) and surfaced as a warning, so a rotated
``LITELLM_SALT_KEY`` can't silently leak ciphertext into ``${NAME}`` headers."""
@ -1700,9 +1610,7 @@ def test_decrypt_global_env_var_drops_undecryptable_value(
)
from litellm.proxy._types import MCPEnvVar
req = _global_env_var_server_request(
[MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")]
)
req = _global_env_var_server_request([MCPEnvVar(name="DB_PASSWORD", value="s3cr3t-p@ss", scope="global")])
entries = json.loads(_prepare_mcp_server_data(req)["env_vars"])
ciphertext = entries[0]["value"]
assert ciphertext != "s3cr3t-p@ss" # encrypted under the original salt key
@ -1749,3 +1657,18 @@ async def test_missing_user_env_vars_error_renders_in_mcp_call_tool():
assert "CorporateDB" in text
assert "CORP_USERNAME" in text
assert "fill_env_vars=srv-99" in text
@pytest.mark.asyncio
async def test_merge_user_env_vars_locks_the_sha256_key_for_user_and_server(env_vars_salt_key):
prisma: Final = _transactional_env_vars_prisma()
await merge_user_env_vars(prisma, "alice", "srv-1", {"TOKEN": "x"}, allowed_names=["TOKEN"])
expected: Final = int.from_bytes(hashlib.sha256(b"alice:srv-1").digest()[:8], "big", signed=True)
assert list(prisma.db._store.locks) == [expected]
def test_mcp_identifier_lock_keys_use_sha256_of_the_lowercased_identifier():
keys: Final = _mcp_identifier_lock_keys("Name", "name", None)
assert keys == (int.from_bytes(hashlib.sha256(b"mcp_identifier:name").digest()[:8], "big", signed=True),)

View file

@ -10,6 +10,7 @@ makes a test fail.
import base64
import datetime
import time
from typing import Final, cast
import pytest
from fastapi import HTTPException, Request
@ -22,11 +23,10 @@ from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.x509.oid import NameOID
from onelogin.saml2.settings import OneLogin_Saml2_Settings
from onelogin.saml2.utils import OneLogin_Saml2_Utils
from starlette.datastructures import URL
from typing import cast
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache
@ -749,3 +749,17 @@ class TestSAMLAuthnCookieSecureFlag:
redirect = await SAMLAuthHandler.build_login_redirect(request, cache)
cookie = redirect.headers["set-cookie"]
assert "Secure" not in cookie
@pytest.mark.asyncio
async def test_effective_sp_settings_use_sha256_signature_and_digest(saml_env):
cache: Final = DualCache()
idp_settings: Final = await SAMLAuthHandler._load_idp_settings(cache)
settings: Final = SAMLAuthHandler._build_settings(_fake_request(), idp_settings)
security: Final = settings["security"]
assert security["signatureAlgorithm"] == "http://www.w3.org/2001/04/xmldsig-more#rsa-sha256"
assert security["digestAlgorithm"] == "http://www.w3.org/2001/04/xmlenc#sha256"
effective: Final = OneLogin_Saml2_Settings(settings, sp_validation_only=True).get_security_data()
assert effective["signatureAlgorithm"] == security["signatureAlgorithm"]
assert effective["digestAlgorithm"] == security["digestAlgorithm"]