feat(proxy): return LiteLLM headers on Google native generateContent routes

Wire build_litellm_proxy_success_headers_from_llm_response for :generateContent
and :streamGenerateContent so x-litellm-*, rate limit, and provider headers
match the OpenAI-style proxy path. Add unit test.

Annotate httpx.HTTPStatusError branch so pyright accepts .response after optional
exception transform. Remove unused variable in streaming tracer test (Ruff F841).

Made-with: Cursor
This commit is contained in:
Sameer Kankute 2026-04-02 10:21:43 +05:30
parent d1df4e838b
commit 3ae14bd9ff
No known key found for this signature in database
3 changed files with 186 additions and 36 deletions

View file

@ -566,6 +566,71 @@ class ProxyBaseLLMRequestProcessing:
verbose_proxy_logger.error(f"Error setting custom headers: {e}")
return {}
@staticmethod
async def build_litellm_proxy_success_headers_from_llm_response(
*,
response: Any,
request_data: dict,
request: Request,
user_api_key_dict: UserAPIKeyAuth,
logging_obj: LiteLLMLoggingObj,
version: Optional[str],
proxy_logging_obj: ProxyLogging,
llm_router: Optional[Router] = None,
) -> dict[str, str]:
"""
Build LiteLLM proxy response headers for routes that call the LLM directly
(e.g. Google native :generateContent) instead of base_process_llm_request.
"""
if isinstance(response, dict):
hidden_params = response.get("_hidden_params") or {}
else:
hidden_params = getattr(response, "_hidden_params", None) or {}
if not isinstance(hidden_params, dict):
hidden_params = {}
model_id = ProxyBaseLLMRequestProcessing._get_model_id_from_response(
hidden_params, request_data
)
cache_key = hidden_params.get("cache_key", None) or ""
api_base = hidden_params.get("api_base", None) or ""
response_cost = hidden_params.get("response_cost", None) or ""
fastest_response_batch_completion = hidden_params.get(
"fastest_response_batch_completion", None
)
additional_headers = hidden_params.get("additional_headers", {}) or {}
if llm_router is not None:
request_data["deployment"] = llm_router.get_deployment(model_id=model_id)
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=logging_obj.litellm_call_id,
model_id=model_id,
cache_key=cache_key,
api_base=api_base,
version=version,
response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
fastest_response_batch_completion=fastest_response_batch_completion,
request_data=request_data,
hidden_params=hidden_params,
litellm_logging_obj=logging_obj,
**additional_headers,
)
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=request_data,
user_api_key_dict=user_api_key_dict,
response=response,
request_headers=dict(request.headers),
)
if callback_headers:
custom_headers.update(callback_headers)
return custom_headers
async def common_processing_pre_call_logic(
self,
request: Request,
@ -813,9 +878,11 @@ class ProxyBaseLLMRequestProcessing:
"Request received by LiteLLM: payload too large to log (%d bytes, limit %d). Keys: %s",
len(_payload_str),
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG,
list(self.data.keys())
if isinstance(self.data, dict)
else type(self.data).__name__,
(
list(self.data.keys())
if isinstance(self.data, dict)
else type(self.data).__name__
),
)
else:
verbose_proxy_logger.debug(
@ -1075,9 +1142,9 @@ class ProxyBaseLLMRequestProcessing:
# aliasing/routing, but the OpenAI-compatible response `model` field should reflect
# what the client sent.
if requested_model_from_client:
self.data[
"_litellm_client_requested_model"
] = requested_model_from_client
self.data["_litellm_client_requested_model"] = (
requested_model_from_client
)
# Streaming: attach a closure that fires after all guardrail
# end-of-stream blocks complete. CSW.__anext__ stores the
@ -1427,9 +1494,7 @@ class ProxyBaseLLMRequestProcessing:
_response = assembled_response
try:
from litellm.proxy.proxy_server import llm_router as _global_llm_router
from litellm.proxy.utils import (
_check_and_merge_model_level_guardrails,
)
from litellm.proxy.utils import _check_and_merge_model_level_guardrails
guardrail_data = _check_and_merge_model_level_guardrails(
data=captured_data, llm_router=_global_llm_router
@ -1599,11 +1664,12 @@ class ProxyBaseLLMRequestProcessing:
elif isinstance(e, httpx.HTTPStatusError):
# Handle httpx.HTTPStatusError - extract actual error from response
# This matches the original behavior before the refactor in commit 511d435f6f
error_body = await e.response.aread()
http_status_error: httpx.HTTPStatusError = e
error_body = await http_status_error.response.aread()
error_text = error_body.decode("utf-8")
raise HTTPException(
status_code=e.response.status_code,
status_code=http_status_error.response.status_code,
detail={"error": error_text},
)
error_msg = f"{str(e)}"
@ -1671,7 +1737,9 @@ class ProxyBaseLLMRequestProcessing:
verbose_proxy_logger.debug("inside generator")
try:
str_so_far = ""
async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook(
async for (
chunk
) in proxy_logging_obj.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=response,
request_data=request_data,
@ -1899,9 +1967,9 @@ class ProxyBaseLLMRequestProcessing:
# Add cache-related fields to **params (handled by Usage.__init__)
if cache_creation_input_tokens is not None:
usage_kwargs[
"cache_creation_input_tokens"
] = cache_creation_input_tokens
usage_kwargs["cache_creation_input_tokens"] = (
cache_creation_input_tokens
)
if cache_read_input_tokens is not None:
usage_kwargs["cache_read_input_tokens"] = cache_read_input_tokens

View file

@ -35,6 +35,7 @@ async def google_generate_content(
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
version,
)
@ -73,6 +74,17 @@ async def google_generate_content(
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
response = await llm_router.agenerate_content(**data)
success_headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
response=response,
request_data=data,
request=request,
user_api_key_dict=user_api_key_dict,
logging_obj=logging_obj,
version=version,
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
)
fastapi_response.headers.update(success_headers)
return response
@ -95,6 +107,7 @@ async def google_stream_generate_content(
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
version,
)
@ -137,9 +150,25 @@ async def google_stream_generate_content(
raise HTTPException(status_code=500, detail="Router not initialized")
response = await llm_router.agenerate_content_stream(**data)
success_headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
response=response,
request_data=data,
request=request,
user_api_key_dict=user_api_key_dict,
logging_obj=logging_obj,
version=version,
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
)
# Check if response is an async iterator (streaming response)
if response is not None and hasattr(response, "__aiter__"):
return StreamingResponse(content=response, media_type="text/event-stream")
return StreamingResponse(
content=response,
media_type="text/event-stream",
headers=success_headers,
)
fastapi_response.headers.update(success_headers)
return response

View file

@ -82,7 +82,9 @@ class TestProxyBaseLLMRequestProcessing:
pytest.fail("litellm_call_id is not a valid UUID")
assert data_passed["litellm_call_id"] == returned_data["litellm_call_id"]
def test_add_dd_apm_tags_for_litellm_call_id_uses_dd_tracing_helper(self, monkeypatch):
def test_add_dd_apm_tags_for_litellm_call_id_uses_dd_tracing_helper(
self, monkeypatch
):
mock_set_active_span_tag = MagicMock(return_value=True)
import litellm.proxy.dd_span_tagger
@ -216,6 +218,63 @@ class TestProxyBaseLLMRequestProcessing:
headers_with_invalid
)
@pytest.mark.asyncio
async def test_build_litellm_proxy_success_headers_from_llm_response(self):
"""
Google native :generateContent uses this helper instead of base_process_llm_request;
ensure x-litellm-* headers and callback hooks merge like the main proxy path.
"""
mock_request = MagicMock(spec=Request)
mock_request.headers = {}
class _FakeGenaiResponse:
_hidden_params = {
"model_id": "deployment-model-id",
"cache_key": "ck-test",
"api_base": "https://generativelanguage.googleapis.com/v1beta",
"response_cost": 0.001,
"additional_headers": {"llm_provider-ratelimit-requests": "1000"},
}
logging_obj = MagicMock()
logging_obj.litellm_call_id = "call-id-test"
mock_user = MagicMock()
mock_user.tpm_limit = None
mock_user.rpm_limit = None
mock_user.max_budget = None
mock_user.spend = 0.0
mock_user.allowed_model_region = None
proxy_logging_obj = MagicMock(spec=ProxyLogging)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(
return_value={"x-ratelimit-remaining-requests": "999"}
)
llm_router = MagicMock()
llm_router.get_deployment.return_value = {"litellm_params": {}}
headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
response=_FakeGenaiResponse(),
request_data={"model": "gemini/gemini-1.5-flash"},
request=mock_request,
user_api_key_dict=mock_user,
logging_obj=logging_obj,
version="9.9.9",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
)
assert headers["x-litellm-call-id"] == "call-id-test"
assert headers["x-litellm-model-id"] == "deployment-model-id"
assert headers["x-litellm-version"] == "9.9.9"
assert headers["llm_provider-ratelimit-requests"] == "1000"
assert headers["x-ratelimit-remaining-requests"] == "999"
proxy_logging_obj.post_call_response_headers_hook.assert_awaited_once()
llm_router.get_deployment.assert_called_once_with(
model_id="deployment-model-id"
)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_with_stream_timeout_header(self):
"""
@ -1009,13 +1068,6 @@ class TestCommonRequestProcessingHelpers:
assert mock_tracer.trace.call_count == 4
# Verify that each call was made with the correct operation name
expected_calls = [
(("streaming.chunk.yield",), {}),
(("streaming.chunk.yield",), {}),
(("streaming.chunk.yield",), {}),
(("streaming.chunk.yield",), {}),
]
actual_calls = mock_tracer.trace.call_args_list
assert len(actual_calls) == 4
@ -1514,7 +1566,10 @@ class TestIsAzureModelRouterRequest:
def test_detects_model_router_with_underscore(self):
assert _is_azure_model_router_request("azure_ai/model_router") is True
assert _is_azure_model_router_request("azure_ai/model_router/my-deployment") is True
assert (
_is_azure_model_router_request("azure_ai/model_router/my-deployment")
is True
)
def test_detects_model_router_with_hyphen(self):
assert _is_azure_model_router_request("azure_ai/model-router") is True
@ -1738,11 +1793,11 @@ class TestDDSpanTaggerTagRequest:
def test_tags_key_alias_and_model(self):
"""key_alias and requested_model are set on the span when present."""
user_key = self._make_user_api_key_dict(key_alias="my-prod-key", token="hashed123")
user_key = self._make_user_api_key_dict(
key_alias="my-prod-key", token="hashed123"
)
with patch(
"litellm.proxy.dd_span_tagger.set_active_span_tag"
) as mock_set_tag:
with patch("litellm.proxy.dd_span_tagger.set_active_span_tag") as mock_set_tag:
DDSpanTagger.tag_request(
user_api_key_dict=user_key,
requested_model="gpt-4o",
@ -1756,9 +1811,7 @@ class TestDDSpanTaggerTagRequest:
"""No key tags are set when key_alias and token are None (e.g. 401 path)."""
user_key = self._make_user_api_key_dict(key_alias=None, token=None)
with patch(
"litellm.proxy.dd_span_tagger.set_active_span_tag"
) as mock_set_tag:
with patch("litellm.proxy.dd_span_tagger.set_active_span_tag") as mock_set_tag:
DDSpanTagger.tag_request(
user_api_key_dict=user_key,
requested_model=None,
@ -1770,15 +1823,15 @@ class TestDDSpanTaggerTagRequest:
"""requested_model is tagged even when there's no key info."""
user_key = self._make_user_api_key_dict(key_alias=None, token=None)
with patch(
"litellm.proxy.dd_span_tagger.set_active_span_tag"
) as mock_set_tag:
with patch("litellm.proxy.dd_span_tagger.set_active_span_tag") as mock_set_tag:
DDSpanTagger.tag_request(
user_api_key_dict=user_key,
requested_model="claude-3-5-sonnet",
)
mock_set_tag.assert_called_once_with("litellm.requested_model", "claude-3-5-sonnet")
mock_set_tag.assert_called_once_with(
"litellm.requested_model", "claude-3-5-sonnet"
)
class TestHasAttributeErrorInChain: