mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
a50f9fa2a8
commit
43890a6d2b
6 changed files with 165 additions and 31 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue