mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 3a23cbb2d4 into d80f8c28ca
This commit is contained in:
commit
2747214749
9 changed files with 315 additions and 214 deletions
14
.github/workflows/codeql.yml
vendored
14
.github/workflows/codeql.yml
vendored
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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),)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue