mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #21430 from BerriAI/litellm_perf_headers_caching
perf: use cached _safe_get_request_headers instead
This commit is contained in:
commit
8020276711
11 changed files with 111 additions and 22 deletions
|
|
@ -10,6 +10,7 @@ from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
|||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import merge_mcp_headers
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
|
@ -685,7 +686,7 @@ if MCP_AVAILABLE:
|
|||
return await _execute_with_mcp_client(
|
||||
new_mcp_server_request,
|
||||
_test_connection_operation,
|
||||
raw_headers=dict(request.headers),
|
||||
raw_headers=_safe_get_request_headers(request),
|
||||
)
|
||||
|
||||
@router.post("/test/tools/list")
|
||||
|
|
@ -744,5 +745,5 @@ if MCP_AVAILABLE:
|
|||
_list_tools_operation,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=dict(request.headers),
|
||||
raw_headers=_safe_get_request_headers(request),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -483,7 +483,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
parent_otel_span = (
|
||||
open_telemetry_logger.create_litellm_proxy_request_started_span(
|
||||
start_time=start_time,
|
||||
headers=dict(request.headers),
|
||||
headers=_safe_get_request_headers(request),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -562,7 +562,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
parent_otel_span=parent_otel_span,
|
||||
request_headers=dict(request.headers),
|
||||
request_headers=_safe_get_request_headers(request),
|
||||
)
|
||||
|
||||
is_proxy_admin = result["is_proxy_admin"]
|
||||
|
|
|
|||
|
|
@ -135,17 +135,29 @@ def _safe_set_request_parsed_body(
|
|||
|
||||
def _safe_get_request_headers(request: Optional[Request]) -> dict:
|
||||
"""
|
||||
[Non-Blocking] Safely get the request headers
|
||||
[Non-Blocking] Safely get the request headers.
|
||||
Caches the result on request.state to avoid re-creating dict(request.headers) per call.
|
||||
|
||||
Warning: Callers must NOT mutate the returned dict — it is shared across
|
||||
all callers within the same request via the cache.
|
||||
"""
|
||||
if request is None:
|
||||
return {}
|
||||
cached = getattr(request.state, "_cached_headers", None)
|
||||
if cached is not None:
|
||||
return cached
|
||||
try:
|
||||
if request is None:
|
||||
return {}
|
||||
return dict(request.headers)
|
||||
headers = dict(request.headers)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"Unexpected error reading request headers - {}".format(e)
|
||||
)
|
||||
return {}
|
||||
headers = {}
|
||||
try:
|
||||
request.state._cached_headers = headers
|
||||
except Exception:
|
||||
pass # request.state may not be available in all contexts
|
||||
return headers
|
||||
|
||||
|
||||
def check_file_size_under_limit(
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from fastapi_sso.sso.base import OpenID
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
|
||||
|
||||
class CustomSSOLoginHandler(CustomLogger):
|
||||
|
|
@ -18,7 +19,7 @@ class CustomSSOLoginHandler(CustomLogger):
|
|||
self,
|
||||
request: Request,
|
||||
) -> OpenID:
|
||||
request_headers_dict = dict(request.headers)
|
||||
request_headers_dict = _safe_get_request_headers(request)
|
||||
verbose_logger.debug("inside custom ui sso sign in hook...")
|
||||
return OpenID(
|
||||
id=request_headers_dict.get("x-litellm-user-id") or "123",
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors,
|
|||
LitellmDataForBackendLLMCall,
|
||||
LitellmUserRoles, SpecialHeaders,
|
||||
TeamCallbackMetadata, UserAPIKeyAuth)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
|
||||
# Cache special headers as a frozenset for O(1) lookup performance
|
||||
_SPECIAL_HEADERS_CACHE = frozenset(
|
||||
|
|
@ -824,7 +825,7 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
from litellm.proxy.proxy_server import llm_router, premium_user
|
||||
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
|
||||
|
||||
_raw_headers: Dict[str, str] = dict(request.headers)
|
||||
_raw_headers: Dict[str, str] = _safe_get_request_headers(request)
|
||||
_headers: Dict[str, str] = clean_headers(
|
||||
request.headers,
|
||||
litellm_key_header_name=(
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from typing_extensions import TypedDict
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -724,7 +725,7 @@ async def get_service_provider_config(request: Request):
|
|||
"SCIM ServiceProviderConfig request: method=%s url=%s headers=%s",
|
||||
request.method,
|
||||
request.url,
|
||||
dict(request.headers),
|
||||
_safe_get_request_headers(request),
|
||||
)
|
||||
meta = {
|
||||
"resourceType": "ServiceProviderConfig",
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from litellm.proxy.auth.route_checks import RouteChecks
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
_safe_get_request_headers,
|
||||
_safe_set_request_parsed_body,
|
||||
get_form_data,
|
||||
get_request_body,
|
||||
|
|
@ -60,7 +61,7 @@ def create_request_copy(request: Request):
|
|||
return {
|
||||
"method": request.method,
|
||||
"url": str(request.url),
|
||||
"headers": dict(request.headers),
|
||||
"headers": _safe_get_request_headers(request).copy(),
|
||||
"cookies": request.cookies,
|
||||
"query_params": dict(request.query_params),
|
||||
}
|
||||
|
|
@ -329,7 +330,7 @@ async def vllm_proxy_route(
|
|||
method=request.method,
|
||||
endpoint=endpoint,
|
||||
request_query_params=request.query_params,
|
||||
request_headers=dict(request.headers),
|
||||
request_headers=_safe_get_request_headers(request),
|
||||
stream=request_body.get("stream", False),
|
||||
content=None,
|
||||
data=None,
|
||||
|
|
@ -1307,7 +1308,7 @@ async def azure_proxy_route(
|
|||
method=request.method,
|
||||
endpoint=endpoint,
|
||||
request_query_params=request.query_params,
|
||||
request_headers=dict(request.headers),
|
||||
request_headers=_safe_get_request_headers(request),
|
||||
stream=request_body.get("stream", False),
|
||||
content=None,
|
||||
data=None,
|
||||
|
|
@ -1505,7 +1506,7 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict:
|
|||
Returns:
|
||||
dict: Headers dictionary with only allowed headers
|
||||
"""
|
||||
incoming_headers = dict(request.headers) or {}
|
||||
incoming_headers = _safe_get_request_headers(request)
|
||||
headers = {}
|
||||
for header_name in ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS:
|
||||
if header_name in incoming_headers:
|
||||
|
|
@ -1621,7 +1622,7 @@ async def _prepare_vertex_auth_headers(
|
|||
if (
|
||||
vertex_credentials is None or vertex_credentials.vertex_project is None
|
||||
) and router_credentials is None:
|
||||
headers = dict(request.headers) or {}
|
||||
headers = _safe_get_request_headers(request).copy()
|
||||
headers_passed_through = True
|
||||
verbose_proxy_logger.debug(
|
||||
"default_vertex_config not set, incoming request headers %s", headers
|
||||
|
|
|
|||
|
|
@ -50,7 +50,10 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
_safe_get_request_headers,
|
||||
)
|
||||
from litellm.proxy.utils import get_server_root_path
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
|
@ -649,7 +652,7 @@ async def pass_through_request( # noqa: PLR0915
|
|||
url = httpx.URL(target)
|
||||
headers = custom_headers
|
||||
headers = HttpPassThroughEndpointHelpers.forward_headers_from_request(
|
||||
request_headers=dict(request.headers),
|
||||
request_headers=_safe_get_request_headers(request).copy(),
|
||||
headers=headers,
|
||||
forward_headers=forward_headers,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -294,6 +294,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
|||
from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
_safe_get_request_headers,
|
||||
check_file_size_under_limit,
|
||||
get_form_data,
|
||||
)
|
||||
|
|
@ -10448,7 +10449,7 @@ async def async_queue_request(
|
|||
data["proxy_server_request"] = {
|
||||
"url": str(request.url),
|
||||
"method": request.method,
|
||||
"headers": dict(request.headers),
|
||||
"headers": _safe_get_request_headers(request).copy(),
|
||||
"body": copy.copy(data), # use copy instead of deepcopy
|
||||
}
|
||||
|
||||
|
|
@ -10469,7 +10470,7 @@ async def async_queue_request(
|
|||
data["metadata"] = {}
|
||||
data["metadata"]["user_api_key"] = user_api_key_dict.api_key
|
||||
data["metadata"]["user_api_key_metadata"] = user_api_key_dict.metadata
|
||||
_headers = dict(request.headers)
|
||||
_headers = _safe_get_request_headers(request).copy()
|
||||
_headers.pop(
|
||||
"authorization", None
|
||||
) # do not store the original `sk-..` api key in the db
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from fastapi import APIRouter, Request, Response
|
|||
import litellm
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
from litellm.proxy.litellm_pre_call_utils import _get_dynamic_logging_metadata
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
create_pass_through_route,
|
||||
|
|
@ -32,7 +33,7 @@ def create_request_copy(request: Request):
|
|||
return {
|
||||
"method": request.method,
|
||||
"url": str(request.url),
|
||||
"headers": dict(request.headers),
|
||||
"headers": _safe_get_request_headers(request).copy(),
|
||||
"cookies": request.cookies,
|
||||
"query_params": dict(request.query_params),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -761,3 +761,70 @@ async def test_request_body_with_html_script_tags():
|
|||
f"Message content with HTML was modified during parsing: "
|
||||
f"expected={msg['content']!r}, got={result['messages'][2]['content']!r}"
|
||||
)
|
||||
|
||||
|
||||
def test_safe_get_request_headers_caches_on_request_state():
|
||||
"""
|
||||
Test that _safe_get_request_headers caches the result on request.state
|
||||
and returns the same object on subsequent calls.
|
||||
"""
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {"content-type": "application/json", "authorization": "Bearer sk-123"}
|
||||
mock_request.state = MagicMock(spec=[]) # empty spec so getattr returns default
|
||||
|
||||
# First call — should create and cache
|
||||
result1 = _safe_get_request_headers(mock_request)
|
||||
assert result1 == {"content-type": "application/json", "authorization": "Bearer sk-123"}
|
||||
assert mock_request.state._cached_headers is result1
|
||||
|
||||
# Second call — should return the cached object (same identity)
|
||||
result2 = _safe_get_request_headers(mock_request)
|
||||
assert result2 is result1
|
||||
|
||||
|
||||
def test_safe_get_request_headers_none_request():
|
||||
"""
|
||||
Test that _safe_get_request_headers returns empty dict for None request.
|
||||
"""
|
||||
result = _safe_get_request_headers(None)
|
||||
assert result == {}
|
||||
|
||||
|
||||
def test_safe_get_request_headers_copy_protects_cache():
|
||||
"""
|
||||
Test that callers using .copy() before mutation do not corrupt the cache.
|
||||
"""
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {"authorization": "Bearer sk-123", "host": "localhost"}
|
||||
mock_request.state = MagicMock(spec=[])
|
||||
|
||||
original = _safe_get_request_headers(mock_request)
|
||||
|
||||
# Simulate what mutation call sites do: copy then pop
|
||||
mutable = _safe_get_request_headers(mock_request).copy()
|
||||
mutable.pop("authorization", None)
|
||||
|
||||
# Cache must be unaffected
|
||||
assert "authorization" in _safe_get_request_headers(mock_request)
|
||||
assert _safe_get_request_headers(mock_request) is original
|
||||
|
||||
|
||||
def test_safe_get_request_headers_state_unavailable():
|
||||
"""
|
||||
Test that _safe_get_request_headers still returns headers when
|
||||
request.state rejects attribute writes (the except path on the cache-write).
|
||||
"""
|
||||
class ReadOnlyState:
|
||||
"""State object that allows reads but raises on writes."""
|
||||
def __setattr__(self, name, value):
|
||||
raise AttributeError("read-only state")
|
||||
|
||||
def __getattr__(self, name):
|
||||
return None # _cached_headers not found → triggers fresh read
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
mock_request.state = ReadOnlyState()
|
||||
|
||||
result = _safe_get_request_headers(mock_request)
|
||||
assert result == {"content-type": "application/json"}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue