From d3c0dd1f8010b41f0b2eddf8a7e0a70f7820fccd Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Tue, 29 Sep 2026 11:53:37 -0700 Subject: [PATCH] fix(proxy): treat marker-prefixed header literals as plaintext A litellm_enc:: value now counts as stored ciphertext only when its payload decodes to at least the size encrypt_value_helper produces (nacl or AES-GCM). Shorter or non-base64 payloads such as litellm_enc::*** are literals: they are encrypted on write and forwarded unchanged, instead of being reported as undecryptable, which made an edit to such a value keep the previously served header. --- .../pass_through_endpoints/common_utils.py | 38 +++++++++++++------ ...test_passthrough_endpoints_common_utils.py | 28 ++++++++++++-- 2 files changed, 51 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/common_utils.py b/litellm/proxy/pass_through_endpoints/common_utils.py index a27bfd4abe7..54fc9b7e252 100644 --- a/litellm/proxy/pass_through_endpoints/common_utils.py +++ b/litellm/proxy/pass_through_endpoints/common_utils.py @@ -1,3 +1,5 @@ +import base64 +import binascii from collections.abc import Mapping from typing import Final @@ -11,6 +13,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helpe # 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:" +# Smallest encrypt_value_helper output for a one-byte value: nonce + tag (AES-GCM), nonce + MAC (nacl). +_MIN_GCM_CIPHERTEXT_BYTES: Final = 29 +_MIN_NACL_CIPHERTEXT_BYTES: Final = 41 def get_litellm_virtual_key(request: Request) -> str: @@ -31,7 +37,7 @@ def get_litellm_virtual_key(request: Request) -> str: 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 value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX): + if _is_marked(value): return value try: return CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(value, new_encryption_key=new_encryption_key) @@ -43,23 +49,33 @@ def _encrypt_header_value(name: str, value: JsonValue, new_encryption_key: str | def _decrypted(name: str, value: str) -> str | None: - """Plaintext of a marked value; None when it does not decrypt or decrypts to an empty string.""" - return ( - decrypt_value_helper( - value=value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX), - key=name, - exception_type="debug", - return_original_value=False, - ) - or 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: + """True for a `litellm_enc::` value whose payload has the shape encrypt_value_helper produces.""" return ( isinstance(value, str) and value.startswith(CALLBACK_VAR_ENCRYPTED_PREFIX) - and value != CALLBACK_VAR_ENCRYPTED_PREFIX + and _has_ciphertext_shape(value.removeprefix(CALLBACK_VAR_ENCRYPTED_PREFIX)) ) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py index 5e579e4d2b1..05c140f14c8 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoints_common_utils.py @@ -223,7 +223,12 @@ def test_undecryptable_pass_through_header_names(salt_key): 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": _ENC + "garbage", "x-plain": "p"} + { + "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"} @@ -232,7 +237,22 @@ def test_decrypt_pass_through_headers_keeps_a_bare_marker_literal(salt_key): assert undecryptable_pass_through_header_names({"x-tag": _ENC}) == frozenset() -@pytest.mark.parametrize("value", [_ENC + "***", _ENC + "!!"]) -def test_decrypt_pass_through_headers_keeps_a_marked_value_that_decrypts_to_nothing(salt_key, value): +@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}) == {"x-tag"} + 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"}