Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions basedpyright-code-budget.json
Original file line number Diff line number Diff line change
Expand Up @@ -105,13 +105,13 @@
"limit": 113
},
"reportUnknownMemberType": {
"limit": 39773
"limit": 39772
},
"reportUnknownParameterType": {
"limit": 20207
},
"reportUnknownVariableType": {
"limit": 31281
"limit": 31280
},
"reportUnnecessaryCast": {
"limit": 122
Expand Down
40 changes: 29 additions & 11 deletions litellm/responses/main.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import asyncio
import contextvars
from collections.abc import Coroutine, Iterable
from collections.abc import Coroutine, Iterable, Mapping
from functools import partial
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast

Expand Down Expand Up @@ -639,6 +639,18 @@ def _pop_use_chat_completions_api_kw(kwargs: dict[str, Any]) -> bool:
return bool(use_cc)


def _merge_forwarded_client_headers(
extra_headers: Mapping[str, Any] | None,
kwargs: Mapping[str, object],
) -> dict[str, Any] | None:
"""Merge the proxy's forwarded client headers (`headers` kwarg) into `extra_headers`."""
client_headers: Final = kwargs.get("headers")
return ResponsesAPIRequestUtils.merge_client_forwarded_headers(
extra_headers=extra_headers,
client_headers=client_headers if isinstance(client_headers, dict) else None,
)


def _resolve_model_provider_for_responses(
model: str,
custom_llm_provider: str | None,
Expand Down Expand Up @@ -904,11 +916,7 @@ def responses(
_is_async: Final = kwargs.pop("aresponses", False) is True
use_chat_completions_api = _pop_use_chat_completions_api_kw(kwargs)

client_headers: Final = kwargs.get("headers")
extra_headers = ResponsesAPIRequestUtils.merge_client_forwarded_headers(
extra_headers=extra_headers,
client_headers=client_headers if isinstance(client_headers, dict) else None,
)
extra_headers = _merge_forwarded_client_headers(extra_headers, kwargs)
local_vars["extra_headers"] = extra_headers

# Convert text_format to text parameter if provided
Expand Down Expand Up @@ -1228,6 +1236,8 @@ def delete_responses(
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
_is_async: Final = kwargs.pop("adelete_responses", False) is True
merged_extra_headers: Final = _merge_forwarded_client_headers(extra_headers, kwargs)
local_vars["extra_headers"] = merged_extra_headers

# get llm provider logic
litellm_params: Final = GenericLiteLLMParams(**kwargs)
Expand Down Expand Up @@ -1275,7 +1285,7 @@ def delete_responses(
responses_api_provider_config=responses_api_provider_config,
litellm_params=litellm_params,
logging_obj=litellm_logging_obj,
extra_headers=extra_headers,
extra_headers=merged_extra_headers,
extra_body=extra_body,
timeout=timeout or request_timeout,
_is_async=_is_async,
Expand Down Expand Up @@ -1399,6 +1409,8 @@ def get_responses(
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
_is_async: Final = kwargs.pop("aget_responses", False) is True
merged_extra_headers: Final = _merge_forwarded_client_headers(extra_headers, kwargs)
local_vars["extra_headers"] = merged_extra_headers

# get llm provider logic
litellm_params: Final = GenericLiteLLMParams(**kwargs)
Expand Down Expand Up @@ -1446,7 +1458,7 @@ def get_responses(
responses_api_provider_config=responses_api_provider_config,
litellm_params=litellm_params,
logging_obj=litellm_logging_obj,
extra_headers=extra_headers,
extra_headers=merged_extra_headers,
extra_body=extra_body,
timeout=timeout or request_timeout,
_is_async=_is_async,
Expand Down Expand Up @@ -1548,6 +1560,8 @@ def list_input_items(
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
_is_async: Final = kwargs.pop("alist_input_items", False) is True
merged_extra_headers: Final = _merge_forwarded_client_headers(extra_headers, kwargs)
local_vars["extra_headers"] = merged_extra_headers

litellm_params: Final = GenericLiteLLMParams(**kwargs)

Expand Down Expand Up @@ -1589,7 +1603,7 @@ def list_input_items(
include=include,
limit=limit,
order=order,
extra_headers=extra_headers,
extra_headers=merged_extra_headers,
timeout=timeout or request_timeout,
_is_async=_is_async,
client=kwargs.get("client"),
Expand Down Expand Up @@ -1692,6 +1706,8 @@ def cancel_responses(
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
_is_async: Final = kwargs.pop("acancel_responses", False) is True
merged_extra_headers: Final = _merge_forwarded_client_headers(extra_headers, kwargs)
local_vars["extra_headers"] = merged_extra_headers

# get llm provider logic
litellm_params: Final = GenericLiteLLMParams(**kwargs)
Expand Down Expand Up @@ -1739,7 +1755,7 @@ def cancel_responses(
responses_api_provider_config=responses_api_provider_config,
litellm_params=litellm_params,
logging_obj=litellm_logging_obj,
extra_headers=extra_headers,
extra_headers=merged_extra_headers,
extra_body=extra_body,
timeout=timeout or request_timeout,
_is_async=_is_async,
Expand Down Expand Up @@ -1864,6 +1880,8 @@ def compact_responses(
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
_is_async: Final = kwargs.pop("acompact_responses", False) is True
merged_extra_headers: Final = _merge_forwarded_client_headers(extra_headers, kwargs)
local_vars["extra_headers"] = merged_extra_headers

# get llm provider logic
litellm_params: Final = GenericLiteLLMParams(**kwargs)
Expand Down Expand Up @@ -1932,7 +1950,7 @@ def compact_responses(
litellm_params=litellm_params,
logging_obj=litellm_logging_obj,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
extra_headers=merged_extra_headers,
extra_body=extra_body,
timeout=timeout or request_timeout,
_is_async=_is_async,
Expand Down
6 changes: 3 additions & 3 deletions litellm/responses/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,8 +93,8 @@ def merge_prompt_management_input(

@staticmethod
def merge_client_forwarded_headers(
extra_headers: dict[str, Any] | None,
client_headers: dict[str, str] | None,
extra_headers: Mapping[str, Any] | None,
client_headers: Mapping[str, str] | None,
) -> dict[str, Any] | None:
"""
Merge headers forwarded by the proxy (`headers` kwarg, set when
Expand All @@ -104,7 +104,7 @@ def merge_client_forwarded_headers(
Header names are compared case-insensitively, as HTTP defines them.
"""
if not client_headers:
return extra_headers
return dict(extra_headers) if extra_headers is not None else None
if not extra_headers:
return dict(client_headers)
explicit_names: Final = frozenset(name.lower() for name in extra_headers)
Expand Down
76 changes: 76 additions & 0 deletions tests/test_litellm/responses/test_responses_api_request_body.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,11 @@ def __init__(self, json_data, status_code=200):
def json(self):
return self._json_data

def raise_for_status(self) -> None:
if self.status_code >= 400:
request = httpx.Request("GET", "http://mock")
raise httpx.HTTPStatusError(self.text, request=request, response=httpx.Response(self.status_code))


def _assert_request_body_matches(request_body: dict, expected_body: dict) -> None:
for key, expected_value in expected_body.items():
Expand Down Expand Up @@ -367,3 +372,74 @@ async def test_aresponses_client_header_conflict_is_case_insensitive():

assert [name for name in request_headers if name.lower() == "x-shared"] == ["x-shared"]
assert request_headers["x-shared"] == "from-caller"


_BY_RESPONSE_ID = {"response_id": "resp_123", "custom_llm_provider": "openai"}

_MANAGEMENT_ROUTES = (
(litellm.aget_responses, "get", _minimal_responses_api_payload("resp_123", "gpt-4o"), _BY_RESPONSE_ID),
(litellm.adelete_responses, "delete", {"id": "resp_123", "object": "response", "deleted": True}, _BY_RESPONSE_ID),
(litellm.acancel_responses, "post", _minimal_responses_api_payload("resp_123", "gpt-4o"), _BY_RESPONSE_ID),
(litellm.alist_input_items, "get", {"object": "list", "data": [], "has_more": False}, _BY_RESPONSE_ID),
(
litellm.acompact_responses,
"post",
_minimal_responses_api_payload("resp_123", "gpt-4o"),
{"model": "openai/gpt-4o", "input": "hi"},
),
)


@pytest.mark.parametrize("route,http_method,payload,route_kwargs", _MANAGEMENT_ROUTES)
@pytest.mark.asyncio
async def test_responses_management_routes_forward_client_headers_to_provider(
route, http_method, payload, route_kwargs
):
"""
The response management routes must forward the proxy's `headers` kwarg
(`forward_client_headers_to_llm_api`) to the provider, like creating a response does.
"""
with patch(
f"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.{http_method}",
new_callable=AsyncMock,
) as mock_request:
mock_request.return_value = MockResponse(payload, 200)

await route(
api_key="fake-api-key",
headers={"x-my-new-header": "hello-from-client"},
**route_kwargs,
)

mock_request.assert_called_once()
request_headers = dict(mock_request.call_args.kwargs["headers"])

assert request_headers["x-my-new-header"] == "hello-from-client"


@pytest.mark.parametrize("route,http_method,payload,route_kwargs", _MANAGEMENT_ROUTES)
@pytest.mark.asyncio
async def test_responses_management_routes_extra_headers_win_over_client_headers(
route, http_method, payload, route_kwargs
):
"""
Explicit `extra_headers` beat the forwarded client headers on the management routes too,
case-insensitively.
"""
with patch(
f"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.{http_method}",
new_callable=AsyncMock,
) as mock_request:
mock_request.return_value = MockResponse(payload, 200)

await route(
api_key="fake-api-key",
headers={"X-Shared": "from-client"},
extra_headers={"x-shared": "from-caller"},
**route_kwargs,
)

request_headers = dict(mock_request.call_args.kwargs["headers"])

assert [name for name in request_headers if name.lower() == "x-shared"] == ["x-shared"]
assert request_headers["x-shared"] == "from-caller"
2 changes: 1 addition & 1 deletion type-discipline-budget.json
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
{
"LIT001": {
"limit": 23149
"limit": 23148
},
"LIT002": {
"limit": 27166
Expand Down
Loading