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)