Merge pull request #25972 from BerriAI/litellm_yj_apr16

[Infra] Merge dev branch
This commit is contained in:
yuneng-jiang 2026-04-17 14:58:52 -07:00 • committed by GitHub
commit 6a9f8f7772
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 128 additions and 21 deletions

View file

@ -3,8 +3,26 @@ from urllib.parse import parse_qs
import httpx
from litellm._logging import verbose_logger
from litellm.constants import PASS_THROUGH_HEADER_PREFIX
# Headers that must not be overwritten via the x-pass- forwarding mechanism.
# Includes standard credential/auth headers and protocol-level headers that
# affect routing or message framing.
_PASS_THROUGH_PROTECTED_HEADERS: frozenset = frozenset(
{
"authorization",
"api-key",
"x-api-key",
"x-goog-api-key",
"host",
"content-length",
}
)
# Header name prefix used to block AWS SigV4 signing headers from being overridden.
_PASS_THROUGH_PROTECTED_HEADER_PREFIXES: tuple = ("x-amz-",)
class BasePassthroughUtils:
@staticmethod
@ -56,11 +74,21 @@ class BasePassthroughUtils:
# Combine request headers with custom headers
headers = {**request_headers, **headers}
# Always process x-pass- prefixed headers (strip prefix and forward)
# Process x-pass- prefixed headers (strip prefix and forward)
# Credential and protocol-level headers are excluded from this mechanism.
for header_name, header_value in request_headers.items():
if header_name.lower().startswith(PASS_THROUGH_HEADER_PREFIX):
# Strip the 'x-pass-' prefix to get the actual header name
actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :]
# Strip the 'x-pass-' prefix and normalize to lowercase
actual_header_name = header_name[len(PASS_THROUGH_HEADER_PREFIX) :].lower()
if actual_header_name in _PASS_THROUGH_PROTECTED_HEADERS or any(
actual_header_name.startswith(p)
for p in _PASS_THROUGH_PROTECTED_HEADER_PREFIXES
):
verbose_logger.debug(
"x-pass- header %s maps to a protected header name; skipping",
header_name,
)
continue
headers[actual_header_name] = header_value
return headers

View file

@ -69,7 +69,8 @@ def check_complete_credentials(request_body: dict) -> bool:
# complex credentials - easier to make a malicious request
return False
if "api_key" in request_body:
api_key_value = request_body.get("api_key")
if api_key_value and isinstance(api_key_value, str) and api_key_value.strip():
return True
return False

View file

@ -19,7 +19,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
import litellm
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import is_request_body_safe
from litellm.proxy.auth.auth_utils import (
check_complete_credentials,
is_request_body_safe,
)
from litellm.proxy.litellm_pre_call_utils import (
_get_dynamic_logging_metadata,
add_litellm_data_to_request,
@ -33,7 +36,9 @@ def mock_request(monkeypatch):
mock_request = Mock(spec=Request)
mock_request.query_params = {} # Set mock query_params to an empty dictionary
mock_request.headers = {"traceparent": "test_traceparent"}
mock_request.state = State() # Real State so _safe_get_request_headers caching works
mock_request.state = (
State()
) # Real State so _safe_get_request_headers caching works
monkeypatch.setattr(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request", mock_request
)
@ -465,6 +470,21 @@ def test_is_request_body_safe_model_enabled(
assert expect_error == error_raised
@pytest.mark.parametrize(
"api_key_value, expect_complete",
[
("sk-real-key", True),
("", False),
(None, False),
(" ", False),
],
)
def test_check_complete_credentials_api_key_values(api_key_value, expect_complete):
request_body = {"model": "gpt-3.5-turbo", "api_key": api_key_value}
result = check_complete_credentials(request_body=request_body)
assert result == expect_complete
def test_reading_openai_org_id_from_headers():
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
@ -735,6 +755,7 @@ def test_get_docs_url(env_vars, expected_url):
result = _get_docs_url()
assert result == expected_url
@pytest.mark.parametrize(
"env_vars, expected_url",
[
@ -1516,7 +1537,7 @@ class MockPrismaClientDB:
mock_key_data,
):
self.db = MockDb(mock_team_data, mock_key_data)
async def get_data(
self,
token: Optional[Union[str, list]] = None,
@ -1534,7 +1555,7 @@ class MockPrismaClientDB:
):
"""Mock get_data method to return user info for admin"""
from litellm.proxy._types import LiteLLM_UserTable
# Return a proper LiteLLM_UserTable object when querying by user_id
if user_id:
return LiteLLM_UserTable(
@ -2072,7 +2093,7 @@ def test_team_alias_stale_bypass_disabled_by_default(monkeypatch):
monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False)
import litellm.proxy.litellm_pre_call_utils as pre_call_utils
from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists
# Reset module-level cache to ensure test isolation
pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
@ -2097,7 +2118,7 @@ def test_team_alias_stale_bypass_disabled_by_default(monkeypatch):
def test_team_alias_stale_bypass_enabled_by_flag(monkeypatch):
import litellm.proxy.litellm_pre_call_utils as pre_call_utils
from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists
# Reset module-level cache to ensure test isolation
pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None
@ -2394,16 +2415,17 @@ async def test_handle_logging_proxy_only_error_syncs_normalized_call_type(
captured_logging_obj["logging_obj"] = logging_obj
return logging_obj, data
with patch(
"litellm.proxy.utils.litellm.utils.function_setup",
side_effect=_capture_function_setup,
), patch.object(
Logging, "async_failure_handler", new=AsyncMock(return_value=None)
), patch.object(
Logging, "failure_handler", return_value=None
), patch(
"litellm.proxy.utils.threading.Thread"
) as mock_thread:
with (
patch(
"litellm.proxy.utils.litellm.utils.function_setup",
side_effect=_capture_function_setup,
),
patch.object(
Logging, "async_failure_handler", new=AsyncMock(return_value=None)
),
patch.object(Logging, "failure_handler", return_value=None),
patch("litellm.proxy.utils.threading.Thread") as mock_thread,
):
mock_thread.return_value.start = Mock()
await proxy_logging._handle_logging_proxy_only_error(
@ -2647,7 +2669,9 @@ async def test_handle_logging_proxy_only_error_skips_handlers_for_pass_through()
"model": "claude-3-5-sonnet",
}
with patch.object(logging_obj, "async_failure_handler", new_callable=AsyncMock) as mock_async:
with patch.object(
logging_obj, "async_failure_handler", new_callable=AsyncMock
) as mock_async:
with patch.object(logging_obj, "failure_handler") as mock_sync:
await proxy_logging._handle_logging_proxy_only_error(
request_data=request_data,

View file

@ -451,6 +451,60 @@ def test_forward_headers_from_request_x_pass_prefix():
assert "x-pass-custom-header" not in result
def test_forward_headers_from_request_protected_headers_not_overwritten():
"""
Test that x-pass- headers whose stripped names resolve to credential or
protocol-level header names are silently dropped and do not overwrite
values already present in the outbound headers dict.
"""
from litellm.passthrough.utils import BasePassthroughUtils
proxy_headers = {
"authorization": "Bearer proxy-upstream-key",
"api-key": "proxy-azure-key",
"x-api-key": "proxy-anthropic-key",
"x-goog-api-key": "proxy-google-key",
}
request_headers = {
"x-pass-authorization": "Bearer attacker-key",
"x-pass-api-key": "attacker-azure-key",
"x-pass-x-api-key": "attacker-anthropic-key",
"x-pass-x-goog-api-key": "attacker-google-key",
"x-pass-host": "evil.example.com",
"x-pass-content-length": "0",
"x-pass-x-amz-security-token": "attacker-aws-token",
# Legitimate x-pass- header that should still be forwarded
"x-pass-anthropic-beta": "context-1m-2025-08-07",
"content-type": "application/json",
}
result = BasePassthroughUtils.forward_headers_from_request(
request_headers=request_headers,
headers=proxy_headers.copy(),
forward_headers=False,
)
# Protected headers must retain the proxy-configured values
assert result["authorization"] == "Bearer proxy-upstream-key"
assert result["api-key"] == "proxy-azure-key"
assert result["x-api-key"] == "proxy-anthropic-key"
assert result["x-goog-api-key"] == "proxy-google-key"
# Protocol headers must not be injected
assert "host" not in result
assert "content-length" not in result
# AWS SigV4 headers must not be injected
assert "x-amz-security-token" not in result
# Legitimate non-protected x-pass- header still forwarded
assert result["anthropic-beta"] == "context-1m-2025-08-07"
# Header name must be normalized to lowercase in output
assert "Anthropic-Beta" not in result
@pytest.mark.asyncio
async def test_vertex_passthrough_custom_model_name_replaced_in_url():
"""