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.
This commit is contained in:
Yucheng He 2026-09-29 01:43:20 -07:00
parent 43890a6d2b
commit 106f52e1c1
2 changed files with 52 additions and 8 deletions

View file

@ -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"]),

View file

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