From a50f9fa2a8865630962d6f01001012ccfbad6c52 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Mon, 28 Sep 2026 15:10:18 -0700 Subject: [PATCH 1/6] fix(proxy): encrypt pass-through endpoint header values at rest Header values on DB-stored pass-through endpoints are written with the litellm_enc:: marker used for callback vars, decrypted where the route resolves its outbound headers, and re-encrypted by in-app master key rotation. Legacy plaintext rows and os.environ/ references keep working. --- .../key_management_endpoints.py | 15 + .../pass_through_endpoints/common_utils.py | 102 +++++ .../pass_through_endpoints.py | 79 +++- litellm/proxy/proxy_server.py | 5 +- .../test_key_management_endpoints.py | 100 +++++ .../test_pass_through_endpoints.py | 399 ++++++++++++++++++ ...test_passthrough_endpoints_common_utils.py | 123 ++++++ tests/test_litellm/proxy/test_proxy_server.py | 61 +++ 8 files changed, 874 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 7e159ec90e7..6517b8b809b 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -66,6 +66,7 @@ from litellm.proxy.auth.auth_utils import ( enforce_batch_enqueued_token_limit_is_admin_only, enforce_output_token_estimates_are_admin_only, ) +from litellm.proxy.auth.master_key_boot_check import SALT_KEY_ENV_VAR from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( evict_and_broadcast, @@ -124,6 +125,7 @@ from litellm.proxy.management_helpers.team_member_permission_checks import ( TeamMemberPermissionChecks, ) from litellm.proxy.management_helpers.utils import management_endpoint_wrapper +from litellm.proxy.pass_through_endpoints.common_utils import reencrypt_general_settings_pass_through from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start from litellm.proxy.spend_tracking.spend_tracking_utils import _is_master_key from litellm.proxy.utils import ( @@ -131,6 +133,7 @@ from litellm.proxy.utils import ( ProxyLogging, _hash_token_if_needed, handle_exception_on_proxy, + invalidate_config_param, is_valid_api_key, ) from litellm.repositories.base_repository import BaseRepository @@ -5301,6 +5304,18 @@ async def _rotate_master_key( data={"param_value": prisma.Json(encrypted_env_vars)}, ) + for c in config: + if ( + c.param_name == "general_settings" + and os.getenv(SALT_KEY_ENV_VAR) is None + and (reencrypted := reencrypt_general_settings_pass_through(c.param_value, new_master_key)) is not None + ): + await _config_table(prisma_client).update( + where={"param_name": "general_settings"}, + data={"param_value": prisma.Json(reencrypted)}, + ) + await invalidate_config_param("general_settings") + # 4. process MCP server table try: await rotate_mcp_server_credentials_master_key( diff --git a/litellm/proxy/pass_through_endpoints/common_utils.py b/litellm/proxy/pass_through_endpoints/common_utils.py index 3ae83cbe906..1550b43a887 100644 --- a/litellm/proxy/pass_through_endpoints/common_utils.py +++ b/litellm/proxy/pass_through_endpoints/common_utils.py @@ -1,6 +1,16 @@ +from collections.abc import Mapping from typing import Final from fastapi import Request +from pydantic import JsonValue + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper + +# Holds the name of the request header that carries the caller's LiteLLM key, which +# user_api_key_auth reads straight from general_settings, so it stays plaintext. +_CALLER_KEY_HEADER_NAME: Final = "litellm_user_api_key" def get_litellm_virtual_key(request: Request) -> str: @@ -16,3 +26,95 @@ def get_litellm_virtual_key(request: Request) -> str: if litellm_api_key: return f"Bearer {litellm_api_key}" return request.headers.get("Authorization", "") + + +def _encrypt_header_value(name: str, value: JsonValue, new_encryption_key: str | None) -> JsonValue: + if not isinstance(value, str) or not value or name == _CALLER_KEY_HEADER_NAME: + return value + if value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): + return value + try: + return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(value, new_encryption_key=new_encryption_key) + except Exception: # noqa: BLE001 # no salt or master key configured: store as before rather than fail the write + return value + + +def _decrypted(name: str, value: str) -> str | None: + return decrypt_value_helper( + value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), + key=name, + exception_type="debug", + return_original_value=False, + ) + + +def _decrypt_header_value(name: str, value: object) -> object: + if not isinstance(value, str) or not value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): + return value + decrypted: Final = _decrypted(name, value) + if decrypted is None: + verbose_proxy_logger.warning( + "Could not decrypt pass-through header %s; check LITELLM_SALT_KEY / master key", name + ) + return value + return decrypted + + +def undecryptable_pass_through_header_names(headers: object) -> frozenset[str]: + """Names of `litellm_enc::` header values that do not decrypt under the current key.""" + if not isinstance(headers, Mapping): + return frozenset() + return frozenset( + str(name) + for name, value in headers.items() + if isinstance(value, str) + and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) + and _decrypted(str(name), value) is None + ) + + +def decrypt_pass_through_headers(headers: Mapping[str, object] | None) -> dict[str, object] | None: + """Decrypt `litellm_enc::` header values; plaintext and os.environ/ values pass through.""" + if headers is None: + return None + return {name: _decrypt_header_value(name, value) for name, value in headers.items()} + + +def _with_encrypted_headers(endpoint: JsonValue, new_encryption_key: str | None, reencrypt: bool) -> JsonValue: + if not isinstance(endpoint, dict) or not isinstance(headers := endpoint.get("headers"), dict): + return endpoint + return { + **endpoint, + "headers": { + name: _encrypt_header_value( + name, _decrypt_header_value(name, value) if reencrypt else value, new_encryption_key + ) + for name, value in headers.items() + }, + } + + +def encrypt_pass_through_endpoints(endpoints: JsonValue) -> JsonValue: + """Encrypt every endpoint's header values for storage; values already carrying the marker are kept.""" + if not isinstance(endpoints, list): + return endpoints + return [_with_encrypted_headers(endpoint, None, reencrypt=False) for endpoint in endpoints] + + +def reencrypt_general_settings_pass_through( + general_settings: object, new_encryption_key: str +) -> dict[str, JsonValue] | None: + """Return general_settings with pass-through header values re-encrypted under new_encryption_key. + + None when the row holds no pass-through endpoint list. + """ + if not isinstance(general_settings, dict) or not isinstance( + endpoints := general_settings.get("pass_through_endpoints"), list + ): + return None + return { + **general_settings, + "pass_through_endpoints": [ + _with_encrypted_headers(endpoint, new_encryption_key, reencrypt=True) for endpoint in endpoints + ], + } diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e0a4184291e..f0d317d74c9 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -101,6 +101,10 @@ from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above ) +from litellm.proxy.pass_through_endpoints.common_utils import ( + decrypt_pass_through_headers, + undecryptable_pass_through_header_names, +) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path from litellm.repositories.team_repository import TeamRepository @@ -161,14 +165,15 @@ async def set_env_variables_in_header(custom_headers: dict | None) -> dict | Non """ if custom_headers is None: return None + stored_headers: Final = decrypt_pass_through_headers(custom_headers) or {} headers: Final = {} - for key, value in custom_headers.items(): + for key, value in stored_headers.items(): # langfuse Api requires base64 encoded headers - it's simpleer to just ask litellm users to set their langfuse public and secret keys # we can then get the b64 encoded keys here if key == "LANGFUSE_PUBLIC_KEY" or key == "LANGFUSE_SECRET_KEY": # langfuse requires b64 encoded headers - we construct that here - _langfuse_public_key = custom_headers["LANGFUSE_PUBLIC_KEY"] - _langfuse_secret_key = custom_headers["LANGFUSE_SECRET_KEY"] + _langfuse_public_key = stored_headers["LANGFUSE_PUBLIC_KEY"] + _langfuse_secret_key = stored_headers["LANGFUSE_SECRET_KEY"] if isinstance(_langfuse_public_key, str) and _langfuse_public_key.startswith("os.environ/"): _langfuse_public_key = get_secret_str(_langfuse_public_key) if isinstance(_langfuse_secret_key, str) and _langfuse_secret_key.startswith("os.environ/"): @@ -196,6 +201,59 @@ async def set_env_variables_in_header(custom_headers: dict | None) -> dict | Non return headers +def _served_custom_headers( + endpoint_id: str, target: str | None, path_without_stored_id: str | None +) -> Mapping[str, object] | None: + by_id: Final[list[Mapping[str, object]]] = [] + by_path: Final[dict[str, Mapping[str, object]]] = {} + for route in _registered_pass_through_routes.values(): + params = route.get("passthrough_params") + if ( + not isinstance(params, Mapping) + or params.get("target") != target + or not isinstance(headers := params.get("custom_headers"), Mapping) + ): + continue + served_headers = cast(Mapping[str, object], headers) + if route["endpoint_id"] == endpoint_id: + by_id.append(served_headers) + elif path_without_stored_id is not None and route.get("path") == path_without_stored_id: + by_path[str(route["endpoint_id"])] = served_headers + if by_id: + return by_id[0] + return next(iter(by_path.values())) if len(by_path) == 1 else None + + +async def _resolve_stored_headers( + endpoint_id: str, + target: str | None, + stored_headers: dict | None, + path_without_stored_id: str | None = None, +) -> dict | None: + """Resolve an endpoint's stored headers for outbound use. + + A header that no longer decrypts under the current key (between an in-app master key + rotation and the restart) keeps the value the endpoint already sends to the same target; + every other header resolves from the stored value. The served endpoint is matched by id, + or, for a stored endpoint without an id (it gets a new one on every reload), by path when + exactly one served endpoint has that path and target. + """ + undecryptable: Final = undecryptable_pass_through_header_names(stored_headers) + registered: Final = _served_custom_headers(endpoint_id, target, path_without_stored_id) if undecryptable else None + if stored_headers is None or registered is None: + return await set_env_variables_in_header(custom_headers=stored_headers) + verbose_proxy_logger.warning( + "Pass-through endpoint %s keeps its current value for headers %s: the stored values do not decrypt " + "under the current key (restart with the new master key after a rotation)", + endpoint_id, + sorted(undecryptable), + ) + if not undecryptable <= registered.keys(): + return dict(registered) + resolved: Final = await set_env_variables_in_header(custom_headers=stored_headers) + return {**(resolved or {}), **{name: registered[name] for name in undecryptable}} + + async def chat_completion_pass_through_endpoint( fastapi_response: Response, request: Request, @@ -3247,7 +3305,8 @@ async def _register_pass_through_endpoint( else: endpoint_data = endpoint - if endpoint_data.get("id") is None: + stored_without_id: Final = endpoint_data.get("id") is None + if stored_without_id: endpoint_data["id"] = str(uuid.uuid4()) endpoint_id: Final = cast(str, endpoint_data["id"]) @@ -3256,7 +3315,9 @@ async def _register_pass_through_endpoint( if path is None: raise ValueError("Path is required for pass-through endpoint") - custom_headers: Final = await set_env_variables_in_header(custom_headers=endpoint_data.get("headers")) + custom_headers: Final = await _resolve_stored_headers( + endpoint_id, target, endpoint_data.get("headers"), path if stored_without_id else None + ) forward_headers: Final = endpoint_data.get("forward_headers") merge_query_params: Final = endpoint_data.get("merge_query_params") default_query_params: Final = endpoint_data.get("default_query_params") @@ -3677,6 +3738,10 @@ async def update_pass_through_endpoints( # Update the list pass_through_endpoint_data[endpoint_index] = endpoint_dict + _custom_headers: Final = await _resolve_stored_headers( + endpoint_id, updated_endpoint.target, updated_endpoint.headers or {} + ) + # Remove old routes from registry before they get re-registered InitPassThroughEndpointHelpers.remove_endpoint_routes(endpoint_id) @@ -3689,10 +3754,6 @@ async def update_pass_through_endpoints( await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict) - # Re-register the route with updated headers - _custom_headers: dict | None = updated_endpoint.headers or {} - _custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers) - route_app: Final = _request_app(request) if updated_endpoint.include_subpath: InitPassThroughEndpointHelpers.add_subpath_route( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ae559f30857..af4dc450346 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -733,6 +733,7 @@ from litellm.proxy.openai_files_endpoints.files_endpoints import ( from litellm.proxy.openai_files_endpoints.files_endpoints import ( set_files_config, ) +from litellm.proxy.pass_through_endpoints.common_utils import encrypt_pass_through_endpoints from litellm.proxy.pass_through_endpoints.openai_passthrough_endpoints import ( router as openai_passthrough_router, ) @@ -17888,7 +17889,7 @@ async def update_config( existing["alerting"] = ["slack"] elif isinstance(existing["alerting"], list) and "slack" not in existing["alerting"]: existing["alerting"].append("slack") - existing[k] = v + existing[k] = encrypt_pass_through_endpoints(v) if k == "pass_through_endpoints" else v await _upsert_section("general_settings", existing) asyncio.create_task( create_config_audit_log( @@ -18149,6 +18150,8 @@ async def update_config_general_settings( field_value = data.field_value if data.field_name == "plugins": field_value = _preserve_redacted_plugin_keys(field_value, general_settings.get("plugins")) + if data.field_name == "pass_through_endpoints": + field_value = encrypt_pass_through_endpoints(field_value) general_settings[data.field_name] = cast(JsonValue, field_value) # cast-ok: ConfigGeneralSettings validated it diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index aa6be328f4a..d2e19b6e2ae 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21093,3 +21093,103 @@ class TestTeamAdminMemberKeyBudgetUpdate: ) assert exc.value.status_code == 403 assert "member_key_budgets" not in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_rotate_master_key_reencrypts_pass_through_endpoint_headers(monkeypatch): + import prisma + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints import key_management_endpoints + from litellm.proxy.pass_through_endpoints.common_utils import ( + decrypt_pass_through_headers, + encrypt_pass_through_endpoints, + ) + + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master-key") + stored_endpoints = encrypt_pass_through_endpoints( + [ + {"path": "/a", "headers": {"x-a": "plain-a"}}, + {"path": "/b", "headers": {"Authorization": "Bearer sk-literal"}}, + ] + ) + general_settings_row = MagicMock( + param_name="general_settings", + param_value={"store_model_in_db": True, "pass_through_endpoints": stored_endpoints}, + ) + mock_prisma_client = AsyncMock() + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[general_settings_row]) + mock_prisma_client.db.litellm_config.update = AsyncMock() + mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[]) + invalidate = AsyncMock() + monkeypatch.setattr(key_management_endpoints, "invalidate_config_param", invalidate) + for rotate_helper in ( + "rotate_mcp_server_credentials_master_key", + "rotate_mcp_user_credentials_master_key", + "rotate_mcp_user_env_vars_master_key", + "rotate_sso_identity_assertions_master_key", + ): + monkeypatch.setattr(key_management_endpoints, rotate_helper, AsyncMock()) + + await key_management_endpoints._rotate_master_key( + prisma_client=mock_prisma_client, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"), + current_master_key="sk-old-master-key", + new_master_key="sk-new-master-key", + ) + + mock_prisma_client.db.litellm_config.update.assert_awaited_once() + update_kwargs = mock_prisma_client.db.litellm_config.update.await_args.kwargs + assert update_kwargs["where"] == {"param_name": "general_settings"} + assert isinstance(update_kwargs["data"]["param_value"], prisma.Json) + rotated = update_kwargs["data"]["param_value"].data + assert rotated["store_model_in_db"] is True + invalidate.assert_awaited_once_with("general_settings") + + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-new-master-key") + assert [decrypt_pass_through_headers(e["headers"]) for e in rotated["pass_through_endpoints"]] == [ + {"x-a": "plain-a"}, + {"Authorization": "Bearer sk-literal"}, + ] + + +@pytest.mark.asyncio +async def test_rotate_master_key_leaves_pass_through_headers_under_salt_key(monkeypatch): + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints import key_management_endpoints + from litellm.proxy.pass_through_endpoints.common_utils import encrypt_pass_through_endpoints + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-stays") + general_settings_row = MagicMock( + param_name="general_settings", + param_value={ + "pass_through_endpoints": encrypt_pass_through_endpoints( + [{"path": "/a", "headers": {"Authorization": "Bearer sk-literal"}}] + ) + }, + ) + mock_prisma_client = AsyncMock() + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[general_settings_row]) + mock_prisma_client.db.litellm_config.update = AsyncMock() + mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[]) + for rotate_helper in ( + "rotate_mcp_server_credentials_master_key", + "rotate_mcp_user_credentials_master_key", + "rotate_mcp_user_env_vars_master_key", + "rotate_sso_identity_assertions_master_key", + ): + monkeypatch.setattr(key_management_endpoints, rotate_helper, AsyncMock()) + + await key_management_endpoints._rotate_master_key( + prisma_client=mock_prisma_client, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"), + current_master_key="sk-old-master-key", + new_master_key="sk-new-master-key", + ) + + mock_prisma_client.db.litellm_config.update.assert_not_awaited() diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 3469df082e0..0579488a34b 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -7622,3 +7622,402 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token() metadata = kwargs["litellm_params"]["metadata"] assert metadata["user_api_key"] == "cli-session-alice" assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice" + + +@pytest.mark.asyncio +async def test_set_env_variables_in_header_decrypts_stored_headers(monkeypatch): + from litellm.proxy.pass_through_endpoints.common_utils import encrypt_pass_through_endpoints + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import set_env_variables_in_header + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-pass-through-tests") + monkeypatch.setenv("UPSTREAM_TOKEN", "resolved-from-env") + [stored] = encrypt_pass_through_endpoints( + [ + { + "path": "/pt", + "headers": { + "x-legacy": "plain-value", + "Authorization": "Bearer sk-literal", + "x-env": "Bearer os.environ/UPSTREAM_TOKEN", + }, + } + ] + ) + headers = {**stored["headers"], "x-legacy": "plain-value"} + + assert await set_env_variables_in_header(custom_headers=headers) == { + "x-legacy": "plain-value", + "Authorization": "Bearer sk-literal", + "x-env": "Bearer resolved-from-env", + } + + +@pytest.mark.asyncio +async def test_register_pass_through_endpoint_keeps_serving_headers_when_stored_headers_do_not_decrypt(monkeypatch): + from fastapi import FastAPI + + from litellm.proxy.pass_through_endpoints.common_utils import ( + encrypt_pass_through_endpoints, + reencrypt_general_settings_pass_through, + ) + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _register_pass_through_endpoint, + _registered_pass_through_routes, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-current-key") + app = FastAPI() + [stored] = encrypt_pass_through_endpoints( + [ + { + "id": "ep-rotate", + "path": "/rotate-pt", + "target": "http://upstream", + "headers": {"Authorization": "Bearer sk-literal"}, + } + ] + ) + rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": [stored]}, "sk-next-key") + assert rotated is not None + [rotated_endpoint] = rotated["pass_through_endpoints"] + try: + await _register_pass_through_endpoint( + endpoint={"id": "ep-rotate-other", "path": "/other-pt", "target": "http://other", "headers": {"x": "y"}}, + app=app, + premium_user=False, + visited_endpoints=set(), + ) + first_visit: set[str] = set() + await _register_pass_through_endpoint( + endpoint=dict(stored), app=app, premium_user=False, visited_endpoints=first_visit + ) + [route_key] = first_visit + second_visit: set[str] = set() + await _register_pass_through_endpoint( + endpoint={ + **rotated_endpoint, + "path": "/rotate-pt-moved", + "headers": {**rotated_endpoint["headers"], "x-org": "acme"}, + }, + app=app, + premium_user=False, + visited_endpoints=second_visit, + ) + + [moved_route_key] = second_visit + assert moved_route_key != route_key + params = _registered_pass_through_routes[moved_route_key]["passthrough_params"] + assert params["custom_headers"] == {"Authorization": "Bearer sk-literal", "x-org": "acme"} + + retargeted_visit: set[str] = set() + await _register_pass_through_endpoint( + endpoint={**rotated_endpoint, "path": "/rotate-pt-moved", "target": "http://upstream-moved"}, + app=app, + premium_user=False, + visited_endpoints=retargeted_visit, + ) + [retargeted_route_key] = retargeted_visit + retargeted = _registered_pass_through_routes[retargeted_route_key]["passthrough_params"] + assert retargeted["target"] == "http://upstream-moved" + assert retargeted["custom_headers"]["Authorization"] != "Bearer sk-literal" + finally: + for key in [k for k in _registered_pass_through_routes if k.startswith("ep-rotate")]: + _registered_pass_through_routes.pop(key, None) + + +@pytest.mark.asyncio +async def test_register_pass_through_endpoint_without_id_keeps_serving_headers_when_stored_headers_do_not_decrypt( + monkeypatch, +): + from fastapi import FastAPI + + from litellm.proxy.pass_through_endpoints.common_utils import ( + encrypt_pass_through_endpoints, + reencrypt_general_settings_pass_through, + ) + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _register_pass_through_endpoint, + _registered_pass_through_routes, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-current-key") + app = FastAPI() + [stored] = encrypt_pass_through_endpoints( + [{"path": "/rotate-no-id-pt", "target": "http://upstream", "headers": {"Authorization": "Bearer sk-no-id"}}] + ) + rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": [stored]}, "sk-next-key") + assert rotated is not None + [rotated_endpoint] = rotated["pass_through_endpoints"] + visited: set[str] = set() + try: + await _register_pass_through_endpoint( + endpoint={"path": "/rotate-no-id-other", "target": "http://other", "headers": {"x": "y"}}, + app=app, + premium_user=False, + visited_endpoints=visited, + ) + await _register_pass_through_endpoint( + endpoint=dict(stored), app=app, premium_user=False, visited_endpoints=visited + ) + reloaded: set[str] = set() + await _register_pass_through_endpoint( + endpoint=dict(rotated_endpoint), app=app, premium_user=False, visited_endpoints=reloaded + ) + + [route_key] = reloaded + params = _registered_pass_through_routes[route_key]["passthrough_params"] + assert params["custom_headers"] == {"Authorization": "Bearer sk-no-id"} + finally: + for key in [k for k in _registered_pass_through_routes if "/rotate-no-id" in k]: + _registered_pass_through_routes.pop(key, None) + + +@pytest.mark.asyncio +async def test_register_pass_through_endpoint_keeps_its_own_headers_when_another_endpoint_shares_the_path( + monkeypatch, +): + from fastapi import FastAPI + + from litellm.proxy.pass_through_endpoints.common_utils import ( + encrypt_pass_through_endpoints, + reencrypt_general_settings_pass_through, + ) + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _register_pass_through_endpoint, + _registered_pass_through_routes, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-current-key") + app = FastAPI() + stored = encrypt_pass_through_endpoints( + [ + { + "id": "ep-shared-a", + "path": "/shared-pt", + "methods": ["GET"], + "target": "http://a", + "headers": {"Authorization": "Bearer sk-a"}, + }, + { + "id": "ep-shared-b", + "path": "/shared-pt", + "methods": ["POST"], + "target": "http://b", + "headers": {"Authorization": "Bearer sk-b"}, + }, + { + "path": "/shared-pt", + "methods": ["PUT"], + "target": "http://c", + "headers": {"Authorization": "Bearer sk-c"}, + }, + ] + ) + rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": stored}, "sk-next-key") + assert rotated is not None + try: + for endpoint in stored: + await _register_pass_through_endpoint( + endpoint=dict(endpoint), app=app, premium_user=False, visited_endpoints=set() + ) + reloaded_b: set[str] = set() + await _register_pass_through_endpoint( + endpoint=dict(rotated["pass_through_endpoints"][1]), + app=app, + premium_user=False, + visited_endpoints=reloaded_b, + ) + reloaded_c: set[str] = set() + await _register_pass_through_endpoint( + endpoint=dict(rotated["pass_through_endpoints"][2]), + app=app, + premium_user=False, + visited_endpoints=reloaded_c, + ) + + [route_b] = reloaded_b + assert _registered_pass_through_routes[route_b]["passthrough_params"]["custom_headers"] == { + "Authorization": "Bearer sk-b" + } + [route_c] = reloaded_c + assert _registered_pass_through_routes[route_c]["passthrough_params"]["custom_headers"] == { + "Authorization": "Bearer sk-c" + } + + twins = encrypt_pass_through_endpoints( + [ + { + "path": "/shared-twin-pt", + "methods": [method], + "target": "http://twin", + "headers": {"Authorization": f"Bearer sk-{method}"}, + } + for method in ("GET", "POST") + ] + ) + rotated_twins = reencrypt_general_settings_pass_through({"pass_through_endpoints": twins}, "sk-next-key") + assert rotated_twins is not None + for twin in twins: + await _register_pass_through_endpoint( + endpoint=dict(twin), app=app, premium_user=False, visited_endpoints=set() + ) + reloaded_twin: set[str] = set() + await _register_pass_through_endpoint( + endpoint=dict(rotated_twins["pass_through_endpoints"][0]), + app=app, + premium_user=False, + visited_endpoints=reloaded_twin, + ) + [route_twin] = reloaded_twin + served_twin = _registered_pass_through_routes[route_twin]["passthrough_params"]["custom_headers"] + assert served_twin["Authorization"] not in ("Bearer sk-GET", "Bearer sk-POST") + + [solo, newcomer] = encrypt_pass_through_endpoints( + [ + {"path": "/shared-solo-pt", "target": "http://d", "headers": {"Authorization": "Bearer sk-d"}}, + {"id": "ep-shared-e", "path": "/shared-solo-pt", "target": "http://d", "headers": {"x-e": "e"}}, + ] + ) + await _register_pass_through_endpoint(endpoint=solo, app=app, premium_user=False, visited_endpoints=set()) + rotated_newcomer = reencrypt_general_settings_pass_through( + {"pass_through_endpoints": [newcomer]}, "sk-next-key" + ) + assert rotated_newcomer is not None + reloaded_e: set[str] = set() + await _register_pass_through_endpoint( + endpoint=dict(rotated_newcomer["pass_through_endpoints"][0]), + app=app, + premium_user=False, + visited_endpoints=reloaded_e, + ) + [route_e] = reloaded_e + assert "Authorization" not in _registered_pass_through_routes[route_e]["passthrough_params"]["custom_headers"] + + [moved] = encrypt_pass_through_endpoints( + [{"path": "/shared-pt", "target": "http://moved", "headers": {"Authorization": "Bearer sk-moved"}}] + ) + rotated_moved = reencrypt_general_settings_pass_through({"pass_through_endpoints": [moved]}, "sk-next-key") + assert rotated_moved is not None + reloaded_moved: set[str] = set() + await _register_pass_through_endpoint( + endpoint=dict(rotated_moved["pass_through_endpoints"][0]), + app=app, + premium_user=False, + visited_endpoints=reloaded_moved, + ) + [route_moved] = reloaded_moved + served_moved = _registered_pass_through_routes[route_moved]["passthrough_params"]["custom_headers"] + assert served_moved["Authorization"] not in ("Bearer sk-a", "Bearer sk-b", "Bearer sk-c", "Bearer sk-moved") + finally: + for key in [k for k in _registered_pass_through_routes if "/shared-" in k]: + _registered_pass_through_routes.pop(key, None) + + +@pytest.mark.asyncio +async def test_register_pass_through_endpoint_reusing_an_id_does_not_get_that_endpoints_headers(monkeypatch): + from fastapi import FastAPI + + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _register_pass_through_endpoint, + _registered_pass_through_routes, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-current-key") + app = FastAPI() + try: + await _register_pass_through_endpoint( + endpoint={ + "id": "ep-reused", + "path": "/reused-victim", + "target": "http://victim", + "headers": {"Authorization": "Bearer sk-victim"}, + }, + app=app, + premium_user=False, + visited_endpoints=set(), + ) + copied: set[str] = set() + await _register_pass_through_endpoint( + endpoint={ + "id": "ep-reused", + "path": "/reused-copy", + "target": "http://elsewhere", + "headers": {"Authorization": "litellm_enc::not-a-ciphertext"}, + }, + app=app, + premium_user=False, + visited_endpoints=copied, + ) + + [copied_route] = copied + assert _registered_pass_through_routes[copied_route]["passthrough_params"]["custom_headers"] == { + "Authorization": "litellm_enc::not-a-ciphertext" + } + finally: + for key in [k for k in _registered_pass_through_routes if k.startswith("ep-reused")]: + _registered_pass_through_routes.pop(key, None) + + +@pytest.mark.asyncio +async def test_update_pass_through_endpoint_keeps_serving_headers_when_stored_headers_do_not_decrypt(monkeypatch): + from fastapi import FastAPI + + from litellm.proxy._types import ConfigFieldInfo, PassThroughGenericEndpoint, UserAPIKeyAuth + from litellm.proxy.pass_through_endpoints.common_utils import ( + encrypt_pass_through_endpoints, + reencrypt_general_settings_pass_through, + ) + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _register_pass_through_endpoint, + _registered_pass_through_routes, + update_pass_through_endpoints, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-current-key") + app = FastAPI() + [stored] = encrypt_pass_through_endpoints( + [ + { + "id": "ep-update-rotate", + "path": "/update-rotate-pt", + "target": "http://upstream", + "headers": {"Authorization": "Bearer sk-update"}, + } + ] + ) + rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": [stored]}, "sk-next-key") + assert rotated is not None + [rotated_endpoint] = rotated["pass_through_endpoints"] + request = MagicMock(spec=Request) + request.app = app + try: + await _register_pass_through_endpoint( + endpoint=dict(stored), app=app, premium_user=False, visited_endpoints=set() + ) + with ( + patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config, + patch("litellm.proxy.proxy_server.update_config_general_settings"), + ): + mock_get_config.return_value = ConfigFieldInfo( + field_name="pass_through_endpoints", field_value=[rotated_endpoint] + ) + await update_pass_through_endpoints( + endpoint_id="ep-update-rotate", + data=PassThroughGenericEndpoint( + path="/update-rotate-pt", + target="http://upstream", + headers={**rotated_endpoint["headers"], "x-org": "acme"}, + ), + request=request, + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + ) + + [served] = [ + route["passthrough_params"] + for route in _registered_pass_through_routes.values() + if route["endpoint_id"] == "ep-update-rotate" + ] + assert served["custom_headers"] == {"Authorization": "Bearer sk-update", "x-org": "acme"} + finally: + for key in [k for k in _registered_pass_through_routes if k.startswith("ep-update-rotate")]: + _registered_pass_through_routes.pop(key, None) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py index 078bd4dd402..6822c6271a1 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py @@ -102,3 +102,126 @@ def test_encode_bedrock_runtime_modelid_arn_partition_arns() -> None: endpoint = "model/arn:aws-us-gov:bedrock:us-gov-west-1:123456789012:inference-profile/test-profile/invoke" expected = "model/arn:aws-us-gov:bedrock:us-gov-west-1:123456789012:inference-profile%2Ftest-profile/invoke" assert CommonUtils.encode_bedrock_runtime_modelid_arn(endpoint) == expected + + +from litellm.proxy.pass_through_endpoints.common_utils import ( + decrypt_pass_through_headers, + encrypt_pass_through_endpoints, + undecryptable_pass_through_header_names, + reencrypt_general_settings_pass_through, +) + +_ENC = "litellm_enc::" + + +@pytest.fixture +def salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-pass-through-tests") + + +def test_encrypt_pass_through_endpoints_encrypts_every_header_of_every_endpoint(salt_key): + stored = encrypt_pass_through_endpoints( + [ + {"path": "/a", "headers": {"x-plain": "value-a"}}, + { + "path": "/b", + "headers": { + "Authorization": "Bearer sk-literal", + "x-env": "Bearer os.environ/UPSTREAM_KEY", + "litellm_user_api_key": "x-my-key", + "x-number": 5, + "x-empty": "", + }, + }, + {"path": "/c"}, + ] + ) + + assert stored[0]["headers"]["x-plain"].startswith(_ENC) + headers_b = stored[1]["headers"] + assert headers_b["Authorization"].startswith(_ENC) + assert "sk-literal" not in headers_b["Authorization"] + assert headers_b["x-env"].startswith(_ENC) + assert headers_b["litellm_user_api_key"] == "x-my-key" + assert headers_b["x-number"] == 5 + assert headers_b["x-empty"] == "" + assert stored[2] == {"path": "/c"} + assert decrypt_pass_through_headers(headers_b) == { + "Authorization": "Bearer sk-literal", + "x-env": "Bearer os.environ/UPSTREAM_KEY", + "litellm_user_api_key": "x-my-key", + "x-number": 5, + "x-empty": "", + } + + +def test_encrypt_pass_through_endpoints_keeps_existing_ciphertext_and_input(salt_key): + endpoints = [{"path": "/a", "headers": {"Authorization": "Bearer sk-literal"}}] + first = encrypt_pass_through_endpoints(endpoints) + second = encrypt_pass_through_endpoints(first) + + assert second == first + assert endpoints == [{"path": "/a", "headers": {"Authorization": "Bearer sk-literal"}}] + assert encrypt_pass_through_endpoints(None) is None + + +def test_decrypt_pass_through_headers_passes_legacy_plaintext_through(salt_key): + legacy = {"Authorization": "Bearer sk-legacy", "x-env": "os.environ/UPSTREAM_KEY"} + + assert decrypt_pass_through_headers(legacy) == legacy + assert decrypt_pass_through_headers(None) is None + + +def test_decrypt_pass_through_headers_keeps_value_it_cannot_decrypt(salt_key, monkeypatch): + stored = encrypt_pass_through_endpoints([{"path": "/a", "headers": {"Authorization": "Bearer sk-literal"}}]) + ciphertext = stored[0]["headers"]["Authorization"] + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-some-other-salt") + + assert decrypt_pass_through_headers({"Authorization": ciphertext}) == {"Authorization": ciphertext} + + +def test_reencrypt_general_settings_pass_through_moves_headers_to_new_key(salt_key, monkeypatch): + stored = encrypt_pass_through_endpoints( + [ + {"path": "/a", "headers": {"x-a": "value-a"}}, + {"path": "/b", "headers": {"Authorization": "Bearer sk-b"}}, + ] + ) + general_settings = { + "store_model_in_db": True, + "pass_through_endpoints": [ + stored[0], + {**stored[1], "headers": {**stored[1]["headers"], "x-legacy": "plain-b"}}, + ], + } + + rotated = reencrypt_general_settings_pass_through(general_settings, "sk-new-master") + + assert rotated is not None + assert rotated["store_model_in_db"] is True + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-new-master") + assert decrypt_pass_through_headers(rotated["pass_through_endpoints"][0]["headers"]) == {"x-a": "value-a"} + assert decrypt_pass_through_headers(rotated["pass_through_endpoints"][1]["headers"]) == { + "Authorization": "Bearer sk-b", + "x-legacy": "plain-b", + } + assert all( + value.startswith(_ENC) + for endpoint in rotated["pass_through_endpoints"] + for value in endpoint["headers"].values() + ) + assert reencrypt_general_settings_pass_through({"store_model_in_db": True}, "sk-new-master") is None + assert reencrypt_general_settings_pass_through(None, "sk-new-master") is None + + +def test_undecryptable_pass_through_header_names(salt_key): + [stored] = encrypt_pass_through_endpoints([{"path": "/a", "headers": {"x-a": "plain-a"}}]) + rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": [stored]}, "sk-next-key") + + assert undecryptable_pass_through_header_names(stored["headers"]) == frozenset() + assert undecryptable_pass_through_header_names({"x-legacy": "plain", "x-n": 5}) == frozenset() + assert undecryptable_pass_through_header_names(None) == frozenset() + assert undecryptable_pass_through_header_names(rotated["pass_through_endpoints"][0]["headers"]) == {"x-a"} + assert undecryptable_pass_through_header_names( + {"x-ok": stored["headers"]["x-a"], "x-bad": _ENC + "garbage", "x-plain": "p"} + ) == {"x-bad"} diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index df05cf0987e..7cdcfa0694d 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -15562,3 +15562,64 @@ async def test_spend_capture_rate_check_job_clears_the_gauge_once_the_setting_is call(api_provider="openai", capture_rate=0.97), call(api_provider="openai", capture_rate=None), ] + + +def test_update_config_encrypts_pass_through_endpoint_headers(_update_config_setup, monkeypatch): + from litellm.proxy.pass_through_endpoints.common_utils import decrypt_pass_through_headers + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-pass-through-tests") + client, prisma, restore = _update_config_setup(initial_rows={"general_settings": {"store_model_in_db": True}}) + try: + resp = client.post( + "/config/update", + json={ + "general_settings": { + "pass_through_endpoints": [ + {"path": "/pt", "target": "http://upstream", "headers": {"Authorization": "Bearer sk-literal"}} + ] + } + }, + ) + assert resp.status_code == 200 + stored = prisma.db.litellm_config.rows["general_settings"] + stored_headers = stored["pass_through_endpoints"][0]["headers"] + assert "sk-literal" not in json.dumps(stored) + assert decrypt_pass_through_headers(stored_headers) == {"Authorization": "Bearer sk-literal"} + assert stored["store_model_in_db"] is True + finally: + restore() + + +@pytest.mark.asyncio +async def test_update_config_general_settings_encrypts_pass_through_endpoint_headers(monkeypatch): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import ConfigFieldUpdate, LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.pass_through_endpoints.common_utils import decrypt_pass_through_headers + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-pass-through-tests") + prisma = _FakePrismaClient(initial_rows={"general_settings": {"store_model_in_db": True}}) + monkeypatch.setattr(ps, "prisma_client", prisma) + monkeypatch.setattr(ps, "invalidate_config_param", AsyncMock(return_value=None)) + monkeypatch.setattr(ps.proxy_config.settings, "apply_db_row", MagicMock()) + monkeypatch.setattr(ps, "create_config_audit_log", AsyncMock(return_value=None)) + endpoints = [ + {"path": "/a", "target": "http://upstream-a", "headers": {"x-a": "plain-a"}}, + {"path": "/b", "target": "http://upstream-b", "headers": {"Authorization": "Bearer sk-literal"}}, + ] + + await ps.update_config_general_settings( + data=ConfigFieldUpdate( + field_name="pass_through_endpoints", field_value=endpoints, config_type="general_settings" + ), + user_api_key_dict=UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + stored = prisma.db.litellm_config.rows["general_settings"] + assert "sk-literal" not in json.dumps(stored) + assert "plain-a" not in json.dumps(stored) + assert [decrypt_pass_through_headers(e["headers"]) for e in stored["pass_through_endpoints"]] == [ + {"x-a": "plain-a"}, + {"Authorization": "Bearer sk-literal"}, + ] + assert endpoints[1]["headers"] == {"Authorization": "Bearer sk-literal"} + ps.proxy_config.settings.apply_db_row.assert_called_once_with("general_settings", stored) From 43890a6d2b0ce7d53d2b2439c934403944228f8f Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 01:05:20 -0700 Subject: [PATCH 2/6] fix(proxy): keep concurrent settings edits during pass-through header rotation Rotation swaps only general_settings.pass_through_endpoints and only when it is unchanged since it was read. The rotation-window fallback is per header, including the Langfuse key pair served as Authorization. A bare litellm_enc:: literal forwards unchanged, and storing a header without an encryption key logs a warning. --- .../key_management_endpoints.py | 46 +++++++--- .../pass_through_endpoints/common_utils.py | 17 +++- .../pass_through_endpoints.py | 16 +++- .../test_key_management_endpoints.py | 28 ++++--- .../test_pass_through_endpoints.py | 84 ++++++++++++++++++- ...test_passthrough_endpoints_common_utils.py | 5 ++ 6 files changed, 165 insertions(+), 31 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 6517b8b809b..86e42d1189e 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -5219,6 +5219,39 @@ async def delete_key_aliases( ) +_PASS_THROUGH_REENCRYPT_ATTEMPTS: Final = 3 +_SWAP_PASS_THROUGH_ENDPOINTS_SQL: Final = ( + 'UPDATE "LiteLLM_Config" ' + "SET param_value = jsonb_set(param_value::jsonb, '{pass_through_endpoints}', $1::jsonb) " + "WHERE param_name = 'general_settings' AND param_value::jsonb -> 'pass_through_endpoints' = $2::jsonb" +) + + +async def _reencrypt_pass_through_endpoint_headers(prisma_client: PrismaClient, new_master_key: str) -> None: + """Re-encrypt pass-through header values in general_settings under new_master_key. + + Only the pass_through_endpoints key is written, and only if it still equals the list that + was read, so a concurrent settings edit is kept; a changed list is re-read and retried. + """ + for _ in range(_PASS_THROUGH_REENCRYPT_ATTEMPTS): + rows: Sequence[ConfigParam] = await _config_table(prisma_client).find_many() + stored = next((row.param_value for row in rows if row.param_name == "general_settings"), None) + reencrypted = reencrypt_general_settings_pass_through(stored, new_master_key) + if not isinstance(stored, dict) or reencrypted is None: + return + swapped = await prisma_client.db.execute_raw( + _SWAP_PASS_THROUGH_ENDPOINTS_SQL, + json.dumps(reencrypted["pass_through_endpoints"]), + json.dumps(stored["pass_through_endpoints"]), + ) + if swapped: + await invalidate_config_param("general_settings") + return + verbose_proxy_logger.warning( + "Pass-through endpoint headers were not re-encrypted: general_settings kept changing during the rotation" + ) + + async def _rotate_master_key( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, @@ -5304,17 +5337,8 @@ async def _rotate_master_key( data={"param_value": prisma.Json(encrypted_env_vars)}, ) - for c in config: - if ( - c.param_name == "general_settings" - and os.getenv(SALT_KEY_ENV_VAR) is None - and (reencrypted := reencrypt_general_settings_pass_through(c.param_value, new_master_key)) is not None - ): - await _config_table(prisma_client).update( - where={"param_name": "general_settings"}, - data={"param_value": prisma.Json(reencrypted)}, - ) - await invalidate_config_param("general_settings") + if os.getenv(SALT_KEY_ENV_VAR) is None: + await _reencrypt_pass_through_endpoint_headers(prisma_client, new_master_key) # 4. process MCP server table try: diff --git a/litellm/proxy/pass_through_endpoints/common_utils.py b/litellm/proxy/pass_through_endpoints/common_utils.py index 1550b43a887..251e2fdac8a 100644 --- a/litellm/proxy/pass_through_endpoints/common_utils.py +++ b/litellm/proxy/pass_through_endpoints/common_utils.py @@ -36,6 +36,9 @@ def _encrypt_header_value(name: str, value: JsonValue, new_encryption_key: str | try: return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(value, new_encryption_key=new_encryption_key) except Exception: # noqa: BLE001 # no salt or master key configured: store as before rather than fail the write + verbose_proxy_logger.warning( + "Pass-through header %s is stored unencrypted: set LITELLM_SALT_KEY or a master key to encrypt it", name + ) return value @@ -48,8 +51,16 @@ def _decrypted(name: str, value: str) -> str | None: ) +def _is_marked(value: object) -> bool: + return ( + isinstance(value, str) + and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) + and value != CALLBACK_VAR_ENCRYPTED_PREFIX + ) + + def _decrypt_header_value(name: str, value: object) -> object: - if not isinstance(value, str) or not value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): + if not isinstance(value, str) or not _is_marked(value): return value decrypted: Final = _decrypted(name, value) if decrypted is None: @@ -67,9 +78,7 @@ def undecryptable_pass_through_header_names(headers: object) -> frozenset[str]: return frozenset( str(name) for name, value in headers.items() - if isinstance(value, str) - and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) - and _decrypted(str(name), value) is None + if isinstance(value, str) and _is_marked(value) and _decrypted(str(name), value) is None ) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index f0d317d74c9..b8c82f6bfc7 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -201,6 +201,11 @@ async def set_env_variables_in_header(custom_headers: dict | None) -> dict | Non return headers +_LANGFUSE_KEY_HEADERS: Final = frozenset({"LANGFUSE_PUBLIC_KEY", "LANGFUSE_SECRET_KEY"}) +# set_env_variables_in_header serves the Langfuse key pair as one Authorization header. +_SERVED_NAME_OF_STORED_HEADER: Final = {name: "Authorization" for name in _LANGFUSE_KEY_HEADERS} + + def _served_custom_headers( endpoint_id: str, target: str | None, path_without_stored_id: str | None ) -> Mapping[str, object] | None: @@ -248,10 +253,13 @@ async def _resolve_stored_headers( endpoint_id, sorted(undecryptable), ) - if not undecryptable <= registered.keys(): - return dict(registered) - resolved: Final = await set_env_variables_in_header(custom_headers=stored_headers) - return {**(resolved or {}), **{name: registered[name] for name in undecryptable}} + replaced: Final = {name for name in undecryptable if _SERVED_NAME_OF_STORED_HEADER.get(name, name) in registered} + dropped: Final = replaced | (_LANGFUSE_KEY_HEADERS if replaced & _LANGFUSE_KEY_HEADERS else frozenset()) + resolved: Final = await set_env_variables_in_header( + custom_headers={name: value for name, value in stored_headers.items() if name not in dropped} + ) + kept: Final = {_SERVED_NAME_OF_STORED_HEADER.get(name, name) for name in replaced} + return {**(resolved or {}), **{name: registered[name] for name in kept}} async def chat_completion_pass_through_endpoint( diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index d2e19b6e2ae..0d3cb60b81d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21097,8 +21097,6 @@ class TestTeamAdminMemberKeyBudgetUpdate: @pytest.mark.asyncio async def test_rotate_master_key_reencrypts_pass_through_endpoint_headers(monkeypatch): - import prisma - from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.management_endpoints import key_management_endpoints from litellm.proxy.pass_through_endpoints.common_utils import ( @@ -21121,8 +21119,15 @@ async def test_rotate_master_key_reencrypts_pass_through_endpoint_headers(monkey mock_prisma_client = AsyncMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) - mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[general_settings_row]) + edited_row = MagicMock( + param_name="general_settings", + param_value={"store_model_in_db": True, "pass_through_endpoints": stored_endpoints[:1]}, + ) + mock_prisma_client.db.litellm_config.find_many = AsyncMock( + side_effect=[[general_settings_row], [general_settings_row], [edited_row]] + ) mock_prisma_client.db.litellm_config.update = AsyncMock() + mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[0, 1]) mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[]) invalidate = AsyncMock() monkeypatch.setattr(key_management_endpoints, "invalidate_config_param", invalidate) @@ -21141,19 +21146,21 @@ async def test_rotate_master_key_reencrypts_pass_through_endpoint_headers(monkey new_master_key="sk-new-master-key", ) - mock_prisma_client.db.litellm_config.update.assert_awaited_once() - update_kwargs = mock_prisma_client.db.litellm_config.update.await_args.kwargs - assert update_kwargs["where"] == {"param_name": "general_settings"} - assert isinstance(update_kwargs["data"]["param_value"], prisma.Json) - rotated = update_kwargs["data"]["param_value"].data - assert rotated["store_model_in_db"] is True + mock_prisma_client.db.litellm_config.update.assert_not_awaited() + [first_swap, retried_swap] = mock_prisma_client.db.execute_raw.await_args_list + assert json.loads(first_swap.args[2]) == stored_endpoints + assert json.loads(retried_swap.args[2]) == stored_endpoints[:1] + assert "jsonb_set" in retried_swap.args[0] invalidate.assert_awaited_once_with("general_settings") monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-new-master-key") - assert [decrypt_pass_through_headers(e["headers"]) for e in rotated["pass_through_endpoints"]] == [ + assert [decrypt_pass_through_headers(e["headers"]) for e in json.loads(first_swap.args[1])] == [ {"x-a": "plain-a"}, {"Authorization": "Bearer sk-literal"}, ] + assert [decrypt_pass_through_headers(e["headers"]) for e in json.loads(retried_swap.args[1])] == [ + {"x-a": "plain-a"} + ] @pytest.mark.asyncio @@ -21193,3 +21200,4 @@ async def test_rotate_master_key_leaves_pass_through_headers_under_salt_key(monk ) mock_prisma_client.db.litellm_config.update.assert_not_awaited() + mock_prisma_client.db.execute_raw.assert_not_called() diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 0579488a34b..f5f88678f10 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -7875,7 +7875,12 @@ async def test_register_pass_through_endpoint_keeps_its_own_headers_when_another [solo, newcomer] = encrypt_pass_through_endpoints( [ {"path": "/shared-solo-pt", "target": "http://d", "headers": {"Authorization": "Bearer sk-d"}}, - {"id": "ep-shared-e", "path": "/shared-solo-pt", "target": "http://d", "headers": {"x-e": "e"}}, + { + "id": "ep-shared-e", + "path": "/shared-solo-pt", + "target": "http://d", + "headers": {"Authorization": "Bearer sk-e"}, + }, ] ) await _register_pass_through_endpoint(endpoint=solo, app=app, premium_user=False, visited_endpoints=set()) @@ -7891,7 +7896,12 @@ async def test_register_pass_through_endpoint_keeps_its_own_headers_when_another visited_endpoints=reloaded_e, ) [route_e] = reloaded_e - assert "Authorization" not in _registered_pass_through_routes[route_e]["passthrough_params"]["custom_headers"] + assert _registered_pass_through_routes[route_e]["passthrough_params"]["custom_headers"][ + "Authorization" + ] not in ( + "Bearer sk-d", + "Bearer sk-e", + ) [moved] = encrypt_pass_through_endpoints( [{"path": "/shared-pt", "target": "http://moved", "headers": {"Authorization": "Bearer sk-moved"}}] @@ -8021,3 +8031,73 @@ async def test_update_pass_through_endpoint_keeps_serving_headers_when_stored_he finally: for key in [k for k in _registered_pass_through_routes if k.startswith("ep-update-rotate")]: _registered_pass_through_routes.pop(key, None) + + +@pytest.mark.asyncio +async def test_register_pass_through_endpoint_keeps_serving_langfuse_keys_and_applies_other_header_edits( + monkeypatch, +): + from fastapi import FastAPI + + from litellm.proxy.pass_through_endpoints.common_utils import ( + encrypt_pass_through_endpoints, + reencrypt_general_settings_pass_through, + ) + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _register_pass_through_endpoint, + _registered_pass_through_routes, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-current-key") + app = FastAPI() + [stored] = encrypt_pass_through_endpoints( + [ + { + "id": "ep-langfuse-rotate", + "path": "/langfuse-rotate-pt", + "target": "http://langfuse", + "headers": {"LANGFUSE_PUBLIC_KEY": "pk-lf", "LANGFUSE_SECRET_KEY": "sk-lf", "x-org": "old"}, + } + ] + ) + rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": [stored]}, "sk-next-key") + assert rotated is not None + [rotated_endpoint] = rotated["pass_through_endpoints"] + try: + await _register_pass_through_endpoint( + endpoint=dict(stored), app=app, premium_user=False, visited_endpoints=set() + ) + [served_before] = [ + route["passthrough_params"]["custom_headers"] + for route in _registered_pass_through_routes.values() + if route["endpoint_id"] == "ep-langfuse-rotate" + ] + reloaded: set[str] = set() + await _register_pass_through_endpoint( + endpoint={**rotated_endpoint, "headers": {**rotated_endpoint["headers"], "x-org": "new"}}, + app=app, + premium_user=False, + visited_endpoints=reloaded, + ) + + [route_key] = reloaded + assert _registered_pass_through_routes[route_key]["passthrough_params"]["custom_headers"] == { + "Authorization": served_before["Authorization"], + "x-org": "new", + } + + half_edited: set[str] = set() + await _register_pass_through_endpoint( + endpoint={**rotated_endpoint, "headers": {**rotated_endpoint["headers"], "LANGFUSE_PUBLIC_KEY": "pk-new"}}, + app=app, + premium_user=False, + visited_endpoints=half_edited, + ) + [half_edited_key] = half_edited + assert _registered_pass_through_routes[half_edited_key]["passthrough_params"]["custom_headers"] == { + "Authorization": served_before["Authorization"], + "x-org": "new", + } + finally: + for key in [k for k in _registered_pass_through_routes if k.startswith("ep-langfuse-rotate")]: + _registered_pass_through_routes.pop(key, None) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py index 6822c6271a1..a96ef88b5aa 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py @@ -225,3 +225,8 @@ def test_undecryptable_pass_through_header_names(salt_key): assert undecryptable_pass_through_header_names( {"x-ok": stored["headers"]["x-a"], "x-bad": _ENC + "garbage", "x-plain": "p"} ) == {"x-bad"} + + +def test_decrypt_pass_through_headers_keeps_a_bare_marker_literal(salt_key): + assert decrypt_pass_through_headers({"x-tag": _ENC}) == {"x-tag": _ENC} + assert undecryptable_pass_through_header_names({"x-tag": _ENC}) == frozenset() From 106f52e1c1e9b71bc4d5d15c8ff40c3d55ab688c Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 01:43:20 -0700 Subject: [PATCH 3/6] fix(proxy): read pass-through settings from the writer during rotation With a read replica configured, the rotation compare-and-swap could read a lagging list and never match the writer. The read and the swap now both go to the writer. --- .../key_management_endpoints.py | 25 ++++++++++--- .../test_key_management_endpoints.py | 35 +++++++++++++++++-- 2 files changed, 52 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 86e42d1189e..d7f9e644bba 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -84,6 +84,7 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_keys from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage from litellm.proxy.management_endpoints.common_utils import ( @@ -5219,6 +5220,16 @@ async def delete_key_aliases( ) +class _ConfigRowFinder(Protocol): + async def find_unique(self, *, where: Mapping[str, object]) -> ConfigParam | None: ... + + +class _PassThroughConfigWriter(Protocol): + litellm_config: _ConfigRowFinder + + async def execute_raw(self, query: str, *args: object) -> int: ... + + _PASS_THROUGH_REENCRYPT_ATTEMPTS: Final = 3 _SWAP_PASS_THROUGH_ENDPOINTS_SQL: Final = ( 'UPDATE "LiteLLM_Config" ' @@ -5230,16 +5241,20 @@ _SWAP_PASS_THROUGH_ENDPOINTS_SQL: Final = ( async def _reencrypt_pass_through_endpoint_headers(prisma_client: PrismaClient, new_master_key: str) -> None: """Re-encrypt pass-through header values in general_settings under new_master_key. - Only the pass_through_endpoints key is written, and only if it still equals the list that - was read, so a concurrent settings edit is kept; a changed list is re-read and retried. + Reads and writes go to the writer. Only the pass_through_endpoints key is written, and only + if it still equals the list that was read, so a concurrent settings edit is kept; a changed + list is re-read and retried. """ + writer: Final = cast( # cast-ok: untyped Prisma client behind the writer pin + "_PassThroughConfigWriter", writer_wrapper(prisma_client.db) + ) for _ in range(_PASS_THROUGH_REENCRYPT_ATTEMPTS): - rows: Sequence[ConfigParam] = await _config_table(prisma_client).find_many() - stored = next((row.param_value for row in rows if row.param_name == "general_settings"), None) + row = await writer.litellm_config.find_unique(where={"param_name": "general_settings"}) + stored = row.param_value if row is not None else None reencrypted = reencrypt_general_settings_pass_through(stored, new_master_key) if not isinstance(stored, dict) or reencrypted is None: return - swapped = await prisma_client.db.execute_raw( + swapped = await writer.execute_raw( _SWAP_PASS_THROUGH_ENDPOINTS_SQL, json.dumps(reencrypted["pass_through_endpoints"]), json.dumps(stored["pass_through_endpoints"]), diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 0d3cb60b81d..4b3fdeac9d7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21123,9 +21123,8 @@ async def test_rotate_master_key_reencrypts_pass_through_endpoint_headers(monkey param_name="general_settings", param_value={"store_model_in_db": True, "pass_through_endpoints": stored_endpoints[:1]}, ) - mock_prisma_client.db.litellm_config.find_many = AsyncMock( - side_effect=[[general_settings_row], [general_settings_row], [edited_row]] - ) + mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[general_settings_row]) + mock_prisma_client.db.litellm_config.find_unique = AsyncMock(side_effect=[general_settings_row, edited_row]) mock_prisma_client.db.litellm_config.update = AsyncMock() mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[0, 1]) mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[]) @@ -21201,3 +21200,33 @@ async def test_rotate_master_key_leaves_pass_through_headers_under_salt_key(monk mock_prisma_client.db.litellm_config.update.assert_not_awaited() mock_prisma_client.db.execute_raw.assert_not_called() + + +@pytest.mark.asyncio +async def test_rotate_master_key_reencrypts_pass_through_headers_on_the_writer(monkeypatch): + from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper + from litellm.proxy.management_endpoints import key_management_endpoints + from litellm.proxy.pass_through_endpoints.common_utils import encrypt_pass_through_endpoints + + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master-key") + row = MagicMock( + param_name="general_settings", + param_value={"pass_through_endpoints": encrypt_pass_through_endpoints([{"path": "/a", "headers": {"x": "v"}}])}, + ) + from unittest.mock import NonCallableMagicMock + + writer = MagicMock() + writer.litellm_config = NonCallableMagicMock(find_unique=AsyncMock(return_value=row)) + writer.execute_raw = AsyncMock(return_value=1) + reader = MagicMock() + reader.litellm_config = NonCallableMagicMock(find_unique=AsyncMock(return_value=None)) + prisma_client = MagicMock() + prisma_client.db = RoutingPrismaWrapper(writer=writer, reader=reader) + monkeypatch.setattr(key_management_endpoints, "invalidate_config_param", AsyncMock()) + + await key_management_endpoints._reencrypt_pass_through_endpoint_headers(prisma_client, "sk-new-master-key") + + writer.litellm_config.find_unique.assert_awaited_once() + writer.execute_raw.assert_awaited_once() + reader.litellm_config.find_unique.assert_not_called() From 9ca392868d677cb9e8d9ee81e7ffbafe7188ae87 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 10:43:27 -0700 Subject: [PATCH 4/6] fix(proxy): keep master key rotation going when pass-through header re-encryption fails Wrap the pass-through header step of _rotate_master_key in the same try/except the MCP, SSO and credentials steps use, so a failed write logs a warning instead of aborting the remaining steps with a 500. Treat a marked header value that decrypts to an empty string as undecryptable, so a literal such as litellm_enc::*** is forwarded unchanged instead of as an empty header. --- .../key_management_endpoints.py | 5 ++- .../pass_through_endpoints/common_utils.py | 14 ++++--- .../test_key_management_endpoints.py | 41 +++++++++++++++++++ ...test_passthrough_endpoints_common_utils.py | 6 +++ 4 files changed, 60 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index d7f9e644bba..5ee35643c54 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -5353,7 +5353,10 @@ async def _rotate_master_key( ) if os.getenv(SALT_KEY_ENV_VAR) is None: - await _reencrypt_pass_through_endpoint_headers(prisma_client, new_master_key) + try: + await _reencrypt_pass_through_endpoint_headers(prisma_client, new_master_key) + except Exception as e: + verbose_proxy_logger.warning("Failed to rotate pass-through endpoint headers: %s", str(e)) # 4. process MCP server table try: diff --git a/litellm/proxy/pass_through_endpoints/common_utils.py b/litellm/proxy/pass_through_endpoints/common_utils.py index 251e2fdac8a..a27bfd4abe7 100644 --- a/litellm/proxy/pass_through_endpoints/common_utils.py +++ b/litellm/proxy/pass_through_endpoints/common_utils.py @@ -43,11 +43,15 @@ def _encrypt_header_value(name: str, value: JsonValue, new_encryption_key: str | def _decrypted(name: str, value: str) -> str | None: - return decrypt_value_helper( - value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), - key=name, - exception_type="debug", - return_original_value=False, + """Plaintext of a marked value; None when it does not decrypt or decrypts to an empty string.""" + return ( + decrypt_value_helper( + value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), + key=name, + exception_type="debug", + return_original_value=False, + ) + or None ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 4b3fdeac9d7..3a9a1537f15 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -21202,6 +21202,47 @@ async def test_rotate_master_key_leaves_pass_through_headers_under_salt_key(monk mock_prisma_client.db.execute_raw.assert_not_called() +@pytest.mark.asyncio +async def test_rotate_master_key_continues_when_pass_through_header_reencryption_fails(monkeypatch): + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints import key_management_endpoints + + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + mock_prisma_client = AsyncMock() + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_config.find_many = AsyncMock( + return_value=[MagicMock(param_name="general_settings", param_value={})] + ) + mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[]) + monkeypatch.setattr( + key_management_endpoints, + "_reencrypt_pass_through_endpoint_headers", + AsyncMock(side_effect=RuntimeError("general_settings write fault")), + ) + later_steps = { + name: AsyncMock() + for name in ( + "rotate_mcp_server_credentials_master_key", + "rotate_mcp_user_credentials_master_key", + "rotate_mcp_user_env_vars_master_key", + "rotate_sso_identity_assertions_master_key", + ) + } + for name, step in later_steps.items(): + monkeypatch.setattr(key_management_endpoints, name, step) + + await key_management_endpoints._rotate_master_key( + prisma_client=mock_prisma_client, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"), + current_master_key="sk-old-master-key", + new_master_key="sk-new-master-key", + ) + + assert all(step.await_count == 1 for step in later_steps.values()) + mock_prisma_client.db.litellm_credentialstable.find_many.assert_awaited_once() + + @pytest.mark.asyncio async def test_rotate_master_key_reencrypts_pass_through_headers_on_the_writer(monkeypatch): from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py index a96ef88b5aa..5e579e4d2b1 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py @@ -230,3 +230,9 @@ def test_undecryptable_pass_through_header_names(salt_key): def test_decrypt_pass_through_headers_keeps_a_bare_marker_literal(salt_key): assert decrypt_pass_through_headers({"x-tag": _ENC}) == {"x-tag": _ENC} assert undecryptable_pass_through_header_names({"x-tag": _ENC}) == frozenset() + + +@pytest.mark.parametrize("value", [_ENC + "***", _ENC + "!!"]) +def test_decrypt_pass_through_headers_keeps_a_marked_value_that_decrypts_to_nothing(salt_key, value): + assert decrypt_pass_through_headers({"x-tag": value}) == {"x-tag": value} + assert undecryptable_pass_through_header_names({"x-tag": value}) == {"x-tag"} From d3c0dd1f8010b41f0b2eddf8a7e0a70f7820fccd Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 11:53:37 -0700 Subject: [PATCH 5/6] fix(proxy): treat marker-prefixed header literals as plaintext A litellm_enc:: value now counts as stored ciphertext only when its payload decodes to at least the size encrypt_value_helper produces (nacl or AES-GCM). Shorter or non-base64 payloads such as litellm_enc::*** are literals: they are encrypted on write and forwarded unchanged, instead of being reported as undecryptable, which made an edit to such a value keep the previously served header. --- .../pass_through_endpoints/common_utils.py | 38 +++++++++++++------ ...test_passthrough_endpoints_common_utils.py | 28 ++++++++++++-- 2 files changed, 51 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/common_utils.py b/litellm/proxy/pass_through_endpoints/common_utils.py index a27bfd4abe7..54fc9b7e252 100644 --- a/litellm/proxy/pass_through_endpoints/common_utils.py +++ b/litellm/proxy/pass_through_endpoints/common_utils.py @@ -1,3 +1,5 @@ +import base64 +import binascii from collections.abc import Mapping from typing import Final @@ -11,6 +13,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helpe # Holds the name of the request header that carries the caller's LiteLLM key, which # user_api_key_auth reads straight from general_settings, so it stays plaintext. _CALLER_KEY_HEADER_NAME: Final = "litellm_user_api_key" +_GCM_CIPHERTEXT_PREFIX: Final = "v2:gcm:" +# Smallest encrypt_value_helper output for a one-byte value: nonce + tag (AES-GCM), nonce + MAC (nacl). +_MIN_GCM_CIPHERTEXT_BYTES: Final = 29 +_MIN_NACL_CIPHERTEXT_BYTES: Final = 41 def get_litellm_virtual_key(request: Request) -> str: @@ -31,7 +37,7 @@ def get_litellm_virtual_key(request: Request) -> str: def _encrypt_header_value(name: str, value: JsonValue, new_encryption_key: str | None) -> JsonValue: if not isinstance(value, str) or not value or name == _CALLER_KEY_HEADER_NAME: return value - if value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): + if _is_marked(value): return value try: return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(value, new_encryption_key=new_encryption_key) @@ -43,23 +49,33 @@ def _encrypt_header_value(name: str, value: JsonValue, new_encryption_key: str | def _decrypted(name: str, value: str) -> str | None: - """Plaintext of a marked value; None when it does not decrypt or decrypts to an empty string.""" - return ( - decrypt_value_helper( - value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), - key=name, - exception_type="debug", - return_original_value=False, - ) - or None + return decrypt_value_helper( + value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), + key=name, + exception_type="debug", + return_original_value=False, ) +def _has_ciphertext_shape(payload: str) -> bool: + gcm: Final = payload.startswith(_GCM_CIPHERTEXT_PREFIX) + encoded: Final = payload.removeprefix(_GCM_CIPHERTEXT_PREFIX) + try: + raw = base64.urlsafe_b64decode(encoded) + except (binascii.Error, ValueError): + try: + raw = base64.b64decode(encoded) + except (binascii.Error, ValueError): + return False + return len(raw) >= (_MIN_GCM_CIPHERTEXT_BYTES if gcm else _MIN_NACL_CIPHERTEXT_BYTES) + + def _is_marked(value: object) -> bool: + """True for a `litellm_enc::` value whose payload has the shape encrypt_value_helper produces.""" return ( isinstance(value, str) and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) - and value != CALLBACK_VAR_ENCRYPTED_PREFIX + and _has_ciphertext_shape(value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX)) ) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py index 5e579e4d2b1..05c140f14c8 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py @@ -223,7 +223,12 @@ def test_undecryptable_pass_through_header_names(salt_key): assert undecryptable_pass_through_header_names(None) == frozenset() assert undecryptable_pass_through_header_names(rotated["pass_through_endpoints"][0]["headers"]) == {"x-a"} assert undecryptable_pass_through_header_names( - {"x-ok": stored["headers"]["x-a"], "x-bad": _ENC + "garbage", "x-plain": "p"} + { + "x-ok": stored["headers"]["x-a"], + "x-bad": rotated["pass_through_endpoints"][0]["headers"]["x-a"], + "x-garbage": _ENC + "garbage", + "x-plain": "p", + } ) == {"x-bad"} @@ -232,7 +237,22 @@ def test_decrypt_pass_through_headers_keeps_a_bare_marker_literal(salt_key): assert undecryptable_pass_through_header_names({"x-tag": _ENC}) == frozenset() -@pytest.mark.parametrize("value", [_ENC + "***", _ENC + "!!"]) -def test_decrypt_pass_through_headers_keeps_a_marked_value_that_decrypts_to_nothing(salt_key, value): +@pytest.mark.parametrize("value", [_ENC + "***", _ENC + "!!", _ENC + "not-a-ciphertext", _ENC + "v2:gcm:abc"]) +def test_marker_prefixed_literal_is_encrypted_and_forwarded_unchanged(salt_key, value): + [stored] = encrypt_pass_through_endpoints([{"path": "/a", "headers": {"x-tag": value}}]) + + assert stored["headers"]["x-tag"] != value + assert decrypt_pass_through_headers(stored["headers"]) == {"x-tag": value} assert decrypt_pass_through_headers({"x-tag": value}) == {"x-tag": value} - assert undecryptable_pass_through_header_names({"x-tag": value}) == {"x-tag"} + assert undecryptable_pass_through_header_names({"x-tag": value}) == frozenset() + + +def test_aes_gcm_ciphertext_from_another_key_is_undecryptable(salt_key, monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"encryption_algorithm": "aes-256-gcm"}) + [stored] = encrypt_pass_through_endpoints([{"path": "/a", "headers": {"x-a": "v"}}]) + assert stored["headers"]["x-a"].startswith(_ENC + "v2:gcm:") + assert decrypt_pass_through_headers(stored["headers"]) == {"x-a": "v"} + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-some-other-salt") + + assert undecryptable_pass_through_header_names(stored["headers"]) == {"x-a"} From 4b2902fa1214af74ad448ae939501f289596d863 Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 12:00:06 -0700 Subject: [PATCH 6/6] chore(proxy): drop comments on the pass-through ciphertext check --- litellm/proxy/pass_through_endpoints/common_utils.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/common_utils.py b/litellm/proxy/pass_through_endpoints/common_utils.py index 54fc9b7e252..a4d09a10d47 100644 --- a/litellm/proxy/pass_through_endpoints/common_utils.py +++ b/litellm/proxy/pass_through_endpoints/common_utils.py @@ -14,7 +14,6 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helpe # user_api_key_auth reads straight from general_settings, so it stays plaintext. _CALLER_KEY_HEADER_NAME: Final = "litellm_user_api_key" _GCM_CIPHERTEXT_PREFIX: Final = "v2:gcm:" -# Smallest encrypt_value_helper output for a one-byte value: nonce + tag (AES-GCM), nonce + MAC (nacl). _MIN_GCM_CIPHERTEXT_BYTES: Final = 29 _MIN_NACL_CIPHERTEXT_BYTES: Final = 41 @@ -71,7 +70,6 @@ def _has_ciphertext_shape(payload: str) -> bool: def _is_marked(value: object) -> bool: - """True for a `litellm_enc::` value whose payload has the shape encrypt_value_helper produces.""" return ( isinstance(value, str) and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX)