mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(proxy): LiteLLM headers on Google native generateContent routes (#25500)
* 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 * fix(proxy): prefill Google GenAI stream _hidden_params for proxy headers - Pass model_id, api_base, and process_response_headers output into streaming iterators so streamGenerateContent gets the same x-litellm-* headers as non-streaming paths. - Drop request_data deployment mutation from build_litellm_proxy_success_headers_from_llm_response. - Avoid logging raw request key names in oversized debug payload (code scanning). - Extend tests for streaming iterator shape, metadata fallback, and helper. Made-with: Cursor * Update litellm/proxy/common_request_processing.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * remove unused key count * Fix greptile review * Update litellm/proxy/common_request_processing.py Co-authored-by: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> --------- Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> Co-authored-by: Mateo Wang <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
9b3cd5ca25
commit
4cecfec9f9
6 changed files with 303 additions and 16 deletions
|
|
@ -1,6 +1,6 @@
|
|||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy.pass_through_endpoints.success_handler import (
|
||||
|
|
@ -29,12 +29,14 @@ class BaseGoogleGenAIGenerateContentStreamingIterator:
|
|||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
request_body: dict,
|
||||
model: str,
|
||||
hidden_params: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.litellm_logging_obj = litellm_logging_obj
|
||||
self.request_body = request_body
|
||||
self.start_time = datetime.now()
|
||||
self.collected_chunks: List[bytes] = []
|
||||
self.model = model
|
||||
self._hidden_params: Dict[str, Any] = hidden_params or {}
|
||||
|
||||
async def _handle_async_streaming_logging(
|
||||
self,
|
||||
|
|
@ -76,11 +78,13 @@ class GoogleGenAIGenerateContentStreamingIterator(
|
|||
litellm_metadata: dict,
|
||||
custom_llm_provider: str,
|
||||
request_body: Optional[dict] = None,
|
||||
hidden_params: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
super().__init__(
|
||||
litellm_logging_obj=logging_obj,
|
||||
request_body=request_body or {},
|
||||
model=model,
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
self.response = response
|
||||
self.model = model
|
||||
|
|
@ -130,11 +134,13 @@ class AsyncGoogleGenAIGenerateContentStreamingIterator(
|
|||
litellm_metadata: dict,
|
||||
custom_llm_provider: str,
|
||||
request_body: Optional[dict] = None,
|
||||
hidden_params: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
super().__init__(
|
||||
litellm_logging_obj=logging_obj,
|
||||
request_body=request_body or {},
|
||||
model=model,
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
self.response = response
|
||||
self.model = model
|
||||
|
|
|
|||
|
|
@ -155,6 +155,30 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
def _google_genai_streaming_hidden_params(
|
||||
*,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
response_headers: httpx.Headers,
|
||||
) -> Dict[str, Any]:
|
||||
"""Pre-stream metadata for proxy response headers (mirrors CustomStreamWrapper._hidden_params)."""
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
|
||||
_model_info: Dict[str, Any] = dict(
|
||||
getattr(litellm_params, "model_info", None) or {}
|
||||
)
|
||||
_raw_id = _model_info.get("id") or logging_obj.get_router_model_id() or ""
|
||||
_model_id = _raw_id if isinstance(_raw_id, str) else str(_raw_id)
|
||||
return {
|
||||
"model_id": _model_id,
|
||||
"api_base": api_base,
|
||||
"cache_key": "",
|
||||
"response_cost": "",
|
||||
"additional_headers": process_response_headers(response_headers),
|
||||
}
|
||||
|
||||
|
||||
class BaseLLMHTTPHandler:
|
||||
async def _make_common_async_call(
|
||||
self,
|
||||
|
|
@ -10425,6 +10449,12 @@ class BaseLLMHTTPHandler:
|
|||
litellm_metadata=litellm_metadata or {},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_body=data,
|
||||
hidden_params=_google_genai_streaming_hidden_params(
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
response_headers=response.headers,
|
||||
),
|
||||
)
|
||||
else:
|
||||
response = sync_httpx_client.post(
|
||||
|
|
@ -10534,6 +10564,12 @@ class BaseLLMHTTPHandler:
|
|||
litellm_metadata=litellm_metadata or {},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_body=data,
|
||||
hidden_params=_google_genai_streaming_hidden_params(
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
response_headers=response.headers,
|
||||
),
|
||||
)
|
||||
else:
|
||||
response = await async_httpx_client.post(
|
||||
|
|
|
|||
|
|
@ -619,6 +619,67 @@ 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,
|
||||
) -> 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 {}
|
||||
|
||||
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,
|
||||
|
|
@ -875,7 +936,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"Request received by LiteLLM:\n%s",
|
||||
json.dumps(self.data, indent=4, default=str),
|
||||
_payload_str,
|
||||
)
|
||||
|
||||
async def base_process_llm_request( # noqa: PLR0915
|
||||
|
|
@ -1511,9 +1572,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
|
||||
|
|
@ -1690,11 +1749,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)}"
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ async def google_generate_content(
|
|||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
version,
|
||||
)
|
||||
|
||||
|
|
@ -73,6 +74,16 @@ 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,
|
||||
)
|
||||
fastapi_response.headers.update(success_headers)
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -95,6 +106,7 @@ async def google_stream_generate_content(
|
|||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
version,
|
||||
)
|
||||
|
||||
|
|
@ -137,9 +149,24 @@ 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,
|
||||
)
|
||||
|
||||
# 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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2,12 +2,16 @@ import os
|
|||
import sys
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import (
|
||||
BaseLLMHTTPHandler,
|
||||
_google_genai_streaming_hidden_params,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
|
|
@ -320,3 +324,29 @@ async def test_async_anthropic_messages_handler_header_priority():
|
|||
assert captured_headers["X-Forwarded-Only"] == "keep"
|
||||
assert captured_headers["X-Extra-Only"] == "also-keep"
|
||||
assert captured_headers["X-Provider-Only"] == "keep-this-too"
|
||||
|
||||
|
||||
def test_google_genai_streaming_hidden_params_model_info_and_router_fallback():
|
||||
logging_obj = Mock()
|
||||
logging_obj.get_router_model_id = Mock(return_value="router-model-id")
|
||||
|
||||
from_model_info = _google_genai_streaming_hidden_params(
|
||||
api_base="https://generativelanguage.googleapis.com/v1beta",
|
||||
litellm_params=GenericLiteLLMParams(model_info={"id": "info-id"}),
|
||||
logging_obj=logging_obj,
|
||||
response_headers=httpx.Headers({"x-ratelimit-remaining": "10"}),
|
||||
)
|
||||
assert from_model_info["model_id"] == "info-id"
|
||||
assert (
|
||||
from_model_info["api_base"]
|
||||
== "https://generativelanguage.googleapis.com/v1beta"
|
||||
)
|
||||
assert isinstance(from_model_info["additional_headers"], dict)
|
||||
|
||||
from_router = _google_genai_streaming_hidden_params(
|
||||
api_base="https://x",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
logging_obj=logging_obj,
|
||||
response_headers=httpx.Headers({}),
|
||||
)
|
||||
assert from_router["model_id"] == "router-model-id"
|
||||
|
|
|
|||
|
|
@ -218,6 +218,141 @@ 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"}
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
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()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_litellm_proxy_success_headers_streaming_style_iterator(self):
|
||||
"""AsyncGoogleGenAIGenerateContentStreamingIterator sets _hidden_params at init; headers must propagate."""
|
||||
|
||||
class _FakeStreamLike:
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise StopAsyncIteration
|
||||
|
||||
_hidden_params = {
|
||||
"model_id": "stream-model-id",
|
||||
"api_base": "https://generativelanguage.googleapis.com/v1beta",
|
||||
"cache_key": "",
|
||||
"response_cost": "",
|
||||
"additional_headers": {"llm_provider-x": "y"},
|
||||
}
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.headers = {}
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.litellm_call_id = "cid-stream"
|
||||
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={})
|
||||
|
||||
headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
|
||||
response=_FakeStreamLike(),
|
||||
request_data={"model": "gemini/gemini-2.0-flash"},
|
||||
request=mock_request,
|
||||
user_api_key_dict=mock_user,
|
||||
logging_obj=logging_obj,
|
||||
version="1.0.0",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert headers["x-litellm-model-id"] == "stream-model-id"
|
||||
assert headers["x-litellm-model-api-base"] == (
|
||||
"https://generativelanguage.googleapis.com/v1beta"
|
||||
)
|
||||
assert headers["llm_provider-x"] == "y"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_litellm_proxy_success_headers_no_hidden_params_metadata_fallback(
|
||||
self,
|
||||
):
|
||||
"""When response has no _hidden_params, model_id can still come from litellm_metadata."""
|
||||
|
||||
class _BareResponse:
|
||||
pass
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.headers = {}
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.litellm_call_id = "cid-meta"
|
||||
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={})
|
||||
|
||||
headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
|
||||
response=_BareResponse(),
|
||||
request_data={
|
||||
"model": "gemini/gemini-1.5-flash",
|
||||
"litellm_metadata": {"model_info": {"id": "meta-model-id"}},
|
||||
},
|
||||
request=mock_request,
|
||||
user_api_key_dict=mock_user,
|
||||
logging_obj=logging_obj,
|
||||
version="1.0.0",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert headers["x-litellm-model-id"] == "meta-model-id"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_with_stream_timeout_header(self):
|
||||
"""
|
||||
|
|
@ -1158,13 +1293,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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue