mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge 0ae5d2c39b into 2dccc0dc79
This commit is contained in:
commit
015e7000e3
4 changed files with 91 additions and 17 deletions
|
|
@ -112,7 +112,10 @@ from litellm.proxy.native_compaction import with_proxy_compaction_executor
|
|||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
|
||||
from litellm.router import Router
|
||||
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
|
||||
from litellm.router_utils.add_retry_fallback_headers import (
|
||||
get_hidden_params_dict,
|
||||
safe_header_value,
|
||||
)
|
||||
from litellm.router_utils.common_utils import resolve_model_group_alias
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.router import RouterRateLimitError
|
||||
|
|
@ -1830,7 +1833,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
headers.update(logging_caching_headers)
|
||||
|
||||
try:
|
||||
return {key: str(value) for key, value in headers.items() if value not in exclude_values}
|
||||
return {key: safe_header_value(str(value)) for key, value in headers.items() if value not in exclude_values}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error setting custom headers: %s", e)
|
||||
return {}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import math
|
|||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Protocol, TypedDict, cast
|
||||
from urllib.parse import quote
|
||||
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
|
|
@ -121,6 +122,33 @@ def response_has_hidden_params(response: object) -> bool:
|
|||
return hasattr(response, "_hidden_params")
|
||||
|
||||
|
||||
def safe_header_value(value: str) -> str:
|
||||
"""
|
||||
HTTP header values must be latin-1 encodable (RFC 7230). Values derived from
|
||||
user-configured strings (e.g. a model name) can contain non-latin-1
|
||||
characters, which crashes response header assembly with a UnicodeEncodeError.
|
||||
Percent-encode such values so they stay ASCII-safe and reversible via
|
||||
``urllib.parse.unquote``, instead of dropping them or crashing the request.
|
||||
|
||||
A value that is already latin-1 safe but contains a literal ``%`` is also
|
||||
percent-encoded: left alone, it would be indistinguishable from a genuinely
|
||||
encoded value once mixed with the encoded case, so a caller applying
|
||||
``unquote`` unconditionally would silently decode the wrong string.
|
||||
"""
|
||||
is_latin1_safe: Final = _is_latin1_safe(value)
|
||||
if is_latin1_safe and "%" not in value:
|
||||
return value
|
||||
return quote(value, safe="")
|
||||
|
||||
|
||||
def _is_latin1_safe(value: str) -> bool:
|
||||
try:
|
||||
value.encode("latin-1")
|
||||
return True
|
||||
except UnicodeEncodeError:
|
||||
return False
|
||||
|
||||
|
||||
def ensure_response_additional_headers(response: object) -> dict[str, object]:
|
||||
hidden_params: Final = get_hidden_params_dict(response, create=isinstance(response, dict))
|
||||
_write_hidden_params(response, hidden_params)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from types import MappingProxyType, SimpleNamespace
|
|||
from typing import AsyncGenerator, Callable, Final, Iterator, Optional, Sequence
|
||||
from urllib.parse import unquote_plus
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from urllib.parse import unquote
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -1353,6 +1354,33 @@ class TestProxyBaseLLMRequestProcessing:
|
|||
assert "x-litellm-response-cost-original" not in headers
|
||||
assert "x-litellm-response-cost-discount-amount" not in headers
|
||||
|
||||
def test_get_custom_headers_percent_encodes_non_latin1_model_name(self):
|
||||
"""
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/39284:
|
||||
a deployment or model-group name outside latin-1 (e.g. Chinese) used to
|
||||
crash header assembly with UnicodeEncodeError when FastAPI/Starlette
|
||||
encoded the response headers.
|
||||
"""
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
mock_user_api_key_dict.tpm_limit = None
|
||||
mock_user_api_key_dict.rpm_limit = None
|
||||
mock_user_api_key_dict.max_budget = None
|
||||
mock_user_api_key_dict.spend = 0
|
||||
|
||||
headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
call_id="test-call-id",
|
||||
model_id="中文模型",
|
||||
)
|
||||
|
||||
encoded_model_id = headers["x-litellm-model-id"]
|
||||
assert encoded_model_id.encode("latin-1")
|
||||
assert unquote(encoded_model_id) == "中文模型"
|
||||
|
||||
# Would raise UnicodeEncodeError before the fix.
|
||||
response = Response(headers=headers)
|
||||
assert response.headers["x-litellm-model-id"] == encoded_model_id
|
||||
|
||||
def test_get_custom_headers_with_margin_info(self):
|
||||
"""
|
||||
Test that margin headers are included when margin is applied.
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Literal
|
||||
from urllib.parse import unquote
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -12,6 +13,7 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
get_fallback_errors_from_headers,
|
||||
get_hidden_params_dict,
|
||||
replace_complexity_router_headers,
|
||||
safe_header_value,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -169,15 +171,8 @@ def test_add_fallback_headers_serializes_fallback_errors():
|
|||
)
|
||||
|
||||
assert result is response
|
||||
assert response._hidden_params["additional_headers"][
|
||||
"x-litellm-attempted-fallbacks"
|
||||
] == 1
|
||||
assert (
|
||||
json.loads(
|
||||
response._hidden_params["additional_headers"]["x-litellm-fallback-errors"]
|
||||
)
|
||||
== fallback_errors
|
||||
)
|
||||
assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1
|
||||
assert json.loads(response._hidden_params["additional_headers"]["x-litellm-fallback-errors"]) == fallback_errors
|
||||
|
||||
|
||||
def test_add_retry_headers_to_streaming_wrapper():
|
||||
|
|
@ -203,9 +198,7 @@ def test_get_hidden_params_dict_with_pydantic_model_hidden_params():
|
|||
|
||||
class Response:
|
||||
def __init__(self):
|
||||
self._hidden_params = InnerHiddenParams(
|
||||
additional_headers={"x-custom": "value"}
|
||||
)
|
||||
self._hidden_params = InnerHiddenParams(additional_headers={"x-custom": "value"})
|
||||
|
||||
result = get_hidden_params_dict(Response())
|
||||
assert result == {"additional_headers": {"x-custom": "value"}}
|
||||
|
|
@ -252,9 +245,7 @@ def test_get_fallback_errors_from_headers_existing_list_passthrough():
|
|||
|
||||
|
||||
def test_get_fallback_errors_from_headers_invalid_json_returns_empty():
|
||||
result = get_fallback_errors_from_headers(
|
||||
{"x-litellm-fallback-errors": "not-valid-json-{"}
|
||||
)
|
||||
result = get_fallback_errors_from_headers({"x-litellm-fallback-errors": "not-valid-json-{"})
|
||||
assert result == []
|
||||
|
||||
|
||||
|
|
@ -283,3 +274,27 @@ def test_add_fallback_headers_to_dict_response():
|
|||
|
||||
assert result is response
|
||||
assert response["_hidden_params"]["additional_headers"]["x-litellm-attempted-fallbacks"] == 1
|
||||
|
||||
|
||||
def test_safe_header_value_passes_through_latin1_values():
|
||||
assert safe_header_value("azure/gpt-4o") == "azure/gpt-4o"
|
||||
|
||||
|
||||
def test_safe_header_value_percent_encodes_non_latin1_values():
|
||||
result = safe_header_value("azure/中文模型名稱")
|
||||
|
||||
assert result.encode("latin-1")
|
||||
assert unquote(result) == "azure/中文模型名稱"
|
||||
|
||||
|
||||
def test_safe_header_value_escapes_a_literal_percent_sign_too():
|
||||
"""
|
||||
A latin-1-safe value that happens to contain a literal "%" (e.g.
|
||||
"model%20name") must not pass through unchanged: it would then be
|
||||
indistinguishable from a genuinely percent-encoded value, and a caller
|
||||
that always unquotes would silently decode it into the wrong string.
|
||||
"""
|
||||
result = safe_header_value("model%20name")
|
||||
|
||||
assert "%" not in result.replace("%25", "")
|
||||
assert unquote(result) == "model%20name"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue