This commit is contained in:
yucheng-berri 2026-09-30 16:54:56 -04:00 • committed by GitHub
commit 953a3afbaa
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 1140 additions and 10 deletions

View file

@ -66,6 +66,7 @@ from litellm.proxy.auth.auth_utils import (
enforce_batch_enqueued_token_limit_is_admin_only,
enforce_output_token_estimates_are_admin_only,
)
from litellm.proxy.auth.master_key_boot_check import SALT_KEY_ENV_VAR
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import (
evict_and_broadcast,
@ -83,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 (
@ -124,6 +126,7 @@ from litellm.proxy.management_helpers.team_member_permission_checks import (
TeamMemberPermissionChecks,
)
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
from litellm.proxy.pass_through_endpoints.common_utils import reencrypt_general_settings_pass_through
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
from litellm.proxy.spend_tracking.spend_tracking_utils import _is_master_key
from litellm.proxy.utils import (
@ -131,6 +134,7 @@ from litellm.proxy.utils import (
ProxyLogging,
_hash_token_if_needed,
handle_exception_on_proxy,
invalidate_config_param,
is_valid_api_key,
)
from litellm.repositories.base_repository import BaseRepository
@ -5218,6 +5222,53 @@ 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" '
"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.
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):
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 writer.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,
@ -5303,6 +5354,12 @@ async def _rotate_master_key(
data={"param_value": prisma.Json(encrypted_env_vars)},
)
if os.getenv(SALT_KEY_ENV_VAR) is None:
try:
await _reencrypt_pass_through_endpoint_headers(prisma_client, new_master_key)
except Exception as e:
verbose_proxy_logger.warning("Failed to rotate pass-through endpoint headers: %s", str(e))
# 4. process MCP server table
try:
await rotate_mcp_server_credentials_master_key(

View file

@ -1,6 +1,21 @@
import base64
import binascii
from collections.abc import Mapping
from typing import Final
from fastapi import Request
from pydantic import JsonValue
from litellm._logging import verbose_proxy_logger
from litellm.proxy.common_utils.callback_utils import CALLBACK_VAR_ENCRYPTED_PREFIX
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper
# Holds the name of the request header that carries the caller's LiteLLM key, which
# user_api_key_auth reads straight from general_settings, so it stays plaintext.
_CALLER_KEY_HEADER_NAME: Final = "litellm_user_api_key"
_GCM_CIPHERTEXT_PREFIX: Final = "v2:gcm:"
_MIN_GCM_CIPHERTEXT_BYTES: Final = 29
_MIN_NACL_CIPHERTEXT_BYTES: Final = 41
def get_litellm_virtual_key(request: Request) -> str:
@ -16,3 +31,117 @@ def get_litellm_virtual_key(request: Request) -> str:
if litellm_api_key:
return f"Bearer {litellm_api_key}"
return request.headers.get("Authorization", "")
def _encrypt_header_value(name: str, value: JsonValue, new_encryption_key: str | None) -> JsonValue:
if not isinstance(value, str) or not value or name == _CALLER_KEY_HEADER_NAME:
return value
if _is_marked(value):
return value
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
def _decrypted(name: str, value: str) -> str | None:
return decrypt_value_helper(
value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX),
key=name,
exception_type="debug",
return_original_value=False,
)
def _has_ciphertext_shape(payload: str) -> bool:
gcm: Final = payload.startswith(_GCM_CIPHERTEXT_PREFIX)
encoded: Final = payload.removeprefix(_GCM_CIPHERTEXT_PREFIX)
try:
raw = base64.urlsafe_b64decode(encoded)
except (binascii.Error, ValueError):
try:
raw = base64.b64decode(encoded)
except (binascii.Error, ValueError):
return False
return len(raw) >= (_MIN_GCM_CIPHERTEXT_BYTES if gcm else _MIN_NACL_CIPHERTEXT_BYTES)
def _is_marked(value: object) -> bool:
return (
isinstance(value, str)
and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX)
and _has_ciphertext_shape(value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX))
)
def _decrypt_header_value(name: str, value: object) -> object:
if not isinstance(value, str) or not _is_marked(value):
return value
decrypted: Final = _decrypted(name, value)
if decrypted is None:
verbose_proxy_logger.warning(
"Could not decrypt pass-through header %s; check LITELLM_SALT_KEY / master key", name
)
return value
return decrypted
def undecryptable_pass_through_header_names(headers: object) -> frozenset[str]:
"""Names of `litellm_enc::` header values that do not decrypt under the current key."""
if not isinstance(headers, Mapping):
return frozenset()
return frozenset(
str(name)
for name, value in headers.items()
if isinstance(value, str) and _is_marked(value) and _decrypted(str(name), value) is None
)
def decrypt_pass_through_headers(headers: Mapping[str, object] | None) -> dict[str, object] | None:
"""Decrypt `litellm_enc::` header values; plaintext and os.environ/ values pass through."""
if headers is None:
return None
return {name: _decrypt_header_value(name, value) for name, value in headers.items()}
def _with_encrypted_headers(endpoint: JsonValue, new_encryption_key: str | None, reencrypt: bool) -> JsonValue:
if not isinstance(endpoint, dict) or not isinstance(headers := endpoint.get("headers"), dict):
return endpoint
return {
**endpoint,
"headers": {
name: _encrypt_header_value(
name, _decrypt_header_value(name, value) if reencrypt else value, new_encryption_key
)
for name, value in headers.items()
},
}
def encrypt_pass_through_endpoints(endpoints: JsonValue) -> JsonValue:
"""Encrypt every endpoint's header values for storage; values already carrying the marker are kept."""
if not isinstance(endpoints, list):
return endpoints
return [_with_encrypted_headers(endpoint, None, reencrypt=False) for endpoint in endpoints]
def reencrypt_general_settings_pass_through(
general_settings: object, new_encryption_key: str
) -> dict[str, JsonValue] | None:
"""Return general_settings with pass-through header values re-encrypted under new_encryption_key.
None when the row holds no pass-through endpoint list.
"""
if not isinstance(general_settings, dict) or not isinstance(
endpoints := general_settings.get("pass_through_endpoints"), list
):
return None
return {
**general_settings,
"pass_through_endpoints": [
_with_encrypted_headers(endpoint, new_encryption_key, reencrypt=True) for endpoint in endpoints
],
}

View file

@ -101,6 +101,10 @@ from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
_get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
)
from litellm.proxy.pass_through_endpoints.common_utils import (
decrypt_pass_through_headers,
undecryptable_pass_through_header_names,
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import normalize_route_for_root_path
from litellm.repositories.team_repository import TeamRepository
@ -161,14 +165,15 @@ async def set_env_variables_in_header(custom_headers: dict | None) -> dict | Non
"""
if custom_headers is None:
return None
stored_headers: Final = decrypt_pass_through_headers(custom_headers) or {}
headers: Final = {}
for key, value in custom_headers.items():
for key, value in stored_headers.items():
# langfuse Api requires base64 encoded headers - it's simpleer to just ask litellm users to set their langfuse public and secret keys
# we can then get the b64 encoded keys here
if key == "LANGFUSE_PUBLIC_KEY" or key == "LANGFUSE_SECRET_KEY":
# langfuse requires b64 encoded headers - we construct that here
_langfuse_public_key = custom_headers["LANGFUSE_PUBLIC_KEY"]
_langfuse_secret_key = custom_headers["LANGFUSE_SECRET_KEY"]
_langfuse_public_key = stored_headers["LANGFUSE_PUBLIC_KEY"]
_langfuse_secret_key = stored_headers["LANGFUSE_SECRET_KEY"]
if isinstance(_langfuse_public_key, str) and _langfuse_public_key.startswith("os.environ/"):
_langfuse_public_key = get_secret_str(_langfuse_public_key)
if isinstance(_langfuse_secret_key, str) and _langfuse_secret_key.startswith("os.environ/"):
@ -196,6 +201,67 @@ 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:
by_id: Final[list[Mapping[str, object]]] = []
by_path: Final[dict[str, Mapping[str, object]]] = {}
for route in _registered_pass_through_routes.values():
params = route.get("passthrough_params")
if (
not isinstance(params, Mapping)
or params.get("target") != target
or not isinstance(headers := params.get("custom_headers"), Mapping)
):
continue
served_headers = cast(Mapping[str, object], headers)
if route["endpoint_id"] == endpoint_id:
by_id.append(served_headers)
elif path_without_stored_id is not None and route.get("path") == path_without_stored_id:
by_path[str(route["endpoint_id"])] = served_headers
if by_id:
return by_id[0]
return next(iter(by_path.values())) if len(by_path) == 1 else None
async def _resolve_stored_headers(
endpoint_id: str,
target: str | None,
stored_headers: dict | None,
path_without_stored_id: str | None = None,
) -> dict | None:
"""Resolve an endpoint's stored headers for outbound use.
A header that no longer decrypts under the current key (between an in-app master key
rotation and the restart) keeps the value the endpoint already sends to the same target;
every other header resolves from the stored value. The served endpoint is matched by id,
or, for a stored endpoint without an id (it gets a new one on every reload), by path when
exactly one served endpoint has that path and target.
"""
undecryptable: Final = undecryptable_pass_through_header_names(stored_headers)
registered: Final = _served_custom_headers(endpoint_id, target, path_without_stored_id) if undecryptable else None
if stored_headers is None or registered is None:
return await set_env_variables_in_header(custom_headers=stored_headers)
verbose_proxy_logger.warning(
"Pass-through endpoint %s keeps its current value for headers %s: the stored values do not decrypt "
"under the current key (restart with the new master key after a rotation)",
endpoint_id,
sorted(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(
fastapi_response: Response,
request: Request,
@ -3247,7 +3313,8 @@ async def _register_pass_through_endpoint(
else:
endpoint_data = endpoint
if endpoint_data.get("id") is None:
stored_without_id: Final = endpoint_data.get("id") is None
if stored_without_id:
endpoint_data["id"] = str(uuid.uuid4())
endpoint_id: Final = cast(str, endpoint_data["id"])
@ -3256,7 +3323,9 @@ async def _register_pass_through_endpoint(
if path is None:
raise ValueError("Path is required for pass-through endpoint")
custom_headers: Final = await set_env_variables_in_header(custom_headers=endpoint_data.get("headers"))
custom_headers: Final = await _resolve_stored_headers(
endpoint_id, target, endpoint_data.get("headers"), path if stored_without_id else None
)
forward_headers: Final = endpoint_data.get("forward_headers")
merge_query_params: Final = endpoint_data.get("merge_query_params")
default_query_params: Final = endpoint_data.get("default_query_params")
@ -3677,6 +3746,10 @@ async def update_pass_through_endpoints(
# Update the list
pass_through_endpoint_data[endpoint_index] = endpoint_dict
_custom_headers: Final = await _resolve_stored_headers(
endpoint_id, updated_endpoint.target, updated_endpoint.headers or {}
)
# Remove old routes from registry before they get re-registered
InitPassThroughEndpointHelpers.remove_endpoint_routes(endpoint_id)
@ -3689,10 +3762,6 @@ async def update_pass_through_endpoints(
await update_config_general_settings(data=updated_data, user_api_key_dict=user_api_key_dict)
# Re-register the route with updated headers
_custom_headers: dict | None = updated_endpoint.headers or {}
_custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers)
route_app: Final = _request_app(request)
if updated_endpoint.include_subpath:
InitPassThroughEndpointHelpers.add_subpath_route(

View file

@ -740,6 +740,7 @@ from litellm.proxy.openai_files_endpoints.files_endpoints import (
from litellm.proxy.openai_files_endpoints.files_endpoints import (
set_files_config,
)
from litellm.proxy.pass_through_endpoints.common_utils import encrypt_pass_through_endpoints
from litellm.proxy.pass_through_endpoints.openai_passthrough_endpoints import (
router as openai_passthrough_router,
)
@ -17988,7 +17989,7 @@ async def update_config(
existing["alerting"] = ["slack"]
elif isinstance(existing["alerting"], list) and "slack" not in existing["alerting"]:
existing["alerting"].append("slack")
existing[k] = v
existing[k] = encrypt_pass_through_endpoints(v) if k == "pass_through_endpoints" else v
await _upsert_section("general_settings", existing)
asyncio.create_task(
create_config_audit_log(
@ -18249,6 +18250,8 @@ async def update_config_general_settings(
field_value = data.field_value
if data.field_name == "plugins":
field_value = _preserve_redacted_plugin_keys(field_value, general_settings.get("plugins"))
if data.field_name == "pass_through_endpoints":
field_value = encrypt_pass_through_endpoints(field_value)
general_settings[data.field_name] = cast(JsonValue, field_value) # cast-ok: ConfigGeneralSettings validated it

View file

@ -21102,3 +21102,181 @@ class TestTeamAdminMemberKeyBudgetUpdate:
)
assert exc.value.status_code == 403
assert "member_key_budgets" not in str(exc.value.detail)
@pytest.mark.asyncio
async def test_rotate_master_key_reencrypts_pass_through_endpoint_headers(monkeypatch):
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 (
decrypt_pass_through_headers,
encrypt_pass_through_endpoints,
)
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master-key")
stored_endpoints = encrypt_pass_through_endpoints(
[
{"path": "/a", "headers": {"x-a": "plain-a"}},
{"path": "/b", "headers": {"Authorization": "Bearer sk-literal"}},
]
)
general_settings_row = MagicMock(
param_name="general_settings",
param_value={"store_model_in_db": True, "pass_through_endpoints": stored_endpoints},
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
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(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=[])
invalidate = AsyncMock()
monkeypatch.setattr(key_management_endpoints, "invalidate_config_param", invalidate)
for rotate_helper in (
"rotate_mcp_server_credentials_master_key",
"rotate_mcp_user_credentials_master_key",
"rotate_mcp_user_env_vars_master_key",
"rotate_sso_identity_assertions_master_key",
):
monkeypatch.setattr(key_management_endpoints, rotate_helper, AsyncMock())
await key_management_endpoints._rotate_master_key(
prisma_client=mock_prisma_client,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"),
current_master_key="sk-old-master-key",
new_master_key="sk-new-master-key",
)
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 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
async def test_rotate_master_key_leaves_pass_through_headers_under_salt_key(monkeypatch):
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 encrypt_pass_through_endpoints
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-stays")
general_settings_row = MagicMock(
param_name="general_settings",
param_value={
"pass_through_endpoints": encrypt_pass_through_endpoints(
[{"path": "/a", "headers": {"Authorization": "Bearer sk-literal"}}]
)
},
)
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])
mock_prisma_client.db.litellm_config.update = AsyncMock()
mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[])
for rotate_helper in (
"rotate_mcp_server_credentials_master_key",
"rotate_mcp_user_credentials_master_key",
"rotate_mcp_user_env_vars_master_key",
"rotate_sso_identity_assertions_master_key",
):
monkeypatch.setattr(key_management_endpoints, rotate_helper, AsyncMock())
await key_management_endpoints._rotate_master_key(
prisma_client=mock_prisma_client,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"),
current_master_key="sk-old-master-key",
new_master_key="sk-new-master-key",
)
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_continues_when_pass_through_header_reencryption_fails(monkeypatch):
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints import key_management_endpoints
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
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=[MagicMock(param_name="general_settings", param_value={})]
)
mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(return_value=[])
monkeypatch.setattr(
key_management_endpoints,
"_reencrypt_pass_through_endpoint_headers",
AsyncMock(side_effect=RuntimeError("general_settings write fault")),
)
later_steps = {
name: AsyncMock()
for name in (
"rotate_mcp_server_credentials_master_key",
"rotate_mcp_user_credentials_master_key",
"rotate_mcp_user_env_vars_master_key",
"rotate_sso_identity_assertions_master_key",
)
}
for name, step in later_steps.items():
monkeypatch.setattr(key_management_endpoints, name, step)
await key_management_endpoints._rotate_master_key(
prisma_client=mock_prisma_client,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin"),
current_master_key="sk-old-master-key",
new_master_key="sk-new-master-key",
)
assert all(step.await_count == 1 for step in later_steps.values())
mock_prisma_client.db.litellm_credentialstable.find_many.assert_awaited_once()
@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()

View file

@ -7683,3 +7683,482 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token()
metadata = kwargs["litellm_params"]["metadata"]
assert metadata["user_api_key"] == "cli-session-alice"
assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice"
@pytest.mark.asyncio
async def test_set_env_variables_in_header_decrypts_stored_headers(monkeypatch):
from litellm.proxy.pass_through_endpoints.common_utils import encrypt_pass_through_endpoints
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import set_env_variables_in_header
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-pass-through-tests")
monkeypatch.setenv("UPSTREAM_TOKEN", "resolved-from-env")
[stored] = encrypt_pass_through_endpoints(
[
{
"path": "/pt",
"headers": {
"x-legacy": "plain-value",
"Authorization": "Bearer sk-literal",
"x-env": "Bearer os.environ/UPSTREAM_TOKEN",
},
}
]
)
headers = {**stored["headers"], "x-legacy": "plain-value"}
assert await set_env_variables_in_header(custom_headers=headers) == {
"x-legacy": "plain-value",
"Authorization": "Bearer sk-literal",
"x-env": "Bearer resolved-from-env",
}
@pytest.mark.asyncio
async def test_register_pass_through_endpoint_keeps_serving_headers_when_stored_headers_do_not_decrypt(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-rotate",
"path": "/rotate-pt",
"target": "http://upstream",
"headers": {"Authorization": "Bearer sk-literal"},
}
]
)
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={"id": "ep-rotate-other", "path": "/other-pt", "target": "http://other", "headers": {"x": "y"}},
app=app,
premium_user=False,
visited_endpoints=set(),
)
first_visit: set[str] = set()
await _register_pass_through_endpoint(
endpoint=dict(stored), app=app, premium_user=False, visited_endpoints=first_visit
)
[route_key] = first_visit
second_visit: set[str] = set()
await _register_pass_through_endpoint(
endpoint={
**rotated_endpoint,
"path": "/rotate-pt-moved",
"headers": {**rotated_endpoint["headers"], "x-org": "acme"},
},
app=app,
premium_user=False,
visited_endpoints=second_visit,
)
[moved_route_key] = second_visit
assert moved_route_key != route_key
params = _registered_pass_through_routes[moved_route_key]["passthrough_params"]
assert params["custom_headers"] == {"Authorization": "Bearer sk-literal", "x-org": "acme"}
retargeted_visit: set[str] = set()
await _register_pass_through_endpoint(
endpoint={**rotated_endpoint, "path": "/rotate-pt-moved", "target": "http://upstream-moved"},
app=app,
premium_user=False,
visited_endpoints=retargeted_visit,
)
[retargeted_route_key] = retargeted_visit
retargeted = _registered_pass_through_routes[retargeted_route_key]["passthrough_params"]
assert retargeted["target"] == "http://upstream-moved"
assert retargeted["custom_headers"]["Authorization"] != "Bearer sk-literal"
finally:
for key in [k for k in _registered_pass_through_routes if k.startswith("ep-rotate")]:
_registered_pass_through_routes.pop(key, None)
@pytest.mark.asyncio
async def test_register_pass_through_endpoint_without_id_keeps_serving_headers_when_stored_headers_do_not_decrypt(
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(
[{"path": "/rotate-no-id-pt", "target": "http://upstream", "headers": {"Authorization": "Bearer sk-no-id"}}]
)
rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": [stored]}, "sk-next-key")
assert rotated is not None
[rotated_endpoint] = rotated["pass_through_endpoints"]
visited: set[str] = set()
try:
await _register_pass_through_endpoint(
endpoint={"path": "/rotate-no-id-other", "target": "http://other", "headers": {"x": "y"}},
app=app,
premium_user=False,
visited_endpoints=visited,
)
await _register_pass_through_endpoint(
endpoint=dict(stored), app=app, premium_user=False, visited_endpoints=visited
)
reloaded: set[str] = set()
await _register_pass_through_endpoint(
endpoint=dict(rotated_endpoint), app=app, premium_user=False, visited_endpoints=reloaded
)
[route_key] = reloaded
params = _registered_pass_through_routes[route_key]["passthrough_params"]
assert params["custom_headers"] == {"Authorization": "Bearer sk-no-id"}
finally:
for key in [k for k in _registered_pass_through_routes if "/rotate-no-id" in k]:
_registered_pass_through_routes.pop(key, None)
@pytest.mark.asyncio
async def test_register_pass_through_endpoint_keeps_its_own_headers_when_another_endpoint_shares_the_path(
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-shared-a",
"path": "/shared-pt",
"methods": ["GET"],
"target": "http://a",
"headers": {"Authorization": "Bearer sk-a"},
},
{
"id": "ep-shared-b",
"path": "/shared-pt",
"methods": ["POST"],
"target": "http://b",
"headers": {"Authorization": "Bearer sk-b"},
},
{
"path": "/shared-pt",
"methods": ["PUT"],
"target": "http://c",
"headers": {"Authorization": "Bearer sk-c"},
},
]
)
rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": stored}, "sk-next-key")
assert rotated is not None
try:
for endpoint in stored:
await _register_pass_through_endpoint(
endpoint=dict(endpoint), app=app, premium_user=False, visited_endpoints=set()
)
reloaded_b: set[str] = set()
await _register_pass_through_endpoint(
endpoint=dict(rotated["pass_through_endpoints"][1]),
app=app,
premium_user=False,
visited_endpoints=reloaded_b,
)
reloaded_c: set[str] = set()
await _register_pass_through_endpoint(
endpoint=dict(rotated["pass_through_endpoints"][2]),
app=app,
premium_user=False,
visited_endpoints=reloaded_c,
)
[route_b] = reloaded_b
assert _registered_pass_through_routes[route_b]["passthrough_params"]["custom_headers"] == {
"Authorization": "Bearer sk-b"
}
[route_c] = reloaded_c
assert _registered_pass_through_routes[route_c]["passthrough_params"]["custom_headers"] == {
"Authorization": "Bearer sk-c"
}
twins = encrypt_pass_through_endpoints(
[
{
"path": "/shared-twin-pt",
"methods": [method],
"target": "http://twin",
"headers": {"Authorization": f"Bearer sk-{method}"},
}
for method in ("GET", "POST")
]
)
rotated_twins = reencrypt_general_settings_pass_through({"pass_through_endpoints": twins}, "sk-next-key")
assert rotated_twins is not None
for twin in twins:
await _register_pass_through_endpoint(
endpoint=dict(twin), app=app, premium_user=False, visited_endpoints=set()
)
reloaded_twin: set[str] = set()
await _register_pass_through_endpoint(
endpoint=dict(rotated_twins["pass_through_endpoints"][0]),
app=app,
premium_user=False,
visited_endpoints=reloaded_twin,
)
[route_twin] = reloaded_twin
served_twin = _registered_pass_through_routes[route_twin]["passthrough_params"]["custom_headers"]
assert served_twin["Authorization"] not in ("Bearer sk-GET", "Bearer sk-POST")
[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": {"Authorization": "Bearer sk-e"},
},
]
)
await _register_pass_through_endpoint(endpoint=solo, app=app, premium_user=False, visited_endpoints=set())
rotated_newcomer = reencrypt_general_settings_pass_through(
{"pass_through_endpoints": [newcomer]}, "sk-next-key"
)
assert rotated_newcomer is not None
reloaded_e: set[str] = set()
await _register_pass_through_endpoint(
endpoint=dict(rotated_newcomer["pass_through_endpoints"][0]),
app=app,
premium_user=False,
visited_endpoints=reloaded_e,
)
[route_e] = reloaded_e
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"}}]
)
rotated_moved = reencrypt_general_settings_pass_through({"pass_through_endpoints": [moved]}, "sk-next-key")
assert rotated_moved is not None
reloaded_moved: set[str] = set()
await _register_pass_through_endpoint(
endpoint=dict(rotated_moved["pass_through_endpoints"][0]),
app=app,
premium_user=False,
visited_endpoints=reloaded_moved,
)
[route_moved] = reloaded_moved
served_moved = _registered_pass_through_routes[route_moved]["passthrough_params"]["custom_headers"]
assert served_moved["Authorization"] not in ("Bearer sk-a", "Bearer sk-b", "Bearer sk-c", "Bearer sk-moved")
finally:
for key in [k for k in _registered_pass_through_routes if "/shared-" in k]:
_registered_pass_through_routes.pop(key, None)
@pytest.mark.asyncio
async def test_register_pass_through_endpoint_reusing_an_id_does_not_get_that_endpoints_headers(monkeypatch):
from fastapi import FastAPI
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()
try:
await _register_pass_through_endpoint(
endpoint={
"id": "ep-reused",
"path": "/reused-victim",
"target": "http://victim",
"headers": {"Authorization": "Bearer sk-victim"},
},
app=app,
premium_user=False,
visited_endpoints=set(),
)
copied: set[str] = set()
await _register_pass_through_endpoint(
endpoint={
"id": "ep-reused",
"path": "/reused-copy",
"target": "http://elsewhere",
"headers": {"Authorization": "litellm_enc::not-a-ciphertext"},
},
app=app,
premium_user=False,
visited_endpoints=copied,
)
[copied_route] = copied
assert _registered_pass_through_routes[copied_route]["passthrough_params"]["custom_headers"] == {
"Authorization": "litellm_enc::not-a-ciphertext"
}
finally:
for key in [k for k in _registered_pass_through_routes if k.startswith("ep-reused")]:
_registered_pass_through_routes.pop(key, None)
@pytest.mark.asyncio
async def test_update_pass_through_endpoint_keeps_serving_headers_when_stored_headers_do_not_decrypt(monkeypatch):
from fastapi import FastAPI
from litellm.proxy._types import ConfigFieldInfo, PassThroughGenericEndpoint, UserAPIKeyAuth
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,
update_pass_through_endpoints,
)
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-current-key")
app = FastAPI()
[stored] = encrypt_pass_through_endpoints(
[
{
"id": "ep-update-rotate",
"path": "/update-rotate-pt",
"target": "http://upstream",
"headers": {"Authorization": "Bearer sk-update"},
}
]
)
rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": [stored]}, "sk-next-key")
assert rotated is not None
[rotated_endpoint] = rotated["pass_through_endpoints"]
request = MagicMock(spec=Request)
request.app = app
try:
await _register_pass_through_endpoint(
endpoint=dict(stored), app=app, premium_user=False, visited_endpoints=set()
)
with (
patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config,
patch("litellm.proxy.proxy_server.update_config_general_settings"),
):
mock_get_config.return_value = ConfigFieldInfo(
field_name="pass_through_endpoints", field_value=[rotated_endpoint]
)
await update_pass_through_endpoints(
endpoint_id="ep-update-rotate",
data=PassThroughGenericEndpoint(
path="/update-rotate-pt",
target="http://upstream",
headers={**rotated_endpoint["headers"], "x-org": "acme"},
),
request=request,
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
)
[served] = [
route["passthrough_params"]
for route in _registered_pass_through_routes.values()
if route["endpoint_id"] == "ep-update-rotate"
]
assert served["custom_headers"] == {"Authorization": "Bearer sk-update", "x-org": "acme"}
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

@ -102,3 +102,157 @@ def test_encode_bedrock_runtime_modelid_arn_partition_arns() -> None:
endpoint = "model/arn:aws-us-gov:bedrock:us-gov-west-1:123456789012:inference-profile/test-profile/invoke"
expected = "model/arn:aws-us-gov:bedrock:us-gov-west-1:123456789012:inference-profile%2Ftest-profile/invoke"
assert CommonUtils.encode_bedrock_runtime_modelid_arn(endpoint) == expected
from litellm.proxy.pass_through_endpoints.common_utils import (
decrypt_pass_through_headers,
encrypt_pass_through_endpoints,
undecryptable_pass_through_header_names,
reencrypt_general_settings_pass_through,
)
_ENC = "litellm_enc::"
@pytest.fixture
def salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-pass-through-tests")
def test_encrypt_pass_through_endpoints_encrypts_every_header_of_every_endpoint(salt_key):
stored = encrypt_pass_through_endpoints(
[
{"path": "/a", "headers": {"x-plain": "value-a"}},
{
"path": "/b",
"headers": {
"Authorization": "Bearer sk-literal",
"x-env": "Bearer os.environ/UPSTREAM_KEY",
"litellm_user_api_key": "x-my-key",
"x-number": 5,
"x-empty": "",
},
},
{"path": "/c"},
]
)
assert stored[0]["headers"]["x-plain"].startswith(_ENC)
headers_b = stored[1]["headers"]
assert headers_b["Authorization"].startswith(_ENC)
assert "sk-literal" not in headers_b["Authorization"]
assert headers_b["x-env"].startswith(_ENC)
assert headers_b["litellm_user_api_key"] == "x-my-key"
assert headers_b["x-number"] == 5
assert headers_b["x-empty"] == ""
assert stored[2] == {"path": "/c"}
assert decrypt_pass_through_headers(headers_b) == {
"Authorization": "Bearer sk-literal",
"x-env": "Bearer os.environ/UPSTREAM_KEY",
"litellm_user_api_key": "x-my-key",
"x-number": 5,
"x-empty": "",
}
def test_encrypt_pass_through_endpoints_keeps_existing_ciphertext_and_input(salt_key):
endpoints = [{"path": "/a", "headers": {"Authorization": "Bearer sk-literal"}}]
first = encrypt_pass_through_endpoints(endpoints)
second = encrypt_pass_through_endpoints(first)
assert second == first
assert endpoints == [{"path": "/a", "headers": {"Authorization": "Bearer sk-literal"}}]
assert encrypt_pass_through_endpoints(None) is None
def test_decrypt_pass_through_headers_passes_legacy_plaintext_through(salt_key):
legacy = {"Authorization": "Bearer sk-legacy", "x-env": "os.environ/UPSTREAM_KEY"}
assert decrypt_pass_through_headers(legacy) == legacy
assert decrypt_pass_through_headers(None) is None
def test_decrypt_pass_through_headers_keeps_value_it_cannot_decrypt(salt_key, monkeypatch):
stored = encrypt_pass_through_endpoints([{"path": "/a", "headers": {"Authorization": "Bearer sk-literal"}}])
ciphertext = stored[0]["headers"]["Authorization"]
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-some-other-salt")
assert decrypt_pass_through_headers({"Authorization": ciphertext}) == {"Authorization": ciphertext}
def test_reencrypt_general_settings_pass_through_moves_headers_to_new_key(salt_key, monkeypatch):
stored = encrypt_pass_through_endpoints(
[
{"path": "/a", "headers": {"x-a": "value-a"}},
{"path": "/b", "headers": {"Authorization": "Bearer sk-b"}},
]
)
general_settings = {
"store_model_in_db": True,
"pass_through_endpoints": [
stored[0],
{**stored[1], "headers": {**stored[1]["headers"], "x-legacy": "plain-b"}},
],
}
rotated = reencrypt_general_settings_pass_through(general_settings, "sk-new-master")
assert rotated is not None
assert rotated["store_model_in_db"] is True
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-new-master")
assert decrypt_pass_through_headers(rotated["pass_through_endpoints"][0]["headers"]) == {"x-a": "value-a"}
assert decrypt_pass_through_headers(rotated["pass_through_endpoints"][1]["headers"]) == {
"Authorization": "Bearer sk-b",
"x-legacy": "plain-b",
}
assert all(
value.startswith(_ENC)
for endpoint in rotated["pass_through_endpoints"]
for value in endpoint["headers"].values()
)
assert reencrypt_general_settings_pass_through({"store_model_in_db": True}, "sk-new-master") is None
assert reencrypt_general_settings_pass_through(None, "sk-new-master") is None
def test_undecryptable_pass_through_header_names(salt_key):
[stored] = encrypt_pass_through_endpoints([{"path": "/a", "headers": {"x-a": "plain-a"}}])
rotated = reencrypt_general_settings_pass_through({"pass_through_endpoints": [stored]}, "sk-next-key")
assert undecryptable_pass_through_header_names(stored["headers"]) == frozenset()
assert undecryptable_pass_through_header_names({"x-legacy": "plain", "x-n": 5}) == frozenset()
assert undecryptable_pass_through_header_names(None) == frozenset()
assert undecryptable_pass_through_header_names(rotated["pass_through_endpoints"][0]["headers"]) == {"x-a"}
assert undecryptable_pass_through_header_names(
{
"x-ok": stored["headers"]["x-a"],
"x-bad": rotated["pass_through_endpoints"][0]["headers"]["x-a"],
"x-garbage": _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()
@pytest.mark.parametrize("value", [_ENC + "***", _ENC + "!!", _ENC + "not-a-ciphertext", _ENC + "v2:gcm:abc"])
def test_marker_prefixed_literal_is_encrypted_and_forwarded_unchanged(salt_key, value):
[stored] = encrypt_pass_through_endpoints([{"path": "/a", "headers": {"x-tag": value}}])
assert stored["headers"]["x-tag"] != value
assert decrypt_pass_through_headers(stored["headers"]) == {"x-tag": value}
assert decrypt_pass_through_headers({"x-tag": value}) == {"x-tag": value}
assert undecryptable_pass_through_header_names({"x-tag": value}) == frozenset()
def test_aes_gcm_ciphertext_from_another_key_is_undecryptable(salt_key, monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"encryption_algorithm": "aes-256-gcm"})
[stored] = encrypt_pass_through_endpoints([{"path": "/a", "headers": {"x-a": "v"}}])
assert stored["headers"]["x-a"].startswith(_ENC + "v2:gcm:")
assert decrypt_pass_through_headers(stored["headers"]) == {"x-a": "v"}
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-some-other-salt")
assert undecryptable_pass_through_header_names(stored["headers"]) == {"x-a"}

View file

@ -15561,3 +15561,64 @@ async def test_spend_capture_rate_check_job_clears_the_gauge_once_the_setting_is
call(api_provider="openai", capture_rate=0.97),
call(api_provider="openai", capture_rate=None),
]
def test_update_config_encrypts_pass_through_endpoint_headers(_update_config_setup, monkeypatch):
from litellm.proxy.pass_through_endpoints.common_utils import decrypt_pass_through_headers
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-pass-through-tests")
client, prisma, restore = _update_config_setup(initial_rows={"general_settings": {"store_model_in_db": True}})
try:
resp = client.post(
"/config/update",
json={
"general_settings": {
"pass_through_endpoints": [
{"path": "/pt", "target": "http://upstream", "headers": {"Authorization": "Bearer sk-literal"}}
]
}
},
)
assert resp.status_code == 200
stored = prisma.db.litellm_config.rows["general_settings"]
stored_headers = stored["pass_through_endpoints"][0]["headers"]
assert "sk-literal" not in json.dumps(stored)
assert decrypt_pass_through_headers(stored_headers) == {"Authorization": "Bearer sk-literal"}
assert stored["store_model_in_db"] is True
finally:
restore()
@pytest.mark.asyncio
async def test_update_config_general_settings_encrypts_pass_through_endpoint_headers(monkeypatch):
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import ConfigFieldUpdate, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.common_utils import decrypt_pass_through_headers
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-pass-through-tests")
prisma = _FakePrismaClient(initial_rows={"general_settings": {"store_model_in_db": True}})
monkeypatch.setattr(ps, "prisma_client", prisma)
monkeypatch.setattr(ps, "invalidate_config_param", AsyncMock(return_value=None))
monkeypatch.setattr(ps.proxy_config.settings, "apply_db_row", MagicMock())
monkeypatch.setattr(ps, "create_config_audit_log", AsyncMock(return_value=None))
endpoints = [
{"path": "/a", "target": "http://upstream-a", "headers": {"x-a": "plain-a"}},
{"path": "/b", "target": "http://upstream-b", "headers": {"Authorization": "Bearer sk-literal"}},
]
await ps.update_config_general_settings(
data=ConfigFieldUpdate(
field_name="pass_through_endpoints", field_value=endpoints, config_type="general_settings"
),
user_api_key_dict=UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN),
)
stored = prisma.db.litellm_config.rows["general_settings"]
assert "sk-literal" not in json.dumps(stored)
assert "plain-a" not in json.dumps(stored)
assert [decrypt_pass_through_headers(e["headers"]) for e in stored["pass_through_endpoints"]] == [
{"x-a": "plain-a"},
{"Authorization": "Bearer sk-literal"},
]
assert endpoints[1]["headers"] == {"Authorization": "Bearer sk-literal"}
ps.proxy_config.settings.apply_db_row.assert_called_once_with("general_settings", stored)