mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 4b2902fa12 into f285229b51
This commit is contained in:
commit
953a3afbaa
8 changed files with 1140 additions and 10 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue