mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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.
This commit is contained in:
parent
9ca392868d
commit
d3c0dd1f80
2 changed files with 51 additions and 15 deletions
|
|
@ -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))
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue