From 106f52e1c1e9b71bc4d5d15c8ff40c3d55ab688c Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 01:43:20 -0700 Subject: [PATCH] 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()