From 67895955f6e9f40e78a96c89a642adc03b5d6a3c Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 23 Sep 2026 07:20:08 +0000 Subject: [PATCH 01/18] feat(proxy): LITELLM_FIPS_MODE startup gate with provider assertion, TLS verify guard and loud password migration Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_utils/fips.py | 151 ++++++++++++++++++ litellm/proxy/proxy_server.py | 31 +++- tests/integration/_support/process.py | 76 +++++++-- .../configuration/test_fips_mode_boot.py | 66 ++++++++ tests/integration/contracts.json | 15 ++ .../proxy/common_utils/test_fips.py | 118 ++++++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 104 ++++++++++++ 7 files changed, 546 insertions(+), 15 deletions(-) create mode 100644 litellm/proxy/common_utils/fips.py create mode 100644 tests/integration/configuration/test_fips_mode_boot.py create mode 100644 tests/test_litellm/proxy/common_utils/test_fips.py diff --git a/litellm/proxy/common_utils/fips.py b/litellm/proxy/common_utils/fips.py new file mode 100644 index 00000000000..33f35388567 --- /dev/null +++ b/litellm/proxy/common_utils/fips.py @@ -0,0 +1,151 @@ +import hashlib +import os +from collections.abc import Callable +from dataclasses import dataclass +from typing import Final + +from typing_extensions import assert_never + +FIPS_MODE_ENV_VAR: Final = "LITELLM_FIPS_MODE" +SSL_VERIFY_ENV_VAR: Final = "SSL_VERIFY" +SSL_VERIFY_SETTING: Final = "litellm_settings.ssl_verify" +REFUSAL_PREFIX: Final = "LiteLLM proxy refused to start" + +_TRUE_VALUES: Final = frozenset({"true", "1", "yes", "on"}) +_FALSE_VALUES: Final = frozenset({"false", "0", "no", "off", ""}) + + +@dataclass(frozen=True, slots=True) +class FipsModeOff: + pass + + +@dataclass(frozen=True, slots=True) +class FipsModeOn: + pass + + +@dataclass(frozen=True, slots=True) +class MalformedFipsMode: + value: str + + +FipsModeSetting = FipsModeOff | FipsModeOn | MalformedFipsMode + + +@dataclass(frozen=True, slots=True) +class ProviderDoesNotEnforceFips: + pass + + +@dataclass(frozen=True, slots=True) +class TlsVerificationDisabled: + sources: tuple[str, ...] + + +FipsBootRefusal = MalformedFipsMode | ProviderDoesNotEnforceFips | TlsVerificationDisabled +FipsBootVerdict = FipsModeOff | FipsModeOn | FipsBootRefusal + + +class FipsModeError(Exception): + pass + + +def parse_fips_mode(raw: str | None) -> FipsModeSetting: + if raw is None: + return FipsModeOff() + normalized: Final = raw.strip().lower() + if normalized in _TRUE_VALUES: + return FipsModeOn() + if normalized in _FALSE_VALUES: + return FipsModeOff() + return MalformedFipsMode(value=raw) + + +def is_fips_mode(environ: Callable[[str], str | None] = os.environ.get) -> bool: + return isinstance(parse_fips_mode(environ(FIPS_MODE_ENV_VAR)), FipsModeOn) + + +def openssl_enforces_fips() -> bool: + """MD5 is not an approved digest, so an enforcing FIPS provider refuses it even when asked for security use.""" + try: + hashlib.md5(b"", usedforsecurity=True) + except ValueError: + return True + return False + + +def fips_boot_verdict( + *, + raw_fips_mode: str | None, + provider_enforces_fips: Callable[[], bool], + ssl_verify_environment: str | None, + ssl_verify_setting: object, +) -> FipsBootVerdict: + setting: Final = parse_fips_mode(raw_fips_mode) + match setting: + case FipsModeOff() | MalformedFipsMode(): + return setting + case FipsModeOn(): + pass + case _: + assert_never(setting) + disabled: Final = tuple( + source + for source, off in ( + (SSL_VERIFY_ENV_VAR, _is_off(ssl_verify_environment)), + (SSL_VERIFY_SETTING, _is_off(ssl_verify_setting)), + ) + if off + ) + if disabled: + return TlsVerificationDisabled(sources=disabled) + if not provider_enforces_fips(): + return ProviderDoesNotEnforceFips() + return setting + + +def enforce_fips_boot_verdict(verdict: FipsBootVerdict, announce: Callable[[str], object]) -> None: + match verdict: + case FipsModeOff() | FipsModeOn(): + return + case MalformedFipsMode() | ProviderDoesNotEnforceFips() | TlsVerificationDisabled(): + message: Final = render_refusal(verdict) + announce(f"\n{message}\n\n") + raise FipsModeError(message) + case _: + assert_never(verdict) + + +def render_refusal(refusal: FipsBootRefusal) -> str: + match refusal: + case MalformedFipsMode(value=value): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR}={value} is not a boolean.\n" + f"Set {FIPS_MODE_ENV_VAR} to true or false, or unset it." + ) + case ProviderDoesNotEnforceFips(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but this Python does not enforce FIPS.\n" + "Its OpenSSL still allows non-approved algorithms (MD5 succeeded), so passwords and keys would be\n" + "protected with algorithms the FIPS 140-3 policy forbids. Run the proxy from a FIPS image whose\n" + f"OpenSSL FIPS provider is enabled, or unset {FIPS_MODE_ENV_VAR} on a non-FIPS runtime." + ) + case TlsVerificationDisabled(sources=sources): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but TLS certificate verification is disabled by " + f"{' and '.join(sources)}.\nFIPS deployments must verify upstream certificates, so remove the " + "override or point ssl_verify at a CA bundle instead." + ) + case _: + assert_never(refusal) + + +def _is_off(value: object) -> bool: + match value: + case bool(): + return value is False + case str(): + return value.strip().lower() in _FALSE_VALUES - {""} + case _: + return False diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index dab7decd4dc..7e2e063ae88 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -404,6 +404,14 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( encrypt_value_helper, ) from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_body_call_id, with_call_id +from litellm.proxy.common_utils.fips import ( + FIPS_MODE_ENV_VAR, + SSL_VERIFY_ENV_VAR, + enforce_fips_boot_verdict, + fips_boot_verdict, + is_fips_mode, + openssl_enforces_fips, +) from litellm.proxy.common_utils.healthy_model_filter import ( get_hidden_unhealthy_model_names, is_healthy_only_listing_default, @@ -1277,6 +1285,16 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: if isinstance(worker_config, dict): await initialize_from_worker_config(worker_config) + enforce_fips_boot_verdict( + fips_boot_verdict( + raw_fips_mode=os.getenv(FIPS_MODE_ENV_VAR), + provider_enforces_fips=openssl_enforces_fips, + ssl_verify_environment=os.getenv(SSL_VERIFY_ENV_VAR), + ssl_verify_setting=litellm.ssl_verify, + ), + announce=announce_on_stderr_at_exit, + ) + enforce_master_key_boot_verdict( await with_stored_secrets_counted( master_key_boot_verdict( @@ -1316,10 +1334,21 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: try: result: Final = await migrate_passwords_to_scrypt_async(prisma_client) verbose_proxy_logger.info("Password migration: %s", result) + except ValueError as e: + verbose_proxy_logger.error( + "Password migration failed, so plaintext passwords stay unhashed in the database: %s. " + "This is what an OpenSSL FIPS provider reports when the hashing algorithm is not approved.", + e, + ) + if is_fips_mode(): + raise except Exception as e: verbose_proxy_logger.warning("Password migration skipped: %s", e) - asyncio.create_task(_run_pw_migration()) + if is_fips_mode(): + await _run_pw_migration() + else: + asyncio.create_task(_run_pw_migration()) async def _run_agent_grant_id_migration() -> None: from litellm.proxy.agent_endpoints.agent_registry import ( diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 5c44beaa570..6dec960097b 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -68,8 +68,15 @@ def owned_proxy( yield owned.gateway +@dataclass(frozen=True, slots=True) +class LaunchedProxy: + port: int + process: subprocess.Popen[bytes] + log: Path + + @contextmanager -def owned_proxy_process( +def launched_proxy( gateway: Gateway, directory: Path, overrides: Mapping[str, str], @@ -77,7 +84,7 @@ def owned_proxy_process( config: Path | None = None, remove_environment: tuple[str, ...] = (), workers: int = 1, -) -> Iterator[OwnedProxy]: +) -> Iterator[LaunchedProxy]: with socket.socket() as reserve: reserve.bind(("127.0.0.1", 0)) port: Final = reserve.getsockname()[1] @@ -116,18 +123,7 @@ def owned_proxy_process( start_new_session=True, ) try: - with httpx.Client(base_url=f"http://127.0.0.1:{port}", timeout=15, trust_env=False) as client: - deadline: Final = time.monotonic() + 70 - while True: - assert process.poll() is None, "Owned proxy exited before readiness" - try: - if client.get("/health/readiness", timeout=2).status_code == 200: - break - except httpx.TransportError: - pass - assert time.monotonic() < deadline, "Owned proxy readiness deadline exceeded" - time.sleep(0.1) - yield OwnedProxy(Gateway(client, gateway.key, gateway.upstream_url), process, log_path) + yield LaunchedProxy(port, process, log_path) finally: root_stopped: Final = stop_root_process(process) residual: Final = group_members(process.pid) @@ -142,3 +138,55 @@ def owned_proxy_process( survivors: Final = group_members(process.pid) assert not survivors, "Owned proxy child survived cleanup" assert root_stopped and not remaining, "Owned proxy required forced cleanup" + + +def refused_boot_log( + gateway: Gateway, + directory: Path, + overrides: Mapping[str, str], + *, + config: Path | None = None, +) -> str: + """Start the proxy and return its log once it exits non-zero instead of becoming ready.""" + with launched_proxy(gateway, directory, overrides, config=config) as launched: + with httpx.Client(base_url=f"http://127.0.0.1:{launched.port}", timeout=15, trust_env=False) as client: + deadline: Final = time.monotonic() + 70 + while launched.process.poll() is None: + try: + ready: Final = client.get("/health/readiness", timeout=2).status_code == 200 + except httpx.TransportError: + ready = False + assert not ready, f"Proxy became ready instead of refusing to boot:\n{launched.log.read_text()}" + assert time.monotonic() < deadline, "Proxy neither exited nor became ready within the deadline" + time.sleep(0.1) + assert launched.process.returncode != 0, ( + f"Proxy exited 0 instead of refusing to boot:\n{launched.log.read_text()}" + ) + return launched.log.read_text() + + +@contextmanager +def owned_proxy_process( + gateway: Gateway, + directory: Path, + overrides: Mapping[str, str], + *, + config: Path | None = None, + remove_environment: tuple[str, ...] = (), + workers: int = 1, +) -> Iterator[OwnedProxy]: + with launched_proxy( + gateway, directory, overrides, config=config, remove_environment=remove_environment, workers=workers + ) as launched: + with httpx.Client(base_url=f"http://127.0.0.1:{launched.port}", timeout=15, trust_env=False) as client: + deadline: Final = time.monotonic() + 70 + while True: + assert launched.process.poll() is None, "Owned proxy exited before readiness" + try: + if client.get("/health/readiness", timeout=2).status_code == 200: + break + except httpx.TransportError: + pass + assert time.monotonic() < deadline, "Owned proxy readiness deadline exceeded" + time.sleep(0.1) + yield OwnedProxy(Gateway(client, gateway.key, gateway.upstream_url), launched.process, launched.log) diff --git a/tests/integration/configuration/test_fips_mode_boot.py b/tests/integration/configuration/test_fips_mode_boot.py new file mode 100644 index 00000000000..5cf69f21ab3 --- /dev/null +++ b/tests/integration/configuration/test_fips_mode_boot.py @@ -0,0 +1,66 @@ +"""LITELLM_FIPS_MODE is a boot gate: the proxy refuses to serve unless the process really enforces FIPS. + +Every leg launches the real proxy binary against the suite's Postgres and asserts on what an operator sees: +exit status and the refusal text in the log. Nothing is patched inside the proxy. +""" + +import hashlib +from pathlib import Path +from typing import Final + +import pytest +import yaml + +from tests.integration._support.client import Gateway +from tests.integration._support.process import owned_proxy, refused_boot_log + +REFUSAL: Final = "LiteLLM proxy refused to start" + + +def _this_python_enforces_fips() -> bool: + try: + hashlib.md5(b"probe", usedforsecurity=True) + except ValueError: + return True + return False + + +@pytest.mark.covers("other.configuration.fips_mode.refuses_boot_when_provider_does_not_enforce_fips") +def test_fips_mode_refuses_to_serve_when_this_python_does_not_enforce_fips(gateway: Gateway, tmp_path: Path) -> None: + if _this_python_enforces_fips(): + pytest.skip("Runner OpenSSL enforces FIPS, so this leg cannot observe the non-enforcing refusal") + log: Final = refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true"}) + assert REFUSAL in log, log + assert "LITELLM_FIPS_MODE" in log and "does not enforce FIPS" in log, log + + +@pytest.mark.covers("other.configuration.fips_mode.refuses_boot_with_tls_verification_disabled") +@pytest.mark.parametrize("source", ("environment", "config")) +def test_fips_mode_refuses_to_serve_with_tls_verification_disabled( + gateway: Gateway, tmp_path: Path, source: str +) -> None: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + path: Final = tmp_path / "ssl_verify_off.yaml" + settings: Final = {**config.get("litellm_settings", {}), "ssl_verify": False} + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings})) + log: Final = ( + refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true", "SSL_VERIFY": "false"}) + if source == "environment" + else refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true"}, config=path) + ) + assert REFUSAL in log, log + assert "TLS certificate verification is disabled" in log, log + assert ("SSL_VERIFY" if source == "environment" else "litellm_settings.ssl_verify") in log, log + + +@pytest.mark.covers("other.configuration.fips_mode.refuses_boot_on_malformed_value") +def test_fips_mode_refuses_to_serve_on_a_value_that_is_not_a_boolean(gateway: Gateway, tmp_path: Path) -> None: + log: Final = refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "enforced"}) + assert REFUSAL in log, log + assert "LITELLM_FIPS_MODE=enforced" in log and "true or false" in log, log + + +@pytest.mark.covers("other.configuration.fips_mode.off_leaves_boot_unchanged") +def test_fips_mode_off_serves_even_with_tls_verification_disabled(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {"LITELLM_FIPS_MODE": "false", "SSL_VERIFY": "false"}) as candidate: + assert candidate.client.get("/health/readiness").status_code == 200 diff --git a/tests/integration/contracts.json b/tests/integration/contracts.json index a8d9cf1df8b..673f62fd731 100644 --- a/tests/integration/contracts.json +++ b/tests/integration/contracts.json @@ -76,6 +76,21 @@ "tests/integration/configuration/test_effective_settings.py::test_credential_value_update_and_model_reload_reach_provider": [ "mgmt.credential.update.saved_value_reaches_wire" ], + "tests/integration/configuration/test_fips_mode_boot.py::test_fips_mode_refuses_to_serve_when_this_python_does_not_enforce_fips": [ + "other.configuration.fips_mode.refuses_boot_when_provider_does_not_enforce_fips" + ], + "tests/integration/configuration/test_fips_mode_boot.py::test_fips_mode_refuses_to_serve_with_tls_verification_disabled[environment]": [ + "other.configuration.fips_mode.refuses_boot_with_tls_verification_disabled" + ], + "tests/integration/configuration/test_fips_mode_boot.py::test_fips_mode_refuses_to_serve_with_tls_verification_disabled[config]": [ + "other.configuration.fips_mode.refuses_boot_with_tls_verification_disabled" + ], + "tests/integration/configuration/test_fips_mode_boot.py::test_fips_mode_refuses_to_serve_on_a_value_that_is_not_a_boolean": [ + "other.configuration.fips_mode.refuses_boot_on_malformed_value" + ], + "tests/integration/configuration/test_fips_mode_boot.py::test_fips_mode_off_serves_even_with_tls_verification_disabled": [ + "other.configuration.fips_mode.off_leaves_boot_unchanged" + ], "tests/integration/management/test_partial_update_sequences.py::test_denied_key_update_preserves_saved_grants_and_serving": [ "mgmt.key.update.denied_request_preserves_effective_state" ], diff --git a/tests/test_litellm/proxy/common_utils/test_fips.py b/tests/test_litellm/proxy/common_utils/test_fips.py new file mode 100644 index 00000000000..e934cafbdda --- /dev/null +++ b/tests/test_litellm/proxy/common_utils/test_fips.py @@ -0,0 +1,118 @@ +import hashlib + +import pytest + +from litellm.proxy.common_utils.fips import ( + FipsModeError, + FipsModeOff, + FipsModeOn, + MalformedFipsMode, + ProviderDoesNotEnforceFips, + TlsVerificationDisabled, + enforce_fips_boot_verdict, + fips_boot_verdict, + is_fips_mode, + openssl_enforces_fips, + parse_fips_mode, +) + + +def _verdict( + raw: str | None, + *, + enforcing: bool = True, + ssl_env: str | None = None, + ssl_setting: object = True, +): + return fips_boot_verdict( + raw_fips_mode=raw, + provider_enforces_fips=lambda: enforcing, + ssl_verify_environment=ssl_env, + ssl_verify_setting=ssl_setting, + ) + + +@pytest.mark.parametrize("raw", [None, "false", "0", "no", "off", "", " False "]) +def test_unset_and_false_spellings_leave_fips_mode_off(raw): + assert parse_fips_mode(raw) == FipsModeOff() + assert is_fips_mode({"LITELLM_FIPS_MODE": raw}.get) is False + + +@pytest.mark.parametrize("raw", ["true", "1", "yes", "on", " TRUE "]) +def test_true_spellings_turn_fips_mode_on(raw): + assert parse_fips_mode(raw) == FipsModeOn() + assert is_fips_mode({"LITELLM_FIPS_MODE": raw}.get) is True + + +@pytest.mark.parametrize("raw", ["enforced", "2", "strict", "yes please"]) +def test_anything_else_is_malformed_and_refused_with_the_offending_value(raw): + assert parse_fips_mode(raw) == MalformedFipsMode(value=raw) + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(_verdict(raw, enforcing=False), announce=lambda _: None) + assert f"LITELLM_FIPS_MODE={raw} is not a boolean" in str(refused.value) + assert "true or false" in str(refused.value) + + +def test_off_never_consults_the_provider_or_tls_settings(): + def explode() -> bool: + raise AssertionError("provider probe must not run while FIPS mode is off") + + verdict = fips_boot_verdict( + raw_fips_mode=None, provider_enforces_fips=explode, ssl_verify_environment="false", ssl_verify_setting=False + ) + assert verdict == FipsModeOff() + enforce_fips_boot_verdict(verdict, announce=lambda _: pytest.fail("nothing to announce when off")) + + +def test_on_with_an_enforcing_provider_and_verified_tls_boots(): + verdict = _verdict("true", enforcing=True, ssl_env="true", ssl_setting="/etc/ssl/certs/ca.pem") + assert verdict == FipsModeOn() + enforce_fips_boot_verdict(verdict, announce=lambda _: pytest.fail("nothing to announce when on")) + + +def test_on_with_a_non_enforcing_provider_is_refused_and_names_the_fix(): + announced = [] + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(_verdict("true", enforcing=False), announce=announced.append) + message = str(refused.value) + assert message.startswith("LiteLLM proxy refused to start") + assert "LITELLM_FIPS_MODE is on but this Python does not enforce FIPS" in message + assert "FIPS image" in message + assert announced == [f"\n{message}\n\n"] + + +@pytest.mark.parametrize( + "ssl_env, ssl_setting, sources", + [ + ("false", True, ("SSL_VERIFY",)), + ("0", True, ("SSL_VERIFY",)), + (None, False, ("litellm_settings.ssl_verify",)), + (None, "False", ("litellm_settings.ssl_verify",)), + ("no", False, ("SSL_VERIFY", "litellm_settings.ssl_verify")), + ], +) +def test_disabled_tls_verification_is_refused_naming_every_source(ssl_env, ssl_setting, sources): + verdict = _verdict("true", enforcing=True, ssl_env=ssl_env, ssl_setting=ssl_setting) + assert verdict == TlsVerificationDisabled(sources=sources) + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(verdict, announce=lambda _: None) + assert "TLS certificate verification is disabled by " + " and ".join(sources) in str(refused.value) + + +@pytest.mark.parametrize("ssl_setting", [True, "true", "/etc/ssl/certs/ca.pem", None, ""]) +def test_verified_or_custom_bundle_tls_settings_are_not_treated_as_disabled(ssl_setting): + assert _verdict("true", enforcing=True, ssl_setting=ssl_setting) == FipsModeOn() + + +def test_disabled_tls_is_reported_before_the_provider_so_operators_see_config_mistakes_first(): + assert _verdict("true", enforcing=False, ssl_env="false") == TlsVerificationDisabled(sources=("SSL_VERIFY",)) + assert _verdict("true", enforcing=False) == ProviderDoesNotEnforceFips() + + +def test_provider_probe_agrees_with_whether_md5_is_usable_for_security_here(): + try: + hashlib.md5(b"", usedforsecurity=True) + except ValueError: + assert openssl_enforces_fips() is True + else: + assert openssl_enforces_fips() is False diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 89fd9c5c9d4..42ad00d6b13 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1681,6 +1681,110 @@ async def test_proxy_startup_refuses_an_unsafe_master_key_even_when_the_database assert ("could not be checked" in announced[0]) == key_can_have_encrypted_the_database +@pytest.mark.asyncio +async def test_proxy_startup_refuses_fips_mode_when_this_python_does_not_enforce_fips(monkeypatch, tmp_path): + from fastapi import FastAPI + + from litellm.proxy.common_utils.fips import FipsModeError + from litellm.proxy.proxy_server import proxy_startup_event + + _, announced = _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: False) + monkeypatch.setenv("LITELLM_FIPS_MODE", "true") + + with pytest.raises(FipsModeError): + async with proxy_startup_event(FastAPI()): + pass + + assert len(announced) == 1 + assert "does not enforce FIPS" in announced[0] + + +@pytest.mark.asyncio +async def test_proxy_startup_refuses_fips_mode_when_the_config_disables_tls_verification(monkeypatch, tmp_path): + import yaml + from fastapi import FastAPI + + from litellm.proxy.common_utils.fips import FipsModeError + from litellm.proxy.proxy_server import proxy_startup_event + + config_path, announced = _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + config_path.write_text( + yaml.dump({"general_settings": {"master_key": "sk-a-safe-master-key"}, "litellm_settings": {"ssl_verify": False}}) + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) + monkeypatch.setattr(litellm, "ssl_verify", True) + monkeypatch.setenv("LITELLM_FIPS_MODE", "true") + + with pytest.raises(FipsModeError): + async with proxy_startup_event(FastAPI()): + pass + + assert "TLS certificate verification is disabled by litellm_settings.ssl_verify" in announced[0] + + +class _PrismaClientWhoseUserTableCannotHash: + class _Table: + async def find_many(self, where): + raise ValueError("[digital envelope routines] unsupported") + + class _Db: + litellm_usertable = None + + def __init__(self): + self.db = self._Db() + self.db.litellm_usertable = self._Table() + self.writer_db = self.db + + async def connect(self): + pass + + async def disconnect(self): + pass + + async def check_view_exists(self): + pass + + async def health_check(self): + pass + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fips_mode", ["true", "false"]) +async def test_proxy_startup_surfaces_a_password_migration_crypto_failure(monkeypatch, tmp_path, caplog, fips_mode): + from fastapi import FastAPI + + from litellm.proxy.proxy_server import proxy_startup_event + + _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setenv("DATABASE_URL", "postgresql://nobody:nothing@127.0.0.1:1/unreachable") + cannot_hash = _PrismaClientWhoseUserTableCannotHash() + + async def connected(**kwargs): + return cannot_hash + + monkeypatch.setattr("litellm.proxy.proxy_server.ProxyStartupEvent._setup_prisma_client", connected) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) + monkeypatch.setenv("LITELLM_FIPS_MODE", fips_mode) + + with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): + if fips_mode == "true": + with pytest.raises(ValueError, match="digital envelope routines"): + async with proxy_startup_event(FastAPI()): + pass + else: + async with proxy_startup_event(FastAPI()): + await asyncio.sleep(0) + + failures = [r.getMessage() for r in caplog.records if "Password migration failed" in r.getMessage()] + assert len(failures) == 1 + assert "plaintext passwords stay unhashed" in failures[0] + assert "digital envelope routines" in failures[0] + + class _DatabaseWithOneStoredCredential: def __init__(self, ciphertext): self._ciphertext = ciphertext From 5c237840f4055a5a0058d09598ec51a66fae8e27 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 23 Sep 2026 07:31:01 +0000 Subject: [PATCH 02/18] refactor(proxy): match ssl_verify off detection to runtime str_to_bool semantics Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_utils/fips.py | 4 +++- tests/integration/_support/process.py | 15 ++++++++++----- .../test_litellm/proxy/common_utils/test_fips.py | 6 +++--- 3 files changed, 16 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/common_utils/fips.py b/litellm/proxy/common_utils/fips.py index 33f35388567..8a3c26b3220 100644 --- a/litellm/proxy/common_utils/fips.py +++ b/litellm/proxy/common_utils/fips.py @@ -6,6 +6,8 @@ from typing import Final from typing_extensions import assert_never +from litellm.secret_managers.main import str_to_bool + FIPS_MODE_ENV_VAR: Final = "LITELLM_FIPS_MODE" SSL_VERIFY_ENV_VAR: Final = "SSL_VERIFY" SSL_VERIFY_SETTING: Final = "litellm_settings.ssl_verify" @@ -146,6 +148,6 @@ def _is_off(value: object) -> bool: case bool(): return value is False case str(): - return value.strip().lower() in _FALSE_VALUES - {""} + return str_to_bool(value) is False case _: return False diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 6dec960097b..f65ab0e8093 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -140,6 +140,13 @@ def launched_proxy( assert root_stopped and not remaining, "Owned proxy required forced cleanup" +def _is_ready(client: httpx.Client) -> bool: + try: + return client.get("/health/readiness", timeout=2).status_code == 200 + except httpx.TransportError: + return False + + def refused_boot_log( gateway: Gateway, directory: Path, @@ -152,11 +159,9 @@ def refused_boot_log( with httpx.Client(base_url=f"http://127.0.0.1:{launched.port}", timeout=15, trust_env=False) as client: deadline: Final = time.monotonic() + 70 while launched.process.poll() is None: - try: - ready: Final = client.get("/health/readiness", timeout=2).status_code == 200 - except httpx.TransportError: - ready = False - assert not ready, f"Proxy became ready instead of refusing to boot:\n{launched.log.read_text()}" + assert not _is_ready(client), ( + f"Proxy became ready instead of refusing to boot:\n{launched.log.read_text()}" + ) assert time.monotonic() < deadline, "Proxy neither exited nor became ready within the deadline" time.sleep(0.1) assert launched.process.returncode != 0, ( diff --git a/tests/test_litellm/proxy/common_utils/test_fips.py b/tests/test_litellm/proxy/common_utils/test_fips.py index e934cafbdda..4b50ae4ad20 100644 --- a/tests/test_litellm/proxy/common_utils/test_fips.py +++ b/tests/test_litellm/proxy/common_utils/test_fips.py @@ -85,10 +85,10 @@ def test_on_with_a_non_enforcing_provider_is_refused_and_names_the_fix(): "ssl_env, ssl_setting, sources", [ ("false", True, ("SSL_VERIFY",)), - ("0", True, ("SSL_VERIFY",)), + (" FALSE ", True, ("SSL_VERIFY",)), (None, False, ("litellm_settings.ssl_verify",)), (None, "False", ("litellm_settings.ssl_verify",)), - ("no", False, ("SSL_VERIFY", "litellm_settings.ssl_verify")), + ("false", False, ("SSL_VERIFY", "litellm_settings.ssl_verify")), ], ) def test_disabled_tls_verification_is_refused_naming_every_source(ssl_env, ssl_setting, sources): @@ -99,7 +99,7 @@ def test_disabled_tls_verification_is_refused_naming_every_source(ssl_env, ssl_s assert "TLS certificate verification is disabled by " + " and ".join(sources) in str(refused.value) -@pytest.mark.parametrize("ssl_setting", [True, "true", "/etc/ssl/certs/ca.pem", None, ""]) +@pytest.mark.parametrize("ssl_setting", [True, "true", "/etc/ssl/certs/ca.pem", None, "", "0", "no"]) def test_verified_or_custom_bundle_tls_settings_are_not_treated_as_disabled(ssl_setting): assert _verdict("true", enforcing=True, ssl_setting=ssl_setting) == FipsModeOn() From 6ae5df66c1d24ba5ac460c2338ea01585753b04a Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 23 Sep 2026 07:45:39 +0000 Subject: [PATCH 03/18] refactor(proxy): drop tautological fips probe test and satisfy CodeQL return checks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_utils/fips.py | 23 ++++++++----------- .../proxy/common_utils/test_fips.py | 12 ---------- 2 files changed, 10 insertions(+), 25 deletions(-) diff --git a/litellm/proxy/common_utils/fips.py b/litellm/proxy/common_utils/fips.py index 8a3c26b3220..07bd4d2462c 100644 --- a/litellm/proxy/common_utils/fips.py +++ b/litellm/proxy/common_utils/fips.py @@ -121,9 +121,9 @@ def enforce_fips_boot_verdict(verdict: FipsBootVerdict, announce: Callable[[str] def render_refusal(refusal: FipsBootRefusal) -> str: match refusal: - case MalformedFipsMode(value=value): + case MalformedFipsMode(): return ( - f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR}={value} is not a boolean.\n" + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR}={refusal.value} is not a boolean.\n" f"Set {FIPS_MODE_ENV_VAR} to true or false, or unset it." ) case ProviderDoesNotEnforceFips(): @@ -133,21 +133,18 @@ def render_refusal(refusal: FipsBootRefusal) -> str: "protected with algorithms the FIPS 140-3 policy forbids. Run the proxy from a FIPS image whose\n" f"OpenSSL FIPS provider is enabled, or unset {FIPS_MODE_ENV_VAR} on a non-FIPS runtime." ) - case TlsVerificationDisabled(sources=sources): + case TlsVerificationDisabled(): return ( f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but TLS certificate verification is disabled by " - f"{' and '.join(sources)}.\nFIPS deployments must verify upstream certificates, so remove the " + f"{' and '.join(refusal.sources)}.\nFIPS deployments must verify upstream certificates, so remove the " "override or point ssl_verify at a CA bundle instead." ) - case _: - assert_never(refusal) + return assert_never(refusal) def _is_off(value: object) -> bool: - match value: - case bool(): - return value is False - case str(): - return str_to_bool(value) is False - case _: - return False + if isinstance(value, bool): + return value is False + if isinstance(value, str): + return str_to_bool(value) is False + return False diff --git a/tests/test_litellm/proxy/common_utils/test_fips.py b/tests/test_litellm/proxy/common_utils/test_fips.py index 4b50ae4ad20..0e8138856c0 100644 --- a/tests/test_litellm/proxy/common_utils/test_fips.py +++ b/tests/test_litellm/proxy/common_utils/test_fips.py @@ -1,5 +1,3 @@ -import hashlib - import pytest from litellm.proxy.common_utils.fips import ( @@ -12,7 +10,6 @@ from litellm.proxy.common_utils.fips import ( enforce_fips_boot_verdict, fips_boot_verdict, is_fips_mode, - openssl_enforces_fips, parse_fips_mode, ) @@ -107,12 +104,3 @@ def test_verified_or_custom_bundle_tls_settings_are_not_treated_as_disabled(ssl_ def test_disabled_tls_is_reported_before_the_provider_so_operators_see_config_mistakes_first(): assert _verdict("true", enforcing=False, ssl_env="false") == TlsVerificationDisabled(sources=("SSL_VERIFY",)) assert _verdict("true", enforcing=False) == ProviderDoesNotEnforceFips() - - -def test_provider_probe_agrees_with_whether_md5_is_usable_for_security_here(): - try: - hashlib.md5(b"", usedforsecurity=True) - except ValueError: - assert openssl_enforces_fips() is True - else: - assert openssl_enforces_fips() is False From e6a5d32cd9925caedbd1eb067b86772031dcf556 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 06:07:49 +0000 Subject: [PATCH 04/18] test(proxy): inject the fake Prisma client through the module boundary instead of patching _setup_prisma_client Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/test_proxy_server.py | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 42ad00d6b13..39e4a417d19 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1733,7 +1733,7 @@ class _PrismaClientWhoseUserTableCannotHash: class _Db: litellm_usertable = None - def __init__(self): + def __init__(self, database_url, proxy_logging_obj): self.db = self._Db() self.db.litellm_usertable = self._Table() self.writer_db = self.db @@ -1744,6 +1744,9 @@ class _PrismaClientWhoseUserTableCannotHash: async def disconnect(self): pass + def start_view_setup_task(self): + pass + async def check_view_exists(self): pass @@ -1761,12 +1764,7 @@ async def test_proxy_startup_surfaces_a_password_migration_crypto_failure(monkey _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) monkeypatch.setenv("DATABASE_URL", "postgresql://nobody:nothing@127.0.0.1:1/unreachable") - cannot_hash = _PrismaClientWhoseUserTableCannotHash() - - async def connected(**kwargs): - return cannot_hash - - monkeypatch.setattr("litellm.proxy.proxy_server.ProxyStartupEvent._setup_prisma_client", connected) + monkeypatch.setattr("litellm.proxy.proxy_server.PrismaClient", _PrismaClientWhoseUserTableCannotHash) monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) monkeypatch.setenv("LITELLM_FIPS_MODE", fips_mode) From f05002acb2187d755f9a7f39732a561b06dd2a8b Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 07:53:12 +0000 Subject: [PATCH 05/18] test(auth): integration cells for JWT algorithm allowlists (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_jwt_algorithm_allowlist.py | 99 +++++++++++++ .../test_mcp_client_assertion_signing_alg.py | 87 ++++++++++++ .../mcp/test_mcp_jwt_signer_jwks_allowlist.py | 133 ++++++++++++++++++ 3 files changed, 319 insertions(+) create mode 100644 tests/integration/authorization/test_jwt_algorithm_allowlist.py create mode 100644 tests/integration/mcp/test_mcp_client_assertion_signing_alg.py create mode 100644 tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py diff --git a/tests/integration/authorization/test_jwt_algorithm_allowlist.py b/tests/integration/authorization/test_jwt_algorithm_allowlist.py new file mode 100644 index 00000000000..210d40a0d6a --- /dev/null +++ b/tests/integration/authorization/test_jwt_algorithm_allowlist.py @@ -0,0 +1,99 @@ +import json +import time +import uuid +from pathlib import Path +from typing import Final + +import jwt +import yaml +from cryptography.hazmat.primitives.asymmetric import ed25519, rsa + +from tests.integration._support.client import Gateway, eventually +from tests.integration._support.process import owned_proxy_process +from tests.integration._support.wire import Reply, Request, wire_server + +KEY_ID: Final = "integration-jwt-allowlist-key" + + +def _jwks_reply(public_jwk: str) -> Reply: + return Reply(body=json.dumps({"keys": [{**json.loads(public_jwk), "kid": KEY_ID}]}).encode()) + + +def _jwt_auth_config(tmp_path: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"] = { + **config["general_settings"], + "enable_jwt_auth": True, + "litellm_jwtauth": {"user_id_jwt_field": "sub", "user_id_upsert": True}, + } + path: Final = tmp_path / "jwt_allowlist.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _token(key: object, algorithm: str, subject: str) -> str: + return jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + key, # pyright: ignore[reportArgumentType] # jwt.encode takes Any key material + algorithm=algorithm, + headers={"kid": KEY_ID}, + ) + + +def test_eddsa_signed_token_is_accepted_and_logged_as_deprecated_outside_fips_mode( + gateway: Gateway, tmp_path: Path +) -> None: + private_key: Final = ed25519.Ed25519PrivateKey.generate() + public_jwk: Final = jwt.algorithms.OKPAlgorithm.to_jwk(private_key.public_key()) + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply(public_jwk) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "EdDSA", subject) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path) + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist control"}]}, + key=token, + ) + assert response.status_code == 200, response.text + deprecation: Final = eventually( + lambda: owned.log.read_text(), + lambda text: "EdDSA" in text and "deprecated" in text and "LITELLM_FIPS_MODE" in text, + seconds=30, + ) + assert "EdDSA" in deprecation and "deprecated" in deprecation and "LITELLM_FIPS_MODE" in deprecation + + +def test_rs256_signed_token_is_accepted_without_deprecation_log(gateway: Gateway, tmp_path: Path) -> None: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key()) + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply(public_jwk) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "RS256", subject) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path) + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist control"}]}, + key=token, + ) + assert response.status_code == 200, response.text + assert "EdDSA" not in owned.log.read_text(), owned.log.read_text() diff --git a/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py b/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py new file mode 100644 index 00000000000..883c335b71d --- /dev/null +++ b/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py @@ -0,0 +1,87 @@ +import json +import uuid +from typing import Final + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway +from integration._support.database import read_rows +from integration._support.mcp import forget_mcp, mcp_peer, register_mcp + + +def _pem_private_key() -> str: + return ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() + ) + + +def _stored_client_assertion_signing_alg(identity: str) -> str: + rows: Final = read_rows('SELECT credentials FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,)) + assert len(rows) == 1, rows + credentials: Final = rows[0]["credentials"] + if credentials is None: + return "RS256" + blob: Final = credentials if isinstance(credentials, dict) else json.loads(credentials) + return str(blob.get("client_assertion_signing_alg") or "RS256") + + +def test_non_approved_client_assertion_signing_alg_is_rejected_on_create_and_update(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "alg" + uuid.uuid4().hex[:8] + created: Final = gateway.request( + "POST", + "/v1/mcp/server", + { + "server_name": alias, + "alias": alias, + **peer.registration(), + "token_exchange_endpoint": "https://idp.integration.invalid/oauth2/token", + "credentials": { + "client_private_key": _pem_private_key(), + "client_assertion_signing_alg": "HS256", + }, + }, + ) + if created.status_code == 201: + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + assert created.status_code in (400, 422), created.text + assert "client_assertion_signing_alg" in created.text, created.text + + identity: Final = register_mcp(scenario, peer, alias + "v") + edited: Final = gateway.request( + "PUT", + "/v1/mcp/server", + {"server_id": identity, "credentials": {"client_assertion_signing_alg": "EdDSA"}}, + ) + assert edited.status_code in (400, 422), edited.text + assert "client_assertion_signing_alg" in edited.text, edited.text + assert _stored_client_assertion_signing_alg(identity) == "RS256" + + +def test_approved_client_assertion_signing_alg_round_trips(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "algok" + uuid.uuid4().hex[:8] + created: Final = gateway.request( + "POST", + "/v1/mcp/server", + { + "server_name": alias, + "alias": alias, + **peer.registration(), + "token_exchange_endpoint": "https://idp.integration.invalid/oauth2/token", + "credentials": { + "client_private_key": _pem_private_key(), + "client_assertion_signing_alg": "ES256", + }, + }, + ) + assert created.status_code == 201, created.text + identity: Final = str(created.json()["server_id"]) + scenario.cleanups.callback(forget_mcp, gateway, identity) + assert _stored_client_assertion_signing_alg(identity) == "ES256" diff --git a/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py b/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py new file mode 100644 index 00000000000..8af03fe0663 --- /dev/null +++ b/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py @@ -0,0 +1,133 @@ +import base64 +import json +import time +import uuid +from collections.abc import Callable, Iterator, Mapping +from contextlib import contextmanager +from typing import Final + +import httpx +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import ed25519, rsa +from integration._support.client import Gateway +from integration._support.mcp import mcp_peer, register_mcp, tool_calls, tool_names +from integration._support.wire import Reply, Request, wire_server + + +@contextmanager +def _jwt_signer_guardrail(gateway: Gateway, discovery_uri: str) -> Iterator[None]: + created: Final = gateway.client.post( + "/guardrails", + headers={"x-litellm-api-key": gateway.key}, + json={ + "guardrail": { + "guardrail_name": "signer" + uuid.uuid4().hex[:8], + "litellm_params": { + "guardrail": "mcp_jwt_signer", + "mode": "pre_mcp_call", + "default_on": True, + "access_token_discovery_uri": discovery_uri, + }, + } + }, + ) + assert created.status_code == 200, created.text + identity: Final = str(created.json()["guardrail_id"]) + try: + yield + finally: + deleted: Final = gateway.client.delete(f"/guardrails/{identity}", headers={"x-litellm-api-key": gateway.key}) + assert deleted.status_code == 200, deleted.text + + +@contextmanager +def _idp_server(jwks_keys: list[Mapping[str, object]]) -> Iterator[str]: + holder: Final = {"url": ""} + + def respond(request: Request) -> Reply: + if request.target == "/.well-known/openid-configuration": + return Reply(body=json.dumps({"jwks_uri": holder["url"] + "/.well-known/jwks.json"}).encode()) + if request.target == "/.well-known/jwks.json": + return Reply(body=json.dumps({"keys": list(jwks_keys)}).encode()) + return Reply(status=404) + + with wire_server(respond) as idp: + holder["url"] = idp.url + yield idp.url + "/.well-known/openid-configuration" + + +def _claims() -> dict[str, object]: + now: Final = int(time.time()) + return {"sub": "integration-mcp-user", "iat": now, "exp": now + 300} + + +def _hs256_key_and_token() -> tuple[dict[str, object], str]: + secret: Final = b"integration-hs256-client-secret-0123456789abcdef" + key: Final = { + "kty": "oct", + "alg": "HS256", + "kid": "sym", + "k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(), + } + token: Final = jwt.encode(_claims(), secret, algorithm="HS256", headers={"kid": "sym"}) + return key, token + + +def _eddsa_key_and_token() -> tuple[dict[str, object], str]: + private_key: Final = ed25519.Ed25519PrivateKey.generate() + public_jwk: Final = json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(private_key.public_key())) + key: Final = {**public_jwk, "alg": "EdDSA", "kid": "ed"} + token: Final = jwt.encode(_claims(), private_key, algorithm="EdDSA", headers={"kid": "ed"}) + return key, token + + +def _rs256_key_and_token() -> tuple[dict[str, object], str]: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = json.loads(jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key())) + key: Final = {**public_jwk, "alg": "RS256", "kid": "rsa"} + token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"}) + return key, token + + +def _call_with_bearer( + gateway: Gateway, key: str, identity: str, name: str, arguments: dict[str, object], bearer: str +) -> httpx.Response: + return gateway.client.post( + "/mcp-rest/tools/call", + headers={"x-litellm-api-key": key, "Authorization": f"Bearer {bearer}"}, + json={"server_id": identity, "name": name, "arguments": arguments}, + ) + + +@pytest.mark.parametrize("key_and_token", (_hs256_key_and_token, _eddsa_key_and_token), ids=("oct-HS256", "OKP-EdDSA")) +def test_jwks_key_with_non_approved_alg_cannot_verify_the_incoming_token( + gateway: Gateway, key_and_token: Callable[[], tuple[dict[str, object], str]] +) -> None: + jwks_key, token = key_and_token() + with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri): + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "jwksalg" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token) + assert response.status_code == 401, response.text + assert "incoming token verification failed" in response.text, response.text + assert tool_calls(peer.drain()) == (), "rejected call reached the peer" + + +def test_jwks_rs256_key_still_verifies_the_incoming_token(gateway: Gateway) -> None: + jwks_key, token = _rs256_key_and_token() + with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri): + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "jwksok" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "9", response.text + assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer" From 6dff6ae5bcd3c5d8d264bb63ebfcb2653a7c9c6d Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 08:09:55 +0000 Subject: [PATCH 06/18] fix(auth): FIPS-aware JWT algorithm allowlists for proxy auth and MCP (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 28 +- .../mcp_server/oauth_identity_binding.py | 15 +- .../mcp_server/outbound_credentials/types.py | 3 +- litellm/proxy/_lazy_openapi_snapshot.json | 44 + litellm/proxy/auth/handle_jwt.py | 83 +- litellm/proxy/auth/jwt_algorithms.py | 48 + .../mcp_jwt_signer/mcp_jwt_signer.py | 7 +- litellm/types/mcp.py | 3 +- .../types/mcp_server/mcp_server_manager.py | 3 +- .../mcp_server/test_mcp_server_manager.py | 680 +++++++++---- .../mcp_server/test_oauth_identity_binding.py | 32 + .../proxy/auth/test_handle_jwt.py | 929 +++++++----------- .../proxy/guardrails/test_mcp_jwt_signer.py | 151 ++- tests/test_litellm/types/test_mcp.py | 34 +- 14 files changed, 1241 insertions(+), 819 deletions(-) create mode 100644 litellm/proxy/auth/jwt_algorithms.py diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 0c520142fb3..8e182755656 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -169,6 +169,7 @@ from litellm.proxy._types import ( is_per_server_oauth_discovery_eligible, ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils +from litellm.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import ( @@ -386,7 +387,7 @@ class MCPServerConfig(TypedDict, total=False): id_jag_resource: str client_private_key: str client_private_key_id: str - client_assertion_signing_alg: str + client_assertion_signing_alg: ApprovedJwtAlgorithm timeout: float max_concurrent_requests: int @@ -1482,6 +1483,18 @@ def _warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None: ) +def _stored_client_assertion_signing_alg(value: object, server_name: str) -> ApprovedJwtAlgorithm: + if value in APPROVED_JWT_ALGORITHMS: + return cast(ApprovedJwtAlgorithm, value) + if value is not None: + verbose_logger.warning( + "MCP server %s: client_assertion_signing_alg %r is not an approved algorithm, using RS256", + server_name, + value, + ) + return "RS256" + + def _warn_legacy_delegate_auth_if_applicable(server: MCPServer, *, source: str) -> None: """Direct legacy delegated OAuth configurations to the admitted replacement.""" if server.auth_type != MCPAuth.oauth2: @@ -2604,7 +2617,10 @@ class MCPServerManager: id_jag_resource=server_config.get("id_jag_resource", None), client_private_key=server_config.get("client_private_key", None), client_private_key_id=server_config.get("client_private_key_id", None), - client_assertion_signing_alg=server_config.get("client_assertion_signing_alg", "RS256"), + client_assertion_signing_alg=_stored_client_assertion_signing_alg( + server_config.get("client_assertion_signing_alg"), + str(server_config.get("alias") or server_config.get("server_name") or server_id), + ), token_exchange_profile=server_config.get("token_exchange_profile", "rfc8693"), allow_sampling=bool(server_config.get("allow_sampling", False)), allow_elicitation=bool(server_config.get("allow_elicitation", False)), @@ -3178,10 +3194,10 @@ class MCPServerManager: credentials_are_encrypted, ), client_private_key_id=(credentials_dict.get("client_private_key_id") if credentials_dict else None), - client_assertion_signing_alg=( - credentials_dict.get("client_assertion_signing_alg") if credentials_dict else None - ) - or "RS256", + client_assertion_signing_alg=_stored_client_assertion_signing_alg( + credentials_dict.get("client_assertion_signing_alg") if credentials_dict else None, + name_for_prefix, + ), token_exchange_profile=mcp_server.token_exchange_profile or (credentials_dict.get("token_exchange_profile") if credentials_dict else None) or "rfc8693", diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index f02e6c85d9b..1f2a492beab 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -23,20 +23,11 @@ from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, jwks_keys_for from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer -_ALLOWED_ID_TOKEN_ALGORITHMS: Final = ( - "RS256", - "RS384", - "RS512", - "ES256", - "ES384", - "ES512", - "PS256", - "PS384", - "PS512", -) +_ALLOWED_ID_TOKEN_ALGORITHMS: Final = APPROVED_JWT_ALGORITHMS _JWKS_CACHE_TTL_SECONDS: Final = 3600 _jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS) @@ -126,7 +117,7 @@ async def _discover_jwks_url(issuer: str) -> str: def _select_signing_key(id_token: str, keys: Sequence[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection": header: Final = jwt.get_unverified_header(id_token) kid: Final = header.get("kid") - for key in keys: + for key in jwks_keys_for(keys, _ALLOWED_ID_TOKEN_ALGORITHMS): if kid is None or key.get("kid") == kid: return jwt.PyJWK(dict(key)) # mutable-ok: PyJWT requires a concrete JWK dictionary return _BindingRejection( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py index 33c3a854058..6c310784370 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py @@ -40,6 +40,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Ok, Result, ) +from litellm.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm from litellm.types.mcp import ( DEFAULT_CREDENTIAL_HEADER, DEFAULT_SUBJECT_TOKEN_TYPE, @@ -326,7 +327,7 @@ class PrivateKeyJwtAuth(BaseModel): source: Literal["private_key_jwt"] = "private_key_jwt" private_key: SecretStr key_id: str | None = None - signing_alg: str = "RS256" + signing_alg: ApprovedJwtAlgorithm = "RS256" class ClientSecretAuth(BaseModel): diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 0b43c3864ab..858a993c4ec 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -29465,6 +29465,17 @@ "client_assertion_signing_alg": { "anyOf": [ { + "enum": [ + "RS256", + "RS384", + "RS512", + "PS256", + "PS384", + "PS512", + "ES256", + "ES384", + "ES512" + ], "type": "string" }, { @@ -32166,6 +32177,17 @@ "client_assertion_signing_alg": { "anyOf": [ { + "enum": [ + "RS256", + "RS384", + "RS512", + "PS256", + "PS384", + "PS512", + "ES256", + "ES384", + "ES512" + ], "type": "string" }, { @@ -35439,6 +35461,17 @@ "client_assertion_signing_alg": { "anyOf": [ { + "enum": [ + "RS256", + "RS384", + "RS512", + "PS256", + "PS384", + "PS512", + "ES256", + "ES384", + "ES512" + ], "type": "string" }, { @@ -39289,6 +39322,17 @@ "client_assertion_signing_alg": { "anyOf": [ { + "enum": [ + "RS256", + "RS384", + "RS512", + "PS256", + "PS384", + "PS512", + "ES256", + "ES384", + "ES512" + ], "type": "string" }, { diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index e07b20fd5d5..6f4a10c6f69 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -16,6 +16,7 @@ import re import time from collections.abc import Awaitable, Callable, Collection, Mapping, Sequence from dataclasses import dataclass +from functools import lru_cache from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast import httpx @@ -53,6 +54,12 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import can_team_access_model +from litellm.proxy.auth.jwt_algorithms import ( + APPROVED_JWT_ALGORITHMS, + LEGACY_JWT_ALGORITHMS, + allowed_jwt_algorithms, + jwks_keys_for, +) from litellm.proxy.auth.model_access_denied import ( ModelAccessDeniedHTTPException, model_access_denied_client_message, @@ -60,6 +67,7 @@ from litellm.proxy.auth.model_access_denied import ( from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canonical_user_id from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.team_grants import team_grants, team_model_aliases +from litellm.proxy.common_utils.fips import is_fips_mode from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, get_management_object_ttl, @@ -104,6 +112,14 @@ UNREACHABLE_CACHE_KEY_PREFIX: Final = "litellm_jwks_unreachable_" _CachedValueT = TypeVar("_CachedValueT", bound=JWKKeyValue | str) +@lru_cache(maxsize=1) +def _log_eddsa_deprecation() -> None: + verbose_proxy_logger.warning( + "JWT Auth: accepted a token signed with EdDSA, which is deprecated and not FIPS 140-3 approved; " + "it is rejected when LITELLM_FIPS_MODE=true. Move the IdP signing key to RS256, PS256 or ES256" + ) + + class _JWTAuthSettings(Protocol): """The JWT auth settings block this handler reads back through ``getattr``, when one is configured.""" @@ -209,18 +225,7 @@ class JWTHandler: # Supported algos: https://pyjwt.readthedocs.io/en/stable/algorithms.html # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret # the key in different ways (e.g. HS* and RS*)." - SUPPORTED_JWT_ALGORITHMS = [ - "RS256", - "RS384", - "RS512", - "PS256", - "PS384", - "PS512", - "ES256", - "ES384", - "ES512", - "EdDSA", - ] + SUPPORTED_JWT_ALGORITHMS = [*APPROVED_JWT_ALGORITHMS, *LEGACY_JWT_ALGORITHMS] LITELLM_JWT_ISSUER_CLAIM = "_litellm_jwt_issuer" LITELLM_USER_ID_CLAIM = "_litellm_user_id" LITELLM_USER_EMAIL_CLAIM = "_litellm_user_email" @@ -240,13 +245,24 @@ class JWTHandler: def __init__( self, + fips_mode: Callable[[], bool] = is_fips_mode, ) -> None: + self._fips_mode: Final = fips_mode self.http_handler = HTTPHandler() self.leeway = 0 # Per-cache-key locks so a TTL lapse triggers one refresh instead of one per in-flight request. self._refresh_locks: dict[str, asyncio.Lock] = {} # mutable-ok: lock registry, keyed by JWKS url self.agent_lookup: AgentLookup = _NoRegisteredAgents() + def allowed_algorithms(self) -> tuple[str, ...]: + return allowed_jwt_algorithms(self._fips_mode()) + + def _warn_deprecated_signing_algorithm(self, token: str) -> None: + if self._fips_mode(): + return + if jwt.get_unverified_header(token).get("alg") == "EdDSA": + _log_eddsa_deprecation() + def bind_agent_lookup(self, agent_lookup: AgentLookup) -> None: self.agent_lookup = agent_lookup @@ -943,7 +959,12 @@ class JWTHandler: log_context=f"kid={kid}", ) - public_key: Final = self.parse_keys(keys=keys, kid=kid) + usable_keys: Final[JWKKeyValue] = ( + list(jwks_keys_for(keys, self.allowed_algorithms())) + if isinstance(keys, list) + else next(iter(jwks_keys_for((keys,), self.allowed_algorithms())), {}) + ) + public_key: Final = self.parse_keys(keys=usable_keys, kid=kid) if public_key is not None: return cast(dict, public_key) @@ -1216,30 +1237,32 @@ class JWTHandler: if isinstance(public_key, dict): public_key_obj: Final = PyJWK.from_dict(self._get_jwk_from_public_key(public_key=public_key)).key - return jwt.decode( + payload: Final = jwt.decode( token, public_key_obj, - algorithms=self.SUPPORTED_JWT_ALGORITHMS, + algorithms=self.allowed_algorithms(), options=decode_options, audience=audience, issuer=issuer, leeway=self.leeway, ) - - cert: Final = x509.load_pem_x509_certificate(public_key.encode(), default_backend()) - key: Final = cert.public_key().public_bytes( - serialization.Encoding.PEM, - serialization.PublicFormat.SubjectPublicKeyInfo, - ) - return jwt.decode( - token, - key, - algorithms=self.SUPPORTED_JWT_ALGORITHMS, - audience=audience, - issuer=issuer, - options=decode_options, - leeway=self.leeway, - ) + else: + cert: Final = x509.load_pem_x509_certificate(public_key.encode(), default_backend()) + key: Final = cert.public_key().public_bytes( + serialization.Encoding.PEM, + serialization.PublicFormat.SubjectPublicKeyInfo, + ) + payload = jwt.decode( + token, + key, + algorithms=self.allowed_algorithms(), + audience=audience, + issuer=issuer, + options=decode_options, + leeway=self.leeway, + ) + self._warn_deprecated_signing_algorithm(token) + return payload async def _auth_jwt_with_issuer(self, token: str, issuer_config: JWTIssuerConfig, kid: str | None) -> dict: try: diff --git a/litellm/proxy/auth/jwt_algorithms.py b/litellm/proxy/auth/jwt_algorithms.py new file mode 100644 index 00000000000..48e03ba9981 --- /dev/null +++ b/litellm/proxy/auth/jwt_algorithms.py @@ -0,0 +1,48 @@ +from collections.abc import Collection, Mapping, Sequence +from types import MappingProxyType +from typing import Final, Literal + +ApprovedJwtAlgorithm = Literal["RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512"] + +APPROVED_JWT_ALGORITHMS: Final[tuple[ApprovedJwtAlgorithm, ...]] = ( + "RS256", + "RS384", + "RS512", + "PS256", + "PS384", + "PS512", + "ES256", + "ES384", + "ES512", +) + +LEGACY_JWT_ALGORITHMS: Final = ("EdDSA",) + +_KEY_TYPE_ALGORITHMS: Final = MappingProxyType( + { + "RSA": frozenset(APPROVED_JWT_ALGORITHMS[:6]), + "EC": frozenset(APPROVED_JWT_ALGORITHMS[6:]), + "OKP": frozenset(LEGACY_JWT_ALGORITHMS), + } +) + + +def allowed_jwt_algorithms(fips_mode: bool) -> tuple[str, ...]: + return APPROVED_JWT_ALGORITHMS if fips_mode else (*APPROVED_JWT_ALGORITHMS, *LEGACY_JWT_ALGORITHMS) + + +def jwks_keys_for( + keys: Sequence[Mapping[str, object]], algorithms: Collection[str] +) -> tuple[Mapping[str, object], ...]: + """Keep keys whose declared alg is allowed; keys without alg are kept only when their kty can sign with an allowed algorithm.""" + return tuple(key for key in keys if _key_allowed(key, frozenset(algorithms))) + + +def _key_allowed(key: Mapping[str, object], algorithms: frozenset[str]) -> bool: + alg: Final = key.get("alg") + if isinstance(alg, str): + return alg in algorithms + key_type: Final = key.get("kty") + if not isinstance(key_type, str): + return False + return bool(_KEY_TYPE_ALGORITHMS.get(key_type, frozenset()) & algorithms) diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py index a269ad31a6b..c1a88f3aa67 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py @@ -89,6 +89,7 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, jwks_keys_for from litellm.types.guardrail_base_init import GuardrailBaseInitKwargs from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import CallTypesLiteral @@ -449,12 +450,16 @@ class MCPJWTSigner(CustomGuardrail): unverified_header: Final = jwt.get_unverified_header(raw_token) kid: Final = unverified_header.get("kid") + approved_keys: Final = jwks_keys_for(jwks_keys, APPROVED_JWT_ALGORITHMS) + if not approved_keys: + raise jwt.exceptions.PyJWKSetError(f"No JWKS key at {jwks_uri!r} uses an approved signing algorithm") + # Build a JWKS object and pick the matching key. # PyJWT's PyJWKSet handles key-type parsing and kid matching correctly. from jwt import PyJWKSet try: - jwks_set: Final = PyJWKSet.from_dict({"keys": jwks_keys}) + jwks_set: Final = PyJWKSet.from_dict({"keys": list(approved_keys)}) except Exception as exc: raise jwt.exceptions.PyJWKSetError(f"Failed to parse JWKS from {jwks_uri!r}: {exc}") from exc diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 83e719810d5..df2d2b979cc 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -11,6 +11,7 @@ import httpx from pydantic import BaseModel, ConfigDict, Field from typing_extensions import TypedDict +from litellm.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm from litellm.types.llms.base import HiddenParams if TYPE_CHECKING: @@ -251,7 +252,7 @@ class MCPCredentials(TypedDict, total=False): Key id (kid) advertised in the client_assertion JWT header """ - client_assertion_signing_alg: str | None + client_assertion_signing_alg: ApprovedJwtAlgorithm | None """ Signing algorithm for the client_assertion JWT. Default: RS256 """ diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index cb32299b143..dc610724485 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -4,6 +4,7 @@ from typing import Any, Final, Literal from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator from typing_extensions import Self +from litellm.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm from litellm.types.mcp import ( DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth, @@ -145,7 +146,7 @@ class MCPServer(BaseModel): id_jag_resource: str | None = None client_private_key: str | None = None client_private_key_id: str | None = None - client_assertion_signing_alg: str = "RS256" + client_assertion_signing_alg: ApprovedJwtAlgorithm = "RS256" # Wire dialect: "rfc8693" (standard token-exchange grant) or "entra_obo" (Microsoft Entra # On-Behalf-Of, the RFC 7523 jwt-bearer grant + requested_token_use extension) token_exchange_profile: str = "rfc8693" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cd5dae1269a..23aa8e4f4cd 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -111,7 +111,6 @@ async def test_manager_sampling_preserves_explicit_headers_without_ambient_conte assert sampling.await_args.kwargs["raw_headers"] == {"x-test-caller": "sampling-caller"} - @pytest.mark.asyncio async def test_sampling_callback_keeps_creation_context_after_caller_switch(): from mcp.server.auth.middleware.auth_context import auth_context_var @@ -219,8 +218,6 @@ def _reload_mcp_manager_module(): return reloaded - - @pytest.fixture(autouse=True) def enable_eager_mcp_oauth_discovery(monkeypatch): monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1") @@ -1394,7 +1391,9 @@ class TestMCPServerManager: assert not any("oauth2_id_jag" in message for message in caplog.messages) @pytest.mark.asyncio - async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso(self, config_only_mcp_manager_factory, monkeypatch, caplog): + async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso( + self, config_only_mcp_manager_factory, monkeypatch, caplog + ): self._clear_sso_env(monkeypatch) monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid") manager = config_only_mcp_manager_factory() @@ -4680,7 +4679,9 @@ class TestMCPServerManager: @pytest.mark.parametrize("auth_type", [MCPAuth.none, MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.oauth2]) @pytest.mark.parametrize("is_byok", [False, True]) @pytest.mark.parametrize("scheme", ["http", "https"]) - async def test_openapi_health_loads_spec_without_mcp_handshake(self, respx_mock, monkeypatch, auth_type, is_byok, scheme): + async def test_openapi_health_loads_spec_without_mcp_handshake( + self, respx_mock, monkeypatch, auth_type, is_byok, scheme + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4730,14 +4731,28 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("failure", "expected_status", "expected_error"), [ - (httpx.Response(401, text="secret response content"), "unhealthy", "OpenAPI specification request failed (HTTP 401)"), + ( + httpx.Response(401, text="secret response content"), + "unhealthy", + "OpenAPI specification request failed (HTTP 401)", + ), (httpx.Response(404), "unhealthy", "OpenAPI specification request failed (HTTP 404)"), (httpx.Response(500), "unhealthy", "OpenAPI specification request failed (HTTP 500)"), - (httpx.ConnectError("secret network details"), "unhealthy", "OpenAPI specification could not be loaded (ConnectError)"), - (httpx.Response(200, text="secret invalid JSON body"), "unhealthy", "OpenAPI specification could not be loaded (JSONDecodeError)"), + ( + httpx.ConnectError("secret network details"), + "unhealthy", + "OpenAPI specification could not be loaded (ConnectError)", + ), + ( + httpx.Response(200, text="secret invalid JSON body"), + "unhealthy", + "OpenAPI specification could not be loaded (JSONDecodeError)", + ), ], ) - async def test_openapi_health_reports_safe_failures(self, respx_mock, monkeypatch, failure, expected_status, expected_error): + async def test_openapi_health_reports_safe_failures( + self, respx_mock, monkeypatch, failure, expected_status, expected_error + ): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -5272,8 +5287,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers captured["server_label"] = server_label @@ -5358,8 +5380,15 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, - auth_type=None, upstream_token_header=None, + path, + method, + operation, + base_url, + headers=None, + server_label=None, + relays_upstream_auth=False, + auth_type=None, + upstream_token_header=None, ): captured["headers"] = headers @@ -12485,7 +12514,9 @@ class TestConfigServerIdPinning: @pytest.mark.asyncio @pytest.mark.parametrize("aliasing_entry_first", [True, False]) - async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected(self, config_only_mcp_manager_factory, aliasing_entry_first: bool): + async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected( + self, config_only_mcp_manager_factory, aliasing_entry_first: bool + ): """A grant naming 'docs_server' reaches both servers unpinned; the pin would narrow it to one.""" manager = config_only_mcp_manager_factory() wiki = ( @@ -12501,7 +12532,9 @@ class TestConfigServerIdPinning: await manager.load_servers_from_config(dict((wiki, docs) if aliasing_entry_first else (docs, wiki))) @pytest.mark.asyncio - async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected(self, config_only_mcp_manager_factory): + async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected( + self, config_only_mcp_manager_factory + ): manager = config_only_mcp_manager_factory() with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"): @@ -12589,7 +12622,9 @@ class TestConfigServerIdPinning: assert second_round == first_round @pytest.mark.asyncio - async def test_shadow_warning_fires_again_when_the_shadowed_set_changes(self, config_only_mcp_manager_factory, caplog): + async def test_shadow_warning_fires_again_when_the_shadowed_set_changes( + self, config_only_mcp_manager_factory, caplog + ): manager = config_only_mcp_manager_factory() await manager.load_servers_from_config(self._config(server_id="docs-prod-1")) @@ -12760,7 +12795,9 @@ class TestConfigServerIdPinning: assert manager.config_mcp_servers["wiki"].url == "https://example.com/mcp" @pytest.mark.asyncio - async def test_a_row_that_shadows_one_id_still_reports_capturing_another(self, config_only_mcp_manager_factory, caplog): + async def test_a_row_that_shadows_one_id_still_reports_capturing_another( + self, config_only_mcp_manager_factory, caplog + ): """Skipping is per identifier, not per row, so the second collision is not lost.""" manager = config_only_mcp_manager_factory() await manager.load_servers_from_config( @@ -13132,7 +13169,8 @@ async def test_pre_call_tool_check_honors_guardrail_attached_to_key(monkeypatch, ("none", {"Authorization": "Bearer injected"}, "extra-headers", "Bearer injected"), ], ) -async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_request_ctx, +async def test_debug_resolution_matches_final_header_conflict_winner( + _mcp_request_ctx, config: Literal["stored", "static", "none"], extra_headers: dict[str, str] | None, expected_source: str, @@ -13203,7 +13241,9 @@ async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_reques @pytest.mark.asyncio @pytest.mark.parametrize("transport", ["http", "stdio"]) -async def test_debug_reports_legacy_signing_and_non_http_transport(_mcp_request_ctx, transport: Literal["http", "stdio"]) -> None: +async def test_debug_reports_legacy_signing_and_non_http_transport( + _mcp_request_ctx, transport: Literal["http", "stdio"] +) -> None: from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from starlette.requests import Request @@ -13247,12 +13287,16 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(_mcp_request_ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + server_id="temporary-oauth-discovery", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, ) manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", ) with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: @@ -13272,13 +13316,18 @@ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publi async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", + server_id="repeated-stale", + name="stale", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + oauth2_flow="authorization_code", ) manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) with ( patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, @@ -13298,13 +13347,20 @@ async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + server_id="resolved-replacement", + name="replacement", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + ) + replacement: Final = original.model_copy( + update={ + "url": "https://new.example.com/mcp", + "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + } ) - replacement: Final = original.model_copy(update={ - "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", - "token_url": "https://new.example.com/token", - }) manager.registry[original.server_id] = replacement assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement @@ -13312,8 +13368,11 @@ async def test_stale_discovery_falls_back_to_resolved_registered_server() -> Non def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="stale-publication", name="publication", url="https://old.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + server_id="stale-publication", + name="publication", + url="https://old.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, ) manager._set_oauth_discovery_deferred(original.server_id, True) original_slot: Final = manager._oauth_discovery_slot(original.server_id) @@ -13329,9 +13388,13 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, - authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + server_id="expiring-session", + name="temporary", + url="https://idp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", ) manager._set_oauth_discovery_deferred(server.server_id, True) resolved: Final = await manager.ensure_oauth_metadata_discovered(server) @@ -13432,7 +13495,9 @@ async def test_openapi_health_reports_size_limit_as_unknown_and_caches_failure(r result = await manager.health_check_server(server.server_id) cached = await manager.health_check_server(server.server_id) assert result.status == "unknown" - assert result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + assert ( + result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" + ) assert cached.health_check_error == result.health_check_error assert cached.last_health_check == result.last_health_check assert route.call_count == 1 @@ -13444,8 +13509,11 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( - server_id="cancelled-cache", name="cancelled-cache", transport=MCPTransport.http, - spec_path="https://93.184.216.34/cancelled-cache.json", auth_type=MCPAuth.none, + server_id="cancelled-cache", + name="cancelled-cache", + transport=MCPTransport.http, + spec_path="https://93.184.216.34/cancelled-cache.json", + auth_type=MCPAuth.none, ) manager.registry = {server.server_id: server} started = asyncio.Event() @@ -13574,7 +13642,9 @@ class _DiscoveryUpstream: def _discovery_server() -> MCPServer: - return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http) + return MCPServer( + server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http + ) @pytest.mark.asyncio @@ -13742,7 +13812,9 @@ async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) assert upstream.initializes == 2 -@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5))) +@pytest.mark.parametrize( + "value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5)) +) def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl @@ -14005,26 +14077,45 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them( class TestProtectedCredentialPreparation: @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,credential", [ - (MCPAuth.bearer_token, None), - (MCPAuth.bearer_token, "Bearer"), - (MCPAuth.api_key, None), - (MCPAuth.basic, "Basic"), - ]) + @pytest.mark.parametrize( + "auth_type,credential", + [ + (MCPAuth.bearer_token, None), + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.api_key, None), + (MCPAuth.basic, "Basic"), + ], + ) @pytest.mark.parametrize("dispatch", ["managed", "local"]) async def test_openapi_dispatch_rejects_unusable_effective_credentials( - self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - auth_type: MCPAuthType, credential: str | None, dispatch: str, + self, + tmp_path: Path, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + auth_type: MCPAuthType, + credential: str | None, + dispatch: str, ) -> None: from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix spec_path: Final = tmp_path / "openapi.json" - spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"}, - "paths": {"/echo": {"get": {"operationId": "echo"}}}})) + spec_path.write_text( + json.dumps( + { + "openapi": "3.0.0", + "info": {"title": "Auth", "version": "1"}, + "paths": {"/echo": {"get": {"operationId": "echo"}}}, + } + ) + ) server: Final = MCPServer( - server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential, + server_id="dispatch-auth", + name="dispatch-auth", + url="https://upstream.example", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=credential, ) manager: Final = MCPServerManager() await manager._register_openapi_tools(str(spec_path), server, server.url) @@ -14047,14 +14138,21 @@ class TestProtectedCredentialPreparation: self, transport: MCPTransport, client_secret: str | None, subject: str | None ) -> None: server = MCPServer( - server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp", - transport=transport, auth_type=MCPAuth.oauth2_token_exchange, - client_id="gateway", client_secret=client_secret, - token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback", + server_id="incomplete-obo", + name="incomplete-obo", + url="https://upstream.example/mcp", + transport=transport, + auth_type=MCPAuth.oauth2_token_exchange, + client_id="gateway", + client_secret=client_secret, + token_exchange_endpoint="https://idp.example/token", + authentication_token="static-fallback", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header="Bearer override", subject_token=subject, + server, + mcp_auth_header="Bearer override", + subject_token=subject, ) assert exc.value.status_code == (401 if subject is None else 500) assert "static-fallback" not in str(exc.value.detail) @@ -14067,8 +14165,11 @@ class TestProtectedCredentialPreparation: self, auth_type: MCPAuthType, credential: str | dict[str, str] | None ) -> None: server = MCPServer( - server_id="empty-static", name="empty-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-static", + name="empty-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential) @@ -14076,16 +14177,22 @@ class TestProtectedCredentialPreparation: assert "credential" in str(exc.value.detail).lower() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,headers", [ - (MCPAuth.api_key, {"X-API-Key": "key"}), - (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), - ]) + @pytest.mark.parametrize( + "auth_type,headers", + [ + (MCPAuth.api_key, {"X-API-Key": "key"}), + (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), + ], + ) async def test_static_auth_accepts_actual_forwarded_credential( self, auth_type: MCPAuthType, headers: dict[str, str] ) -> None: server = MCPServer( - server_id="header-static", name="header-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="header-static", + name="header-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, ) client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers) assert client._get_auth_headers() == headers @@ -14094,29 +14201,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange]) async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None: server = MCPServer( - server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="openapi-empty", + name="openapi-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, token_exchange_endpoint="https://idp.example/token", ) with pytest.raises(HTTPException) as exc: await MCPServerManager().resolve_openapi_upstream_auth( - mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None, - user_api_key_auth=None, forwarded_headers=None, + mcp_server=server, + oauth2_headers=None, + raw_headers=None, + mcp_auth_header=None, + user_api_key_auth=None, + forwarded_headers=None, ) assert exc.value.status_code in (401, 500) @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,slot,value", [ - (MCPAuth.api_key, "X-API-Key", "token"), - (MCPAuth.authorization, "Authorization", "opaque-secret-value"), - (MCPAuth.authorization, "Authorization", "Bearer abc"), - (MCPAuth.authorization, "Authorization", "Custom abc"), - ]) + @pytest.mark.parametrize( + "auth_type,slot,value", + [ + (MCPAuth.api_key, "X-API-Key", "token"), + (MCPAuth.authorization, "Authorization", "opaque-secret-value"), + (MCPAuth.authorization, "Authorization", "Bearer abc"), + (MCPAuth.authorization, "Authorization", "Custom abc"), + ], + ) async def test_raw_static_credentials_are_forwarded_unchanged( - self, auth_type: MCPAuthType, slot: str, value: str, + self, + auth_type: MCPAuthType, + slot: str, + value: str, ) -> None: - server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value) + server = MCPServer( + server_id="raw-key", + name="raw-key", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, + ) client = await MCPServerManager()._create_mcp_client(server) assert client._resolved_auth is not None request = httpx.Request("GET", server.url) @@ -14130,17 +14256,24 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"]) @pytest.mark.parametrize("source", ["configured", "caller", "forwarded"]) async def test_raw_authorization_rejects_bare_schemes_before_dispatch( - self, respx_mock: MockRouter, value: str, source: str, + self, + respx_mock: MockRouter, + value: str, + source: str, ) -> None: server: Final = MCPServer( - server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.authorization, + server_id="raw-empty", + name="raw-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.authorization, authentication_token=value if source == "configured" else None, ) destination: Final = respx_mock.route().respond(200) with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc: await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, + server, + mcp_auth_header=value if source == "caller" else None, extra_headers={"Authorization": value} if source == "forwarded" else None, ) assert exc.value.status_code == 500 @@ -14148,9 +14281,15 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None: - server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True, - token_exchange_endpoint="https://idp.example/token") + server = MCPServer( + server_id="obo-byok", + name="obo-byok", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + is_byok=True, + token_exchange_endpoint="https://idp.example/token", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override") assert exc.value.status_code == 401 @@ -14158,41 +14297,66 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")]) async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None: - server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured) + server = MCPServer( + server_id="override", + name="override", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=configured, + ) client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override) assert client._get_auth_headers()["Authorization"] == override @pytest.mark.asyncio @pytest.mark.parametrize("token", [None, "shared"]) async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None: - server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token) + server = MCPServer( + server_id="empty-header", + name="empty-header", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.bearer_token, + authentication_token=token, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "}) assert exc.value.status_code == 500 @pytest.mark.asyncio async def test_custom_slot_uses_its_actual_credential(self) -> None: - server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, - upstream_token_header="X-Custom", authentication_token="key") + server = MCPServer( + server_id="custom", + name="custom", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", + authentication_token="key", + ) client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"}) assert client._credential_slot == "X-Custom" assert await client.discovery_auth_fingerprint() @pytest.mark.asyncio - @pytest.mark.parametrize("static_headers,accepted", [ - ({"apikey": "static-key"}, True), - ({"apikey": ""}, False), - ({"X-Tenant": "tenant"}, True), - ]) + @pytest.mark.parametrize( + "static_headers,accepted", + [ + ({"apikey": "static-key"}, True), + ({"apikey": ""}, False), + ({"X-Tenant": "tenant"}, True), + ], + ) async def test_api_key_carried_by_static_header_passes_fail_closed_check( self, static_headers: dict[str, str], accepted: bool ) -> None: server: Final = MCPServer( - server_id="static-slot", name="static-slot", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, static_headers=static_headers, + server_id="static-slot", + name="static-slot", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + static_headers=static_headers, ) if not accepted: with pytest.raises(HTTPException) as exc: @@ -14204,21 +14368,36 @@ class TestProtectedCredentialPreparation: assert all(request.headers[name] == value for name, value in static_headers.items()) @pytest.mark.asyncio - @pytest.mark.parametrize("static,forwarded,caller", [ - ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), - ({}, {"X-API-Key": "forwarded"}, None), - ({}, None, "ApiKey caller"), - ({"X-API-Key": "static"}, {"Authorization": ""}, None), - ]) + @pytest.mark.parametrize( + "static,forwarded,caller", + [ + ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), + ({}, {"X-API-Key": "forwarded"}, None), + ({}, None, "ApiKey caller"), + ({"X-API-Key": "static"}, {"Authorization": ""}, None), + ], + ) async def test_openapi_static_credentials_remain_supported( - self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, - static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None + self, + respx_mock: MockRouter, + monkeypatch: pytest.MonkeyPatch, + static: dict[str, str], + forwarded: dict[str, str] | None, + caller: str | None, ) -> None: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, _request_extra_headers, create_tool_function, + _request_auth_header, + _request_extra_headers, + create_tool_function, ) + tool: Final = create_tool_function( - "/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key, + "/echo", + "get", + {}, + "https://upstream.example", + headers=static, + auth_type=MCPAuth.api_key, ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") @@ -14252,8 +14431,13 @@ class TestProtectedCredentialPreparation: self.closed = True auth = CancelledAuth() - server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key) + server = MCPServer( + server_id="cancel", + name="cancel", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + ) client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth) with pytest.raises(asyncio.CancelledError): await prepare_mcp_client(server, client) @@ -14262,8 +14446,14 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization]) async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None: - server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ") + server = MCPServer( + server_id="blank-static", + name="blank-static", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=" ", + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server) assert exc.value.status_code == 500 @@ -14271,8 +14461,13 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="]) async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None: - server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic) + server = MCPServer( + server_id="bad-basic", + name="bad-basic", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header}) assert exc.value.status_code == 500 @@ -14281,34 +14476,48 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"]) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None: - server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, - authentication_token=value if source == "configured" else None) + server = MCPServer( + server_id="basic-scheme", + name="basic-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value if source == "configured" else None, + ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None) assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,default_slot", [ - (MCPAuth.api_key, "fixture-key", "X-API-Key"), - (MCPAuth.bearer_token, "fixture-key", "Authorization"), - (MCPAuth.basic, "user:pass", "Authorization"), - (MCPAuth.token, "fixture-key", "Authorization"), - (MCPAuth.authorization, "fixture-key", "Authorization"), - ]) + @pytest.mark.parametrize( + "auth_type,value,default_slot", + [ + (MCPAuth.api_key, "fixture-key", "X-API-Key"), + (MCPAuth.bearer_token, "fixture-key", "Authorization"), + (MCPAuth.basic, "user:pass", "Authorization"), + (MCPAuth.token, "fixture-key", "Authorization"), + (MCPAuth.authorization, "fixture-key", "Authorization"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_usable_credential_survives_an_empty_alternate_header( self, auth_type: MCPAuthType, value: str, default_slot: str, source: str ) -> None: server: Final = MCPServer( - server_id="alternate", name="alternate", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom", + server_id="alternate", + name="alternate", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + upstream_token_header="X-Custom", authentication_token=value if source == "configured" else None, ) empty_slot: Final = default_slot if source == "configured" else "X-Custom" selected_slot: Final = "X-Custom" if source == "configured" else default_slot client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""}, + server, + mcp_auth_header=value if source == "caller" else None, + extra_headers={empty_slot: ""}, ) request: Final = await client.prepare_request_auth() assert request.headers[selected_slot] @@ -14317,8 +14526,12 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None: server: Final = MCPServer( - server_id="both-empty", name="both-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom", + server_id="both-empty", + name="both-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""}) @@ -14331,12 +14544,17 @@ class TestProtectedCredentialPreparation: self, custom_slot: str | None, source: str ) -> None: server: Final = MCPServer( - server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot, + server_id="caller-auth", + name="caller-auth", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + upstream_token_header=custom_slot, ) headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""} client: Final = await MCPServerManager()._create_mcp_client( - server, mcp_auth_header=headers if source == "caller" else None, + server, + mcp_auth_header=headers if source == "caller" else None, extra_headers=headers if source == "forwarded" else None, ) request: Final = await client.prepare_request_auth() @@ -14345,14 +14563,29 @@ class TestProtectedCredentialPreparation: assert custom_slot is None or custom_slot not in request.headers @pytest.mark.asyncio - @pytest.mark.parametrize("value", [ - "", " ", "Bearer", "Basic", "token", "ApiKey", - "Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY", - ]) + @pytest.mark.parametrize( + "value", + [ + "", + " ", + "Bearer", + "Basic", + "token", + "ApiKey", + "Bearer Bearer", + "ApiKey ApiKey", + "token token", + "bEaReR BEARER", + "aPiKeY\tAPIKEY", + ], + ) async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None: server: Final = MCPServer( - server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.api_key, + server_id="caller-empty", + name="caller-empty", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value}) @@ -14363,8 +14596,11 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None: server: Final = MCPServer( - server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, + server_id="basic-pair", + name="basic-pair", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14377,8 +14613,12 @@ class TestProtectedCredentialPreparation: import base64 server: Final = MCPServer( - server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value, + server_id="basic-valid", + name="basic-valid", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.basic, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -14387,17 +14627,27 @@ class TestProtectedCredentialPreparation: assert base64.b64decode(encoded) == value.encode() @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value", [ - (MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"), - (MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"), - ]) + @pytest.mark.parametrize( + "auth_type,value", + [ + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.bearer_token, "Bearer "), + (MCPAuth.bearer_token, "bearer"), + (MCPAuth.token, "token"), + (MCPAuth.token, "token "), + (MCPAuth.token, "TOKEN"), + ], + ) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix( self, auth_type: MCPAuthType, value: str, source: str ) -> None: server: Final = MCPServer( - server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, + server_id="empty-scheme", + name="empty-scheme", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14405,17 +14655,24 @@ class TestProtectedCredentialPreparation: assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize("auth_type,value,expected", [ - (MCPAuth.bearer_token, "token", "Bearer token"), - (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), - (MCPAuth.token, "tokenish", "token tokenish"), - ]) + @pytest.mark.parametrize( + "auth_type,value,expected", + [ + (MCPAuth.bearer_token, "token", "Bearer token"), + (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), + (MCPAuth.token, "tokenish", "token tokenish"), + ], + ) async def test_static_credentials_that_resemble_schemes_remain_usable( self, auth_type: MCPAuthType, value: str, expected: str ) -> None: server: Final = MCPServer( - server_id="real-token", name="real-token", url="https://upstream.example/mcp", - transport=MCPTransport.http, auth_type=auth_type, authentication_token=value, + server_id="real-token", + name="real-token", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -14454,16 +14711,31 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon registry.register_tool("observer-execute", "Execute", {"type": "object"}, upstream) monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) manager = MCPServerManager() - manager.registry = {"observer": MCPServer( - server_id="observer", name="observer", server_name="observer", transport="http", - url="https://observer.example/mcp", spec_path="observer.json", auth_type="none", - )} + manager.registry = { + "observer": MCPServer( + server_id="observer", + name="observer", + server_name="observer", + transport="http", + url="https://observer.example/mcp", + spec_path="observer.json", + auth_type="none", + ) + } manager.tool_name_to_mcp_server_name_mapping = {"observer-execute": "observer"} - result = await asyncio.wait_for(manager.call_tool( - server_name="observer", name="execute", arguments={"text": "hello"}, - user_api_key_auth=UserAPIKeyAuth(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), - guardrail_context=MCPRequestContext.resolve_guardrail_context({"metadata": {"guardrails": ["observe"] if selected else []}}), - ), timeout=5) + result = await asyncio.wait_for( + manager.call_tool( + server_name="observer", + name="execute", + arguments={"text": "hello"}, + user_api_key_auth=UserAPIKeyAuth(), + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + guardrail_context=MCPRequestContext.resolve_guardrail_context( + {"metadata": {"guardrails": ["observe"] if selected else []}} + ), + ), + timeout=5, + ) assert tool_started.is_set() assert guardrail_started.is_set() is selected assert result.is_error is False @@ -14492,11 +14764,21 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie from litellm.proxy._experimental.mcp_server import server as legacy_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback - upstream = MCPServer(server_id="explicit-empty", name="explicit_empty", url="https://example.invalid/mcp", transport=MCPTransport.http, allow_sampling=True) + upstream = MCPServer( + server_id="explicit-empty", + name="explicit_empty", + url="https://example.invalid/mcp", + transport=MCPTransport.http, + allow_sampling=True, + ) token = auth_context_var.set(None) sampling = AsyncMock() try: - legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99") + legacy_server.set_auth_context( + UserAPIKeyAuth(user_id="unrelated"), + raw_headers={"authorization": "unrelated-credential"}, + client_ip="192.0.2.99", + ) with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory, patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), @@ -14504,7 +14786,9 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie if legacy_factory: callback = _create_sampling_callback(user_api_key_auth=UserAPIKeyAuth(user_id="explicit")) else: - await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) + await MCPServerManager()._create_mcp_client( + upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None + ) callback = factory.call_args.kwargs["sampling_callback"] await callback(None, None) captured = sampling.await_args.kwargs @@ -14527,16 +14811,28 @@ class TestSharedIdentifierPrefixWarning: manager = MCPServerManager() rows = [ LiteLLM_MCPServerTable( - server_id="srv-a", server_name="alpha", alias="shared", url="https://a.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-a", + server_name="alpha", + alias="shared", + url="https://a.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-b", server_name="beta", alias="Shared", url="https://b.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-b", + server_name="beta", + alias="Shared", + url="https://b.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-c", server_name="gamma", alias="lonely", url="https://c.example.com/mcp", - transport=MCPTransport.http, updated_at=datetime.now(), + server_id="srv-c", + server_name="gamma", + alias="lonely", + url="https://c.example.com/mcp", + transport=MCPTransport.http, + updated_at=datetime.now(), ), ] raw_rows = [MagicMock(model_dump=lambda row=row: row.model_dump()) for row in rows] @@ -14575,3 +14871,37 @@ class TestSharedIdentifierPrefixWarning: assert "srv-b" in shared_warnings[0] assert "srv-c" not in shared_warnings[0] assert "'shared'" in shared_warnings[0] + + +def test_stored_client_assertion_signing_alg_falls_back_to_rs256_for_non_approved(caplog): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _stored_client_assertion_signing_alg, + ) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert _stored_client_assertion_signing_alg("HS256", "srv") == "RS256" + assert "client_assertion_signing_alg" in caplog.text and "HS256" in caplog.text and "srv" in caplog.text + + +def test_stored_client_assertion_signing_alg_passes_through_approved(caplog): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _stored_client_assertion_signing_alg, + ) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert _stored_client_assertion_signing_alg("PS384", "srv") == "PS384" + assert _stored_client_assertion_signing_alg(None, "srv") == "RS256" + assert "approved algorithm" not in caplog.text + + +def test_mcp_server_model_rejects_non_approved_client_assertion_signing_alg(): + from pydantic import ValidationError + + with pytest.raises(ValidationError) as exc: + MCPServer( + server_id="srv", + name="srv", + transport="http", + client_assertion_signing_alg="EdDSA", + ) + assert "client_assertion_signing_alg" in str(exc.value) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py index 0036035f448..ee1623fe22a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_identity_binding.py @@ -665,3 +665,35 @@ async def test_audit_matching_login_without_nonce_does_not_report_failure(caplog ) assert result is None assert "oauth_identity_binding audit" not in caplog.text + + +def test_select_signing_key_rejects_kid_matching_key_with_non_approved_alg() -> None: + import json + + from cryptography.hazmat.primitives.asymmetric import ed25519 + + from litellm.proxy._experimental.mcp_server.oauth_identity_binding import _BindingRejection + + okp_jwk: Final = { + **json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(ed25519.Ed25519PrivateKey.generate().public_key())), + "kid": KID, + "alg": "EdDSA", + } + result: Final = _select_signing_key(_sign_id_token({}), [okp_jwk]) + assert isinstance(result, _BindingRejection) + + +def test_select_signing_key_picks_approved_key_when_kid_is_shared() -> None: + import json + + from cryptography.hazmat.primitives.asymmetric import ed25519 + + okp_jwk: Final = { + **json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(ed25519.Ed25519PrivateKey.generate().public_key())), + "kid": KID, + "alg": "EdDSA", + } + hs_jwk: Final = {"kty": "oct", "kid": KID, "alg": "HS256", "k": "c2VjcmV0"} + result: Final = _select_signing_key(_sign_id_token({}), [okp_jwk, hs_jwk, _PUBLIC_JWK]) + assert isinstance(result, jwt.PyJWK) + assert result.key_id == KID diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index f8b9043a23f..585f4e0bca1 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -106,9 +106,7 @@ async def test_map_user_to_teams_handles_already_in_team_exception(): ) as mock_add: with patch("litellm.proxy.auth.handle_jwt.verbose_proxy_logger") as mock_logger: # This should not raise an exception - result = await JWTAuthManager.map_user_to_teams( - user_object=user, team_object=team - ) + result = await JWTAuthManager.map_user_to_teams(user_object=user, team_object=team) # Verify the method completed successfully assert result is None @@ -145,14 +143,10 @@ async def test_map_user_to_teams_reraises_other_proxy_exceptions(): async def test_map_user_to_teams_null_inputs(): """Test that method handles null inputs gracefully""" # Test with null user - await JWTAuthManager.map_user_to_teams( - user_object=None, team_object=LiteLLM_TeamTable(team_id="test_team_1") - ) + await JWTAuthManager.map_user_to_teams(user_object=None, team_object=LiteLLM_TeamTable(team_id="test_team_1")) # Test with null team - await JWTAuthManager.map_user_to_teams( - user_object=LiteLLM_UserTable(user_id="test_user_1"), team_object=None - ) + await JWTAuthManager.map_user_to_teams(user_object=LiteLLM_UserTable(user_id="test_user_1"), team_object=None) # Test with both null await JWTAuthManager.map_user_to_teams(user_object=None, team_object=None) @@ -208,9 +202,7 @@ async def test_find_team_with_model_access_reports_passthrough_allowlist_denial( assert exc_info.value.status_code == 403 assert "allowed_passthrough_routes" in exc_info.value.detail assert "requested model" not in exc_info.value.detail - mock_is_auth_enforced_pass_through_route.assert_called_once_with( - route="/my-pass-through", method="POST" - ) + mock_is_auth_enforced_pass_through_route.assert_called_once_with(route="/my-pass-through", method="POST") user_api_key_dict = mock_passthrough_check.call_args.kwargs["user_api_key_dict"] assert user_api_key_dict.metadata == {} @@ -403,9 +395,7 @@ async def test_auth_builder_proxy_admin_user_role(): route = "/chat/completions" # Create user object with PROXY_ADMIN role - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.PROXY_ADMIN - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.PROXY_ADMIN) # Create mock JWT handler jwt_handler = JWTHandler() @@ -414,14 +404,10 @@ async def test_auth_builder_proxy_admin_user_role(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -429,9 +415,7 @@ async def test_auth_builder_proxy_admin_user_role(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -444,9 +428,7 @@ async def test_auth_builder_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -459,12 +441,8 @@ async def test_auth_builder_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -498,9 +476,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): route = "/chat/completions" # Create user object with regular USER role - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create mock JWT handler jwt_handler = JWTHandler() @@ -509,14 +485,10 @@ async def test_auth_builder_non_proxy_admin_user_role(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -524,9 +496,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -539,9 +509,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -554,12 +522,8 @@ async def test_auth_builder_non_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -708,11 +672,7 @@ async def test_sync_user_role_and_teams(): prisma_client=None, user_api_key_cache=mock_user_api_key_cache, litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -721,9 +681,7 @@ async def test_sync_user_role_and_teams(): token = {"roles": ["ADMIN"], "my_id_teams": ["team1", "team2"]} - user = LiteLLM_UserTable( - user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER.value, teams=["team2"] - ) + user = LiteLLM_UserTable(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER.value, teams=["team2"]) prisma = AsyncMock() prisma.db.litellm_usertable.update = AsyncMock() @@ -750,11 +708,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -771,9 +725,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): prisma = AsyncMock() prisma.db.litellm_usertable.update = AsyncMock() - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args @@ -793,11 +745,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -818,9 +766,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ): - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args @@ -840,11 +786,7 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -860,9 +802,7 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): prisma = AsyncMock() - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_not_called() @@ -887,9 +827,7 @@ def test_get_all_jwt_team_ids_unions_singular_and_plural(): assert jwt_handler.get_all_jwt_team_ids({"teams": ["a", "b"]}) == ["a", "b"] # both populated, no overlap - assert jwt_handler.get_all_jwt_team_ids( - {"team_id": "primary", "teams": ["a", "b"]} - ) == ["a", "b", "primary"] + assert jwt_handler.get_all_jwt_team_ids({"team_id": "primary", "teams": ["a", "b"]}) == ["a", "b", "primary"] # both populated with overlap — singular dedup'd assert jwt_handler.get_all_jwt_team_ids({"team_id": "a", "teams": ["a", "b"]}) == [ @@ -898,9 +836,7 @@ def test_get_all_jwt_team_ids_unions_singular_and_plural(): ] # singular field as multi-element list (some IdPs) — merge all, preserve plural-first order - assert jwt_handler.get_all_jwt_team_ids( - {"team_id": ["primary", "secondary"], "teams": ["a"]} - ) == [ + assert jwt_handler.get_all_jwt_team_ids({"team_id": ["primary", "secondary"], "teams": ["a"]}) == [ "a", "primary", "secondary", @@ -960,19 +896,11 @@ async def test_map_jwt_role_to_litellm_role(): litellm_jwtauth=LiteLLM_JWTAuth( jwt_litellm_role_map=[ # Exact match - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ), + JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN), # Wildcard patterns - JWTLiteLLMRoleMap( - jwt_role="user_*", litellm_role=LitellmUserRoles.INTERNAL_USER - ), - JWTLiteLLMRoleMap( - jwt_role="team_?", litellm_role=LitellmUserRoles.TEAM - ), - JWTLiteLLMRoleMap( - jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER - ), + JWTLiteLLMRoleMap(jwt_role="user_*", litellm_role=LitellmUserRoles.INTERNAL_USER), + JWTLiteLLMRoleMap(jwt_role="team_?", litellm_role=LitellmUserRoles.TEAM), + JWTLiteLLMRoleMap(jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER), ], roles_jwt_field="roles", ), @@ -1044,9 +972,7 @@ async def test_map_jwt_role_to_litellm_role(): # Test patterns that don't match character classes jwt_handler.litellm_jwtauth.jwt_litellm_role_map = [ - JWTLiteLLMRoleMap( - jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER - ), + JWTLiteLLMRoleMap(jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER), ] token = {"roles": ["dev_4"]} # 4 is not in [123] result = jwt_handler.map_jwt_role_to_litellm_role(token) @@ -1141,25 +1067,19 @@ async def test_nested_jwt_field_access(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="profile.object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(nested_token, None) == "obj789" # Test 5b: object_id_jwt_field with flat access (backward compatibility) jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(flat_token, None) == "obj789" # Test 6: end_user_id_jwt_field with nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - end_user_id_jwt_field="customer.end_user_id" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") assert jwt_handler.get_end_user_id(nested_token, None) == "customer123" # Test 6b: end_user_id_jwt_field with flat access (backward compatibility) @@ -1175,9 +1095,7 @@ async def test_nested_jwt_field_access(): assert jwt_handler.get_team_id(flat_token, None) == "team456" # Test 8: roles_jwt_field with deeply nested access (already supported, but testing) - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - roles_jwt_field="resource_access.my-client.roles" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") assert jwt_handler.get_jwt_role(nested_token, []) == ["admin", "user"] # Test 9: user_roles_jwt_field with nested access (already supported, but testing) @@ -1223,10 +1141,7 @@ async def test_nested_jwt_field_missing_paths(): # Test 2: Missing user.email should return default jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="user.email") - assert ( - jwt_handler.get_user_email(incomplete_token, "default@example.com") - == "default@example.com" - ) + assert jwt_handler.get_user_email(incomplete_token, "default@example.com") == "default@example.com" # Test 3: Missing groups should return empty list jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups") @@ -1241,43 +1156,28 @@ async def test_nested_jwt_field_missing_paths(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="profile.object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(incomplete_token, "default_obj") == "default_obj" # Test 6: Missing customer.end_user_id should return default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - end_user_id_jwt_field="customer.end_user_id" - ) - assert ( - jwt_handler.get_end_user_id(incomplete_token, "default_customer") - == "default_customer" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") + assert jwt_handler.get_end_user_id(incomplete_token, "default_customer") == "default_customer" # Test 7: Missing tenant.team_id should use team_id_default fallback - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_id_jwt_field="tenant.team_id", team_id_default="fallback_team" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="tenant.team_id", team_id_default="fallback_team") assert jwt_handler.get_team_id(incomplete_token, "default_team") == "fallback_team" # Test 8: Missing resource_access.my-client.roles should return default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - roles_jwt_field="resource_access.my-client.roles" - ) - assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == [ - "default_role" - ] + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") + assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == ["default_role"] # Test 9: Missing nested user roles should return default jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( user_roles_jwt_field="resource_access.my-client.roles", user_allowed_roles=["admin", "user"], ) - assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == [ - "default_user_role" - ] + assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == ["default_user_role"] @pytest.mark.asyncio @@ -1302,9 +1202,7 @@ async def test_metadata_prefix_handling_in_nested_fields(): } # Test 1: metadata.user.email should access user.email after prefix removal - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - user_email_jwt_field="metadata.user.email" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="metadata.user.email") # The get_nested_value function removes "metadata." prefix, so "metadata.user.email" becomes "user.email" assert jwt_handler.get_user_email(token, None) == "user@example.com" @@ -1340,9 +1238,7 @@ async def test_find_team_with_model_access_model_group(monkeypatch): async def mock_get_team_object(*args, **kwargs): # type: ignore return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -1450,9 +1346,7 @@ async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatc async def mock_get_team_object(*args, **kwargs): # type: ignore return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -1518,14 +1412,10 @@ async def test_auth_builder_returns_team_membership_object(): team_id=_team_id, budget_id="budget_123", spend=10.5, - litellm_budget_table=LiteLLM_BudgetTable( - budget_id="budget_123", rpm_limit=100, tpm_limit=5000 - ), + litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget_123", rpm_limit=100, tpm_limit=5000), ) - user_object = LiteLLM_UserTable( - user_id=_user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id=_user_id, user_role=LitellmUserRoles.INTERNAL_USER) team_object = LiteLLM_TeamTable(team_id=_team_id) @@ -1536,14 +1426,10 @@ async def test_auth_builder_returns_team_membership_object(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1551,9 +1437,7 @@ async def test_auth_builder_returns_team_membership_object(): return_value=(_user_id, "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1566,9 +1450,7 @@ async def test_auth_builder_returns_team_membership_object(): new_callable=AsyncMock, return_value=(_team_id, team_object), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1587,15 +1469,9 @@ async def test_auth_builder_returns_team_membership_object(): user_object.user_id, ), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": _user_id, "scope": ""} @@ -1614,24 +1490,12 @@ async def test_auth_builder_returns_team_membership_object(): ) # Verify that team_membership_object is returned - assert result["team_membership"] is not None, ( - "team_membership should be present" - ) - assert result["team_membership"] == mock_team_membership, ( - "team_membership should match the mock object" - ) - assert result["team_membership"].user_id == _user_id, ( - "team_membership user_id should match" - ) - assert result["team_membership"].team_id == _team_id, ( - "team_membership team_id should match" - ) - assert result["team_membership"].budget_id == "budget_123", ( - "team_membership budget_id should match" - ) - assert result["team_membership"].spend == 10.5, ( - "team_membership spend should match" - ) + assert result["team_membership"] is not None, "team_membership should be present" + assert result["team_membership"] == mock_team_membership, "team_membership should match the mock object" + assert result["team_membership"].user_id == _user_id, "team_membership user_id should match" + assert result["team_membership"].team_id == _team_id, "team_membership team_id should match" + assert result["team_membership"].budget_id == "budget_123", "team_membership budget_id should match" + assert result["team_membership"].spend == 10.5, "team_membership spend should match" @pytest.mark.asyncio @@ -1648,9 +1512,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create JWT handler with OIDC UserInfo enabled jwt_handler = JWTHandler() @@ -1677,18 +1539,12 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): # Mock all the dependencies with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1696,9 +1552,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1711,9 +1565,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1726,15 +1578,9 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, ): # Set up mock return values mock_get_userinfo.return_value = userinfo_response @@ -1775,9 +1621,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create JWT handler with OIDC UserInfo disabled jwt_handler = JWTHandler() @@ -1801,18 +1645,12 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): # Mock all the dependencies with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1820,9 +1658,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): return_value=("test_user_1", None, None), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1835,9 +1671,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1850,15 +1684,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, ): # Set up mock return values mock_auth_jwt.return_value = jwt_response @@ -1904,9 +1732,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) jwt_handler = JWTHandler() user_api_key_cache = DualCache() @@ -1925,9 +1751,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() jwt_response = {"sub": "test_user_1", "scope": ""} with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), @@ -1968,9 +1792,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), ): mock_auth_jwt.return_value = jwt_response @@ -2052,9 +1874,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): ) team_object = LiteLLM_TeamTable(team_id="team-2") - user_object = LiteLLM_UserTable( - user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER) with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, @@ -2065,9 +1885,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): new_callable=AsyncMock, return_value=None, ), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, patch.object( JWTAuthManager, "get_objects", @@ -2075,9 +1893,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): return_value=(user_object, None, None, None, user_object.user_id), ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), ): mock_auth_jwt.return_value = { "sub": "user-1", @@ -2294,9 +2110,7 @@ async def test_auth_builder_rbac_team_loads_team_for_passthrough_allowlist(): return_value=(None, None, None, None, None), ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.RouteChecks.is_auth_enforced_pass_through_route", return_value=True, @@ -2325,9 +2139,7 @@ async def test_auth_builder_rbac_team_loads_team_for_passthrough_allowlist(): mock_get_team.assert_awaited_once() assert mock_get_team.await_args.kwargs["team_id"] == "team-rbac" user_api_key_dict = mock_passthrough_check.call_args.kwargs["user_api_key_dict"] - assert user_api_key_dict.team_metadata == { - "allowed_passthrough_routes": ["/my-pass-through"] - } + assert user_api_key_dict.team_metadata == {"allowed_passthrough_routes": ["/my-pass-through"]} @pytest.mark.asyncio @@ -2421,9 +2233,7 @@ async def test_auth_builder_admin_on_llm_route_honors_team_header(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2475,9 +2285,7 @@ async def test_auth_builder_admin_on_mgmt_route_ignores_team_header(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2531,9 +2339,7 @@ async def test_auth_builder_admin_on_llm_route_without_header_unchanged(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2577,9 +2383,7 @@ async def test_get_team_alias_with_nested_fields(): } # Test nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_alias_jwt_field="organization.team.name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_alias_jwt_field="organization.team.name") assert jwt_handler.get_team_alias(nested_token, None) == "engineering-team" # Test flat access (backward compatibility) @@ -2587,9 +2391,7 @@ async def test_get_team_alias_with_nested_fields(): assert jwt_handler.get_team_alias(nested_token, None) == "flat-team" # Test missing field returns default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_alias_jwt_field="nonexistent.field" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_alias_jwt_field="nonexistent.field") assert jwt_handler.get_team_alias(nested_token, "default-team") == "default-team" # Test with team_alias_jwt_field not configured @@ -2620,9 +2422,7 @@ async def test_is_required_team_id_with_team_alias_field(): assert jwt_handler.is_required_team_id() is True # Both fields set - should return True - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_id_jwt_field="team_id", team_alias_jwt_field="team_name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_name") assert jwt_handler.is_required_team_id() is True @@ -2654,9 +2454,7 @@ async def test_find_and_validate_specific_team_id_with_team_alias(): # Mock team object returned by get_team_object_by_alias team_object = LiteLLM_TeamTable(team_id="resolved-team-id", team_alias="my-team") - with patch( - "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock - ) as mock_get_by_alias: + with patch("litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias: mock_get_by_alias.return_value = team_object team_id, result_team = await JWTAuthManager.find_and_validate_specific_team_id( @@ -2699,9 +2497,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): jwt_handler.update_environment( prisma_client=None, user_api_key_cache=user_api_key_cache, - litellm_jwtauth=LiteLLM_JWTAuth( - team_id_jwt_field="team_id", team_alias_jwt_field="team_alias" - ), + litellm_jwtauth=LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"), ) # Token with both team_id and team name @@ -2711,9 +2507,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): team_object = LiteLLM_TeamTable(team_id="direct-team-id") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -2793,9 +2587,7 @@ async def test_get_org_alias_with_nested_fields(): } # Test nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - org_alias_jwt_field="company.organization.name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_alias_jwt_field="company.organization.name") assert jwt_handler.get_org_alias(nested_token, None) == "acme-corp" # Test flat access @@ -2803,9 +2595,7 @@ async def test_get_org_alias_with_nested_fields(): assert jwt_handler.get_org_alias(nested_token, None) == "flat-org" # Test missing field returns default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - org_alias_jwt_field="nonexistent.field" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_alias_jwt_field="nonexistent.field") assert jwt_handler.get_org_alias(nested_token, "default-org") == "default-org" # Test with org_alias_jwt_field not configured @@ -2843,9 +2633,7 @@ async def test_get_objects_resolves_org_by_name(): models=[], ) - with patch( - "litellm.proxy.auth.handle_jwt.get_org_object_by_alias", new_callable=AsyncMock - ) as mock_get_by_alias: + with patch("litellm.proxy.auth.handle_jwt.get_org_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias: mock_get_by_alias.return_value = org_object ( @@ -2923,9 +2711,7 @@ async def test_resolve_jwks_url_resolves_oidc_discovery_document(): litellm_jwtauth=LiteLLM_JWTAuth(), ) - discovery_url = ( - "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" - ) + discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" mock_response = MagicMock() @@ -2956,9 +2742,7 @@ async def test_resolve_jwks_url_caches_resolved_jwks_uri(): litellm_jwtauth=LiteLLM_JWTAuth(), ) - discovery_url = ( - "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" - ) + discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" mock_response = MagicMock() @@ -3102,9 +2886,7 @@ async def test_find_and_validate_specific_team_id_hints_bracket_notation(): error_msg = str(exc_info.value) # Should mention the bad field name and suggest the fix assert "roles.0" in error_msg, f"Expected field name in: {error_msg}" - assert "roles" in error_msg and "list" in error_msg, ( - f"Expected hint about using 'roles' instead: {error_msg}" - ) + assert "roles" in error_msg and "list" in error_msg, f"Expected hint about using 'roles' instead: {error_msg}" @pytest.mark.asyncio @@ -3132,9 +2914,7 @@ async def test_find_and_validate_specific_team_id_hints_bracket_index_notation() error_msg = str(exc_info.value) assert "roles[0]" in error_msg, f"Expected field name in: {error_msg}" - assert "roles" in error_msg and "list" in error_msg, ( - f"Expected hint about using 'roles' instead: {error_msg}" - ) + assert "roles" in error_msg and "list" in error_msg, f"Expected hint about using 'roles' instead: {error_msg}" @pytest.mark.asyncio @@ -3244,9 +3024,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( if len(user_teams) == 1 and get_team_object_return == "resolved_row": only = user_teams[0] team_table = LiteLLM_TeamTable(team_id=only) - membership = LiteLLM_TeamMembership( - user_id=user_id, team_id=only, litellm_budget_table=None - ) + membership = LiteLLM_TeamMembership(user_id=user_id, team_id=only, litellm_budget_table=None) get_team_return_value = team_table membership_return_value = membership else: @@ -3305,9 +3083,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -3324,9 +3100,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( code = 404 if get_team_object_return == "http_404" else 500 mock_get_team.side_effect = HTTPException( status_code=code, - detail={ - "error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call." - }, + detail={"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."}, ) else: mock_get_team.return_value = get_team_return_value @@ -3435,9 +3209,7 @@ async def test_auth_builder_single_team_fallback_membership_outage_raises_instea ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -3478,9 +3250,7 @@ def _reset_unscoped_warning_flag(): JWTHandler._unscoped_jwt_warning_emitted = False -def test_build_decode_kwargs_no_env_disables_both_verifications( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_no_env_disables_both_verifications(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.delenv("JWT_AUDIENCE", raising=False) monkeypatch.delenv("JWT_ISSUER", raising=False) @@ -3491,9 +3261,7 @@ def test_build_decode_kwargs_no_env_disables_both_verifications( assert kwargs["options"] == {"verify_aud": False, "verify_iss": False} -def test_build_decode_kwargs_audience_only_enables_aud_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_audience_only_enables_aud_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") monkeypatch.delenv("JWT_ISSUER", raising=False) @@ -3505,9 +3273,7 @@ def test_build_decode_kwargs_audience_only_enables_aud_verification( assert kwargs["options"] == {"verify_iss": False} -def test_build_decode_kwargs_issuer_only_enables_iss_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_issuer_only_enables_iss_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.delenv("JWT_AUDIENCE", raising=False) monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") @@ -3518,9 +3284,7 @@ def test_build_decode_kwargs_issuer_only_enables_iss_verification( assert kwargs["options"] == {"verify_aud": False} -def test_build_decode_kwargs_both_set_enables_full_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_both_set_enables_full_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") @@ -3532,9 +3296,7 @@ def test_build_decode_kwargs_both_set_enables_full_verification( assert kwargs["options"] is None -def test_build_decode_kwargs_warns_once_when_unscoped( - monkeypatch, _reset_unscoped_warning_flag, caplog -): +def test_build_decode_kwargs_warns_once_when_unscoped(monkeypatch, _reset_unscoped_warning_flag, caplog): """The warning about unscoped JWT auth should fire on the first call but not on every subsequent decode.""" import logging @@ -3550,17 +3312,12 @@ def test_build_decode_kwargs_warns_once_when_unscoped( matching = [ r for r in caplog.records - if "JWT auth is enabled" in r.getMessage() - and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + if "JWT auth is enabled" in r.getMessage() and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() ] - assert len(matching) == 1, ( - f"Expected exactly one warning across 3 calls, got {len(matching)}" - ) + assert len(matching) == 1, f"Expected exactly one warning across 3 calls, got {len(matching)}" -def test_build_decode_kwargs_no_warning_when_scoped( - monkeypatch, _reset_unscoped_warning_flag, caplog -): +def test_build_decode_kwargs_no_warning_when_scoped(monkeypatch, _reset_unscoped_warning_flag, caplog): import logging monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") @@ -3569,11 +3326,7 @@ def test_build_decode_kwargs_no_warning_when_scoped( JWTHandler._build_decode_kwargs() - matching = [ - r - for r in caplog.records - if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() - ] + matching = [r for r in caplog.records if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()] assert matching == [] @@ -3631,11 +3384,7 @@ async def test_find_team_with_model_access_unresolved_group_claim_returns_none( from litellm.router import Router - router = Router( - model_list=[ - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} - ] - ) + router = Router(model_list=[{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}]) proxy_server_module = types.ModuleType("proxy_server") proxy_server_module.llm_router = router monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) @@ -3707,9 +3456,7 @@ async def test_find_and_validate_specific_team_id_non_404_http_exception_propaga "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, ) as mock_get_team: - mock_get_team.side_effect = HTTPException( - status_code=status_code, detail="non-404 failure" - ) + mock_get_team.side_effect = HTTPException(status_code=status_code, detail="non-404 failure") with pytest.raises(HTTPException) as exc_info: await JWTAuthManager.find_and_validate_specific_team_id( @@ -3782,9 +3529,7 @@ async def test_find_team_with_model_access_resolved_team_without_model_still_rai async def mock_get_team_object(*_args, **_kwargs): return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -3849,11 +3594,7 @@ async def test_find_team_with_model_access_unresolved_group_claim_default_raises from litellm.router import Router - router = Router( - model_list=[ - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} - ] - ) + router = Router(model_list=[{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}]) proxy_server_module = types.ModuleType("proxy_server") proxy_server_module.llm_router = router monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) @@ -3890,12 +3631,7 @@ def test_canonical_user_id_rebinds_to_legacy_uuid(): jwt_email = "matt@example.com" user_object = LiteLLM_UserTable(user_id=legacy_uuid, user_email=jwt_email) - assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id=jwt_email, user_object=user_object - ) - == legacy_uuid - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=jwt_email, user_object=user_object) == legacy_uuid def test_canonical_user_id_no_change_when_ids_match(): @@ -3903,28 +3639,20 @@ def test_canonical_user_id_no_change_when_ids_match(): same = "alice@example.com" user_object = LiteLLM_UserTable(user_id=same, user_email=same) - assert ( - JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object) - == same - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object) == same def test_canonical_user_id_returns_claim_when_no_user_object(): """No resolved row (e.g. upsert disabled / brand new) -> keep the claim.""" assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id="newcomer@example.com", user_object=None - ) + JWTAuthManager._canonical_user_id_from_db(user_id="newcomer@example.com", user_object=None) == "newcomer@example.com" ) def test_canonical_user_id_returns_none_when_claim_none_and_no_object(): """Defensive: no claim and no row -> stays None, never invents an id.""" - assert ( - JWTAuthManager._canonical_user_id_from_db(user_id=None, user_object=None) - is None - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=None, user_object=None) is None def test_canonical_user_id_no_change_when_db_user_id_falsy(): @@ -3934,10 +3662,7 @@ def test_canonical_user_id_no_change_when_db_user_id_falsy(): user_id = "" assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id="jwt@example.com", user_object=_Stub() - ) - == "jwt@example.com" + JWTAuthManager._canonical_user_id_from_db(user_id="jwt@example.com", user_object=_Stub()) == "jwt@example.com" ) @@ -3953,9 +3678,7 @@ async def test_auth_jwt_expired_token_raises_401_jwk_path(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() with ( - patch.object( - jwt_handler, "get_public_key", new_callable=AsyncMock - ) as mock_get_public_key, + patch.object(jwt_handler, "get_public_key", new_callable=AsyncMock) as mock_get_public_key, patch( "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", return_value={"kid": "test-kid"}, @@ -3991,9 +3714,7 @@ async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): mock_cert.public_key.return_value.public_bytes.return_value = b"fake-key" with ( - patch.object( - jwt_handler, "get_public_key", new_callable=AsyncMock - ) as mock_get_public_key, + patch.object(jwt_handler, "get_public_key", new_callable=AsyncMock) as mock_get_public_key, patch( "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", return_value={"kid": "test-kid"}, @@ -4007,9 +3728,7 @@ async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): side_effect=jwt_lib.ExpiredSignatureError("Signature has expired"), ), ): - mock_get_public_key.return_value = ( - "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" - ) + mock_get_public_key.return_value = "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" with pytest.raises(ProxyException) as exc_info: await jwt_handler.auth_jwt(token="expired.jwt.token") @@ -4123,9 +3842,7 @@ async def test_get_public_key_fetches_and_caches_jwks_response(): ) assert public_key == jwk - cached_keys = await cache.async_get_cache( - key="litellm_jwt_auth_keys_https://issuer.example.com/keys" - ) + cached_keys = await cache.async_get_cache(key="litellm_jwt_auth_keys_https://issuer.example.com/keys") assert cached_keys == [jwk] @@ -4370,9 +4087,7 @@ async def test_lowering_public_key_stale_ttl_stops_serving_a_copy_cached_under_t # The operator tightens the window and restarts; the cache, and its long-lived copy, survive. endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) - tightened = _get_jwt_handler_with_scripted_endpoint( - cache, endpoint, public_key_stale_ttl=lowered_stale_ttl - ) + tightened = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_stale_ttl=lowered_stale_ttl) await cache.async_set_cache( key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}", value=time.time() - 7200, @@ -4720,9 +4435,7 @@ def test_get_jwks_url_for_issuer_falls_back_to_discovery_document(): jwks_url = jwt_handler._get_jwks_url_for_issuer(issuer_config=issuer_config) - assert ( - jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" - ) + assert jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" @pytest.mark.asyncio @@ -4746,9 +4459,7 @@ async def test_get_objects_team_membership_uses_rebound_user_id(): return None jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - user_id_jwt_field="email", user_id_upsert=True - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="email", user_id_upsert=True) with ( patch( @@ -4844,9 +4555,7 @@ async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims( assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-org" - assert jwt_handler.get_team_id(token=claims, default_value=None) == ( - "example-org/litellm-fork" - ) + assert jwt_handler.get_team_id(token=claims, default_value=None) == ("example-org/litellm-fork") @pytest.mark.asyncio @@ -4949,9 +4658,7 @@ async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): claims = await jwt_handler.auth_jwt(token=token) - assert ( - jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" - ) + assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" @pytest.mark.asyncio @@ -4984,7 +4691,7 @@ async def test_multi_issuer_jwt_unknown_issuer_falls_back_to_global_jwks(monkeyp kid="issuer-key", ) - with pytest.raises(Exception, match='Missing JWT Public Key URL from environment\\.') as exc: + with pytest.raises(Exception, match="Missing JWT Public Key URL from environment\\.") as exc: await jwt_handler.auth_jwt(token=token) assert "Missing JWT Public Key URL from environment." in str(exc.value) @@ -5058,7 +4765,7 @@ async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch) kid=shared_kid, ) - with pytest.raises(Exception, match='Validation fails: Signature verification failed') as exc: + with pytest.raises(Exception, match="Validation fails: Signature verification failed") as exc: await jwt_handler.auth_jwt(token=token) assert "Validation fails" in str(exc.value) @@ -5113,7 +4820,7 @@ def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( issuer = "https://issuer.example.com" jwks_url = f"{issuer}/keys" - with pytest.raises(Exception, match='must configure audience or set') as exc: + with pytest.raises(Exception, match="must configure audience or set") as exc: LiteLLM_JWTAuth( issuers=[ { @@ -5130,7 +4837,7 @@ def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): issuer = "https://issuer.example.com" jwks_url = f"{issuer}/keys" - with pytest.raises(Exception, match='cannot set audience and disable_audience_validation=True') as exc: + with pytest.raises(Exception, match="cannot set audience and disable_audience_validation=True") as exc: LiteLLM_JWTAuth( issuers=[ { @@ -5142,9 +4849,7 @@ def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): ] ) - assert "cannot set audience and disable_audience_validation=True together" in str( - exc.value - ) + assert "cannot set audience and disable_audience_validation=True together" in str(exc.value) @pytest.mark.asyncio @@ -5197,21 +4902,15 @@ async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): claims = await jwt_handler.auth_jwt(token=token) - assert jwt_handler.get_user_id(token=claims, default_value=None) == ( - "real-user@example.com" - ) - assert jwt_handler.get_user_email(token=claims, default_value=None) == ( - "real-user@example.com" - ) + assert jwt_handler.get_user_id(token=claims, default_value=None) == ("real-user@example.com") + assert jwt_handler.get_user_email(token=claims, default_value=None) == ("real-user@example.com") assert jwt_handler.get_team_id(token=claims, default_value=None) == "real-team" assert jwt_handler.get_team_ids_from_jwt(token=claims) == [ "real-team", "secondary-team", ] assert jwt_handler.get_org_id(token=claims, default_value=None) == "real-org" - assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ( - "real-end-user" - ) + assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ("real-end-user") @pytest.mark.asyncio @@ -5251,15 +4950,11 @@ async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): assert JWTHandler.LITELLM_TEAM_ID_CLAIM not in claims assert jwt_handler.get_user_id(token=claims, default_value=None) is None assert jwt_handler.get_team_id(token=claims, default_value=None) is None - assert jwt_handler.get_user_email(token=claims, default_value=None) == ( - "real-user@example.com" - ) + assert jwt_handler.get_user_email(token=claims, default_value=None) == ("real-user@example.com") @pytest.mark.asyncio -async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( - monkeypatch, caplog -): +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning(monkeypatch, caplog): import logging monkeypatch.delenv("JWT_AUDIENCE", raising=False) @@ -5309,11 +5004,7 @@ def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deploym JWTHandler._build_decode_kwargs() - matching = [ - r - for r in caplog.records - if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() - ] + matching = [r for r in caplog.records if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()] assert len(matching) == 1 @@ -5441,9 +5132,7 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership(): "expect_403", ), [ - pytest.param( - True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team" - ), + pytest.param(True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"), pytest.param( True, ["team_a", "team_b"], @@ -5521,9 +5210,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( async def call_auth_builder(): with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock - ) as mock_auth_jwt, + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -5563,9 +5250,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -5669,9 +5354,7 @@ async def test_resolve_db_team_fallback_skips_team_without_model_access(): teams=["restricted_team", "allowed_team"], ) teams = { - "restricted_team": LiteLLM_TeamTable( - team_id="restricted_team", models=["claude-3"] - ), + "restricted_team": LiteLLM_TeamTable(team_id="restricted_team", models=["claude-3"]), "allowed_team": LiteLLM_TeamTable(team_id="allowed_team", models=["gpt-4"]), } @@ -5853,9 +5536,7 @@ async def _run_auth_builder_with_header_team( jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = jwt_auth_config with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value=token - ), + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value=token), patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -5874,9 +5555,7 @@ async def _run_auth_builder_with_header_team( new_callable=AsyncMock, return_value=None, ), - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=allowed_team_ids - ), + patch.object(JWTAuthManager, "get_all_team_ids", return_value=allowed_team_ids), patch.object( JWTAuthManager, "get_objects", @@ -5885,9 +5564,7 @@ async def _run_auth_builder_with_header_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -5914,9 +5591,7 @@ async def _run_auth_builder_with_header_team( @pytest.mark.asyncio -async def test_auth_builder_header_team_not_found_matches_non_membership_denial() -> ( - None -): +async def test_auth_builder_header_team_not_found_matches_non_membership_denial() -> None: """A provisional x-litellm-team-id naming a nonexistent team must produce the exact same 403 shape as one naming an existing team outside the caller's memberships. Letting get_team_object's 404 surface would give any @@ -5936,19 +5611,15 @@ async def test_auth_builder_header_team_not_found_matches_non_membership_denial( return LiteLLM_TeamTable(team_id=team_id) with pytest.raises(HTTPException) as missing_exc: - await _run_auth_builder_with_header_team( - config, token, "team_ghost", user_object, _team_lookup_404, set() - ) + await _run_auth_builder_with_header_team(config, token, "team_ghost", user_object, _team_lookup_404, set()) with pytest.raises(HTTPException) as outsider_exc: - await _run_auth_builder_with_header_team( - config, token, "team_other", user_object, team_exists, set() - ) + await _run_auth_builder_with_header_team(config, token, "team_other", user_object, team_exists, set()) assert missing_exc.value.status_code == 403 assert outsider_exc.value.status_code == 403 - assert missing_exc.value.detail.replace( - "team_ghost", "" - ) == outsider_exc.value.detail.replace("team_other", "") + assert missing_exc.value.detail.replace("team_ghost", "") == outsider_exc.value.detail.replace( + "team_other", "" + ) assert "exist" not in missing_exc.value.detail @@ -6145,9 +5816,7 @@ async def test_auth_builder_db_fallback_does_not_validate_rbac_team_against_db_m ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6299,9 +5968,7 @@ async def test_auth_builder_db_fallback_runs_when_only_team_id_default_set(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6383,9 +6050,7 @@ async def test_auth_builder_alias_only_token_resolves_alias_not_db_fallback(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6442,14 +6107,10 @@ async def test_find_and_validate_specific_team_id_alias_wins_over_team_id_defaul ) jwt_token = {"sub": "user-1", "team_alias": "my-team"} - alias_team = LiteLLM_TeamTable( - team_id="alias_resolved_team", team_alias="my-team" - ) + alias_team = LiteLLM_TeamTable(team_id="alias_resolved_team", team_alias="my-team") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -6496,9 +6157,7 @@ async def test_find_and_validate_specific_team_id_team_id_default_used_without_a default_team = LiteLLM_TeamTable(team_id="config_default_team") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -6567,9 +6226,7 @@ async def test_auth_builder_db_fallback_enforces_passthrough_route_access(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6723,18 +6380,14 @@ async def test_sync_user_role_and_teams_singular_claim_reconciles_memberships(): "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ) as mock_patch: - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, AsyncMock() - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, AsyncMock()) mock_patch.assert_awaited_once() assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == { "team_stale_a", "team_stale_b", } - assert set(mock_patch.call_args.kwargs["teams_ids_to_add_user_to"]) == { - "team_primary" - } + assert set(mock_patch.call_args.kwargs["teams_ids_to_add_user_to"]) == {"team_primary"} assert user.teams == ["team_primary"] @@ -6793,9 +6446,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6870,9 +6521,7 @@ async def test_auth_builder_header_cannot_override_rbac_team_under_db_fallback() ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6922,9 +6571,7 @@ async def test_auth_builder_header_team_enforces_team_allowed_routes_under_db_fa async def call(route: str): with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock - ) as mock_auth_jwt, + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -6952,9 +6599,7 @@ async def test_auth_builder_header_team_enforces_team_allowed_routes_under_db_fa ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -7312,14 +6957,10 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ) as mock_patch: - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, AsyncMock() - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, AsyncMock()) mock_patch.assert_awaited_once() - assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == { - "team_existing" - } + assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == {"team_existing"} assert mock_patch.call_args.kwargs["teams_ids_to_add_user_to"] == [] assert user.teams == [] @@ -7544,8 +7185,12 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a if identity_only: identity = await JWTAuthManager.resolve_identity( - api_key=token, jwt_handler=jwt_handler, prisma_client=None, - user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, + api_key=token, + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, ) assert identity.agent_id == "canonical-agent-id" return @@ -7580,8 +7225,12 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch if identity_only: with pytest.raises(HTTPException) as denial: await JWTAuthManager.resolve_identity( - api_key=token, jwt_handler=jwt_handler, prisma_client=None, - user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, + api_key=token, + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, ) assert denial.value.status_code == 403 return @@ -7648,7 +7297,9 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc from litellm.proxy.management_endpoints import team_endpoints handler, token = _entra_signed_app_token( - monkeypatch, azp="canonical-agent-id", scope=LiteLLM_JWTAuth().admin_jwt_scope, + monkeypatch, + azp="canonical-agent-id", + scope=LiteLLM_JWTAuth().admin_jwt_scope, ) handler.bind_agent_lookup(_entra_agent_registry()) handler.litellm_jwtauth.team_id_upsert = True @@ -7660,10 +7311,16 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc resolve = JWTAuthManager.auth_builder if admission else JWTAuthManager.authorize_jwt result = await resolve( - api_key=token, jwt_handler=handler, request_data={}, general_settings={}, - route="/chat/completions", prisma_client=database, - user_api_key_cache=handler.user_api_key_cache, parent_otel_span=None, - proxy_logging_obj=MagicMock(), request_headers={"x-litellm-team-id": "new-team"}, + api_key=token, + jwt_handler=handler, + request_data={}, + general_settings={}, + route="/chat/completions", + prisma_client=database, + user_api_key_cache=handler.user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + request_headers={"x-litellm-team-id": "new-team"}, ) assert result["is_proxy_admin"] is True @@ -7688,7 +7345,12 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning cache: Final = UserApiKeyCache() cache.set_cache("litellm_jwt_auth_keys_https://admin.example/jwks", [jwk]) user_id: Final = f"admin-status-{existing_user}-{warm_cache}-{email}" - user: Final = LiteLLM_UserTable(user_id=user_id, user_email="admin@allowed.example", metadata={"scim_active": False}, organization_memberships=[]) + user: Final = LiteLLM_UserTable( + user_id=user_id, + user_email="admin@allowed.example", + metadata={"scim_active": False}, + organization_memberships=[], + ) if existing_user and warm_cache: cache.set_cache(user_id, user) database: Final = MagicMock() @@ -7701,7 +7363,9 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning prisma_client=database, user_api_key_cache=cache, litellm_jwtauth=LiteLLM_JWTAuth( - user_id_jwt_field="sub", user_id_upsert=True, user_email_jwt_field="email", + user_id_jwt_field="sub", + user_id_upsert=True, + user_email_jwt_field="email", user_allowed_email_domain="allowed.example", ), ) @@ -7709,12 +7373,22 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning monkeypatch.setenv("JWT_ISSUER", "https://admin.example") monkeypatch.setenv("JWT_AUDIENCE", "gateway") token: Final = _encode_rsa_jwt( - private_key, "https://admin.example", "gateway", "admin-status", + private_key, + "https://admin.example", + "gateway", + "admin-status", {"sub": user_id, "scope": "litellm_proxy_admin", **({"email": email} if email else {})}, ) result: Final = await JWTAuthManager.auth_builder( - api_key=token, jwt_handler=handler, prisma_client=database, user_api_key_cache=cache, - parent_otel_span=None, proxy_logging_obj=MagicMock(), request_data={}, general_settings={}, route="/user/info", + api_key=token, + jwt_handler=handler, + prisma_client=database, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + request_data={}, + general_settings={}, + route="/user/info", ) assert result["is_proxy_admin"] is True assert result["user_id"] == user_id @@ -7722,3 +7396,162 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning users.create.assert_not_awaited() if existing_user: assert users.find_unique.await_count == (0 if warm_cache else 1) + + +# --- B5: JWT algorithm allowlists (LIT-8429) --- + + +def _okp_keypair_and_jwk() -> "tuple[object, dict]": + import json + + from cryptography.hazmat.primitives.asymmetric import ed25519 + + private_key = ed25519.Ed25519PrivateKey.generate() + jwk = { + **json.loads(__import__("jwt").algorithms.OKPAlgorithm.to_jwk(private_key.public_key())), + "kid": "ed", + "alg": "EdDSA", + "use": "sig", + } + return private_key, jwk + + +def _eddsa_jwt(private_key: object, kid: str = "ed") -> str: + import jwt + + current_time = int(time.time()) + return jwt.encode( + {"sub": "test-subject", "iat": current_time, "exp": current_time + 300}, + private_key, # pyright: ignore[reportArgumentType] + algorithm="EdDSA", + headers={"kid": kid}, + ) + + +def test_approved_jwt_algorithm_literal_matches_tuple(): + from typing import get_args + + from litellm.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm + + assert set(get_args(ApprovedJwtAlgorithm)) == set(APPROVED_JWT_ALGORITHMS) + + +def test_allowed_jwt_algorithms_drops_legacy_only_in_fips_mode(): + from litellm.proxy.auth.jwt_algorithms import allowed_jwt_algorithms + + assert "EdDSA" not in allowed_jwt_algorithms(True) + assert list(allowed_jwt_algorithms(False)) == JWTHandler.SUPPORTED_JWT_ALGORITHMS + + +@pytest.mark.parametrize( + "key,algorithms,expected", + [ + ({"kty": "RSA", "alg": "RS256", "kid": "a"}, ("RS256",), True), + ({"kty": "RSA", "alg": "HS256", "kid": "a"}, ("RS256", "ES256"), False), + ({"kty": "OKP", "alg": "EdDSA", "kid": "a"}, ("RS256", "ES256"), False), + ({"kty": "RSA", "kid": "a"}, ("RS256", "ES256"), True), + ({"kty": "EC", "kid": "a"}, ("RS256", "ES256"), True), + ({"kty": "EC", "kid": "a"}, ("RS256",), False), + ({"kty": "OKP", "kid": "a"}, ("RS256", "EdDSA"), True), + ({"kty": "OKP", "kid": "a"}, ("RS256", "ES256"), False), + ({"kty": "oct", "kid": "a"}, ("RS256", "HS256"), False), + ({"kid": "a"}, ("RS256",), False), + ], +) +def test_jwks_keys_for_filters_by_declared_or_inferred_algorithm(key, algorithms, expected): + from litellm.proxy.auth.jwt_algorithms import jwks_keys_for + + assert list(jwks_keys_for([key], algorithms)) == ([key] if expected else []) + + +def test_fips_mode_rejects_eddsa_token_but_accepts_rs256(): + import jwt + + ed_private, ed_jwk = _okp_keypair_and_jwk() + rsa_private, rsa_jwk = _rsa_keypair_and_jwk() + handler = JWTHandler(fips_mode=lambda: True) + + with pytest.raises(jwt.exceptions.InvalidAlgorithmError): + handler._decode_jwt_with_public_key( + token=_eddsa_jwt(ed_private), + public_key=ed_jwk, + audience=None, + disable_audience_validation=True, + ) + claims: Final = handler._decode_jwt_with_public_key( + token=_encode_rsa_jwt(rsa_private, "iss", "aud", "rsa"), + public_key=rsa_jwk, + audience="aud", + issuer="iss", + ) + assert claims["sub"] == "test-subject" + + +def test_eddsa_token_accepted_with_deprecation_log_outside_fips(caplog): + import logging + + from litellm.proxy.auth.handle_jwt import _log_eddsa_deprecation + + ed_private, ed_jwk = _okp_keypair_and_jwk() + handler = JWTHandler(fips_mode=lambda: False) + + _log_eddsa_deprecation.cache_clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + claims: Final = handler._decode_jwt_with_public_key( + token=_eddsa_jwt(ed_private), + public_key=ed_jwk, + audience=None, + disable_audience_validation=True, + ) + assert claims["sub"] == "test-subject" + assert "EdDSA" in caplog.text and "deprecated" in caplog.text and "LITELLM_FIPS_MODE" in caplog.text + + _log_eddsa_deprecation.cache_clear() + caplog.clear() + rsa_private, rsa_jwk = _rsa_keypair_and_jwk() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + handler._decode_jwt_with_public_key( + token=_encode_rsa_jwt(rsa_private, "iss", "aud", "rsa"), + public_key=rsa_jwk, + audience="aud", + issuer="iss", + ) + assert "deprecated" not in caplog.text + _log_eddsa_deprecation.cache_clear() + + +def _rsa_keypair_and_jwk() -> "tuple[object, dict]": + import json + + import jwt as jwt_lib + from cryptography.hazmat.primitives.asymmetric import rsa + + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + jwk = { + **json.loads(jwt_lib.algorithms.RSAAlgorithm.to_jwk(private_key.public_key())), + "kid": "rsa", + "alg": "RS256", + "use": "sig", + } + return private_key, jwk + + +@pytest.mark.asyncio +async def test_jwks_url_kid_matching_only_a_filtered_key_raises(): + _, ed_jwk = _okp_keypair_and_jwk() + jwks_url = "https://idp.example.com/jwks" + + fips_handler = JWTHandler(fips_mode=lambda: True) + fips_handler.user_api_key_cache = DualCache() + await fips_handler.user_api_key_cache.async_set_cache( + key=f"litellm_jwt_auth_keys_{jwks_url}", value=[ed_jwk], ttl=600 + ) + with pytest.raises(NoMatchingJWTPublicKeyError): + await fips_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="ed") + + normal_handler = JWTHandler(fips_mode=lambda: False) + normal_handler.user_api_key_cache = DualCache() + await normal_handler.user_api_key_cache.async_set_cache( + key=f"litellm_jwt_auth_keys_{jwks_url}", value=[ed_jwk], ttl=600 + ) + assert (await normal_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="ed"))["kid"] == "ed" diff --git a/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py b/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py index cb2276ab39d..1150fab69b4 100644 --- a/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py +++ b/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py @@ -174,9 +174,7 @@ def test_jwks_public_key_can_verify_signed_jwt(): def test_build_claims_standard_fields(): """_build_claims() populates iss, aud, iat, exp, nbf.""" - signer = _make_signer( - issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300 - ) + signer = _make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300) user_dict = _make_user_api_key_dict() data = {"mcp_tool_name": "get_weather"} @@ -347,17 +345,13 @@ async def test_hook_skips_non_mcp_call_types(): data=original_data, call_type=call_type, # type: ignore[arg-type] ) - assert "extra_headers" not in ( - result or {} - ), f"extra_headers should not be set for {call_type}" + assert "extra_headers" not in (result or {}), f"extra_headers should not be set for {call_type}" @pytest.mark.asyncio async def test_hook_signs_list_mcp_tools(): """async_pre_call_hook() signs JWT for list_mcp_tools with list scope.""" - signer = _make_signer( - issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300 - ) + signer = _make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300) user_dict = _make_user_api_key_dict(user_id="alice", team_id="backend") data = {"mcp_tool_name": "should_be_cleared"} @@ -382,9 +376,7 @@ async def test_hook_signs_list_mcp_tools(): @pytest.mark.asyncio async def test_signed_token_is_verifiable(): """The JWT injected by the hook can be verified against the JWKS public key.""" - signer = _make_signer( - issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300 - ) + signer = _make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300) user_dict = _make_user_api_key_dict(user_id="alice", team_id="backend") data = {"mcp_tool_name": "search"} @@ -600,9 +592,7 @@ async def test_channel_token_injected_when_configured(): assert isinstance(result, dict) assert "x-mcp-channel-token" in result["extra_headers"] - channel_token = result["extra_headers"]["x-mcp-channel-token"].removeprefix( - "Bearer " - ) + channel_token = result["extra_headers"]["x-mcp-channel-token"].removeprefix("Bearer ") channel_payload = _decode_unverified(channel_token) assert channel_payload["aud"] == "bedrock-gateway" @@ -818,9 +808,7 @@ def test_initialize_guardrail_passes_all_params(): litellm_params.issuer = "https://litellm.example.com" litellm_params.audience = "mcp-test" litellm_params.ttl_seconds = 120 - litellm_params.access_token_discovery_uri = ( - "https://idp.example.com/.well-known/openid-configuration" - ) + litellm_params.access_token_discovery_uri = "https://idp.example.com/.well-known/openid-configuration" litellm_params.token_introspection_endpoint = "https://idp.example.com/introspect" litellm_params.verify_issuer = "https://idp.example.com" litellm_params.verify_audience = "api://test" @@ -843,10 +831,7 @@ def test_initialize_guardrail_passes_all_params(): assert signer.issuer == "https://litellm.example.com" assert signer.audience == "mcp-test" assert signer.ttl_seconds == 120 - assert ( - signer.access_token_discovery_uri - == "https://idp.example.com/.well-known/openid-configuration" - ) + assert signer.access_token_discovery_uri == "https://idp.example.com/.well-known/openid-configuration" assert signer.token_introspection_endpoint == "https://idp.example.com/introspect" assert signer.verify_issuer == "https://idp.example.com" assert signer.verify_audience == "api://test" @@ -879,9 +864,7 @@ def _make_httpx_response(json_body: dict, status_code: int = 200): if status_code >= 400: from httpx import HTTPStatusError, Request, Response - mock_resp.raise_for_status.side_effect = HTTPStatusError( - "error", request=MagicMock(), response=MagicMock() - ) + mock_resp.raise_for_status.side_effect = HTTPStatusError("error", request=MagicMock(), response=MagicMock()) return mock_resp @@ -940,9 +923,7 @@ async def test_fetch_jwks_uses_cache_on_second_call(): @pytest.mark.asyncio async def test_get_oidc_discovery_caches_when_jwks_uri_present(): """_get_oidc_discovery caches the doc when jwks_uri is in the response.""" - signer = _make_signer( - access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration" - ) + signer = _make_signer(access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration") signer._oidc_discovery_doc = None # ensure fresh discovery_doc = { @@ -964,9 +945,7 @@ async def test_get_oidc_discovery_caches_when_jwks_uri_present(): @pytest.mark.asyncio async def test_get_oidc_discovery_does_not_cache_when_jwks_uri_absent(): """_get_oidc_discovery does NOT cache a doc that is missing jwks_uri.""" - signer = _make_signer( - access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration" - ) + signer = _make_signer(access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration") signer._oidc_discovery_doc = None bad_doc = {"issuer": "https://idp.example.com"} # no jwks_uri @@ -1069,9 +1048,7 @@ async def test_verify_incoming_jwt_raises_on_expired_token(): @pytest.mark.asyncio async def test_introspect_opaque_token_returns_claims_when_active(): """_introspect_opaque_token returns the introspection payload for active tokens.""" - signer = _make_signer( - token_introspection_endpoint="https://idp.example.com/introspect" - ) + signer = _make_signer(token_introspection_endpoint="https://idp.example.com/introspect") introspection_response = { "active": True, @@ -1095,9 +1072,7 @@ async def test_introspect_opaque_token_returns_claims_when_active(): @pytest.mark.asyncio async def test_introspect_opaque_token_raises_on_inactive_token(): """_introspect_opaque_token raises ExpiredSignatureError when active=false.""" - signer = _make_signer( - token_introspection_endpoint="https://idp.example.com/introspect" - ) + signer = _make_signer(token_introspection_endpoint="https://idp.example.com/introspect") fake_resp = _make_httpx_response({"active": False}) mock_client = MagicMock() @@ -1128,9 +1103,7 @@ async def test_hook_raises_401_when_jwt_verification_fails(): """async_pre_call_hook raises HTTP 401 when incoming JWT verification fails.""" from fastapi import HTTPException - signer = _make_signer( - access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration" - ) + signer = _make_signer(access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration") with patch.object( signer, @@ -1269,3 +1242,101 @@ async def test_inject_mcp_jwt_signs_for_tool_call_path(): scopes = set(decoded["scope"].split()) assert "mcp:tools/call" in scopes assert "mcp:tools/search_web:call" in scopes + + +# --------------------------------------------------------------------------- +# B5: incoming-JWT JWKS allowlist (LIT-8429) +# --------------------------------------------------------------------------- + + +def _okp_idp_key_and_token(now: int): + import json + + from cryptography.hazmat.primitives.asymmetric import ed25519 + + private_key = ed25519.Ed25519PrivateKey.generate() + jwk = { + **json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(private_key.public_key())), + "kid": "ed", + "alg": "EdDSA", + } + token = jwt.encode( + {"sub": "idp-user", "iat": now, "exp": now + 300}, + private_key, + algorithm="EdDSA", + headers={"kid": "ed"}, + ) + return jwk, token + + +def _oct_idp_key_and_token(now: int): + secret = b"integration-hs256-client-secret-0123456789abcdef" + jwk = { + "kty": "oct", + "kid": "sym", + "alg": "HS256", + "k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(), + } + token = jwt.encode( + {"sub": "idp-user", "iat": now, "exp": now + 300}, + secret, + algorithm="HS256", + headers={"kid": "sym"}, + ) + return jwk, token + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "key_and_token", (_oct_idp_key_and_token, _okp_idp_key_and_token), ids=("oct-HS256", "OKP-EdDSA") +) +async def test_verify_incoming_jwt_rejects_jwks_without_approved_algorithms(key_and_token): + """A JWKS whose only keys use non-approved algorithms can never verify the incoming token.""" + jwks_key, incoming_token = key_and_token(int(time.time())) + signer = _make_signer( + access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration", + ) + with patch.object( + signer, + "_get_oidc_discovery", + new_callable=AsyncMock, + return_value={"jwks_uri": "https://idp.example.com/jwks"}, + ): + with patch( + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer._fetch_jwks", + new_callable=AsyncMock, + return_value=[jwks_key], + ): + with pytest.raises(jwt.exceptions.PyJWKSetError, match="approved signing algorithm"): + await signer._verify_incoming_jwt(incoming_token) + + +@pytest.mark.asyncio +async def test_verify_incoming_jwt_ignores_filtered_keys_sharing_kid(): + """An approved RS256 key still verifies when the JWKS also carries a non-approved key with the same kid.""" + signer = _make_signer( + access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration", + ) + now = int(time.time()) + oct_jwk, _ = _oct_idp_key_and_token(now) + oct_jwk["kid"] = signer._kid + incoming_token = jwt.encode( + {"sub": "idp-user", "iat": now, "exp": now + 300}, + signer._private_key, + algorithm="RS256", + headers={"kid": signer._kid}, + ) + jwks = [oct_jwk, *signer.get_jwks()["keys"]] + with patch.object( + signer, + "_get_oidc_discovery", + new_callable=AsyncMock, + return_value={"jwks_uri": "https://idp.example.com/jwks"}, + ): + with patch( + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer._fetch_jwks", + new_callable=AsyncMock, + return_value=jwks, + ): + payload = await signer._verify_incoming_jwt(incoming_token) + assert payload["sub"] == "idp-user" diff --git a/tests/test_litellm/types/test_mcp.py b/tests/test_litellm/types/test_mcp.py index 5450ec4aa48..de138bf662c 100644 --- a/tests/test_litellm/types/test_mcp.py +++ b/tests/test_litellm/types/test_mcp.py @@ -53,12 +53,12 @@ def test_has_header_matches_any_casing() -> None: @pytest.mark.parametrize( "target,expected", [ - ("https://upstream.example.com/other", False), # same origin + ("https://upstream.example.com/other", False), # same origin ("https://upstream.example.com:443/other", False), # explicit default port - ("https://attacker.example.com/collect", True), # different host - ("http://upstream.example.com/collect", True), # scheme downgrade, same host + ("https://attacker.example.com/collect", True), # different host + ("http://upstream.example.com/collect", True), # scheme downgrade, same host ("https://upstream.example.com:8443/other", True), # different port, same host - ("https://sub.upstream.example.com/x", True), # different host + ("https://sub.upstream.example.com/x", True), # different host ], ) def test_origin_is_scheme_host_and_port_not_host_alone(target: str, expected: bool) -> None: @@ -85,3 +85,29 @@ async def test_the_hook_drops_the_slot_only_once_the_origin_changes() -> None: foreign = httpx.Request("GET", "https://attacker.example.com/x", headers={"esb-oauth": "Bearer x"}) await hook(foreign) assert "esb-oauth" not in foreign.headers + + +def test_new_mcp_server_request_rejects_non_approved_client_assertion_signing_alg() -> None: + from pydantic import ValidationError + + from litellm.proxy._types import NewMCPServerRequest + + with pytest.raises(ValidationError) as exc: + NewMCPServerRequest( + transport="http", + url="https://mcp.example.com", + credentials={"client_assertion_signing_alg": "HS256"}, + ) + assert "client_assertion_signing_alg" in str(exc.value) + + +def test_new_mcp_server_request_accepts_approved_client_assertion_signing_alg() -> None: + from litellm.proxy._types import NewMCPServerRequest + + request = NewMCPServerRequest( + transport="http", + url="https://mcp.example.com", + credentials={"client_assertion_signing_alg": "ES256"}, + ) + assert request.credentials is not None + assert request.credentials["client_assertion_signing_alg"] == "ES256" From 505eb9d80816ae0b407aa36514a73b8ce554708d Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 08:13:06 +0000 Subject: [PATCH 07/18] refactor(types): keep JWT algorithm literals under litellm/types (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 2 +- .../mcp_server/oauth_identity_binding.py | 3 ++- .../mcp_server/outbound_credentials/types.py | 2 +- litellm/proxy/auth/handle_jwt.py | 8 ++------ litellm/proxy/auth/jwt_algorithms.py | 18 ++---------------- .../mcp_jwt_signer/mcp_jwt_signer.py | 3 ++- litellm/types/mcp.py | 2 +- litellm/types/mcp_server/mcp_server_manager.py | 2 +- litellm/types/proxy/auth/jwt_algorithms.py | 17 +++++++++++++++++ .../test_litellm/proxy/auth/test_handle_jwt.py | 2 +- 10 files changed, 30 insertions(+), 29 deletions(-) create mode 100644 litellm/types/proxy/auth/jwt_algorithms.py diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8e182755656..2ff52dd5edc 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -169,7 +169,6 @@ from litellm.proxy._types import ( is_per_server_oauth_discovery_eligible, ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils -from litellm.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl from litellm.proxy.management_endpoints.sso.id_jag_assertion_capture import ( @@ -194,6 +193,7 @@ from litellm.types.mcp_server.mcp_server_manager import ( MCPOAuthMetadata, MCPServer, ) +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm from litellm.types.utils import CallTypes if TYPE_CHECKING: diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index 1f2a492beab..ee2ee79e697 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -23,9 +23,10 @@ from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache from litellm.llms.custom_httpx.http_handler import get_async_httpx_client -from litellm.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, jwks_keys_for +from litellm.proxy.auth.jwt_algorithms import jwks_keys_for from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS _ALLOWED_ID_TOKEN_ALGORITHMS: Final = APPROVED_JWT_ALGORITHMS _JWKS_CACHE_TTL_SECONDS: Final = 3600 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py index 6c310784370..8a7f2653971 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py @@ -40,12 +40,12 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Ok, Result, ) -from litellm.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm from litellm.types.mcp import ( DEFAULT_CREDENTIAL_HEADER, DEFAULT_SUBJECT_TOKEN_TYPE, normalize_upstream_header_name, ) +from litellm.types.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm class AuthResolution(str, Enum): diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 6f4a10c6f69..0e0d0376578 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -54,12 +54,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_checks import can_team_access_model -from litellm.proxy.auth.jwt_algorithms import ( - APPROVED_JWT_ALGORITHMS, - LEGACY_JWT_ALGORITHMS, - allowed_jwt_algorithms, - jwks_keys_for, -) +from litellm.proxy.auth.jwt_algorithms import allowed_jwt_algorithms, jwks_keys_for from litellm.proxy.auth.model_access_denied import ( ModelAccessDeniedHTTPException, model_access_denied_client_message, @@ -76,6 +71,7 @@ from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.user_repository import UserRepository from litellm.types.agents import AgentResponse from litellm.types.proxy.auth.auth_checks import UserNotFoundError +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, LEGACY_JWT_ALGORITHMS from .auth_checks import ( TeamNotFoundError, diff --git a/litellm/proxy/auth/jwt_algorithms.py b/litellm/proxy/auth/jwt_algorithms.py index 48e03ba9981..a3f7ec0c4f9 100644 --- a/litellm/proxy/auth/jwt_algorithms.py +++ b/litellm/proxy/auth/jwt_algorithms.py @@ -1,22 +1,8 @@ from collections.abc import Collection, Mapping, Sequence from types import MappingProxyType -from typing import Final, Literal +from typing import Final -ApprovedJwtAlgorithm = Literal["RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512"] - -APPROVED_JWT_ALGORITHMS: Final[tuple[ApprovedJwtAlgorithm, ...]] = ( - "RS256", - "RS384", - "RS512", - "PS256", - "PS384", - "PS512", - "ES256", - "ES384", - "ES512", -) - -LEGACY_JWT_ALGORITHMS: Final = ("EdDSA",) +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, LEGACY_JWT_ALGORITHMS _KEY_TYPE_ALGORITHMS: Final = MappingProxyType( { diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py index c1a88f3aa67..dd0611176d7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py @@ -89,9 +89,10 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, jwks_keys_for +from litellm.proxy.auth.jwt_algorithms import jwks_keys_for from litellm.types.guardrail_base_init import GuardrailBaseInitKwargs from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS from litellm.types.utils import CallTypesLiteral if TYPE_CHECKING: diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index df2d2b979cc..7960abcd083 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -11,8 +11,8 @@ import httpx from pydantic import BaseModel, ConfigDict, Field from typing_extensions import TypedDict -from litellm.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm from litellm.types.llms.base import HiddenParams +from litellm.types.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm if TYPE_CHECKING: import httpx2 diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index dc610724485..ee4f08cfa47 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -4,7 +4,6 @@ from typing import Any, Final, Literal from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator from typing_extensions import Self -from litellm.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm from litellm.types.mcp import ( DEFAULT_SUBJECT_TOKEN_TYPE, MCPAuth, @@ -13,6 +12,7 @@ from litellm.types.mcp import ( MCPTransportType, normalize_upstream_header_name, ) +from litellm.types.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm # MCPInfo now allows arbitrary additional fields for custom metadata MCPInfo = dict[str, Any] diff --git a/litellm/types/proxy/auth/jwt_algorithms.py b/litellm/types/proxy/auth/jwt_algorithms.py new file mode 100644 index 00000000000..73c383dd9a1 --- /dev/null +++ b/litellm/types/proxy/auth/jwt_algorithms.py @@ -0,0 +1,17 @@ +from typing import Final, Literal + +ApprovedJwtAlgorithm = Literal["RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512"] + +APPROVED_JWT_ALGORITHMS: Final[tuple[ApprovedJwtAlgorithm, ...]] = ( + "RS256", + "RS384", + "RS512", + "PS256", + "PS384", + "PS512", + "ES256", + "ES384", + "ES512", +) + +LEGACY_JWT_ALGORITHMS: Final = ("EdDSA",) diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 585f4e0bca1..b3e0795d070 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -7431,7 +7431,7 @@ def _eddsa_jwt(private_key: object, kid: str = "ed") -> str: def test_approved_jwt_algorithm_literal_matches_tuple(): from typing import get_args - from litellm.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm + from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm assert set(get_args(ApprovedJwtAlgorithm)) == set(APPROVED_JWT_ALGORITHMS) From 2a726aa659ebd4218f7b2b4c9d950ca9c935a02f Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 08:47:29 +0000 Subject: [PATCH 08/18] fix(auth): satisfy type gates and adapt unknown-alg test for approved JWT algorithms (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 2 +- litellm/proxy/auth/handle_jwt.py | 49 +++++++++---------- .../mcp_jwt_signer/mcp_jwt_signer.py | 4 +- .../test_token_endpoint.py | 24 +++++---- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +- 5 files changed, 43 insertions(+), 38 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 2ff52dd5edc..e5ff4f72867 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1485,7 +1485,7 @@ def _warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None: def _stored_client_assertion_signing_alg(value: object, server_name: str) -> ApprovedJwtAlgorithm: if value in APPROVED_JWT_ALGORITHMS: - return cast(ApprovedJwtAlgorithm, value) + return cast(ApprovedJwtAlgorithm, value) # cast-ok: value was checked against APPROVED_JWT_ALGORITHMS if value is not None: verbose_logger.warning( "MCP server %s: client_assertion_signing_alg %r is not an approved algorithm, using RS256", diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 0e0d0376578..4fff3b99469 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -221,7 +221,10 @@ class JWTHandler: # Supported algos: https://pyjwt.readthedocs.io/en/stable/algorithms.html # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret # the key in different ways (e.g. HS* and RS*)." - SUPPORTED_JWT_ALGORITHMS = [*APPROVED_JWT_ALGORITHMS, *LEGACY_JWT_ALGORITHMS] + SUPPORTED_JWT_ALGORITHMS = [ # mutable-ok: list kept for backward compatibility + *APPROVED_JWT_ALGORITHMS, + *LEGACY_JWT_ALGORITHMS, + ] LITELLM_JWT_ISSUER_CLAIM = "_litellm_jwt_issuer" LITELLM_USER_ID_CLAIM = "_litellm_user_id" LITELLM_USER_EMAIL_CLAIM = "_litellm_user_email" @@ -955,10 +958,11 @@ class JWTHandler: log_context=f"kid={kid}", ) + allowed: Final = self.allowed_algorithms() usable_keys: Final[JWKKeyValue] = ( - list(jwks_keys_for(keys, self.allowed_algorithms())) + list(jwks_keys_for(keys, allowed)) # mutable-ok: parse_keys consumes a JWKKeyValue list if isinstance(keys, list) - else next(iter(jwks_keys_for((keys,), self.allowed_algorithms())), {}) + else next(iter(jwks_keys_for((keys,), allowed)), {}) # mutable-ok: single-key dict is a JWKKeyValue ) public_key: Final = self.parse_keys(keys=usable_keys, kid=kid) if public_key is not None: @@ -1231,32 +1235,25 @@ class JWTHandler: ) ) - if isinstance(public_key, dict): - public_key_obj: Final = PyJWK.from_dict(self._get_jwk_from_public_key(public_key=public_key)).key - payload: Final = jwt.decode( - token, - public_key_obj, - algorithms=self.allowed_algorithms(), - options=decode_options, - audience=audience, - issuer=issuer, - leeway=self.leeway, - ) - else: - cert: Final = x509.load_pem_x509_certificate(public_key.encode(), default_backend()) - key: Final = cert.public_key().public_bytes( + key_obj: Final = ( + PyJWK.from_dict(self._get_jwk_from_public_key(public_key=public_key)).key + if isinstance(public_key, dict) + else x509.load_pem_x509_certificate(public_key.encode(), default_backend()) + .public_key() + .public_bytes( serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo, ) - payload = jwt.decode( - token, - key, - algorithms=self.allowed_algorithms(), - audience=audience, - issuer=issuer, - options=decode_options, - leeway=self.leeway, - ) + ) + payload: Final = jwt.decode( + token, + key_obj, + algorithms=self.allowed_algorithms(), + audience=audience, + issuer=issuer, + options=decode_options, + leeway=self.leeway, + ) self._warn_deprecated_signing_algorithm(token) return payload diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py index dd0611176d7..1a3fab21b7c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py @@ -460,7 +460,9 @@ class MCPJWTSigner(CustomGuardrail): from jwt import PyJWKSet try: - jwks_set: Final = PyJWKSet.from_dict({"keys": list(approved_keys)}) + jwks_set: Final = PyJWKSet.from_dict( + {"keys": list(approved_keys)} # mutable-ok: PyJWKSet.from_dict takes a dict + ) except Exception as exc: raise jwt.exceptions.PyJWKSetError(f"Failed to parse JWKS from {jwks_uri!r}: {exc}") from exc diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py index db3a1a386a3..b306e53c0b2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py @@ -13,11 +13,12 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import jwt -import litellm import pytest from cryptography.hazmat.primitives import serialization -from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.hazmat.primitives.asymmetric import ec, rsa +from pydantic import SecretStr +import litellm from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Error, Ok, @@ -34,12 +35,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( CredError, PrivateKeyJwtAuth, ) -from pydantic import SecretStr -_PATCH_TARGET = ( - "litellm.proxy._experimental.mcp_server.outbound_credentials." - "token_endpoint.get_async_httpx_client" -) +_PATCH_TARGET = "litellm.proxy._experimental.mcp_server.outbound_credentials.token_endpoint.get_async_httpx_client" _ENDPOINT = "https://idp.example.com/oauth2/token" _CLIENT_ID = "litellm-client-id" @@ -50,6 +47,15 @@ _PRIVATE_PEM = _RSA_KEY.private_bytes( serialization.PrivateFormat.PKCS8, serialization.NoEncryption(), ).decode() +_EC_PRIVATE_PEM = ( + ec.generate_private_key(ec.SECP256R1()) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() +) _PUBLIC_PEM = ( _RSA_KEY.public_key() .public_bytes( @@ -185,9 +191,9 @@ async def test_fetch_network_error_maps_to_upstream_unavailable(raised): "auth", [ PrivateKeyJwtAuth(private_key=SecretStr("not-a-pem-key"), signing_alg="RS256"), - PrivateKeyJwtAuth(private_key=SecretStr(_PRIVATE_PEM), signing_alg="XX999"), + PrivateKeyJwtAuth(private_key=SecretStr(_EC_PRIVATE_PEM), signing_alg="RS256"), ], - ids=["garbage-key", "unknown-alg"], + ids=["garbage-key", "key-alg-mismatch"], ) async def test_fetch_unsignable_client_assertion_is_misconfigured_not_a_crash(auth): client = AsyncMock() diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 2938e65cfde..af40e742905 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -34841,7 +34841,7 @@ export interface components { /** Aws Session Token */ aws_session_token?: string | null; /** Client Assertion Signing Alg */ - client_assertion_signing_alg?: string | null; + client_assertion_signing_alg?: ("RS256" | "RS384" | "RS512" | "PS256" | "PS384" | "PS512" | "ES256" | "ES384" | "ES512") | null; /** Client Id */ client_id?: string | null; /** Client Private Key */ From 5745a2a7313fc2faf9a14af5dea21e333e6ef21d Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 08:52:40 +0000 Subject: [PATCH 09/18] test: keep existing test files unformatted, add only the new cells (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_token_endpoint.py | 9 +- .../mcp_server/test_mcp_server_manager.py | 646 ++++----------- .../proxy/auth/test_handle_jwt.py | 770 +++++++++++++----- .../proxy/guardrails/test_mcp_jwt_signer.py | 53 +- tests/test_litellm/types/test_mcp.py | 8 +- 5 files changed, 773 insertions(+), 713 deletions(-) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py index b306e53c0b2..29d1f6b693b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py @@ -13,12 +13,11 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import jwt +import litellm import pytest from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import ec, rsa -from pydantic import SecretStr -import litellm from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Error, Ok, @@ -35,8 +34,12 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( CredError, PrivateKeyJwtAuth, ) +from pydantic import SecretStr -_PATCH_TARGET = "litellm.proxy._experimental.mcp_server.outbound_credentials.token_endpoint.get_async_httpx_client" +_PATCH_TARGET = ( + "litellm.proxy._experimental.mcp_server.outbound_credentials." + "token_endpoint.get_async_httpx_client" +) _ENDPOINT = "https://idp.example.com/oauth2/token" _CLIENT_ID = "litellm-client-id" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 23aa8e4f4cd..47bd5b1a8c4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -111,6 +111,7 @@ async def test_manager_sampling_preserves_explicit_headers_without_ambient_conte assert sampling.await_args.kwargs["raw_headers"] == {"x-test-caller": "sampling-caller"} + @pytest.mark.asyncio async def test_sampling_callback_keeps_creation_context_after_caller_switch(): from mcp.server.auth.middleware.auth_context import auth_context_var @@ -218,6 +219,8 @@ def _reload_mcp_manager_module(): return reloaded + + @pytest.fixture(autouse=True) def enable_eager_mcp_oauth_discovery(monkeypatch): monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1") @@ -1391,9 +1394,7 @@ class TestMCPServerManager: assert not any("oauth2_id_jag" in message for message in caplog.messages) @pytest.mark.asyncio - async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso( - self, config_only_mcp_manager_factory, monkeypatch, caplog - ): + async def test_load_servers_from_config_does_not_warn_for_api_key_with_google_sso(self, config_only_mcp_manager_factory, monkeypatch, caplog): self._clear_sso_env(monkeypatch) monkeypatch.setenv("GOOGLE_CLIENT_ID", "google-cid") manager = config_only_mcp_manager_factory() @@ -4679,9 +4680,7 @@ class TestMCPServerManager: @pytest.mark.parametrize("auth_type", [MCPAuth.none, MCPAuth.bearer_token, MCPAuth.api_key, MCPAuth.oauth2]) @pytest.mark.parametrize("is_byok", [False, True]) @pytest.mark.parametrize("scheme", ["http", "https"]) - async def test_openapi_health_loads_spec_without_mcp_handshake( - self, respx_mock, monkeypatch, auth_type, is_byok, scheme - ): + async def test_openapi_health_loads_spec_without_mcp_handshake(self, respx_mock, monkeypatch, auth_type, is_byok, scheme): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -4731,28 +4730,14 @@ class TestMCPServerManager: @pytest.mark.parametrize( ("failure", "expected_status", "expected_error"), [ - ( - httpx.Response(401, text="secret response content"), - "unhealthy", - "OpenAPI specification request failed (HTTP 401)", - ), + (httpx.Response(401, text="secret response content"), "unhealthy", "OpenAPI specification request failed (HTTP 401)"), (httpx.Response(404), "unhealthy", "OpenAPI specification request failed (HTTP 404)"), (httpx.Response(500), "unhealthy", "OpenAPI specification request failed (HTTP 500)"), - ( - httpx.ConnectError("secret network details"), - "unhealthy", - "OpenAPI specification could not be loaded (ConnectError)", - ), - ( - httpx.Response(200, text="secret invalid JSON body"), - "unhealthy", - "OpenAPI specification could not be loaded (JSONDecodeError)", - ), + (httpx.ConnectError("secret network details"), "unhealthy", "OpenAPI specification could not be loaded (ConnectError)"), + (httpx.Response(200, text="secret invalid JSON body"), "unhealthy", "OpenAPI specification could not be loaded (JSONDecodeError)"), ], ) - async def test_openapi_health_reports_safe_failures( - self, respx_mock, monkeypatch, failure, expected_status, expected_error - ): + async def test_openapi_health_reports_safe_failures(self, respx_mock, monkeypatch, failure, expected_status, expected_error): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( @@ -5287,15 +5272,8 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, - method, - operation, - base_url, - headers=None, - server_label=None, - relays_upstream_auth=False, - auth_type=None, - upstream_token_header=None, + path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, + auth_type=None, upstream_token_header=None, ): captured["headers"] = headers captured["server_label"] = server_label @@ -5380,15 +5358,8 @@ class TestMCPServerManager: captured: dict = {} def fake_create_tool_function( - path, - method, - operation, - base_url, - headers=None, - server_label=None, - relays_upstream_auth=False, - auth_type=None, - upstream_token_header=None, + path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False, + auth_type=None, upstream_token_header=None, ): captured["headers"] = headers @@ -12514,9 +12485,7 @@ class TestConfigServerIdPinning: @pytest.mark.asyncio @pytest.mark.parametrize("aliasing_entry_first", [True, False]) - async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected( - self, config_only_mcp_manager_factory, aliasing_entry_first: bool - ): + async def test_pinning_own_name_that_is_another_entrys_alias_is_rejected(self, config_only_mcp_manager_factory, aliasing_entry_first: bool): """A grant naming 'docs_server' reaches both servers unpinned; the pin would narrow it to one.""" manager = config_only_mcp_manager_factory() wiki = ( @@ -12532,9 +12501,7 @@ class TestConfigServerIdPinning: await manager.load_servers_from_config(dict((wiki, docs) if aliasing_entry_first else (docs, wiki))) @pytest.mark.asyncio - async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected( - self, config_only_mcp_manager_factory - ): + async def test_pinning_own_name_that_is_another_entrys_mapped_alias_is_rejected(self, config_only_mcp_manager_factory): manager = config_only_mcp_manager_factory() with pytest.raises(ValueError, match="server_name or alias of MCP server 'wiki_server'"): @@ -12622,9 +12589,7 @@ class TestConfigServerIdPinning: assert second_round == first_round @pytest.mark.asyncio - async def test_shadow_warning_fires_again_when_the_shadowed_set_changes( - self, config_only_mcp_manager_factory, caplog - ): + async def test_shadow_warning_fires_again_when_the_shadowed_set_changes(self, config_only_mcp_manager_factory, caplog): manager = config_only_mcp_manager_factory() await manager.load_servers_from_config(self._config(server_id="docs-prod-1")) @@ -12795,9 +12760,7 @@ class TestConfigServerIdPinning: assert manager.config_mcp_servers["wiki"].url == "https://example.com/mcp" @pytest.mark.asyncio - async def test_a_row_that_shadows_one_id_still_reports_capturing_another( - self, config_only_mcp_manager_factory, caplog - ): + async def test_a_row_that_shadows_one_id_still_reports_capturing_another(self, config_only_mcp_manager_factory, caplog): """Skipping is per identifier, not per row, so the second collision is not lost.""" manager = config_only_mcp_manager_factory() await manager.load_servers_from_config( @@ -13169,8 +13132,7 @@ async def test_pre_call_tool_check_honors_guardrail_attached_to_key(monkeypatch, ("none", {"Authorization": "Bearer injected"}, "extra-headers", "Bearer injected"), ], ) -async def test_debug_resolution_matches_final_header_conflict_winner( - _mcp_request_ctx, +async def test_debug_resolution_matches_final_header_conflict_winner(_mcp_request_ctx, config: Literal["stored", "static", "none"], extra_headers: dict[str, str] | None, expected_source: str, @@ -13241,9 +13203,7 @@ async def test_debug_resolution_matches_final_header_conflict_winner( @pytest.mark.asyncio @pytest.mark.parametrize("transport", ["http", "stdio"]) -async def test_debug_reports_legacy_signing_and_non_http_transport( - _mcp_request_ctx, transport: Literal["http", "stdio"] -) -> None: +async def test_debug_reports_legacy_signing_and_non_http_transport(_mcp_request_ctx, transport: Literal["http", "stdio"]) -> None: from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from starlette.requests import Request @@ -13287,16 +13247,12 @@ async def test_debug_reports_legacy_signing_and_non_http_transport( async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="temporary-oauth-discovery", - name="temporary", - url="https://idp.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.true_passthrough, + server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, ) manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", ) with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: @@ -13316,18 +13272,13 @@ async def test_temporary_server_discovery_reuses_resolved_metadata_without_publi async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="repeated-stale", - name="stale", - url="https://idp.example.com/mcp", - transport=MCPTransport.http, - auth_type=auth_type, - oauth2_flow="authorization_code", + server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", ) manager.registry[server.server_id] = server manager._set_oauth_discovery_deferred(server.server_id, True) metadata: Final = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", ) with ( patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, @@ -13347,20 +13298,13 @@ async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="resolved-replacement", - name="replacement", - url="https://old.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - oauth2_flow="authorization_code", - ) - replacement: Final = original.model_copy( - update={ - "url": "https://new.example.com/mcp", - "authorization_url": "https://new.example.com/authorize", - "token_url": "https://new.example.com/token", - } + server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", ) + replacement: Final = original.model_copy(update={ + "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + }) manager.registry[original.server_id] = replacement assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement @@ -13368,11 +13312,8 @@ async def test_stale_discovery_falls_back_to_resolved_registered_server() -> Non def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: manager: Final = MCPServerManager() original: Final = MCPServer( - server_id="stale-publication", - name="publication", - url="https://old.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, + server_id="stale-publication", name="publication", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, ) manager._set_oauth_discovery_deferred(original.server_id, True) original_slot: Final = manager._oauth_discovery_slot(original.server_id) @@ -13388,13 +13329,9 @@ def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: manager: Final = MCPServerManager() server: Final = MCPServer( - server_id="expiring-session", - name="temporary", - url="https://idp.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.true_passthrough, - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", + server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", ) manager._set_oauth_discovery_deferred(server.server_id, True) resolved: Final = await manager.ensure_oauth_metadata_discovered(server) @@ -13495,9 +13432,7 @@ async def test_openapi_health_reports_size_limit_as_unknown_and_caches_failure(r result = await manager.health_check_server(server.server_id) cached = await manager.health_check_server(server.server_id) assert result.status == "unknown" - assert ( - result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" - ) + assert result.health_check_error == "OpenAPI specification probe refused: Response exceeds the configured size limit" assert cached.health_check_error == result.health_check_error assert cached.last_health_check == result.last_health_check assert route.call_count == 1 @@ -13509,11 +13444,8 @@ async def test_openapi_health_cancellation_does_not_poison_cache(respx_mock, mon monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") manager = MCPServerManager() server = MCPServer( - server_id="cancelled-cache", - name="cancelled-cache", - transport=MCPTransport.http, - spec_path="https://93.184.216.34/cancelled-cache.json", - auth_type=MCPAuth.none, + server_id="cancelled-cache", name="cancelled-cache", transport=MCPTransport.http, + spec_path="https://93.184.216.34/cancelled-cache.json", auth_type=MCPAuth.none, ) manager.registry = {server.server_id: server} started = asyncio.Event() @@ -13642,9 +13574,7 @@ class _DiscoveryUpstream: def _discovery_server() -> MCPServer: - return MCPServer( - server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http - ) + return MCPServer(server_id="discovery", name="discovery", url="https://discovery.example/mcp", transport=MCPTransport.http) @pytest.mark.asyncio @@ -13812,9 +13742,7 @@ async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) assert upstream.initializes == 2 -@pytest.mark.parametrize( - "value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5)) -) +@pytest.mark.parametrize("value,expected", (("invalid", 60.0), ("nan", 60.0), ("inf", 60.0), ("-1", 60.0), ("12.5", 12.5))) def test_discovery_cache_ttl_validation(value: str, expected: float, monkeypatch: pytest.MonkeyPatch) -> None: from litellm.proxy._experimental.mcp_server.mcp_server_manager import _mcp_discovery_cache_ttl @@ -14077,45 +14005,26 @@ async def test_discovery_cache_returns_oversized_results_without_retaining_them( class TestProtectedCredentialPreparation: @pytest.mark.asyncio - @pytest.mark.parametrize( - "auth_type,credential", - [ - (MCPAuth.bearer_token, None), - (MCPAuth.bearer_token, "Bearer"), - (MCPAuth.api_key, None), - (MCPAuth.basic, "Basic"), - ], - ) + @pytest.mark.parametrize("auth_type,credential", [ + (MCPAuth.bearer_token, None), + (MCPAuth.bearer_token, "Bearer"), + (MCPAuth.api_key, None), + (MCPAuth.basic, "Basic"), + ]) @pytest.mark.parametrize("dispatch", ["managed", "local"]) async def test_openapi_dispatch_rejects_unusable_effective_credentials( - self, - tmp_path: Path, - respx_mock: MockRouter, - monkeypatch: pytest.MonkeyPatch, - auth_type: MCPAuthType, - credential: str | None, - dispatch: str, + self, tmp_path: Path, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, + auth_type: MCPAuthType, credential: str | None, dispatch: str, ) -> None: from litellm.proxy._experimental.mcp_server.server import _handle_local_mcp_tool from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_name, get_server_prefix spec_path: Final = tmp_path / "openapi.json" - spec_path.write_text( - json.dumps( - { - "openapi": "3.0.0", - "info": {"title": "Auth", "version": "1"}, - "paths": {"/echo": {"get": {"operationId": "echo"}}}, - } - ) - ) + spec_path.write_text(json.dumps({"openapi": "3.0.0", "info": {"title": "Auth", "version": "1"}, + "paths": {"/echo": {"get": {"operationId": "echo"}}}})) server: Final = MCPServer( - server_id="dispatch-auth", - name="dispatch-auth", - url="https://upstream.example", - transport=MCPTransport.http, - auth_type=auth_type, - authentication_token=credential, + server_id="dispatch-auth", name="dispatch-auth", url="https://upstream.example", + transport=MCPTransport.http, auth_type=auth_type, authentication_token=credential, ) manager: Final = MCPServerManager() await manager._register_openapi_tools(str(spec_path), server, server.url) @@ -14138,21 +14047,14 @@ class TestProtectedCredentialPreparation: self, transport: MCPTransport, client_secret: str | None, subject: str | None ) -> None: server = MCPServer( - server_id="incomplete-obo", - name="incomplete-obo", - url="https://upstream.example/mcp", - transport=transport, - auth_type=MCPAuth.oauth2_token_exchange, - client_id="gateway", - client_secret=client_secret, - token_exchange_endpoint="https://idp.example/token", - authentication_token="static-fallback", + server_id="incomplete-obo", name="incomplete-obo", url="https://upstream.example/mcp", + transport=transport, auth_type=MCPAuth.oauth2_token_exchange, + client_id="gateway", client_secret=client_secret, + token_exchange_endpoint="https://idp.example/token", authentication_token="static-fallback", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client( - server, - mcp_auth_header="Bearer override", - subject_token=subject, + server, mcp_auth_header="Bearer override", subject_token=subject, ) assert exc.value.status_code == (401 if subject is None else 500) assert "static-fallback" not in str(exc.value.detail) @@ -14165,11 +14067,8 @@ class TestProtectedCredentialPreparation: self, auth_type: MCPAuthType, credential: str | dict[str, str] | None ) -> None: server = MCPServer( - server_id="empty-static", - name="empty-static", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=auth_type, + server_id="empty-static", name="empty-static", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=auth_type, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=credential) @@ -14177,22 +14076,16 @@ class TestProtectedCredentialPreparation: assert "credential" in str(exc.value.detail).lower() @pytest.mark.asyncio - @pytest.mark.parametrize( - "auth_type,headers", - [ - (MCPAuth.api_key, {"X-API-Key": "key"}), - (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), - ], - ) + @pytest.mark.parametrize("auth_type,headers", [ + (MCPAuth.api_key, {"X-API-Key": "key"}), + (MCPAuth.bearer_token, {"Authorization": "Bearer token"}), + ]) async def test_static_auth_accepts_actual_forwarded_credential( self, auth_type: MCPAuthType, headers: dict[str, str] ) -> None: server = MCPServer( - server_id="header-static", - name="header-static", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=auth_type, + server_id="header-static", name="header-static", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=auth_type, ) client = await MCPServerManager()._create_mcp_client(server, extra_headers=headers) assert client._get_auth_headers() == headers @@ -14201,48 +14094,29 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("auth_type", [MCPAuth.oauth2_token_exchange]) async def test_openapi_protected_auth_rejects_missing_credentials(self, auth_type: MCPAuthType) -> None: server = MCPServer( - server_id="openapi-empty", - name="openapi-empty", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=auth_type, + server_id="openapi-empty", name="openapi-empty", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=auth_type, token_exchange_endpoint="https://idp.example/token", ) with pytest.raises(HTTPException) as exc: await MCPServerManager().resolve_openapi_upstream_auth( - mcp_server=server, - oauth2_headers=None, - raw_headers=None, - mcp_auth_header=None, - user_api_key_auth=None, - forwarded_headers=None, + mcp_server=server, oauth2_headers=None, raw_headers=None, mcp_auth_header=None, + user_api_key_auth=None, forwarded_headers=None, ) assert exc.value.status_code in (401, 500) @pytest.mark.asyncio - @pytest.mark.parametrize( - "auth_type,slot,value", - [ - (MCPAuth.api_key, "X-API-Key", "token"), - (MCPAuth.authorization, "Authorization", "opaque-secret-value"), - (MCPAuth.authorization, "Authorization", "Bearer abc"), - (MCPAuth.authorization, "Authorization", "Custom abc"), - ], - ) + @pytest.mark.parametrize("auth_type,slot,value", [ + (MCPAuth.api_key, "X-API-Key", "token"), + (MCPAuth.authorization, "Authorization", "opaque-secret-value"), + (MCPAuth.authorization, "Authorization", "Bearer abc"), + (MCPAuth.authorization, "Authorization", "Custom abc"), + ]) async def test_raw_static_credentials_are_forwarded_unchanged( - self, - auth_type: MCPAuthType, - slot: str, - value: str, + self, auth_type: MCPAuthType, slot: str, value: str, ) -> None: - server = MCPServer( - server_id="raw-key", - name="raw-key", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=auth_type, - authentication_token=value, - ) + server = MCPServer(server_id="raw-key", name="raw-key", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=auth_type, authentication_token=value) client = await MCPServerManager()._create_mcp_client(server) assert client._resolved_auth is not None request = httpx.Request("GET", server.url) @@ -14256,24 +14130,17 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Bearer", "basic", "token", "ApiKey", " bEaReR ", "\tTOKEN\t"]) @pytest.mark.parametrize("source", ["configured", "caller", "forwarded"]) async def test_raw_authorization_rejects_bare_schemes_before_dispatch( - self, - respx_mock: MockRouter, - value: str, - source: str, + self, respx_mock: MockRouter, value: str, source: str, ) -> None: server: Final = MCPServer( - server_id="raw-empty", - name="raw-empty", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.authorization, + server_id="raw-empty", name="raw-empty", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.authorization, authentication_token=value if source == "configured" else None, ) destination: Final = respx_mock.route().respond(200) with pytest.raises(HTTPException, match="requires a usable upstream credential") as exc: await MCPServerManager()._create_mcp_client( - server, - mcp_auth_header=value if source == "caller" else None, + server, mcp_auth_header=value if source == "caller" else None, extra_headers={"Authorization": value} if source == "forwarded" else None, ) assert exc.value.status_code == 500 @@ -14281,15 +14148,9 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_byok_flag_cannot_bypass_incomplete_obo(self) -> None: - server = MCPServer( - server_id="obo-byok", - name="obo-byok", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - is_byok=True, - token_exchange_endpoint="https://idp.example/token", - ) + server = MCPServer(server_id="obo-byok", name="obo-byok", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2_token_exchange, is_byok=True, + token_exchange_endpoint="https://idp.example/token") with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header="Bearer override") assert exc.value.status_code == 401 @@ -14297,66 +14158,41 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("configured,override", [(None, "Bearer usable"), ("shared", "Bearer usable")]) async def test_bearer_override_remains_usable(self, configured: str | None, override: str) -> None: - server = MCPServer( - server_id="override", - name="override", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.bearer_token, - authentication_token=configured, - ) + server = MCPServer(server_id="override", name="override", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=configured) client = await MCPServerManager()._create_mcp_client(server, mcp_auth_header=override) assert client._get_auth_headers()["Authorization"] == override @pytest.mark.asyncio @pytest.mark.parametrize("token", [None, "shared"]) async def test_empty_injected_header_cannot_satisfy_protected_auth(self, token: str | None) -> None: - server = MCPServer( - server_id="empty-header", - name="empty-header", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.bearer_token, - authentication_token=token, - ) + server = MCPServer(server_id="empty-header", name="empty-header", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.bearer_token, authentication_token=token) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"authorization": " "}) assert exc.value.status_code == 500 @pytest.mark.asyncio async def test_custom_slot_uses_its_actual_credential(self) -> None: - server = MCPServer( - server_id="custom", - name="custom", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.api_key, - upstream_token_header="X-Custom", - authentication_token="key", - ) + server = MCPServer(server_id="custom", name="custom", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.api_key, + upstream_token_header="X-Custom", authentication_token="key") client = await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Trace": "trace"}) assert client._credential_slot == "X-Custom" assert await client.discovery_auth_fingerprint() @pytest.mark.asyncio - @pytest.mark.parametrize( - "static_headers,accepted", - [ - ({"apikey": "static-key"}, True), - ({"apikey": ""}, False), - ({"X-Tenant": "tenant"}, True), - ], - ) + @pytest.mark.parametrize("static_headers,accepted", [ + ({"apikey": "static-key"}, True), + ({"apikey": ""}, False), + ({"X-Tenant": "tenant"}, True), + ]) async def test_api_key_carried_by_static_header_passes_fail_closed_check( self, static_headers: dict[str, str], accepted: bool ) -> None: server: Final = MCPServer( - server_id="static-slot", - name="static-slot", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.api_key, - static_headers=static_headers, + server_id="static-slot", name="static-slot", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.api_key, static_headers=static_headers, ) if not accepted: with pytest.raises(HTTPException) as exc: @@ -14368,36 +14204,21 @@ class TestProtectedCredentialPreparation: assert all(request.headers[name] == value for name, value in static_headers.items()) @pytest.mark.asyncio - @pytest.mark.parametrize( - "static,forwarded,caller", - [ - ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), - ({}, {"X-API-Key": "forwarded"}, None), - ({}, None, "ApiKey caller"), - ({"X-API-Key": "static"}, {"Authorization": ""}, None), - ], - ) + @pytest.mark.parametrize("static,forwarded,caller", [ + ({"X-API-Key": "static"}, {"x-api-key": "forwarded"}, None), + ({}, {"X-API-Key": "forwarded"}, None), + ({}, None, "ApiKey caller"), + ({"X-API-Key": "static"}, {"Authorization": ""}, None), + ]) async def test_openapi_static_credentials_remain_supported( - self, - respx_mock: MockRouter, - monkeypatch: pytest.MonkeyPatch, - static: dict[str, str], - forwarded: dict[str, str] | None, - caller: str | None, + self, respx_mock: MockRouter, monkeypatch: pytest.MonkeyPatch, + static: dict[str, str], forwarded: dict[str, str] | None, caller: str | None ) -> None: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - create_tool_function, + _request_auth_header, _request_extra_headers, create_tool_function, ) - tool: Final = create_tool_function( - "/echo", - "get", - {}, - "https://upstream.example", - headers=static, - auth_type=MCPAuth.api_key, + "/echo", "get", {}, "https://upstream.example", headers=static, auth_type=MCPAuth.api_key, ) monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") destination: Final = respx_mock.get("https://upstream.example/echo").respond(200, text="authenticated") @@ -14431,13 +14252,8 @@ class TestProtectedCredentialPreparation: self.closed = True auth = CancelledAuth() - server = MCPServer( - server_id="cancel", - name="cancel", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.api_key, - ) + server = MCPServer(server_id="cancel", name="cancel", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.api_key) client = MCPClient(server_url=server.url, auth_type=MCPAuth.api_key, resolved_auth=auth) with pytest.raises(asyncio.CancelledError): await prepare_mcp_client(server, client) @@ -14446,14 +14262,8 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("auth_type", [MCPAuth.basic, MCPAuth.token, MCPAuth.authorization]) async def test_other_static_schemes_reject_whitespace_credentials(self, auth_type: MCPAuthType) -> None: - server = MCPServer( - server_id="blank-static", - name="blank-static", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=auth_type, - authentication_token=" ", - ) + server = MCPServer(server_id="blank-static", name="blank-static", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=auth_type, authentication_token=" ") with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server) assert exc.value.status_code == 500 @@ -14461,13 +14271,8 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("header", ["Basic", "Basic @@@", "Other abc", "Basic QmFzaWM=", "Basic bm8tY29sb24="]) async def test_basic_headers_without_usable_credentials_reject(self, header: str) -> None: - server = MCPServer( - server_id="bad-basic", - name="bad-basic", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.basic, - ) + server = MCPServer(server_id="bad-basic", name="bad-basic", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.basic) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"Authorization": header}) assert exc.value.status_code == 500 @@ -14476,48 +14281,34 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("value", ["Basic", "Basic ", "basic"]) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_scheme_alone_is_not_a_credential(self, value: str, source: str) -> None: - server = MCPServer( - server_id="basic-scheme", - name="basic-scheme", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.basic, - authentication_token=value if source == "configured" else None, - ) + server = MCPServer(server_id="basic-scheme", name="basic-scheme", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.basic, + authentication_token=value if source == "configured" else None) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header=value if source == "caller" else None) assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize( - "auth_type,value,default_slot", - [ - (MCPAuth.api_key, "fixture-key", "X-API-Key"), - (MCPAuth.bearer_token, "fixture-key", "Authorization"), - (MCPAuth.basic, "user:pass", "Authorization"), - (MCPAuth.token, "fixture-key", "Authorization"), - (MCPAuth.authorization, "fixture-key", "Authorization"), - ], - ) + @pytest.mark.parametrize("auth_type,value,default_slot", [ + (MCPAuth.api_key, "fixture-key", "X-API-Key"), + (MCPAuth.bearer_token, "fixture-key", "Authorization"), + (MCPAuth.basic, "user:pass", "Authorization"), + (MCPAuth.token, "fixture-key", "Authorization"), + (MCPAuth.authorization, "fixture-key", "Authorization"), + ]) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_usable_credential_survives_an_empty_alternate_header( self, auth_type: MCPAuthType, value: str, default_slot: str, source: str ) -> None: server: Final = MCPServer( - server_id="alternate", - name="alternate", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=auth_type, - upstream_token_header="X-Custom", + server_id="alternate", name="alternate", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=auth_type, upstream_token_header="X-Custom", authentication_token=value if source == "configured" else None, ) empty_slot: Final = default_slot if source == "configured" else "X-Custom" selected_slot: Final = "X-Custom" if source == "configured" else default_slot client: Final = await MCPServerManager()._create_mcp_client( - server, - mcp_auth_header=value if source == "caller" else None, - extra_headers={empty_slot: ""}, + server, mcp_auth_header=value if source == "caller" else None, extra_headers={empty_slot: ""}, ) request: Final = await client.prepare_request_auth() assert request.headers[selected_slot] @@ -14526,12 +14317,8 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_empty_custom_and_default_headers_do_not_satisfy_auth(self) -> None: server: Final = MCPServer( - server_id="both-empty", - name="both-empty", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.api_key, - upstream_token_header="X-Custom", + server_id="both-empty", name="both-empty", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header="X-Custom", ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, extra_headers={"X-Custom": "", "X-API-Key": ""}) @@ -14544,17 +14331,12 @@ class TestProtectedCredentialPreparation: self, custom_slot: str | None, source: str ) -> None: server: Final = MCPServer( - server_id="caller-auth", - name="caller-auth", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.api_key, - upstream_token_header=custom_slot, + server_id="caller-auth", name="caller-auth", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.api_key, upstream_token_header=custom_slot, ) headers: Final = {"Authorization": "Bearer caller-credential", "X-API-Key": ""} client: Final = await MCPServerManager()._create_mcp_client( - server, - mcp_auth_header=headers if source == "caller" else None, + server, mcp_auth_header=headers if source == "caller" else None, extra_headers=headers if source == "forwarded" else None, ) request: Final = await client.prepare_request_auth() @@ -14563,29 +14345,14 @@ class TestProtectedCredentialPreparation: assert custom_slot is None or custom_slot not in request.headers @pytest.mark.asyncio - @pytest.mark.parametrize( - "value", - [ - "", - " ", - "Bearer", - "Basic", - "token", - "ApiKey", - "Bearer Bearer", - "ApiKey ApiKey", - "token token", - "bEaReR BEARER", - "aPiKeY\tAPIKEY", - ], - ) + @pytest.mark.parametrize("value", [ + "", " ", "Bearer", "Basic", "token", "ApiKey", + "Bearer Bearer", "ApiKey ApiKey", "token token", "bEaReR BEARER", "aPiKeY\tAPIKEY", + ]) async def test_api_key_rejects_authorization_without_a_credential(self, value: str) -> None: server: Final = MCPServer( - server_id="caller-empty", - name="caller-empty", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.api_key, + server_id="caller-empty", name="caller-empty", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.api_key, ) with pytest.raises(HTTPException) as exc: await MCPServerManager()._create_mcp_client(server, mcp_auth_header={"Authorization": value}) @@ -14596,11 +14363,8 @@ class TestProtectedCredentialPreparation: @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_basic_requires_a_username_password_separator(self, value: str, source: str) -> None: server: Final = MCPServer( - server_id="basic-pair", - name="basic-pair", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.basic, + server_id="basic-pair", name="basic-pair", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14613,12 +14377,8 @@ class TestProtectedCredentialPreparation: import base64 server: Final = MCPServer( - server_id="basic-valid", - name="basic-valid", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.basic, - authentication_token=value, + server_id="basic-valid", name="basic-valid", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.basic, authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -14627,27 +14387,17 @@ class TestProtectedCredentialPreparation: assert base64.b64decode(encoded) == value.encode() @pytest.mark.asyncio - @pytest.mark.parametrize( - "auth_type,value", - [ - (MCPAuth.bearer_token, "Bearer"), - (MCPAuth.bearer_token, "Bearer "), - (MCPAuth.bearer_token, "bearer"), - (MCPAuth.token, "token"), - (MCPAuth.token, "token "), - (MCPAuth.token, "TOKEN"), - ], - ) + @pytest.mark.parametrize("auth_type,value", [ + (MCPAuth.bearer_token, "Bearer"), (MCPAuth.bearer_token, "Bearer "), (MCPAuth.bearer_token, "bearer"), + (MCPAuth.token, "token"), (MCPAuth.token, "token "), (MCPAuth.token, "TOKEN"), + ]) @pytest.mark.parametrize("source", ["configured", "caller"]) async def test_static_scheme_only_input_cannot_hide_behind_rendered_prefix( self, auth_type: MCPAuthType, value: str, source: str ) -> None: server: Final = MCPServer( - server_id="empty-scheme", - name="empty-scheme", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=auth_type, + server_id="empty-scheme", name="empty-scheme", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=auth_type, authentication_token=value if source == "configured" else None, ) with pytest.raises(HTTPException) as exc: @@ -14655,24 +14405,17 @@ class TestProtectedCredentialPreparation: assert exc.value.status_code == 500 @pytest.mark.asyncio - @pytest.mark.parametrize( - "auth_type,value,expected", - [ - (MCPAuth.bearer_token, "token", "Bearer token"), - (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), - (MCPAuth.token, "tokenish", "token tokenish"), - ], - ) + @pytest.mark.parametrize("auth_type,value,expected", [ + (MCPAuth.bearer_token, "token", "Bearer token"), + (MCPAuth.bearer_token, "Bearertoken", "Bearer Bearertoken"), + (MCPAuth.token, "tokenish", "token tokenish"), + ]) async def test_static_credentials_that_resemble_schemes_remain_usable( self, auth_type: MCPAuthType, value: str, expected: str ) -> None: server: Final = MCPServer( - server_id="real-token", - name="real-token", - url="https://upstream.example/mcp", - transport=MCPTransport.http, - auth_type=auth_type, - authentication_token=value, + server_id="real-token", name="real-token", url="https://upstream.example/mcp", + transport=MCPTransport.http, auth_type=auth_type, authentication_token=value, ) client: Final = await MCPServerManager()._create_mcp_client(server) request: Final = await client.prepare_request_auth() @@ -14711,31 +14454,16 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon registry.register_tool("observer-execute", "Execute", {"type": "object"}, upstream) monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) manager = MCPServerManager() - manager.registry = { - "observer": MCPServer( - server_id="observer", - name="observer", - server_name="observer", - transport="http", - url="https://observer.example/mcp", - spec_path="observer.json", - auth_type="none", - ) - } + manager.registry = {"observer": MCPServer( + server_id="observer", name="observer", server_name="observer", transport="http", + url="https://observer.example/mcp", spec_path="observer.json", auth_type="none", + )} manager.tool_name_to_mcp_server_name_mapping = {"observer-execute": "observer"} - result = await asyncio.wait_for( - manager.call_tool( - server_name="observer", - name="execute", - arguments={"text": "hello"}, - user_api_key_auth=UserAPIKeyAuth(), - proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), - guardrail_context=MCPRequestContext.resolve_guardrail_context( - {"metadata": {"guardrails": ["observe"] if selected else []}} - ), - ), - timeout=5, - ) + result = await asyncio.wait_for(manager.call_tool( + server_name="observer", name="execute", arguments={"text": "hello"}, + user_api_key_auth=UserAPIKeyAuth(), proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + guardrail_context=MCPRequestContext.resolve_guardrail_context({"metadata": {"guardrails": ["observe"] if selected else []}}), + ), timeout=5) assert tool_started.is_set() assert guardrail_started.is_set() is selected assert result.is_error is False @@ -14764,21 +14492,11 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie from litellm.proxy._experimental.mcp_server import server as legacy_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback - upstream = MCPServer( - server_id="explicit-empty", - name="explicit_empty", - url="https://example.invalid/mcp", - transport=MCPTransport.http, - allow_sampling=True, - ) + upstream = MCPServer(server_id="explicit-empty", name="explicit_empty", url="https://example.invalid/mcp", transport=MCPTransport.http, allow_sampling=True) token = auth_context_var.set(None) sampling = AsyncMock() try: - legacy_server.set_auth_context( - UserAPIKeyAuth(user_id="unrelated"), - raw_headers={"authorization": "unrelated-credential"}, - client_ip="192.0.2.99", - ) + legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99") with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory, patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), @@ -14786,9 +14504,7 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie if legacy_factory: callback = _create_sampling_callback(user_api_key_auth=UserAPIKeyAuth(user_id="explicit")) else: - await MCPServerManager()._create_mcp_client( - upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None - ) + await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) callback = factory.call_args.kwargs["sampling_callback"] await callback(None, None) captured = sampling.await_args.kwargs @@ -14811,28 +14527,16 @@ class TestSharedIdentifierPrefixWarning: manager = MCPServerManager() rows = [ LiteLLM_MCPServerTable( - server_id="srv-a", - server_name="alpha", - alias="shared", - url="https://a.example.com/mcp", - transport=MCPTransport.http, - updated_at=datetime.now(), + server_id="srv-a", server_name="alpha", alias="shared", url="https://a.example.com/mcp", + transport=MCPTransport.http, updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-b", - server_name="beta", - alias="Shared", - url="https://b.example.com/mcp", - transport=MCPTransport.http, - updated_at=datetime.now(), + server_id="srv-b", server_name="beta", alias="Shared", url="https://b.example.com/mcp", + transport=MCPTransport.http, updated_at=datetime.now(), ), LiteLLM_MCPServerTable( - server_id="srv-c", - server_name="gamma", - alias="lonely", - url="https://c.example.com/mcp", - transport=MCPTransport.http, - updated_at=datetime.now(), + server_id="srv-c", server_name="gamma", alias="lonely", url="https://c.example.com/mcp", + transport=MCPTransport.http, updated_at=datetime.now(), ), ] raw_rows = [MagicMock(model_dump=lambda row=row: row.model_dump()) for row in rows] diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index b3e0795d070..5c34bd66b17 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -106,7 +106,9 @@ async def test_map_user_to_teams_handles_already_in_team_exception(): ) as mock_add: with patch("litellm.proxy.auth.handle_jwt.verbose_proxy_logger") as mock_logger: # This should not raise an exception - result = await JWTAuthManager.map_user_to_teams(user_object=user, team_object=team) + result = await JWTAuthManager.map_user_to_teams( + user_object=user, team_object=team + ) # Verify the method completed successfully assert result is None @@ -143,10 +145,14 @@ async def test_map_user_to_teams_reraises_other_proxy_exceptions(): async def test_map_user_to_teams_null_inputs(): """Test that method handles null inputs gracefully""" # Test with null user - await JWTAuthManager.map_user_to_teams(user_object=None, team_object=LiteLLM_TeamTable(team_id="test_team_1")) + await JWTAuthManager.map_user_to_teams( + user_object=None, team_object=LiteLLM_TeamTable(team_id="test_team_1") + ) # Test with null team - await JWTAuthManager.map_user_to_teams(user_object=LiteLLM_UserTable(user_id="test_user_1"), team_object=None) + await JWTAuthManager.map_user_to_teams( + user_object=LiteLLM_UserTable(user_id="test_user_1"), team_object=None + ) # Test with both null await JWTAuthManager.map_user_to_teams(user_object=None, team_object=None) @@ -202,7 +208,9 @@ async def test_find_team_with_model_access_reports_passthrough_allowlist_denial( assert exc_info.value.status_code == 403 assert "allowed_passthrough_routes" in exc_info.value.detail assert "requested model" not in exc_info.value.detail - mock_is_auth_enforced_pass_through_route.assert_called_once_with(route="/my-pass-through", method="POST") + mock_is_auth_enforced_pass_through_route.assert_called_once_with( + route="/my-pass-through", method="POST" + ) user_api_key_dict = mock_passthrough_check.call_args.kwargs["user_api_key_dict"] assert user_api_key_dict.metadata == {} @@ -395,7 +403,9 @@ async def test_auth_builder_proxy_admin_user_role(): route = "/chat/completions" # Create user object with PROXY_ADMIN role - user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.PROXY_ADMIN) + user_object = LiteLLM_UserTable( + user_id="test_user_1", user_role=LitellmUserRoles.PROXY_ADMIN + ) # Create mock JWT handler jwt_handler = JWTHandler() @@ -404,10 +414,14 @@ async def test_auth_builder_proxy_admin_user_role(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, + patch.object( + JWTAuthManager, "check_rbac_role", new_callable=AsyncMock + ) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, + patch.object( + jwt_handler, "get_object_id", return_value=None + ) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -415,7 +429,9 @@ async def test_auth_builder_proxy_admin_user_role(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, + patch.object( + jwt_handler, "get_end_user_id", return_value=None + ) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -428,7 +444,9 @@ async def test_auth_builder_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, + patch.object( + JWTAuthManager, "get_all_team_ids", return_value=set() + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -441,8 +459,12 @@ async def test_auth_builder_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, - patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object( + JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock + ) as mock_map_user, + patch.object( + JWTAuthManager, "validate_object_id", return_value=True + ) as mock_validate_object, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -476,7 +498,9 @@ async def test_auth_builder_non_proxy_admin_user_role(): route = "/chat/completions" # Create user object with regular USER role - user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) + user_object = LiteLLM_UserTable( + user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER + ) # Create mock JWT handler jwt_handler = JWTHandler() @@ -485,10 +509,14 @@ async def test_auth_builder_non_proxy_admin_user_role(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, + patch.object( + JWTAuthManager, "check_rbac_role", new_callable=AsyncMock + ) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, + patch.object( + jwt_handler, "get_object_id", return_value=None + ) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -496,7 +524,9 @@ async def test_auth_builder_non_proxy_admin_user_role(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, + patch.object( + jwt_handler, "get_end_user_id", return_value=None + ) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -509,7 +539,9 @@ async def test_auth_builder_non_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, + patch.object( + JWTAuthManager, "get_all_team_ids", return_value=set() + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -522,8 +554,12 @@ async def test_auth_builder_non_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, - patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object( + JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock + ) as mock_map_user, + patch.object( + JWTAuthManager, "validate_object_id", return_value=True + ) as mock_validate_object, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -672,7 +708,11 @@ async def test_sync_user_role_and_teams(): prisma_client=None, user_api_key_cache=mock_user_api_key_cache, litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], + jwt_litellm_role_map=[ + JWTLiteLLMRoleMap( + jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN + ) + ], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -681,7 +721,9 @@ async def test_sync_user_role_and_teams(): token = {"roles": ["ADMIN"], "my_id_teams": ["team1", "team2"]} - user = LiteLLM_UserTable(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER.value, teams=["team2"]) + user = LiteLLM_UserTable( + user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER.value, teams=["team2"] + ) prisma = AsyncMock() prisma.db.litellm_usertable.update = AsyncMock() @@ -708,7 +750,11 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], + jwt_litellm_role_map=[ + JWTLiteLLMRoleMap( + jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN + ) + ], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -725,7 +771,9 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): prisma = AsyncMock() prisma.db.litellm_usertable.update = AsyncMock() - await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) + await JWTAuthManager.sync_user_role_and_teams( + jwt_handler, token, user, prisma, user_api_key_cache=mock_cache + ) mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args @@ -745,7 +793,11 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], + jwt_litellm_role_map=[ + JWTLiteLLMRoleMap( + jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN + ) + ], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -766,7 +818,9 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ): - await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) + await JWTAuthManager.sync_user_role_and_teams( + jwt_handler, token, user, prisma, user_api_key_cache=mock_cache + ) mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args @@ -786,7 +840,11 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], + jwt_litellm_role_map=[ + JWTLiteLLMRoleMap( + jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN + ) + ], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -802,7 +860,9 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): prisma = AsyncMock() - await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) + await JWTAuthManager.sync_user_role_and_teams( + jwt_handler, token, user, prisma, user_api_key_cache=mock_cache + ) mock_cache.async_set_cache.assert_not_called() @@ -827,7 +887,9 @@ def test_get_all_jwt_team_ids_unions_singular_and_plural(): assert jwt_handler.get_all_jwt_team_ids({"teams": ["a", "b"]}) == ["a", "b"] # both populated, no overlap - assert jwt_handler.get_all_jwt_team_ids({"team_id": "primary", "teams": ["a", "b"]}) == ["a", "b", "primary"] + assert jwt_handler.get_all_jwt_team_ids( + {"team_id": "primary", "teams": ["a", "b"]} + ) == ["a", "b", "primary"] # both populated with overlap — singular dedup'd assert jwt_handler.get_all_jwt_team_ids({"team_id": "a", "teams": ["a", "b"]}) == [ @@ -836,7 +898,9 @@ def test_get_all_jwt_team_ids_unions_singular_and_plural(): ] # singular field as multi-element list (some IdPs) — merge all, preserve plural-first order - assert jwt_handler.get_all_jwt_team_ids({"team_id": ["primary", "secondary"], "teams": ["a"]}) == [ + assert jwt_handler.get_all_jwt_team_ids( + {"team_id": ["primary", "secondary"], "teams": ["a"]} + ) == [ "a", "primary", "secondary", @@ -896,11 +960,19 @@ async def test_map_jwt_role_to_litellm_role(): litellm_jwtauth=LiteLLM_JWTAuth( jwt_litellm_role_map=[ # Exact match - JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN), + JWTLiteLLMRoleMap( + jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN + ), # Wildcard patterns - JWTLiteLLMRoleMap(jwt_role="user_*", litellm_role=LitellmUserRoles.INTERNAL_USER), - JWTLiteLLMRoleMap(jwt_role="team_?", litellm_role=LitellmUserRoles.TEAM), - JWTLiteLLMRoleMap(jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER), + JWTLiteLLMRoleMap( + jwt_role="user_*", litellm_role=LitellmUserRoles.INTERNAL_USER + ), + JWTLiteLLMRoleMap( + jwt_role="team_?", litellm_role=LitellmUserRoles.TEAM + ), + JWTLiteLLMRoleMap( + jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER + ), ], roles_jwt_field="roles", ), @@ -972,7 +1044,9 @@ async def test_map_jwt_role_to_litellm_role(): # Test patterns that don't match character classes jwt_handler.litellm_jwtauth.jwt_litellm_role_map = [ - JWTLiteLLMRoleMap(jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER), + JWTLiteLLMRoleMap( + jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER + ), ] token = {"roles": ["dev_4"]} # 4 is not in [123] result = jwt_handler.map_jwt_role_to_litellm_role(token) @@ -1067,19 +1141,25 @@ async def test_nested_jwt_field_access(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="profile.object_id", - role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], + role_mappings=[ + RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) + ], ) assert jwt_handler.get_object_id(nested_token, None) == "obj789" # Test 5b: object_id_jwt_field with flat access (backward compatibility) jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="object_id", - role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], + role_mappings=[ + RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) + ], ) assert jwt_handler.get_object_id(flat_token, None) == "obj789" # Test 6: end_user_id_jwt_field with nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + end_user_id_jwt_field="customer.end_user_id" + ) assert jwt_handler.get_end_user_id(nested_token, None) == "customer123" # Test 6b: end_user_id_jwt_field with flat access (backward compatibility) @@ -1095,7 +1175,9 @@ async def test_nested_jwt_field_access(): assert jwt_handler.get_team_id(flat_token, None) == "team456" # Test 8: roles_jwt_field with deeply nested access (already supported, but testing) - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + roles_jwt_field="resource_access.my-client.roles" + ) assert jwt_handler.get_jwt_role(nested_token, []) == ["admin", "user"] # Test 9: user_roles_jwt_field with nested access (already supported, but testing) @@ -1141,7 +1223,10 @@ async def test_nested_jwt_field_missing_paths(): # Test 2: Missing user.email should return default jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="user.email") - assert jwt_handler.get_user_email(incomplete_token, "default@example.com") == "default@example.com" + assert ( + jwt_handler.get_user_email(incomplete_token, "default@example.com") + == "default@example.com" + ) # Test 3: Missing groups should return empty list jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups") @@ -1156,28 +1241,43 @@ async def test_nested_jwt_field_missing_paths(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="profile.object_id", - role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], + role_mappings=[ + RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) + ], ) assert jwt_handler.get_object_id(incomplete_token, "default_obj") == "default_obj" # Test 6: Missing customer.end_user_id should return default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") - assert jwt_handler.get_end_user_id(incomplete_token, "default_customer") == "default_customer" + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + end_user_id_jwt_field="customer.end_user_id" + ) + assert ( + jwt_handler.get_end_user_id(incomplete_token, "default_customer") + == "default_customer" + ) # Test 7: Missing tenant.team_id should use team_id_default fallback - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="tenant.team_id", team_id_default="fallback_team") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_id_jwt_field="tenant.team_id", team_id_default="fallback_team" + ) assert jwt_handler.get_team_id(incomplete_token, "default_team") == "fallback_team" # Test 8: Missing resource_access.my-client.roles should return default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") - assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == ["default_role"] + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + roles_jwt_field="resource_access.my-client.roles" + ) + assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == [ + "default_role" + ] # Test 9: Missing nested user roles should return default jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( user_roles_jwt_field="resource_access.my-client.roles", user_allowed_roles=["admin", "user"], ) - assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == ["default_user_role"] + assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == [ + "default_user_role" + ] @pytest.mark.asyncio @@ -1202,7 +1302,9 @@ async def test_metadata_prefix_handling_in_nested_fields(): } # Test 1: metadata.user.email should access user.email after prefix removal - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="metadata.user.email") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_email_jwt_field="metadata.user.email" + ) # The get_nested_value function removes "metadata." prefix, so "metadata.user.email" becomes "user.email" assert jwt_handler.get_user_email(token, None) == "user@example.com" @@ -1238,7 +1340,9 @@ async def test_find_team_with_model_access_model_group(monkeypatch): async def mock_get_team_object(*args, **kwargs): # type: ignore return team - monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -1346,7 +1450,9 @@ async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatc async def mock_get_team_object(*args, **kwargs): # type: ignore return team - monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -1412,10 +1518,14 @@ async def test_auth_builder_returns_team_membership_object(): team_id=_team_id, budget_id="budget_123", spend=10.5, - litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget_123", rpm_limit=100, tpm_limit=5000), + litellm_budget_table=LiteLLM_BudgetTable( + budget_id="budget_123", rpm_limit=100, tpm_limit=5000 + ), ) - user_object = LiteLLM_UserTable(user_id=_user_id, user_role=LitellmUserRoles.INTERNAL_USER) + user_object = LiteLLM_UserTable( + user_id=_user_id, user_role=LitellmUserRoles.INTERNAL_USER + ) team_object = LiteLLM_TeamTable(team_id=_team_id) @@ -1426,10 +1536,14 @@ async def test_auth_builder_returns_team_membership_object(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, + patch.object( + JWTAuthManager, "check_rbac_role", new_callable=AsyncMock + ) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, + patch.object( + jwt_handler, "get_object_id", return_value=None + ) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1437,7 +1551,9 @@ async def test_auth_builder_returns_team_membership_object(): return_value=(_user_id, "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, + patch.object( + jwt_handler, "get_end_user_id", return_value=None + ) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1450,7 +1566,9 @@ async def test_auth_builder_returns_team_membership_object(): new_callable=AsyncMock, return_value=(_team_id, team_object), ) as mock_find_team, - patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, + patch.object( + JWTAuthManager, "get_all_team_ids", return_value=set() + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1469,9 +1587,15 @@ async def test_auth_builder_returns_team_membership_object(): user_object.user_id, ), ) as mock_get_objects, - patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, - patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, + patch.object( + JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock + ) as mock_map_user, + patch.object( + JWTAuthManager, "validate_object_id", return_value=True + ) as mock_validate_object, + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ) as mock_sync_user, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": _user_id, "scope": ""} @@ -1490,12 +1614,24 @@ async def test_auth_builder_returns_team_membership_object(): ) # Verify that team_membership_object is returned - assert result["team_membership"] is not None, "team_membership should be present" - assert result["team_membership"] == mock_team_membership, "team_membership should match the mock object" - assert result["team_membership"].user_id == _user_id, "team_membership user_id should match" - assert result["team_membership"].team_id == _team_id, "team_membership team_id should match" - assert result["team_membership"].budget_id == "budget_123", "team_membership budget_id should match" - assert result["team_membership"].spend == 10.5, "team_membership spend should match" + assert result["team_membership"] is not None, ( + "team_membership should be present" + ) + assert result["team_membership"] == mock_team_membership, ( + "team_membership should match the mock object" + ) + assert result["team_membership"].user_id == _user_id, ( + "team_membership user_id should match" + ) + assert result["team_membership"].team_id == _team_id, ( + "team_membership team_id should match" + ) + assert result["team_membership"].budget_id == "budget_123", ( + "team_membership budget_id should match" + ) + assert result["team_membership"].spend == 10.5, ( + "team_membership spend should match" + ) @pytest.mark.asyncio @@ -1512,7 +1648,9 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) + user_object = LiteLLM_UserTable( + user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER + ) # Create JWT handler with OIDC UserInfo enabled jwt_handler = JWTHandler() @@ -1539,12 +1677,18 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): # Mock all the dependencies with ( - patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, + patch.object( + jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock + ) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, + patch.object( + JWTAuthManager, "check_rbac_role", new_callable=AsyncMock + ) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, + patch.object( + jwt_handler, "get_object_id", return_value=None + ) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1552,7 +1696,9 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, + patch.object( + jwt_handler, "get_end_user_id", return_value=None + ) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1565,7 +1711,9 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, + patch.object( + JWTAuthManager, "get_all_team_ids", return_value=set() + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1578,9 +1726,15 @@ async def test_auth_builder_with_oidc_userinfo_enabled(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, - patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, + patch.object( + JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock + ) as mock_map_user, + patch.object( + JWTAuthManager, "validate_object_id", return_value=True + ) as mock_validate_object, + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ) as mock_sync_user, ): # Set up mock return values mock_get_userinfo.return_value = userinfo_response @@ -1621,7 +1775,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) + user_object = LiteLLM_UserTable( + user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER + ) # Create JWT handler with OIDC UserInfo disabled jwt_handler = JWTHandler() @@ -1645,12 +1801,18 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): # Mock all the dependencies with ( - patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, + patch.object( + jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock + ) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, + patch.object( + JWTAuthManager, "check_rbac_role", new_callable=AsyncMock + ) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, + patch.object( + jwt_handler, "get_object_id", return_value=None + ) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1658,7 +1820,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): return_value=("test_user_1", None, None), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, + patch.object( + jwt_handler, "get_end_user_id", return_value=None + ) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1671,7 +1835,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, + patch.object( + JWTAuthManager, "get_all_team_ids", return_value=set() + ) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1684,9 +1850,15 @@ async def test_auth_builder_with_oidc_userinfo_disabled(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, - patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, + patch.object( + JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock + ) as mock_map_user, + patch.object( + JWTAuthManager, "validate_object_id", return_value=True + ) as mock_validate_object, + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ) as mock_sync_user, ): # Set up mock return values mock_auth_jwt.return_value = jwt_response @@ -1732,7 +1904,9 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) + user_object = LiteLLM_UserTable( + user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER + ) jwt_handler = JWTHandler() user_api_key_cache = DualCache() @@ -1751,7 +1925,9 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() jwt_response = {"sub": "test_user_1", "scope": ""} with ( - patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, + patch.object( + jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock + ) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), @@ -1792,7 +1968,9 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), ): mock_auth_jwt.return_value = jwt_response @@ -1874,7 +2052,9 @@ async def test_auth_builder_uses_team_from_header_e2e(): ) team_object = LiteLLM_TeamTable(team_id="team-2") - user_object = LiteLLM_UserTable(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER) + user_object = LiteLLM_UserTable( + user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER + ) with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, @@ -1885,7 +2065,9 @@ async def test_auth_builder_uses_team_from_header_e2e(): new_callable=AsyncMock, return_value=None, ), - patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_team, patch.object( JWTAuthManager, "get_objects", @@ -1893,7 +2075,9 @@ async def test_auth_builder_uses_team_from_header_e2e(): return_value=(user_object, None, None, None, user_object.user_id), ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), ): mock_auth_jwt.return_value = { "sub": "user-1", @@ -2110,7 +2294,9 @@ async def test_auth_builder_rbac_team_loads_team_for_passthrough_allowlist(): return_value=(None, None, None, None, None), ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.RouteChecks.is_auth_enforced_pass_through_route", return_value=True, @@ -2139,7 +2325,9 @@ async def test_auth_builder_rbac_team_loads_team_for_passthrough_allowlist(): mock_get_team.assert_awaited_once() assert mock_get_team.await_args.kwargs["team_id"] == "team-rbac" user_api_key_dict = mock_passthrough_check.call_args.kwargs["user_api_key_dict"] - assert user_api_key_dict.team_metadata == {"allowed_passthrough_routes": ["/my-pass-through"]} + assert user_api_key_dict.team_metadata == { + "allowed_passthrough_routes": ["/my-pass-through"] + } @pytest.mark.asyncio @@ -2233,7 +2421,9 @@ async def test_auth_builder_admin_on_llm_route_honors_team_header(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2285,7 +2475,9 @@ async def test_auth_builder_admin_on_mgmt_route_ignores_team_header(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2339,7 +2531,9 @@ async def test_auth_builder_admin_on_llm_route_without_header_unchanged(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2383,7 +2577,9 @@ async def test_get_team_alias_with_nested_fields(): } # Test nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_alias_jwt_field="organization.team.name") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_alias_jwt_field="organization.team.name" + ) assert jwt_handler.get_team_alias(nested_token, None) == "engineering-team" # Test flat access (backward compatibility) @@ -2391,7 +2587,9 @@ async def test_get_team_alias_with_nested_fields(): assert jwt_handler.get_team_alias(nested_token, None) == "flat-team" # Test missing field returns default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_alias_jwt_field="nonexistent.field") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_alias_jwt_field="nonexistent.field" + ) assert jwt_handler.get_team_alias(nested_token, "default-team") == "default-team" # Test with team_alias_jwt_field not configured @@ -2422,7 +2620,9 @@ async def test_is_required_team_id_with_team_alias_field(): assert jwt_handler.is_required_team_id() is True # Both fields set - should return True - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_name") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_id_jwt_field="team_id", team_alias_jwt_field="team_name" + ) assert jwt_handler.is_required_team_id() is True @@ -2454,7 +2654,9 @@ async def test_find_and_validate_specific_team_id_with_team_alias(): # Mock team object returned by get_team_object_by_alias team_object = LiteLLM_TeamTable(team_id="resolved-team-id", team_alias="my-team") - with patch("litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias: + with patch( + "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock + ) as mock_get_by_alias: mock_get_by_alias.return_value = team_object team_id, result_team = await JWTAuthManager.find_and_validate_specific_team_id( @@ -2497,7 +2699,9 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): jwt_handler.update_environment( prisma_client=None, user_api_key_cache=user_api_key_cache, - litellm_jwtauth=LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"), + litellm_jwtauth=LiteLLM_JWTAuth( + team_id_jwt_field="team_id", team_alias_jwt_field="team_alias" + ), ) # Token with both team_id and team name @@ -2507,7 +2711,9 @@ async def test_find_and_validate_team_id_takes_precedence_over_name(): team_object = LiteLLM_TeamTable(team_id="direct-team-id") with ( - patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -2587,7 +2793,9 @@ async def test_get_org_alias_with_nested_fields(): } # Test nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_alias_jwt_field="company.organization.name") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + org_alias_jwt_field="company.organization.name" + ) assert jwt_handler.get_org_alias(nested_token, None) == "acme-corp" # Test flat access @@ -2595,7 +2803,9 @@ async def test_get_org_alias_with_nested_fields(): assert jwt_handler.get_org_alias(nested_token, None) == "flat-org" # Test missing field returns default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_alias_jwt_field="nonexistent.field") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + org_alias_jwt_field="nonexistent.field" + ) assert jwt_handler.get_org_alias(nested_token, "default-org") == "default-org" # Test with org_alias_jwt_field not configured @@ -2633,7 +2843,9 @@ async def test_get_objects_resolves_org_by_name(): models=[], ) - with patch("litellm.proxy.auth.handle_jwt.get_org_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias: + with patch( + "litellm.proxy.auth.handle_jwt.get_org_object_by_alias", new_callable=AsyncMock + ) as mock_get_by_alias: mock_get_by_alias.return_value = org_object ( @@ -2711,7 +2923,9 @@ async def test_resolve_jwks_url_resolves_oidc_discovery_document(): litellm_jwtauth=LiteLLM_JWTAuth(), ) - discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" + discovery_url = ( + "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" + ) jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" mock_response = MagicMock() @@ -2742,7 +2956,9 @@ async def test_resolve_jwks_url_caches_resolved_jwks_uri(): litellm_jwtauth=LiteLLM_JWTAuth(), ) - discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" + discovery_url = ( + "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" + ) jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" mock_response = MagicMock() @@ -2886,7 +3102,9 @@ async def test_find_and_validate_specific_team_id_hints_bracket_notation(): error_msg = str(exc_info.value) # Should mention the bad field name and suggest the fix assert "roles.0" in error_msg, f"Expected field name in: {error_msg}" - assert "roles" in error_msg and "list" in error_msg, f"Expected hint about using 'roles' instead: {error_msg}" + assert "roles" in error_msg and "list" in error_msg, ( + f"Expected hint about using 'roles' instead: {error_msg}" + ) @pytest.mark.asyncio @@ -2914,7 +3132,9 @@ async def test_find_and_validate_specific_team_id_hints_bracket_index_notation() error_msg = str(exc_info.value) assert "roles[0]" in error_msg, f"Expected field name in: {error_msg}" - assert "roles" in error_msg and "list" in error_msg, f"Expected hint about using 'roles' instead: {error_msg}" + assert "roles" in error_msg and "list" in error_msg, ( + f"Expected hint about using 'roles' instead: {error_msg}" + ) @pytest.mark.asyncio @@ -3024,7 +3244,9 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( if len(user_teams) == 1 and get_team_object_return == "resolved_row": only = user_teams[0] team_table = LiteLLM_TeamTable(team_id=only) - membership = LiteLLM_TeamMembership(user_id=user_id, team_id=only, litellm_budget_table=None) + membership = LiteLLM_TeamMembership( + user_id=user_id, team_id=only, litellm_budget_table=None + ) get_team_return_value = team_table membership_return_value = membership else: @@ -3083,7 +3305,9 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -3100,7 +3324,9 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team( code = 404 if get_team_object_return == "http_404" else 500 mock_get_team.side_effect = HTTPException( status_code=code, - detail={"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."}, + detail={ + "error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call." + }, ) else: mock_get_team.return_value = get_team_return_value @@ -3209,7 +3435,9 @@ async def test_auth_builder_single_team_fallback_membership_outage_raises_instea ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -3250,7 +3478,9 @@ def _reset_unscoped_warning_flag(): JWTHandler._unscoped_jwt_warning_emitted = False -def test_build_decode_kwargs_no_env_disables_both_verifications(monkeypatch, _reset_unscoped_warning_flag): +def test_build_decode_kwargs_no_env_disables_both_verifications( + monkeypatch, _reset_unscoped_warning_flag +): monkeypatch.delenv("JWT_AUDIENCE", raising=False) monkeypatch.delenv("JWT_ISSUER", raising=False) @@ -3261,7 +3491,9 @@ def test_build_decode_kwargs_no_env_disables_both_verifications(monkeypatch, _re assert kwargs["options"] == {"verify_aud": False, "verify_iss": False} -def test_build_decode_kwargs_audience_only_enables_aud_verification(monkeypatch, _reset_unscoped_warning_flag): +def test_build_decode_kwargs_audience_only_enables_aud_verification( + monkeypatch, _reset_unscoped_warning_flag +): monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") monkeypatch.delenv("JWT_ISSUER", raising=False) @@ -3273,7 +3505,9 @@ def test_build_decode_kwargs_audience_only_enables_aud_verification(monkeypatch, assert kwargs["options"] == {"verify_iss": False} -def test_build_decode_kwargs_issuer_only_enables_iss_verification(monkeypatch, _reset_unscoped_warning_flag): +def test_build_decode_kwargs_issuer_only_enables_iss_verification( + monkeypatch, _reset_unscoped_warning_flag +): monkeypatch.delenv("JWT_AUDIENCE", raising=False) monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") @@ -3284,7 +3518,9 @@ def test_build_decode_kwargs_issuer_only_enables_iss_verification(monkeypatch, _ assert kwargs["options"] == {"verify_aud": False} -def test_build_decode_kwargs_both_set_enables_full_verification(monkeypatch, _reset_unscoped_warning_flag): +def test_build_decode_kwargs_both_set_enables_full_verification( + monkeypatch, _reset_unscoped_warning_flag +): monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") @@ -3296,7 +3532,9 @@ def test_build_decode_kwargs_both_set_enables_full_verification(monkeypatch, _re assert kwargs["options"] is None -def test_build_decode_kwargs_warns_once_when_unscoped(monkeypatch, _reset_unscoped_warning_flag, caplog): +def test_build_decode_kwargs_warns_once_when_unscoped( + monkeypatch, _reset_unscoped_warning_flag, caplog +): """The warning about unscoped JWT auth should fire on the first call but not on every subsequent decode.""" import logging @@ -3312,12 +3550,17 @@ def test_build_decode_kwargs_warns_once_when_unscoped(monkeypatch, _reset_unscop matching = [ r for r in caplog.records - if "JWT auth is enabled" in r.getMessage() and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + if "JWT auth is enabled" in r.getMessage() + and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() ] - assert len(matching) == 1, f"Expected exactly one warning across 3 calls, got {len(matching)}" + assert len(matching) == 1, ( + f"Expected exactly one warning across 3 calls, got {len(matching)}" + ) -def test_build_decode_kwargs_no_warning_when_scoped(monkeypatch, _reset_unscoped_warning_flag, caplog): +def test_build_decode_kwargs_no_warning_when_scoped( + monkeypatch, _reset_unscoped_warning_flag, caplog +): import logging monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") @@ -3326,7 +3569,11 @@ def test_build_decode_kwargs_no_warning_when_scoped(monkeypatch, _reset_unscoped JWTHandler._build_decode_kwargs() - matching = [r for r in caplog.records if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()] + matching = [ + r + for r in caplog.records + if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + ] assert matching == [] @@ -3384,7 +3631,11 @@ async def test_find_team_with_model_access_unresolved_group_claim_returns_none( from litellm.router import Router - router = Router(model_list=[{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}]) + router = Router( + model_list=[ + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} + ] + ) proxy_server_module = types.ModuleType("proxy_server") proxy_server_module.llm_router = router monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) @@ -3456,7 +3707,9 @@ async def test_find_and_validate_specific_team_id_non_404_http_exception_propaga "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, ) as mock_get_team: - mock_get_team.side_effect = HTTPException(status_code=status_code, detail="non-404 failure") + mock_get_team.side_effect = HTTPException( + status_code=status_code, detail="non-404 failure" + ) with pytest.raises(HTTPException) as exc_info: await JWTAuthManager.find_and_validate_specific_team_id( @@ -3529,7 +3782,9 @@ async def test_find_team_with_model_access_resolved_team_without_model_still_rai async def mock_get_team_object(*_args, **_kwargs): return team - monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -3594,7 +3849,11 @@ async def test_find_team_with_model_access_unresolved_group_claim_default_raises from litellm.router import Router - router = Router(model_list=[{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}]) + router = Router( + model_list=[ + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} + ] + ) proxy_server_module = types.ModuleType("proxy_server") proxy_server_module.llm_router = router monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) @@ -3631,7 +3890,12 @@ def test_canonical_user_id_rebinds_to_legacy_uuid(): jwt_email = "matt@example.com" user_object = LiteLLM_UserTable(user_id=legacy_uuid, user_email=jwt_email) - assert JWTAuthManager._canonical_user_id_from_db(user_id=jwt_email, user_object=user_object) == legacy_uuid + assert ( + JWTAuthManager._canonical_user_id_from_db( + user_id=jwt_email, user_object=user_object + ) + == legacy_uuid + ) def test_canonical_user_id_no_change_when_ids_match(): @@ -3639,20 +3903,28 @@ def test_canonical_user_id_no_change_when_ids_match(): same = "alice@example.com" user_object = LiteLLM_UserTable(user_id=same, user_email=same) - assert JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object) == same + assert ( + JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object) + == same + ) def test_canonical_user_id_returns_claim_when_no_user_object(): """No resolved row (e.g. upsert disabled / brand new) -> keep the claim.""" assert ( - JWTAuthManager._canonical_user_id_from_db(user_id="newcomer@example.com", user_object=None) + JWTAuthManager._canonical_user_id_from_db( + user_id="newcomer@example.com", user_object=None + ) == "newcomer@example.com" ) def test_canonical_user_id_returns_none_when_claim_none_and_no_object(): """Defensive: no claim and no row -> stays None, never invents an id.""" - assert JWTAuthManager._canonical_user_id_from_db(user_id=None, user_object=None) is None + assert ( + JWTAuthManager._canonical_user_id_from_db(user_id=None, user_object=None) + is None + ) def test_canonical_user_id_no_change_when_db_user_id_falsy(): @@ -3662,7 +3934,10 @@ def test_canonical_user_id_no_change_when_db_user_id_falsy(): user_id = "" assert ( - JWTAuthManager._canonical_user_id_from_db(user_id="jwt@example.com", user_object=_Stub()) == "jwt@example.com" + JWTAuthManager._canonical_user_id_from_db( + user_id="jwt@example.com", user_object=_Stub() + ) + == "jwt@example.com" ) @@ -3678,7 +3953,9 @@ async def test_auth_jwt_expired_token_raises_401_jwk_path(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() with ( - patch.object(jwt_handler, "get_public_key", new_callable=AsyncMock) as mock_get_public_key, + patch.object( + jwt_handler, "get_public_key", new_callable=AsyncMock + ) as mock_get_public_key, patch( "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", return_value={"kid": "test-kid"}, @@ -3714,7 +3991,9 @@ async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): mock_cert.public_key.return_value.public_bytes.return_value = b"fake-key" with ( - patch.object(jwt_handler, "get_public_key", new_callable=AsyncMock) as mock_get_public_key, + patch.object( + jwt_handler, "get_public_key", new_callable=AsyncMock + ) as mock_get_public_key, patch( "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", return_value={"kid": "test-kid"}, @@ -3728,7 +4007,9 @@ async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): side_effect=jwt_lib.ExpiredSignatureError("Signature has expired"), ), ): - mock_get_public_key.return_value = "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" + mock_get_public_key.return_value = ( + "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" + ) with pytest.raises(ProxyException) as exc_info: await jwt_handler.auth_jwt(token="expired.jwt.token") @@ -3842,7 +4123,9 @@ async def test_get_public_key_fetches_and_caches_jwks_response(): ) assert public_key == jwk - cached_keys = await cache.async_get_cache(key="litellm_jwt_auth_keys_https://issuer.example.com/keys") + cached_keys = await cache.async_get_cache( + key="litellm_jwt_auth_keys_https://issuer.example.com/keys" + ) assert cached_keys == [jwk] @@ -4087,7 +4370,9 @@ async def test_lowering_public_key_stale_ttl_stops_serving_a_copy_cached_under_t # The operator tightens the window and restarts; the cache, and its long-lived copy, survive. endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) - tightened = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_stale_ttl=lowered_stale_ttl) + tightened = _get_jwt_handler_with_scripted_endpoint( + cache, endpoint, public_key_stale_ttl=lowered_stale_ttl + ) await cache.async_set_cache( key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}", value=time.time() - 7200, @@ -4435,7 +4720,9 @@ def test_get_jwks_url_for_issuer_falls_back_to_discovery_document(): jwks_url = jwt_handler._get_jwks_url_for_issuer(issuer_config=issuer_config) - assert jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" + assert ( + jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" + ) @pytest.mark.asyncio @@ -4459,7 +4746,9 @@ async def test_get_objects_team_membership_uses_rebound_user_id(): return None jwt_handler = JWTHandler() - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="email", user_id_upsert=True) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + user_id_jwt_field="email", user_id_upsert=True + ) with ( patch( @@ -4555,7 +4844,9 @@ async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims( assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-org" - assert jwt_handler.get_team_id(token=claims, default_value=None) == ("example-org/litellm-fork") + assert jwt_handler.get_team_id(token=claims, default_value=None) == ( + "example-org/litellm-fork" + ) @pytest.mark.asyncio @@ -4658,7 +4949,9 @@ async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): claims = await jwt_handler.auth_jwt(token=token) - assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" + assert ( + jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" + ) @pytest.mark.asyncio @@ -4691,7 +4984,7 @@ async def test_multi_issuer_jwt_unknown_issuer_falls_back_to_global_jwks(monkeyp kid="issuer-key", ) - with pytest.raises(Exception, match="Missing JWT Public Key URL from environment\\.") as exc: + with pytest.raises(Exception, match='Missing JWT Public Key URL from environment\\.') as exc: await jwt_handler.auth_jwt(token=token) assert "Missing JWT Public Key URL from environment." in str(exc.value) @@ -4765,7 +5058,7 @@ async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch) kid=shared_kid, ) - with pytest.raises(Exception, match="Validation fails: Signature verification failed") as exc: + with pytest.raises(Exception, match='Validation fails: Signature verification failed') as exc: await jwt_handler.auth_jwt(token=token) assert "Validation fails" in str(exc.value) @@ -4820,7 +5113,7 @@ def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( issuer = "https://issuer.example.com" jwks_url = f"{issuer}/keys" - with pytest.raises(Exception, match="must configure audience or set") as exc: + with pytest.raises(Exception, match='must configure audience or set') as exc: LiteLLM_JWTAuth( issuers=[ { @@ -4837,7 +5130,7 @@ def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): issuer = "https://issuer.example.com" jwks_url = f"{issuer}/keys" - with pytest.raises(Exception, match="cannot set audience and disable_audience_validation=True") as exc: + with pytest.raises(Exception, match='cannot set audience and disable_audience_validation=True') as exc: LiteLLM_JWTAuth( issuers=[ { @@ -4849,7 +5142,9 @@ def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): ] ) - assert "cannot set audience and disable_audience_validation=True together" in str(exc.value) + assert "cannot set audience and disable_audience_validation=True together" in str( + exc.value + ) @pytest.mark.asyncio @@ -4902,15 +5197,21 @@ async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): claims = await jwt_handler.auth_jwt(token=token) - assert jwt_handler.get_user_id(token=claims, default_value=None) == ("real-user@example.com") - assert jwt_handler.get_user_email(token=claims, default_value=None) == ("real-user@example.com") + assert jwt_handler.get_user_id(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) assert jwt_handler.get_team_id(token=claims, default_value=None) == "real-team" assert jwt_handler.get_team_ids_from_jwt(token=claims) == [ "real-team", "secondary-team", ] assert jwt_handler.get_org_id(token=claims, default_value=None) == "real-org" - assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ("real-end-user") + assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ( + "real-end-user" + ) @pytest.mark.asyncio @@ -4950,11 +5251,15 @@ async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): assert JWTHandler.LITELLM_TEAM_ID_CLAIM not in claims assert jwt_handler.get_user_id(token=claims, default_value=None) is None assert jwt_handler.get_team_id(token=claims, default_value=None) is None - assert jwt_handler.get_user_email(token=claims, default_value=None) == ("real-user@example.com") + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) @pytest.mark.asyncio -async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning(monkeypatch, caplog): +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( + monkeypatch, caplog +): import logging monkeypatch.delenv("JWT_AUDIENCE", raising=False) @@ -5004,7 +5309,11 @@ def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deploym JWTHandler._build_decode_kwargs() - matching = [r for r in caplog.records if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()] + matching = [ + r + for r in caplog.records + if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + ] assert len(matching) == 1 @@ -5132,7 +5441,9 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership(): "expect_403", ), [ - pytest.param(True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"), + pytest.param( + True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team" + ), pytest.param( True, ["team_a", "team_b"], @@ -5210,7 +5521,9 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( async def call_auth_builder(): with ( - patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, + patch.object( + jwt_handler, "auth_jwt", new_callable=AsyncMock + ) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -5250,7 +5563,9 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -5354,7 +5669,9 @@ async def test_resolve_db_team_fallback_skips_team_without_model_access(): teams=["restricted_team", "allowed_team"], ) teams = { - "restricted_team": LiteLLM_TeamTable(team_id="restricted_team", models=["claude-3"]), + "restricted_team": LiteLLM_TeamTable( + team_id="restricted_team", models=["claude-3"] + ), "allowed_team": LiteLLM_TeamTable(team_id="allowed_team", models=["gpt-4"]), } @@ -5536,7 +5853,9 @@ async def _run_auth_builder_with_header_team( jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = jwt_auth_config with ( - patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value=token), + patch.object( + jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value=token + ), patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -5555,7 +5874,9 @@ async def _run_auth_builder_with_header_team( new_callable=AsyncMock, return_value=None, ), - patch.object(JWTAuthManager, "get_all_team_ids", return_value=allowed_team_ids), + patch.object( + JWTAuthManager, "get_all_team_ids", return_value=allowed_team_ids + ), patch.object( JWTAuthManager, "get_objects", @@ -5564,7 +5885,9 @@ async def _run_auth_builder_with_header_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -5591,7 +5914,9 @@ async def _run_auth_builder_with_header_team( @pytest.mark.asyncio -async def test_auth_builder_header_team_not_found_matches_non_membership_denial() -> None: +async def test_auth_builder_header_team_not_found_matches_non_membership_denial() -> ( + None +): """A provisional x-litellm-team-id naming a nonexistent team must produce the exact same 403 shape as one naming an existing team outside the caller's memberships. Letting get_team_object's 404 surface would give any @@ -5611,15 +5936,19 @@ async def test_auth_builder_header_team_not_found_matches_non_membership_denial( return LiteLLM_TeamTable(team_id=team_id) with pytest.raises(HTTPException) as missing_exc: - await _run_auth_builder_with_header_team(config, token, "team_ghost", user_object, _team_lookup_404, set()) + await _run_auth_builder_with_header_team( + config, token, "team_ghost", user_object, _team_lookup_404, set() + ) with pytest.raises(HTTPException) as outsider_exc: - await _run_auth_builder_with_header_team(config, token, "team_other", user_object, team_exists, set()) + await _run_auth_builder_with_header_team( + config, token, "team_other", user_object, team_exists, set() + ) assert missing_exc.value.status_code == 403 assert outsider_exc.value.status_code == 403 - assert missing_exc.value.detail.replace("team_ghost", "") == outsider_exc.value.detail.replace( - "team_other", "" - ) + assert missing_exc.value.detail.replace( + "team_ghost", "" + ) == outsider_exc.value.detail.replace("team_other", "") assert "exist" not in missing_exc.value.detail @@ -5816,7 +6145,9 @@ async def test_auth_builder_db_fallback_does_not_validate_rbac_team_against_db_m ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -5968,7 +6299,9 @@ async def test_auth_builder_db_fallback_runs_when_only_team_id_default_set(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6050,7 +6383,9 @@ async def test_auth_builder_alias_only_token_resolves_alias_not_db_fallback(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6107,10 +6442,14 @@ async def test_find_and_validate_specific_team_id_alias_wins_over_team_id_defaul ) jwt_token = {"sub": "user-1", "team_alias": "my-team"} - alias_team = LiteLLM_TeamTable(team_id="alias_resolved_team", team_alias="my-team") + alias_team = LiteLLM_TeamTable( + team_id="alias_resolved_team", team_alias="my-team" + ) with ( - patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -6157,7 +6496,9 @@ async def test_find_and_validate_specific_team_id_team_id_default_used_without_a default_team = LiteLLM_TeamTable(team_id="config_default_team") with ( - patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, + patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -6226,7 +6567,9 @@ async def test_auth_builder_db_fallback_enforces_passthrough_route_access(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6380,14 +6723,18 @@ async def test_sync_user_role_and_teams_singular_claim_reconciles_memberships(): "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ) as mock_patch: - await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, AsyncMock()) + await JWTAuthManager.sync_user_role_and_teams( + jwt_handler, token, user, AsyncMock() + ) mock_patch.assert_awaited_once() assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == { "team_stale_a", "team_stale_b", } - assert set(mock_patch.call_args.kwargs["teams_ids_to_add_user_to"]) == {"team_primary"} + assert set(mock_patch.call_args.kwargs["teams_ids_to_add_user_to"]) == { + "team_primary" + } assert user.teams == ["team_primary"] @@ -6446,7 +6793,9 @@ async def test_auth_builder_provisional_header_team_is_not_upserted(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6521,7 +6870,9 @@ async def test_auth_builder_header_cannot_override_rbac_team_under_db_fallback() ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6571,7 +6922,9 @@ async def test_auth_builder_header_team_enforces_team_allowed_routes_under_db_fa async def call(route: str): with ( - patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, + patch.object( + jwt_handler, "auth_jwt", new_callable=AsyncMock + ) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -6599,7 +6952,9 @@ async def test_auth_builder_header_team_enforces_team_allowed_routes_under_db_fa ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), + patch.object( + JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock + ), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6957,10 +7312,14 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ) as mock_patch: - await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, AsyncMock()) + await JWTAuthManager.sync_user_role_and_teams( + jwt_handler, token, user, AsyncMock() + ) mock_patch.assert_awaited_once() - assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == {"team_existing"} + assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == { + "team_existing" + } assert mock_patch.call_args.kwargs["teams_ids_to_add_user_to"] == [] assert user.teams == [] @@ -7185,12 +7544,8 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a if identity_only: identity = await JWTAuthManager.resolve_identity( - api_key=token, - jwt_handler=jwt_handler, - prisma_client=None, - user_api_key_cache=None, - parent_otel_span=None, - proxy_logging_obj=None, + api_key=token, jwt_handler=jwt_handler, prisma_client=None, + user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, ) assert identity.agent_id == "canonical-agent-id" return @@ -7225,12 +7580,8 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch if identity_only: with pytest.raises(HTTPException) as denial: await JWTAuthManager.resolve_identity( - api_key=token, - jwt_handler=jwt_handler, - prisma_client=None, - user_api_key_cache=None, - parent_otel_span=None, - proxy_logging_obj=None, + api_key=token, jwt_handler=jwt_handler, prisma_client=None, + user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, ) assert denial.value.status_code == 403 return @@ -7297,9 +7648,7 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc from litellm.proxy.management_endpoints import team_endpoints handler, token = _entra_signed_app_token( - monkeypatch, - azp="canonical-agent-id", - scope=LiteLLM_JWTAuth().admin_jwt_scope, + monkeypatch, azp="canonical-agent-id", scope=LiteLLM_JWTAuth().admin_jwt_scope, ) handler.bind_agent_lookup(_entra_agent_registry()) handler.litellm_jwtauth.team_id_upsert = True @@ -7311,16 +7660,10 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc resolve = JWTAuthManager.auth_builder if admission else JWTAuthManager.authorize_jwt result = await resolve( - api_key=token, - jwt_handler=handler, - request_data={}, - general_settings={}, - route="/chat/completions", - prisma_client=database, - user_api_key_cache=handler.user_api_key_cache, - parent_otel_span=None, - proxy_logging_obj=MagicMock(), - request_headers={"x-litellm-team-id": "new-team"}, + api_key=token, jwt_handler=handler, request_data={}, general_settings={}, + route="/chat/completions", prisma_client=database, + user_api_key_cache=handler.user_api_key_cache, parent_otel_span=None, + proxy_logging_obj=MagicMock(), request_headers={"x-litellm-team-id": "new-team"}, ) assert result["is_proxy_admin"] is True @@ -7345,12 +7688,7 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning cache: Final = UserApiKeyCache() cache.set_cache("litellm_jwt_auth_keys_https://admin.example/jwks", [jwk]) user_id: Final = f"admin-status-{existing_user}-{warm_cache}-{email}" - user: Final = LiteLLM_UserTable( - user_id=user_id, - user_email="admin@allowed.example", - metadata={"scim_active": False}, - organization_memberships=[], - ) + user: Final = LiteLLM_UserTable(user_id=user_id, user_email="admin@allowed.example", metadata={"scim_active": False}, organization_memberships=[]) if existing_user and warm_cache: cache.set_cache(user_id, user) database: Final = MagicMock() @@ -7363,9 +7701,7 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning prisma_client=database, user_api_key_cache=cache, litellm_jwtauth=LiteLLM_JWTAuth( - user_id_jwt_field="sub", - user_id_upsert=True, - user_email_jwt_field="email", + user_id_jwt_field="sub", user_id_upsert=True, user_email_jwt_field="email", user_allowed_email_domain="allowed.example", ), ) @@ -7373,22 +7709,12 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning monkeypatch.setenv("JWT_ISSUER", "https://admin.example") monkeypatch.setenv("JWT_AUDIENCE", "gateway") token: Final = _encode_rsa_jwt( - private_key, - "https://admin.example", - "gateway", - "admin-status", + private_key, "https://admin.example", "gateway", "admin-status", {"sub": user_id, "scope": "litellm_proxy_admin", **({"email": email} if email else {})}, ) result: Final = await JWTAuthManager.auth_builder( - api_key=token, - jwt_handler=handler, - prisma_client=database, - user_api_key_cache=cache, - parent_otel_span=None, - proxy_logging_obj=MagicMock(), - request_data={}, - general_settings={}, - route="/user/info", + api_key=token, jwt_handler=handler, prisma_client=database, user_api_key_cache=cache, + parent_otel_span=None, proxy_logging_obj=MagicMock(), request_data={}, general_settings={}, route="/user/info", ) assert result["is_proxy_admin"] is True assert result["user_id"] == user_id diff --git a/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py b/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py index 1150fab69b4..9ec4aac4692 100644 --- a/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py +++ b/tests/test_litellm/proxy/guardrails/test_mcp_jwt_signer.py @@ -174,7 +174,9 @@ def test_jwks_public_key_can_verify_signed_jwt(): def test_build_claims_standard_fields(): """_build_claims() populates iss, aud, iat, exp, nbf.""" - signer = _make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300) + signer = _make_signer( + issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300 + ) user_dict = _make_user_api_key_dict() data = {"mcp_tool_name": "get_weather"} @@ -345,13 +347,17 @@ async def test_hook_skips_non_mcp_call_types(): data=original_data, call_type=call_type, # type: ignore[arg-type] ) - assert "extra_headers" not in (result or {}), f"extra_headers should not be set for {call_type}" + assert "extra_headers" not in ( + result or {} + ), f"extra_headers should not be set for {call_type}" @pytest.mark.asyncio async def test_hook_signs_list_mcp_tools(): """async_pre_call_hook() signs JWT for list_mcp_tools with list scope.""" - signer = _make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300) + signer = _make_signer( + issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300 + ) user_dict = _make_user_api_key_dict(user_id="alice", team_id="backend") data = {"mcp_tool_name": "should_be_cleared"} @@ -376,7 +382,9 @@ async def test_hook_signs_list_mcp_tools(): @pytest.mark.asyncio async def test_signed_token_is_verifiable(): """The JWT injected by the hook can be verified against the JWKS public key.""" - signer = _make_signer(issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300) + signer = _make_signer( + issuer="https://litellm.example.com", audience="mcp", ttl_seconds=300 + ) user_dict = _make_user_api_key_dict(user_id="alice", team_id="backend") data = {"mcp_tool_name": "search"} @@ -592,7 +600,9 @@ async def test_channel_token_injected_when_configured(): assert isinstance(result, dict) assert "x-mcp-channel-token" in result["extra_headers"] - channel_token = result["extra_headers"]["x-mcp-channel-token"].removeprefix("Bearer ") + channel_token = result["extra_headers"]["x-mcp-channel-token"].removeprefix( + "Bearer " + ) channel_payload = _decode_unverified(channel_token) assert channel_payload["aud"] == "bedrock-gateway" @@ -808,7 +818,9 @@ def test_initialize_guardrail_passes_all_params(): litellm_params.issuer = "https://litellm.example.com" litellm_params.audience = "mcp-test" litellm_params.ttl_seconds = 120 - litellm_params.access_token_discovery_uri = "https://idp.example.com/.well-known/openid-configuration" + litellm_params.access_token_discovery_uri = ( + "https://idp.example.com/.well-known/openid-configuration" + ) litellm_params.token_introspection_endpoint = "https://idp.example.com/introspect" litellm_params.verify_issuer = "https://idp.example.com" litellm_params.verify_audience = "api://test" @@ -831,7 +843,10 @@ def test_initialize_guardrail_passes_all_params(): assert signer.issuer == "https://litellm.example.com" assert signer.audience == "mcp-test" assert signer.ttl_seconds == 120 - assert signer.access_token_discovery_uri == "https://idp.example.com/.well-known/openid-configuration" + assert ( + signer.access_token_discovery_uri + == "https://idp.example.com/.well-known/openid-configuration" + ) assert signer.token_introspection_endpoint == "https://idp.example.com/introspect" assert signer.verify_issuer == "https://idp.example.com" assert signer.verify_audience == "api://test" @@ -864,7 +879,9 @@ def _make_httpx_response(json_body: dict, status_code: int = 200): if status_code >= 400: from httpx import HTTPStatusError, Request, Response - mock_resp.raise_for_status.side_effect = HTTPStatusError("error", request=MagicMock(), response=MagicMock()) + mock_resp.raise_for_status.side_effect = HTTPStatusError( + "error", request=MagicMock(), response=MagicMock() + ) return mock_resp @@ -923,7 +940,9 @@ async def test_fetch_jwks_uses_cache_on_second_call(): @pytest.mark.asyncio async def test_get_oidc_discovery_caches_when_jwks_uri_present(): """_get_oidc_discovery caches the doc when jwks_uri is in the response.""" - signer = _make_signer(access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration") + signer = _make_signer( + access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration" + ) signer._oidc_discovery_doc = None # ensure fresh discovery_doc = { @@ -945,7 +964,9 @@ async def test_get_oidc_discovery_caches_when_jwks_uri_present(): @pytest.mark.asyncio async def test_get_oidc_discovery_does_not_cache_when_jwks_uri_absent(): """_get_oidc_discovery does NOT cache a doc that is missing jwks_uri.""" - signer = _make_signer(access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration") + signer = _make_signer( + access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration" + ) signer._oidc_discovery_doc = None bad_doc = {"issuer": "https://idp.example.com"} # no jwks_uri @@ -1048,7 +1069,9 @@ async def test_verify_incoming_jwt_raises_on_expired_token(): @pytest.mark.asyncio async def test_introspect_opaque_token_returns_claims_when_active(): """_introspect_opaque_token returns the introspection payload for active tokens.""" - signer = _make_signer(token_introspection_endpoint="https://idp.example.com/introspect") + signer = _make_signer( + token_introspection_endpoint="https://idp.example.com/introspect" + ) introspection_response = { "active": True, @@ -1072,7 +1095,9 @@ async def test_introspect_opaque_token_returns_claims_when_active(): @pytest.mark.asyncio async def test_introspect_opaque_token_raises_on_inactive_token(): """_introspect_opaque_token raises ExpiredSignatureError when active=false.""" - signer = _make_signer(token_introspection_endpoint="https://idp.example.com/introspect") + signer = _make_signer( + token_introspection_endpoint="https://idp.example.com/introspect" + ) fake_resp = _make_httpx_response({"active": False}) mock_client = MagicMock() @@ -1103,7 +1128,9 @@ async def test_hook_raises_401_when_jwt_verification_fails(): """async_pre_call_hook raises HTTP 401 when incoming JWT verification fails.""" from fastapi import HTTPException - signer = _make_signer(access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration") + signer = _make_signer( + access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration" + ) with patch.object( signer, diff --git a/tests/test_litellm/types/test_mcp.py b/tests/test_litellm/types/test_mcp.py index de138bf662c..1e997a3a892 100644 --- a/tests/test_litellm/types/test_mcp.py +++ b/tests/test_litellm/types/test_mcp.py @@ -53,12 +53,12 @@ def test_has_header_matches_any_casing() -> None: @pytest.mark.parametrize( "target,expected", [ - ("https://upstream.example.com/other", False), # same origin + ("https://upstream.example.com/other", False), # same origin ("https://upstream.example.com:443/other", False), # explicit default port - ("https://attacker.example.com/collect", True), # different host - ("http://upstream.example.com/collect", True), # scheme downgrade, same host + ("https://attacker.example.com/collect", True), # different host + ("http://upstream.example.com/collect", True), # scheme downgrade, same host ("https://upstream.example.com:8443/other", True), # different port, same host - ("https://sub.upstream.example.com/x", True), # different host + ("https://sub.upstream.example.com/x", True), # different host ], ) def test_origin_is_scheme_host_and_port_not_host_alone(target: str, expected: bool) -> None: From 19675c0f5c4ab9d857faaf35c33ee6a74a131625 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 09:15:26 +0000 Subject: [PATCH 10/18] fix(mcp): validate client_assertion_signing_alg at the REST boundary, keep stored rows readable (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_lazy_openapi_snapshot.json | 44 ------------------- litellm/proxy/_types.py | 29 ++++++++++++ litellm/types/mcp.py | 3 +- .../test_mcp_client_assertion_signing_alg.py | 35 ++++++++++++++- .../mcp_server/test_db_credentials.py | 17 +++++++ tests/test_litellm/proxy/test__types.py | 30 +++++++++++++ tests/test_litellm/types/test_mcp.py | 25 ----------- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +- 8 files changed, 112 insertions(+), 73 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 858a993c4ec..0b43c3864ab 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -29465,17 +29465,6 @@ "client_assertion_signing_alg": { "anyOf": [ { - "enum": [ - "RS256", - "RS384", - "RS512", - "PS256", - "PS384", - "PS512", - "ES256", - "ES384", - "ES512" - ], "type": "string" }, { @@ -32177,17 +32166,6 @@ "client_assertion_signing_alg": { "anyOf": [ { - "enum": [ - "RS256", - "RS384", - "RS512", - "PS256", - "PS384", - "PS512", - "ES256", - "ES384", - "ES512" - ], "type": "string" }, { @@ -35461,17 +35439,6 @@ "client_assertion_signing_alg": { "anyOf": [ { - "enum": [ - "RS256", - "RS384", - "RS512", - "PS256", - "PS384", - "PS512", - "ES256", - "ES384", - "ES512" - ], "type": "string" }, { @@ -39322,17 +39289,6 @@ "client_assertion_signing_alg": { "anyOf": [ { - "enum": [ - "RS256", - "RS384", - "RS512", - "PS256", - "PS384", - "PS512", - "ES256", - "ES384", - "ES512" - ], "type": "string" }, { diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 4affa55f903..036dcb01759 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -15,6 +15,8 @@ from pydantic import ( Json, JsonValue, PositiveInt, + TypeAdapter, + ValidationError, field_validator, model_validator, ) @@ -45,6 +47,10 @@ from litellm.types.mcp import ( MCPTransportType, ) from litellm.types.mcp_server.mcp_server_manager import MCPInfo +from litellm.types.proxy.auth.jwt_algorithms import ( + APPROVED_JWT_ALGORITHMS, + ApprovedJwtAlgorithm, +) from litellm.types.proxy.carried_budget_state import ( OrgBudgetSnapshot, TeamBudgetSnapshot, @@ -1502,6 +1508,19 @@ def _reject_unsupported_per_server_oauth_discovery(values: object, require_auth_ raise _per_server_oauth_discovery_error() +_APPROVED_JWT_ALGORITHM_ADAPTER: Final = TypeAdapter(ApprovedJwtAlgorithm) + + +def _validated_client_assertion_signing_alg(credentials: MCPCredentials | None) -> MCPCredentials | None: + if credentials is None or credentials.get("client_assertion_signing_alg") is None: + return credentials + try: + _APPROVED_JWT_ALGORITHM_ADAPTER.validate_python(credentials["client_assertion_signing_alg"]) + except ValidationError as exc: + raise ValueError(f"client_assertion_signing_alg must be one of {', '.join(APPROVED_JWT_ALGORITHMS)}") from exc + return credentials + + class NewMCPServerRequest(LiteLLMPydanticObjectBase): server_id: str | None = None server_name: str | None = None @@ -1565,6 +1584,11 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): description="Server-managed: set by the endpoint; caller values are overridden.", ) + @field_validator("credentials") + @classmethod + def check_client_assertion_signing_alg(cls, credentials): + return _validated_client_assertion_signing_alg(credentials) + @model_validator(mode="before") @classmethod def validate_transport_fields(cls, values): @@ -1664,6 +1688,11 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): timeout: float | None = None max_concurrent_requests: int | None = None + @field_validator("credentials") + @classmethod + def check_client_assertion_signing_alg(cls, credentials): + return _validated_client_assertion_signing_alg(credentials) + @model_validator(mode="before") @classmethod def validate_transport_fields(cls, values): diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 7960abcd083..83e719810d5 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -12,7 +12,6 @@ from pydantic import BaseModel, ConfigDict, Field from typing_extensions import TypedDict from litellm.types.llms.base import HiddenParams -from litellm.types.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm if TYPE_CHECKING: import httpx2 @@ -252,7 +251,7 @@ class MCPCredentials(TypedDict, total=False): Key id (kid) advertised in the client_assertion JWT header """ - client_assertion_signing_alg: ApprovedJwtAlgorithm | None + client_assertion_signing_alg: str | None """ Signing algorithm for the client_assertion JWT. Default: RS256 """ diff --git a/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py b/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py index 883c335b71d..c2c32f9866b 100644 --- a/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py +++ b/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py @@ -1,10 +1,12 @@ import json +import os import uuid from typing import Final +import psycopg from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import rsa -from integration._support.client import Gateway +from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.mcp import forget_mcp, mcp_peer, register_mcp @@ -85,3 +87,34 @@ def test_approved_client_assertion_signing_alg_round_trips(gateway: Gateway) -> identity: Final = str(created.json()["server_id"]) scenario.cleanups.callback(forget_mcp, gateway, identity) assert _stored_client_assertion_signing_alg(identity) == "ES256" + + +def test_server_row_with_stale_client_assertion_signing_alg_still_loads(gateway: Gateway) -> None: + """A row persisted before the allowlist (credentials.client_assertion_signing_alg = "HS256") + must keep loading with the RS256 fallback instead of disappearing on upgrade.""" + identity: Final = "stalealg" + uuid.uuid4().hex[:8] + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute( + 'INSERT INTO "LiteLLM_MCPServerTable" (server_id, server_name, url, transport, credentials,' + " created_at, updated_at) VALUES (%s, %s, %s, %s, %s::jsonb, NOW(), NOW())", + ( + identity, + "stalealg" + uuid.uuid4().hex[:8], + "https://mcp.integration.invalid", + "http", + json.dumps({"client_assertion_signing_alg": "HS256"}), + ), + ) + try: + with mcp_peer() as peer, gateway.scenario() as scenario: + register_mcp(scenario, peer, "staletrigger" + uuid.uuid4().hex[:8]) + listed: Final = eventually( + lambda: gateway.request("GET", "/v1/mcp/server"), + lambda response: response.status_code == 200 + and any(server.get("server_id") == identity for server in response.json()), + seconds=30, + ) + assert listed.status_code == 200, listed.text + finally: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute('DELETE FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,)) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index cfcff73b857..ec145bded45 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -1751,3 +1751,20 @@ async def test_unverified_legacy_cache_cannot_bypass_enforcement(monkeypatch): await mcp_per_user_token_cache.set("alice", "srv", "bob", 60) assert await module.resolve_user_oauth_access_token("alice", server) is None assert await mcp_per_user_token_cache.get("alice", "srv") is None + + +def test_server_table_row_with_stale_client_assertion_signing_alg_still_validates() -> None: + """Rows written before the approved-algorithm allowlist carry values like HS256 in the + credentials blob; the stored shape stays lenient so reload_servers_from_database still + parses the row and _stored_client_assertion_signing_alg falls back to RS256 instead of + the server silently disappearing on upgrade.""" + row: Final = { + "server_id": "srv-stale-alg", + "server_name": "stale_alg_server", + "url": "https://mcp.example.com", + "transport": "http", + "credentials": {"client_assertion_signing_alg": "HS256"}, + } + parsed: Final = LiteLLM_MCPServerTable.model_validate(row) + assert parsed.credentials is not None + assert parsed.credentials["client_assertion_signing_alg"] == "HS256" diff --git a/tests/test_litellm/proxy/test__types.py b/tests/test_litellm/proxy/test__types.py index f5abe0561db..9f17ddb12a3 100644 --- a/tests/test_litellm/proxy/test__types.py +++ b/tests/test_litellm/proxy/test__types.py @@ -377,3 +377,33 @@ def test_change_password_request_passwords_hidden_from_repr(): for rendered in (repr(request), str(request)): assert "hunter2hunter2" not in rendered assert "NewP@ssw0rd-2026" not in rendered + + +def test_mcp_server_requests_reject_non_approved_client_assertion_signing_alg() -> None: + """The REST boundary is strict even though the stored blob stays lenient: + a write carrying HS256/hs256/"" must 422 naming the field.""" + from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest + + for cls, base in ( + (NewMCPServerRequest, {"transport": "http", "url": "https://mcp.example.com"}), + (UpdateMCPServerRequest, {"server_id": "srv-1", "transport": "http", "url": "https://mcp.example.com"}), + ): + for alg in ("HS256", "hs256", "", "EdDSA"): + with pytest.raises(ValidationError) as exc: + cls(**base, credentials={"client_assertion_signing_alg": alg}) + assert "client_assertion_signing_alg" in str(exc.value), f"{cls.__name__} accepted {alg!r}" + + +def test_mcp_server_requests_accept_approved_client_assertion_signing_alg() -> None: + from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest + + for alg in ("ES256", "PS384", "RS256", None): + request = NewMCPServerRequest( + transport="http", + url="https://mcp.example.com", + credentials={"client_assertion_signing_alg": alg}, + ) + assert request.credentials is not None + assert request.credentials["client_assertion_signing_alg"] == alg + update = UpdateMCPServerRequest(server_id="srv-1") + assert update.credentials is None diff --git a/tests/test_litellm/types/test_mcp.py b/tests/test_litellm/types/test_mcp.py index 1e997a3a892..e714f89be03 100644 --- a/tests/test_litellm/types/test_mcp.py +++ b/tests/test_litellm/types/test_mcp.py @@ -86,28 +86,3 @@ async def test_the_hook_drops_the_slot_only_once_the_origin_changes() -> None: await hook(foreign) assert "esb-oauth" not in foreign.headers - -def test_new_mcp_server_request_rejects_non_approved_client_assertion_signing_alg() -> None: - from pydantic import ValidationError - - from litellm.proxy._types import NewMCPServerRequest - - with pytest.raises(ValidationError) as exc: - NewMCPServerRequest( - transport="http", - url="https://mcp.example.com", - credentials={"client_assertion_signing_alg": "HS256"}, - ) - assert "client_assertion_signing_alg" in str(exc.value) - - -def test_new_mcp_server_request_accepts_approved_client_assertion_signing_alg() -> None: - from litellm.proxy._types import NewMCPServerRequest - - request = NewMCPServerRequest( - transport="http", - url="https://mcp.example.com", - credentials={"client_assertion_signing_alg": "ES256"}, - ) - assert request.credentials is not None - assert request.credentials["client_assertion_signing_alg"] == "ES256" diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index af40e742905..2938e65cfde 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -34841,7 +34841,7 @@ export interface components { /** Aws Session Token */ aws_session_token?: string | null; /** Client Assertion Signing Alg */ - client_assertion_signing_alg?: ("RS256" | "RS384" | "RS512" | "PS256" | "PS384" | "PS512" | "ES256" | "ES384" | "ES512") | null; + client_assertion_signing_alg?: string | null; /** Client Id */ client_id?: string | null; /** Client Private Key */ From a7c83a14e20109d77f679aaa73497e9aa0813789 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 09:24:27 +0000 Subject: [PATCH 11/18] refactor: type the MCP credentials validators and drop stray whitespace (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 4 ++-- litellm/proxy/auth/jwt_algorithms.py | 3 ++- tests/test_litellm/types/test_mcp.py | 1 - 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 036dcb01759..fdf58e50ede 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1586,7 +1586,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): @field_validator("credentials") @classmethod - def check_client_assertion_signing_alg(cls, credentials): + def check_client_assertion_signing_alg(cls, credentials: MCPCredentials | None) -> MCPCredentials | None: return _validated_client_assertion_signing_alg(credentials) @model_validator(mode="before") @@ -1690,7 +1690,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): @field_validator("credentials") @classmethod - def check_client_assertion_signing_alg(cls, credentials): + def check_client_assertion_signing_alg(cls, credentials: MCPCredentials | None) -> MCPCredentials | None: return _validated_client_assertion_signing_alg(credentials) @model_validator(mode="before") diff --git a/litellm/proxy/auth/jwt_algorithms.py b/litellm/proxy/auth/jwt_algorithms.py index a3f7ec0c4f9..0e1839ec200 100644 --- a/litellm/proxy/auth/jwt_algorithms.py +++ b/litellm/proxy/auth/jwt_algorithms.py @@ -20,7 +20,8 @@ def allowed_jwt_algorithms(fips_mode: bool) -> tuple[str, ...]: def jwks_keys_for( keys: Sequence[Mapping[str, object]], algorithms: Collection[str] ) -> tuple[Mapping[str, object], ...]: - """Keep keys whose declared alg is allowed; keys without alg are kept only when their kty can sign with an allowed algorithm.""" + """Keep keys whose declared alg is allowed; keys without alg are kept only when their kty can sign with an + allowed algorithm.""" return tuple(key for key in keys if _key_allowed(key, frozenset(algorithms))) diff --git a/tests/test_litellm/types/test_mcp.py b/tests/test_litellm/types/test_mcp.py index e714f89be03..5450ec4aa48 100644 --- a/tests/test_litellm/types/test_mcp.py +++ b/tests/test_litellm/types/test_mcp.py @@ -85,4 +85,3 @@ async def test_the_hook_drops_the_slot_only_once_the_origin_changes() -> None: foreign = httpx.Request("GET", "https://attacker.example.com/x", headers={"esb-oauth": "Bearer x"}) await hook(foreign) assert "esb-oauth" not in foreign.headers - From 93939f2a5357f5b3ed72b9babcc2f7caaa62a36f Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 11:50:28 +0000 Subject: [PATCH 12/18] test(integration): audit cells for JWT and MCP algorithm allowlists (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_jwt_algorithm_allowlist.py | 600 +++++++++++++++++- .../test_mcp_client_assertion_signing_alg.py | 161 ++++- .../mcp/test_mcp_jwt_signer_jwks_allowlist.py | 160 ++++- 3 files changed, 907 insertions(+), 14 deletions(-) diff --git a/tests/integration/authorization/test_jwt_algorithm_allowlist.py b/tests/integration/authorization/test_jwt_algorithm_allowlist.py index 210d40a0d6a..93dc98c17c4 100644 --- a/tests/integration/authorization/test_jwt_algorithm_allowlist.py +++ b/tests/integration/authorization/test_jwt_algorithm_allowlist.py @@ -1,22 +1,65 @@ +import asyncio +import base64 import json +import os +import socket import time import uuid +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import Final +import httpx import jwt +import psutil import yaml -from cryptography.hazmat.primitives.asymmetric import ed25519, rsa +from anthropic import Anthropic +from cryptography.hazmat.primitives.asymmetric import ec, ed25519, rsa +from jwt.algorithms import ECAlgorithm, OKPAlgorithm, RSAAlgorithm +from openai import AsyncOpenAI, OpenAI from tests.integration._support.client import Gateway, eventually +from tests.integration._support.database import read_rows from tests.integration._support.process import owned_proxy_process from tests.integration._support.wire import Reply, Request, wire_server KEY_ID: Final = "integration-jwt-allowlist-key" +APPROVED_WARNING_PARTS: Final = ("EdDSA", "deprecated", "LITELLM_FIPS_MODE") def _jwks_reply(public_jwk: str) -> Reply: - return Reply(body=json.dumps({"keys": [{**json.loads(public_jwk), "kid": KEY_ID}]}).encode()) + return _jwks_reply_keys([{**json.loads(public_jwk), "kid": KEY_ID}]) + + +def _jwks_reply_keys(keys: list[dict[str, object]]) -> Reply: + return Reply(body=json.dumps({"keys": keys}).encode()) + + +def _eddsa_keypair() -> tuple[ed25519.Ed25519PrivateKey, dict[str, object]]: + private_key: Final = ed25519.Ed25519PrivateKey.generate() + return private_key, {**json.loads(OKPAlgorithm.to_jwk(private_key.public_key())), "kid": "ed"} + + +def _rsa_keypair() -> tuple[rsa.RSAPrivateKey, dict[str, object]]: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + return private_key, {**json.loads(RSAAlgorithm.to_jwk(private_key.public_key())), "kid": "rsa"} + + +def _ec_keypair() -> tuple[ec.EllipticCurvePrivateKey, dict[str, object]]: + private_key: Final = ec.generate_private_key(ec.SECP256R1()) + return private_key, {**json.loads(ECAlgorithm.to_jwk(private_key.public_key())), "kid": "ec"} + + +def _free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] + + +def _observed_bodies() -> tuple[str, ...]: + response: Final = httpx.get(f"{os.environ['INTEGRATION_UPSTREAM_URL']}/__observations", timeout=15) + assert response.status_code == 200, response.text + return tuple(json.dumps(entry["body"]) for entry in response.json()["requests"]) def _jwt_auth_config(tmp_path: Path) -> Path: @@ -31,12 +74,12 @@ def _jwt_auth_config(tmp_path: Path) -> Path: return path -def _token(key: object, algorithm: str, subject: str) -> str: +def _token(key: object, algorithm: str, subject: str, kid: str = KEY_ID) -> str: return jwt.encode( {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, key, # pyright: ignore[reportArgumentType] # jwt.encode takes Any key material algorithm=algorithm, - headers={"kid": KEY_ID}, + headers={"kid": kid}, ) @@ -44,7 +87,7 @@ def test_eddsa_signed_token_is_accepted_and_logged_as_deprecated_outside_fips_mo gateway: Gateway, tmp_path: Path ) -> None: private_key: Final = ed25519.Ed25519PrivateKey.generate() - public_jwk: Final = jwt.algorithms.OKPAlgorithm.to_jwk(private_key.public_key()) + public_jwk: Final = OKPAlgorithm.to_jwk(private_key.public_key()) def respond(request: Request) -> Reply: assert request.method == "GET", request @@ -75,7 +118,7 @@ def test_eddsa_signed_token_is_accepted_and_logged_as_deprecated_outside_fips_mo def test_rs256_signed_token_is_accepted_without_deprecation_log(gateway: Gateway, tmp_path: Path) -> None: private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) - public_jwk: Final = jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key()) + public_jwk: Final = RSAAlgorithm.to_jwk(private_key.public_key()) def respond(request: Request) -> Reply: assert request.method == "GET", request @@ -97,3 +140,548 @@ def test_rs256_signed_token_is_accepted_without_deprecation_log(gateway: Gateway ) assert response.status_code == 200, response.text assert "EdDSA" not in owned.log.read_text(), owned.log.read_text() + + +def _deprecation_lines(log_text: str) -> list[str]: + return [line for line in log_text.splitlines() if all(part in line for part in APPROVED_WARNING_PARTS)] + + +def _alg_lied_token(private_key: ed25519.Ed25519PrivateKey, kid: str, subject: str) -> str: + def segment(value: bytes) -> str: + return base64.urlsafe_b64encode(value).rstrip(b"=").decode() + + signing_input: Final = ( + f"{segment(json.dumps({'alg': 'RS256', 'typ': 'JWT', 'kid': kid}).encode())}." + f"{segment(json.dumps({'sub': subject, 'iat': int(time.time()), 'exp': int(time.time()) + 300}).encode())}" + ) + signature: Final = segment(private_key.sign(signing_input.encode())) + return f"{signing_input}.{signature}" + + +def test_eddsa_non_stream_request_reaches_upstream_and_logs_deprecation(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "EdDSA", subject, kid="ed") + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + marker: Final = f"jwt-allowlist-a1-{uuid.uuid4().hex}" + _observed_bodies() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": marker}]}, + key=token, + ) + assert response.status_code == 200, response.text + deprecation: Final = eventually( + lambda: owned.log.read_text(), + lambda text: _deprecation_lines(text) != [], + seconds=30, + ) + assert _deprecation_lines(deprecation) != [] + observed: Final = _observed_bodies() + assert any(marker in body for body in observed), observed + + +def test_eddsa_streamed_request_through_openai_sdk_reaches_upstream_and_logs_deprecation( + gateway: Gateway, tmp_path: Path +) -> None: + private_key, jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "EdDSA", subject, kid="ed") + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + marker: Final = f"jwt-allowlist-a3-{uuid.uuid4().hex}" + _observed_bodies() + sdk: Final = OpenAI(base_url=f"{owned.gateway.client.base_url}/v1", api_key=token, timeout=30) + chunks: Final = [ + chunk + for chunk in sdk.chat.completions.create( + model=model, messages=[{"role": "user", "content": marker}], stream=True + ) + ] + assert chunks, "stream produced no chunks" + deprecation: Final = eventually( + lambda: owned.log.read_text(), + lambda text: _deprecation_lines(text) != [], + seconds=30, + ) + assert _deprecation_lines(deprecation) != [] + observed: Final = _observed_bodies() + assert any(marker in body for body in observed), observed + + +def test_rs256_messages_request_through_anthropic_sdk_reaches_upstream(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _rsa_keypair() + completion: Final = json.dumps( + { + "id": "msg_allowlist_a4", + "type": "message", + "role": "assistant", + "model": "claude-3-5-haiku-20241022", + "content": [{"type": "text", "text": "a4-reply"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 5, "output_tokens": 2}, + } + ).encode() + + def jwks(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + def provider(request: Request) -> Reply: + return Reply(body=completion) + + with wire_server(jwks) as keys_server, wire_server(provider) as provider_server, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "RS256", subject, kid="rsa") + with owned_proxy_process( + gateway, + tmp_path, + {"JWT_PUBLIC_KEY_URL": keys_server.url}, + config=_jwt_auth_config(tmp_path), + workers=2, + ) as owned: + model: Final = scenario.model(model="anthropic/claude-3-5-haiku-latest", api_base=provider_server.url) + scenario.cleanups.callback(scenario.delete_user, subject) + sdk: Final = Anthropic( + base_url=str(owned.gateway.client.base_url).rstrip("/"), + auth_token=token, + timeout=30, + max_retries=0, + ) + message: Final = sdk.messages.create( + model=model, max_tokens=8, messages=[{"role": "user", "content": "a4"}] + ) + text: Final = message.content[0].text + assert "a4-reply" in text, message + received: Final = provider_server.drain() + assert received, "upstream never received the request" + + +def test_rs256_responses_request_through_openai_async_sdk_reaches_upstream(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _rsa_keypair() + response_object: Final = json.dumps( + { + "id": "resp_allowlist_a5", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": "msg_a5", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "a5-reply", "annotations": []}], + } + ], + "usage": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7}, + } + ).encode() + + def jwks(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + def provider(request: Request) -> Reply: + return Reply(body=response_object) + + with wire_server(jwks) as keys_server, wire_server(provider) as provider_server, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "RS256", subject, kid="rsa") + with owned_proxy_process( + gateway, + tmp_path, + {"JWT_PUBLIC_KEY_URL": keys_server.url}, + config=_jwt_auth_config(tmp_path), + workers=2, + ) as owned: + model: Final = scenario.model(api_base=provider_server.url) + scenario.cleanups.callback(scenario.delete_user, subject) + + async def call() -> str: + sdk: Final = AsyncOpenAI(base_url=f"{owned.gateway.client.base_url}/v1", api_key=token, timeout=30) + response: Final = await sdk.responses.create(model=model, input="a5") + return response.id + + identity: Final = asyncio.run(call()) + assert identity.startswith("resp_"), identity + received: Final = provider_server.drain() + assert received, "upstream never received the request" + + +def test_es256_signed_token_is_accepted(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _ec_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + private_key, + algorithm="ES256", + headers={"kid": "ec"}, + ) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist es256"}]}, + key=token, + ) + assert response.status_code == 200, response.text + + +def test_hs256_token_against_oct_jwks_key_is_rejected(gateway: Gateway, tmp_path: Path) -> None: + secret: Final = b"integration-hs256-jwt-secret-0123456789abcdef" + oct_key: Final = { + "kty": "oct", + "kid": "sym", + "alg": "HS256", + "k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(), + } + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([oct_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + secret, + algorithm="HS256", + headers={"kid": "sym"}, + ) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist hs256"}]}, + key=token, + ) + assert response.status_code == 401, response.text + assert "error" in response.text, response.text + + +def test_token_signed_eddsa_with_rs256_header_is_rejected(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _alg_lied_token(private_key, "ed", subject) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist lied"}]}, + key=token, + ) + assert response.status_code == 401, response.text + + +def test_request_without_authorization_header_is_rejected(gateway: Gateway, tmp_path: Path) -> None: + _, jwks_key = _rsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.client.post( + "/v1/chat/completions", + json={"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist noauth"}]}, + ) + assert response.status_code == 401, response.text + + +def test_empty_jwks_document_rejects_tokens(gateway: Gateway, tmp_path: Path) -> None: + private_key, _jwks_key = _rsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + private_key, + algorithm="RS256", + headers={"kid": "rsa"}, + ) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist empty"}]}, + key=token, + ) + assert response.status_code == 401, response.text + + +def test_five_eddsa_requests_warn_at_most_once_per_worker(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "EdDSA", subject, kid="ed") + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + for _ in range(5): + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist burst"}]}, + key=token, + ) + assert response.status_code == 200, response.text + log_text: Final = eventually( + lambda: owned.log.read_text(), + lambda text: _deprecation_lines(text) != [], + seconds=30, + ) + warnings: Final = _deprecation_lines(log_text) + assert 1 <= len(warnings) <= 2, warnings + + +def test_master_key_request_still_works_on_jwt_enabled_proxy(gateway: Gateway, tmp_path: Path) -> None: + _, jwks_key = _rsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist master"}]}, + ) + assert response.status_code == 200, response.text + + +def _spend_row_counts(identifiers: set[str]) -> dict[str, int]: + if not identifiers: + return {} + rows: Final = read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s::text[])', + ("{" + ",".join(sorted(identifiers)) + "}",), + ) + counts: Final = {identifier: 0 for identifier in identifiers} + for row in rows: + identifier = str(row["request_id"]) + if identifier in counts: + counts[identifier] += 1 + return counts + + +def test_mixed_jwt_burst_survives_jwks_restart_and_writes_spend_rows(gateway: Gateway, tmp_path: Path) -> None: + rsa_private, rsa_jwks_key = _rsa_keypair() + ed_private, ed_jwks_key = _eddsa_keypair() + port: Final = _free_port() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([rsa_jwks_key, ed_jwks_key]) + + with gateway.scenario() as scenario: + rs_token: Final = jwt.encode( + {"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900}, + rsa_private, + algorithm="RS256", + headers={"kid": "rsa"}, + ) + ed_token: Final = jwt.encode( + {"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900}, + ed_private, + algorithm="EdDSA", + headers={"kid": "ed"}, + ) + with owned_proxy_process( + gateway, + tmp_path, + {"JWT_PUBLIC_KEY_URL": f"http://127.0.0.1:{port}/jwks"}, + config=_jwt_auth_config(tmp_path), + workers=2, + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback( + scenario.delete_user, str(jwt.decode(rs_token, options={"verify_signature": False})["sub"]) + ) + scenario.cleanups.callback( + scenario.delete_user, str(jwt.decode(ed_token, options={"verify_signature": False})["sub"]) + ) + + def chat(token: str, index: int, stream: bool = False) -> httpx.Response: + with owned.gateway.client.stream( + "POST", + "/v1/chat/completions", + json={ + "model": model, + "messages": [{"role": "user", "content": f"d1-{index}"}], + **({"stream": True} if stream else {}), + }, + headers={"Authorization": f"Bearer {token}"}, + ) as response: + return httpx.Response( + status_code=response.status_code, + headers=dict(response.headers), + content=response.read(), + ) + + with wire_server(respond, port=port): + warmup_rs: Final = chat(rs_token, 0) + warmup_ed: Final = chat(ed_token, 1) + assert warmup_rs.status_code == 200, warmup_rs.text + assert warmup_ed.status_code == 200, warmup_ed.text + + statuses: Final = [] + bodies: Final = [] + with ThreadPoolExecutor(max_workers=15) as pool: + futures = [ + pool.submit(chat, rs_token if index % 2 == 0 else ed_token, index, index % 5 == 0) + for index in range(30) + ] + for future in futures: + result = future.result(timeout=60) + statuses.append(result.status_code) + bodies.append(result.text) + assert all(status in (200, 401) for status in statuses), (statuses, bodies[:3]) + assert any(status == 200 for status in statuses), statuses + + with wire_server(respond, port=port): + fresh_rs: Final = chat(rs_token, 100) + fresh_ed: Final = chat(ed_token, 101) + assert fresh_rs.status_code == 200, fresh_rs.text + assert fresh_ed.status_code == 200, fresh_ed.text + + succeeded: Final = { + json.loads(text)["id"] + for status, text in zip(statuses, bodies) + if status == 200 and "id" in text and text.strip().startswith("{") + } + succeeded.update(json.loads(text)["id"] for text in (fresh_rs.text, fresh_ed.text) if "id" in text) + counts: Final = eventually( + lambda: _spend_row_counts(succeeded), + lambda value: value == {identifier: 1 for identifier in succeeded}, + seconds=70, + ) + assert all(count == 1 for count in counts.values()), counts + + +def test_burst_survives_killed_worker_and_stays_alive(gateway: Gateway, tmp_path: Path) -> None: + rsa_private, rsa_jwks_key = _rsa_keypair() + ed_private, ed_jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([rsa_jwks_key, ed_jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + rs_token: Final = jwt.encode( + {"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900}, + rsa_private, + algorithm="RS256", + headers={"kid": "rsa"}, + ) + ed_token: Final = jwt.encode( + {"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900}, + ed_private, + algorithm="EdDSA", + headers={"kid": "ed"}, + ) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback( + scenario.delete_user, str(jwt.decode(rs_token, options={"verify_signature": False})["sub"]) + ) + scenario.cleanups.callback( + scenario.delete_user, str(jwt.decode(ed_token, options={"verify_signature": False})["sub"]) + ) + + def chat(token: str, index: int) -> httpx.Response: + return owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"d2-{index}"}]}, + key=token, + ) + + with ThreadPoolExecutor(max_workers=12) as pool: + futures = [pool.submit(chat, rs_token if index % 2 == 0 else ed_token, index) for index in range(24)] + children: Final = psutil.Process(owned.process.pid).children(recursive=True) + assert children, "owned proxy reported no worker children" + children[0].kill() + results: Final = [] + for future in futures: + try: + result = future.result(timeout=60) + results.append(result.status_code) + except httpx.TransportError: + results.append(0) + assert all(status in (200, 401, 0) for status in results), results + assert any(status == 200 for status in results), results + liveliness: Final = eventually( + lambda: owned.gateway.client.get("/health/liveliness"), + lambda response: response.status_code == 200, + seconds=45, + ) + assert liveliness.status_code == 200, liveliness.text diff --git a/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py b/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py index c2c32f9866b..d00c51d1156 100644 --- a/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py +++ b/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py @@ -1,14 +1,18 @@ import json import os import uuid +from pathlib import Path from typing import Final +import httpx import psycopg +import pytest from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import rsa from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.mcp import forget_mcp, mcp_peer, register_mcp +from integration._support.process import owned_proxy_process def _pem_private_key() -> str: @@ -110,11 +114,164 @@ def test_server_row_with_stale_client_assertion_signing_alg_still_loads(gateway: register_mcp(scenario, peer, "staletrigger" + uuid.uuid4().hex[:8]) listed: Final = eventually( lambda: gateway.request("GET", "/v1/mcp/server"), - lambda response: response.status_code == 200 - and any(server.get("server_id") == identity for server in response.json()), + lambda response: ( + response.status_code == 200 + and any(server.get("server_id") == identity for server in response.json()) + ), seconds=30, ) assert listed.status_code == 200, listed.text finally: with psycopg.connect(os.environ["DATABASE_URL"]) as connection: connection.execute('DELETE FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,)) + + +APPROVED_ALGS: Final = ("RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512") + + +def _post_server(gateway: Gateway, name: str, credentials: object, omit_alg: bool = False) -> httpx.Response: + blob: Final = {} if omit_alg else {"client_assertion_signing_alg": credentials} + return gateway.request( + "POST", + "/v1/mcp/server", + { + "server_name": name, + "alias": name, + "url": "https://mcp.integration.invalid", + "transport": "http", + "credentials": blob, + }, + ) + + +def test_post_with_lowercase_hs256_is_rejected_naming_the_field(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + created: Final = _post_server(gateway, "b1hs" + uuid.uuid4().hex[:8], "hs256") + if created.status_code == 201: + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + assert created.status_code == 422, created.text + assert "client_assertion_signing_alg" in created.text, created.text + for algorithm in APPROVED_ALGS: + assert algorithm in created.text, created.text + + +def test_put_with_eddsa_on_existing_server_is_rejected(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "b2ed" + uuid.uuid4().hex[:8]) + edited: Final = gateway.request( + "PUT", + "/v1/mcp/server", + {"server_id": identity, "credentials": {"client_assertion_signing_alg": "EdDSA"}}, + ) + assert edited.status_code == 422, edited.text + assert "client_assertion_signing_alg" in edited.text, edited.text + + +@pytest.mark.parametrize("algorithm", APPROVED_ALGS) +def test_each_approved_client_assertion_signing_alg_is_accepted_and_stored(gateway: Gateway, algorithm: str) -> None: + with gateway.scenario() as scenario: + name: Final = "b3" + algorithm.lower() + uuid.uuid4().hex[:6] + created: Final = _post_server(gateway, name, algorithm) + assert created.status_code == 201, created.text + identity: Final = str(created.json()["server_id"]) + scenario.cleanups.callback(forget_mcp, gateway, identity) + listed: Final = gateway.request("GET", "/v1/mcp/server") + assert listed.status_code == 200, listed.text + assert any(server.get("server_id") == identity for server in listed.json()), listed.text + assert _stored_client_assertion_signing_alg(identity) == algorithm + + +@pytest.mark.parametrize( + ("label", "value", "expected"), + ( + ("empty", "", 422), + ("five-kb", "x" * 5120, 422), + ("integer", 7, 422), + ("list", ["RS256"], 422), + ("null", None, 201), + ), +) +def test_client_assertion_signing_alg_payload_variants( + gateway: Gateway, label: str, value: object, expected: int +) -> None: + with gateway.scenario() as scenario: + created: Final = _post_server(gateway, f"b{label}" + uuid.uuid4().hex[:8], value) + if created.status_code == 201: + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + assert created.status_code == expected, created.text + if expected == 422 and label in ("empty", "five-kb"): + assert "client_assertion_signing_alg" in created.text, created.text + + +def test_server_post_without_credentials_alg_key_is_accepted(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + created: Final = _post_server(gateway, "b9absent" + uuid.uuid4().hex[:8], None, omit_alg=True) + assert created.status_code == 201, created.text + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + + +def test_unauthenticated_server_post_with_hs256_is_rejected(gateway: Gateway) -> None: + response: Final = gateway.client.post( + "/v1/mcp/server", + json={ + "server_name": "b10unauth" + uuid.uuid4().hex[:8], + "credentials": {"client_assertion_signing_alg": "hs256"}, + }, + ) + assert response.status_code == 401, response.text + + +def _seeded_row_loads_with_fallback_warning(gateway: Gateway, tmp_path: Path, algorithm: str) -> None: + identity: Final = "dualread" + uuid.uuid4().hex[:8] + name: Final = "dualread" + uuid.uuid4().hex[:8] + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute( + 'INSERT INTO "LiteLLM_MCPServerTable" (server_id, server_name, url, transport, credentials,' + " created_at, updated_at) VALUES (%s, %s, %s, %s, %s::jsonb, NOW(), NOW())", + ( + identity, + name, + "https://mcp.integration.invalid", + "http", + json.dumps({"client_assertion_signing_alg": algorithm}), + ), + ) + try: + with owned_proxy_process(gateway, tmp_path, {}) as owned: + listed: Final = eventually( + lambda: owned.gateway.request("GET", "/v1/mcp/server"), + lambda response: ( + response.status_code == 200 + and any(server.get("server_id") == identity for server in response.json()) + ), + seconds=30, + ) + assert listed.status_code == 200, listed.text + log_text: Final = eventually( + lambda: owned.log.read_text(), + lambda text: name in text and "not an approved algorithm" in text and "using RS256" in text, + seconds=30, + ) + assert name in log_text and "not an approved algorithm" in log_text and "using RS256" in log_text + finally: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute('DELETE FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,)) + + +def test_seeded_hs256_row_loads_with_rs256_fallback_warning(gateway: Gateway, tmp_path: Path) -> None: + _seeded_row_loads_with_fallback_warning(gateway, tmp_path, "HS256") + + +def test_seeded_eddsa_row_loads_with_rs256_fallback_warning(gateway: Gateway, tmp_path: Path) -> None: + _seeded_row_loads_with_fallback_warning(gateway, tmp_path, "EdDSA") + + +def test_chat_still_works_after_alg_rejection(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + created: Final = _post_server(gateway, "b13chat" + uuid.uuid4().hex[:8], "hs256") + assert created.status_code in (201, 422), created.text + if created.status_code == 201: + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + model: Final = scenario.model() + body: Final = gateway.chat(model) + assert body["id"], body diff --git a/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py b/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py index 8af03fe0663..2d1d48a7a95 100644 --- a/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py +++ b/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py @@ -1,22 +1,31 @@ import base64 import json +import os +import signal +import socket +import subprocess +import sys import time import uuid -from collections.abc import Callable, Iterator, Mapping +from collections.abc import Callable, Generator, Mapping +from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager +from pathlib import Path from typing import Final import httpx import jwt import pytest from cryptography.hazmat.primitives.asymmetric import ed25519, rsa -from integration._support.client import Gateway +from integration._support.client import Gateway, eventually from integration._support.mcp import mcp_peer, register_mcp, tool_calls, tool_names +from integration._support.process import owned_proxy_process from integration._support.wire import Reply, Request, wire_server +from jwt.algorithms import OKPAlgorithm, RSAAlgorithm @contextmanager -def _jwt_signer_guardrail(gateway: Gateway, discovery_uri: str) -> Iterator[None]: +def _jwt_signer_guardrail(gateway: Gateway, discovery_uri: str) -> Generator[None]: created: Final = gateway.client.post( "/guardrails", headers={"x-litellm-api-key": gateway.key}, @@ -42,7 +51,7 @@ def _jwt_signer_guardrail(gateway: Gateway, discovery_uri: str) -> Iterator[None @contextmanager -def _idp_server(jwks_keys: list[Mapping[str, object]]) -> Iterator[str]: +def _idp_server(jwks_keys: list[Mapping[str, object]]) -> Generator[str]: holder: Final = {"url": ""} def respond(request: Request) -> Reply: @@ -76,7 +85,7 @@ def _hs256_key_and_token() -> tuple[dict[str, object], str]: def _eddsa_key_and_token() -> tuple[dict[str, object], str]: private_key: Final = ed25519.Ed25519PrivateKey.generate() - public_jwk: Final = json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(private_key.public_key())) + public_jwk: Final = json.loads(OKPAlgorithm.to_jwk(private_key.public_key())) key: Final = {**public_jwk, "alg": "EdDSA", "kid": "ed"} token: Final = jwt.encode(_claims(), private_key, algorithm="EdDSA", headers={"kid": "ed"}) return key, token @@ -84,7 +93,7 @@ def _eddsa_key_and_token() -> tuple[dict[str, object], str]: def _rs256_key_and_token() -> tuple[dict[str, object], str]: private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) - public_jwk: Final = json.loads(jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key())) + public_jwk: Final = json.loads(RSAAlgorithm.to_jwk(private_key.public_key())) key: Final = {**public_jwk, "alg": "RS256", "kid": "rsa"} token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"}) return key, token @@ -131,3 +140,142 @@ def test_jwks_rs256_key_still_verifies_the_incoming_token(gateway: Gateway) -> N assert response.status_code == 200, response.text assert response.json()["content"][0]["text"] == "9", response.text assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer" + + +def _rs256_keypair() -> tuple[rsa.RSAPrivateKey, dict[str, object]]: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + return private_key, json.loads(RSAAlgorithm.to_jwk(private_key.public_key())) + + +def _rs256_key_and_token_no_alg() -> tuple[dict[str, object], str]: + private_key, public_jwk = _rs256_keypair() + key: Final = {**public_jwk, "kid": "rsa"} + token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"}) + return key, token + + +def _signer_call(gateway: Gateway, key: str, identity: str, name: str, bearer: str) -> tuple[int, str]: + try: + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, bearer) + return response.status_code, response.text + except httpx.HTTPError as error: + return 0, repr(error) + + +def test_jwks_rs256_key_without_alg_field_still_verifies_the_token(gateway: Gateway) -> None: + jwks_key, token = _rs256_key_and_token_no_alg() + with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri): + with mcp_peer() as peer, gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "jwksc4" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "9", response.text + assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer" + + +def test_rs256_token_verifies_when_jwks_also_carries_a_non_approved_key(gateway: Gateway) -> None: + okp_key, _eddsa_token = _eddsa_key_and_token() + jwks_key, token = _rs256_key_and_token() + with _idp_server([okp_key, jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri): + with mcp_peer() as peer, gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "jwksc5" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "9", response.text + assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer" + + +@pytest.mark.parametrize( + ("label", "jwks_reply"), + ( + ("jwks-404", Reply(status=404, body=b'{"error": "missing"}')), + ("jwks-not-json", Reply(body=b"this is not json")), + ), + ids=("jwks-404", "jwks-not-json"), +) +def test_unusable_jwks_document_rejects_the_incoming_token(gateway: Gateway, label: str, jwks_reply: Reply) -> None: + _, token = _rs256_key_and_token() + + def respond(request: Request) -> Reply: + if request.target == "/.well-known/openid-configuration": + return Reply(body=json.dumps({"jwks_uri": holder["url"] + "/.well-known/jwks.json"}).encode()) + return jwks_reply + + holder: Final = {"url": ""} + with wire_server(respond) as idp: + holder["url"] = idp.url + with _jwt_signer_guardrail(gateway, idp.url + "/.well-known/openid-configuration"): + with mcp_peer() as peer, gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "jwks" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + status, text = _signer_call(gateway, key, identity, name, token) + assert status == 401, (status, text) + assert "incoming token verification failed" in text, text + assert tool_calls(peer.drain()) == (), "rejected call reached the peer" + alive: Final = gateway.client.get("/health/liveliness") + assert alive.status_code == 200, alive.text + + +def test_signer_jwks_pause_rejects_then_recovers(gateway: Gateway, tmp_path: Path) -> None: + private_key, public_jwk = _rs256_keypair() + jwks_dir: Final = tmp_path / "jwks" / ".well-known" + jwks_dir.mkdir(parents=True) + port: Final = _reserve_port() + (jwks_dir / "jwks.json").write_text(json.dumps({"keys": [{**public_jwk, "kid": "rsa", "alg": "RS256"}]})) + (jwks_dir / "openid-configuration").write_text( + json.dumps({"jwks_uri": f"http://127.0.0.1:{port}/.well-known/jwks.json"}) + ) + token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"}) + server: Final = subprocess.Popen( + [sys.executable, "-m", "http.server", str(port), "--bind", "127.0.0.1", "--directory", str(tmp_path / "jwks")], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + try: + with owned_proxy_process(gateway, tmp_path, {}) as owned: + discovery: Final = f"http://127.0.0.1:{port}/.well-known/openid-configuration" + with _jwt_signer_guardrail(owned.gateway, discovery): + with mcp_peer() as peer, owned.gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "jwksd3" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(owned.gateway, key, identity)["add"] + peer.drain() + os.kill(server.pid, signal.SIGSTOP) + outcomes: Final = [] + with ThreadPoolExecutor(max_workers=8) as pool: + futures = [ + pool.submit(_signer_call, owned.gateway, key, identity, name, token) for _ in range(8) + ] + for future in futures: + outcomes.append(future.result(timeout=80)) + assert all(status != 200 for status, _text in outcomes), outcomes + assert any(status in (0, 401, 500) for status, _text in outcomes), outcomes + os.kill(server.pid, signal.SIGCONT) + recovered: Final = eventually( + lambda: _signer_call(owned.gateway, key, identity, name, token), + lambda outcome: outcome[0] == 200, + seconds=70, + ) + assert recovered[0] == 200, recovered + assert tool_calls(peer.drain()) != (), "resumed call never reached the peer" + finally: + try: + os.kill(server.pid, signal.SIGCONT) + except ProcessLookupError: + pass + server.terminate() + server.wait(timeout=10) + + +def _reserve_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] From f53a005254c81fadcea58f4dd5022544c05d8001 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 12:52:31 +0000 Subject: [PATCH 13/18] test(integration): isolate the http.server child interpreter in the JWKS pause cell (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp/test_mcp_jwt_signer_jwks_allowlist.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py b/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py index 2d1d48a7a95..4d75ba60d71 100644 --- a/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py +++ b/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py @@ -235,7 +235,17 @@ def test_signer_jwks_pause_rejects_then_recovers(gateway: Gateway, tmp_path: Pat ) token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"}) server: Final = subprocess.Popen( - [sys.executable, "-m", "http.server", str(port), "--bind", "127.0.0.1", "--directory", str(tmp_path / "jwks")], + [ + sys.executable, + "-I", + "-m", + "http.server", + str(port), + "--bind", + "127.0.0.1", + "--directory", + str(tmp_path / "jwks"), + ], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) From 27737bd38cf97b5b24bb7f60da0ccb00e75c5e32 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 14:15:34 +0000 Subject: [PATCH 14/18] test(integration): tighten JWT chaos cells (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_jwt_algorithm_allowlist.py | 54 +++++++++++++++---- 1 file changed, 44 insertions(+), 10 deletions(-) diff --git a/tests/integration/authorization/test_jwt_algorithm_allowlist.py b/tests/integration/authorization/test_jwt_algorithm_allowlist.py index 93dc98c17c4..44a5df0505e 100644 --- a/tests/integration/authorization/test_jwt_algorithm_allowlist.py +++ b/tests/integration/authorization/test_jwt_algorithm_allowlist.py @@ -516,6 +516,23 @@ def test_master_key_request_still_works_on_jwt_enabled_proxy(gateway: Gateway, t assert response.status_code == 200, response.text +def _response_ids(status: int, text: str) -> tuple[str, ...]: + if status != 200: + return () + stripped: Final = text.strip() + if stripped.startswith("{"): + body: Final = json.loads(stripped) + return (str(body["id"]),) if "id" in body else () + ids: Final = { + str(chunk["id"]) + for line in stripped.splitlines() + if line.startswith("data: ") and line[len("data: ") :].strip() != "[DONE]" + for chunk in (json.loads(line[len("data: ") :]),) + if "id" in chunk + } + return tuple(ids) + + def _spend_row_counts(identifiers: set[str]) -> dict[str, int]: if not identifiers: return {} @@ -591,8 +608,8 @@ def test_mixed_jwt_burst_survives_jwks_restart_and_writes_spend_rows(gateway: Ga assert warmup_rs.status_code == 200, warmup_rs.text assert warmup_ed.status_code == 200, warmup_ed.text - statuses: Final = [] - bodies: Final = [] + statuses: Final[list[int]] = [] + bodies: Final[list[str]] = [] with ThreadPoolExecutor(max_workers=15) as pool: futures = [ pool.submit(chat, rs_token if index % 2 == 0 else ed_token, index, index % 5 == 0) @@ -612,11 +629,13 @@ def test_mixed_jwt_burst_survives_jwks_restart_and_writes_spend_rows(gateway: Ga assert fresh_ed.status_code == 200, fresh_ed.text succeeded: Final = { - json.loads(text)["id"] - for status, text in zip(statuses, bodies) - if status == 200 and "id" in text and text.strip().startswith("{") + identifier for status, text in zip(statuses, bodies) for identifier in _response_ids(status, text) } - succeeded.update(json.loads(text)["id"] for text in (fresh_rs.text, fresh_ed.text) if "id" in text) + succeeded.update(_response_ids(200, fresh_rs.text) + _response_ids(200, fresh_ed.text)) + assert len(succeeded) == sum(1 for status in statuses if status == 200) + 2, ( + statuses, + bodies[:3], + ) counts: Final = eventually( lambda: _spend_row_counts(succeeded), lambda value: value == {identifier: 1 for identifier in succeeded}, @@ -667,10 +686,19 @@ def test_burst_survives_killed_worker_and_stays_alive(gateway: Gateway, tmp_path with ThreadPoolExecutor(max_workers=12) as pool: futures = [pool.submit(chat, rs_token if index % 2 == 0 else ed_token, index) for index in range(24)] - children: Final = psutil.Process(owned.process.pid).children(recursive=True) - assert children, "owned proxy reported no worker children" - children[0].kill() - results: Final = [] + owned_port: Final = owned.gateway.client.base_url.port + listeners: Final = [ + child + for child in psutil.Process(owned.process.pid).children(recursive=True) + if any( + connection.status == psutil.CONN_LISTEN and connection.laddr.port == owned_port + for connection in child.net_connections(kind="inet") + ) + ] + assert len(listeners) >= 2, listeners + killed_pid: Final = listeners[0].pid + listeners[0].kill() + results: Final[list[int]] = [] for future in futures: try: result = future.result(timeout=60) @@ -679,6 +707,12 @@ def test_burst_survives_killed_worker_and_stays_alive(gateway: Gateway, tmp_path results.append(0) assert all(status in (200, 401, 0) for status in results), results assert any(status == 200 for status in results), results + dead: Final = eventually( + lambda: psutil.pid_exists(killed_pid), + lambda exists: not exists, + seconds=30, + ) + assert not dead, f"killed worker pid {killed_pid} still exists" liveliness: Final = eventually( lambda: owned.gateway.client.get("/health/liveliness"), lambda response: response.status_code == 200, From a5e79564cb9a0cfe2a86bd1ea3a0a1f05872dd68 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 15:59:08 +0000 Subject: [PATCH 15/18] test(integration): apply the premium entitlement patch at import so spawned workers inherit it (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/_support/proxy.py | 17 +++++++++++------ 1 file changed, 11 insertions(+), 6 deletions(-) diff --git a/tests/integration/_support/proxy.py b/tests/integration/_support/proxy.py index a444b93757d..19d479085a9 100644 --- a/tests/integration/_support/proxy.py +++ b/tests/integration/_support/proxy.py @@ -1,11 +1,19 @@ -"""Run the normal single-process CLI with the existing behavior-suite test entitlement.""" +"""Run the normal single-process CLI with the behavior-suite test entitlement, applied at import so spawned workers inherit it.""" import signal import sys from types import FrameType +from typing import Final from unittest.mock import patch -from litellm import run_server +_PREMIUM: Final = ( + patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts + "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True + ) +) +_PREMIUM.start() + +from litellm import run_server # noqa: E402 def _exit_on_reraised_term(signum: int, frame: FrameType | None) -> None: @@ -14,10 +22,7 @@ def _exit_on_reraised_term(signum: int, frame: FrameType | None) -> None: def main() -> None: signal.signal(signal.SIGTERM, _exit_on_reraised_term) - with patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts - "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True - ): - run_server() + run_server() if __name__ == "__main__": From 1c795e36e7950ec9f1283a5b5e83136139fe9cdd Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 15:59:45 +0000 Subject: [PATCH 16/18] test(integration): keep the run_server import at the top of the entitlement shim (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/_support/proxy.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/integration/_support/proxy.py b/tests/integration/_support/proxy.py index 19d479085a9..c6ae4f30375 100644 --- a/tests/integration/_support/proxy.py +++ b/tests/integration/_support/proxy.py @@ -6,6 +6,8 @@ from types import FrameType from typing import Final from unittest.mock import patch +from litellm import run_server + _PREMIUM: Final = ( patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True @@ -13,8 +15,6 @@ _PREMIUM: Final = ( ) _PREMIUM.start() -from litellm import run_server # noqa: E402 - def _exit_on_reraised_term(signum: int, frame: FrameType | None) -> None: sys.exit(0) From b20805b719af6892696f09bbc50a1a84dc8edd5e Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 5 Oct 2026 08:15:39 +0000 Subject: [PATCH 17/18] fix(auth): drop unused mutable-ok suppressions flagged by LIT013 (LIT-8429) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/handle_jwt.py | 6 +++--- .../guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py | 4 +--- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 1b360035bf7..8dac077fe3d 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -233,7 +233,7 @@ class JWTHandler: # Supported algos: https://pyjwt.readthedocs.io/en/stable/algorithms.html # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret # the key in different ways (e.g. HS* and RS*)." - SUPPORTED_JWT_ALGORITHMS = [ # mutable-ok: list kept for backward compatibility + SUPPORTED_JWT_ALGORITHMS = [ *APPROVED_JWT_ALGORITHMS, *LEGACY_JWT_ALGORITHMS, ] @@ -975,9 +975,9 @@ class JWTHandler: allowed: Final = self.allowed_algorithms() usable_keys: Final[JWKKeyValue] = ( - list(jwks_keys_for(keys, allowed)) # mutable-ok: parse_keys consumes a JWKKeyValue list + list(jwks_keys_for(keys, allowed)) if isinstance(keys, list) - else next(iter(jwks_keys_for((keys,), allowed)), {}) # mutable-ok: single-key dict is a JWKKeyValue + else next(iter(jwks_keys_for((keys,), allowed)), {}) ) public_key: Final = self.parse_keys(keys=usable_keys, kid=kid) if public_key is not None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py index 9256d0fe5f3..75ce60ebacf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py @@ -463,9 +463,7 @@ class MCPJWTSigner(CustomGuardrail): from jwt import PyJWKSet try: - jwks_set: Final = PyJWKSet.from_dict( - {"keys": list(approved_keys)} # mutable-ok: PyJWKSet.from_dict takes a dict - ) + jwks_set: Final = PyJWKSet.from_dict({"keys": list(approved_keys)}) except Exception as exc: raise jwt.exceptions.PyJWKSetError(f"Failed to parse JWKS from {jwks_uri!r}: {exc}") from exc From 4713920478de00ef8611a5570251245b9441785e Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 5 Oct 2026 12:46:03 +0000 Subject: [PATCH 18/18] test(integration): take main's license forwarding for owned proxies instead of the entitlement patch Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/_support/proxy.py | 11 ----------- 1 file changed, 11 deletions(-) diff --git a/tests/integration/_support/proxy.py b/tests/integration/_support/proxy.py index c6ae4f30375..610d4724349 100644 --- a/tests/integration/_support/proxy.py +++ b/tests/integration/_support/proxy.py @@ -1,20 +1,9 @@ -"""Run the normal single-process CLI with the behavior-suite test entitlement, applied at import so spawned workers inherit it.""" - import signal import sys from types import FrameType -from typing import Final -from unittest.mock import patch from litellm import run_server -_PREMIUM: Final = ( - patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts - "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True - ) -) -_PREMIUM.start() - def _exit_on_reraised_term(signum: int, frame: FrameType | None) -> None: sys.exit(0)