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()