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.
This commit is contained in:
Yucheng He 2026-09-29 01:05:20 -07:00
parent a50f9fa2a8
commit 43890a6d2b
6 changed files with 165 additions and 31 deletions

View file

@ -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:

View file

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

View file

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

View file

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

View file

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

View file

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