Merge pull request #24964 from BerriAI/litellm_google_generate_content_response_headers

feat(proxy): LiteLLM headers on Google native generateContent routes
This commit is contained in:
Sameer Kankute 2026-04-02 18:34:36 +05:30 committed by GitHub
commit 34eff472a7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 374 additions and 69 deletions

View file

@ -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

View file

@ -150,6 +150,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,
@ -4497,9 +4521,9 @@ class BaseLLMHTTPHandler:
# Second: Execute agentic loop
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider[
"custom_llm_provider"
] = custom_llm_provider
kwargs_with_provider["custom_llm_provider"] = (
custom_llm_provider
)
agentic_response = await callback.async_run_agentic_loop(
tools=tool_calls,
model=model,
@ -4615,9 +4639,9 @@ class BaseLLMHTTPHandler:
# Second: Execute agentic loop
# Add custom_llm_provider to kwargs so the agentic loop can reconstruct the full model name
kwargs_with_provider = kwargs.copy() if kwargs else {}
kwargs_with_provider[
"custom_llm_provider"
] = custom_llm_provider
kwargs_with_provider["custom_llm_provider"] = (
custom_llm_provider
)
agentic_response = (
await callback.async_run_chat_completion_agentic_loop(
tools=tool_calls,
@ -5101,7 +5125,10 @@ class BaseLLMHTTPHandler:
_is_async: bool = False,
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]:
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
]:
"""
Handles image edit requests.
@ -5313,7 +5340,10 @@ class BaseLLMHTTPHandler:
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
) -> Union[ImageResponse, Coroutine[Any, Any, ImageResponse],]:
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
]:
"""
Handles image generation requests.
When _is_async=True, returns a coroutine instead of making the call directly.
@ -5553,7 +5583,10 @@ class BaseLLMHTTPHandler:
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
) -> Union[VideoObject, Coroutine[Any, Any, VideoObject],]:
) -> Union[
VideoObject,
Coroutine[Any, Any, VideoObject],
]:
"""
Handles video generation requests.
When _is_async=True, returns a coroutine instead of making the call directly.
@ -10033,6 +10066,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(
@ -10142,6 +10181,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(

View file

@ -566,6 +566,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,
@ -809,18 +870,18 @@ class ProxyBaseLLMRequestProcessing:
return
_payload_str = json.dumps(self.data, default=str)
if len(_payload_str) > MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG:
_key_count = len(self.data) if isinstance(self.data, dict) else 0
verbose_proxy_logger.debug(
"Request received by LiteLLM: payload too large to log (%d bytes, limit %d). Keys: %s",
"Request received by LiteLLM: payload too large to log (%d bytes, limit %d). Dict key count: %d; type: %s",
len(_payload_str),
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG,
list(self.data.keys())
if isinstance(self.data, dict)
else type(self.data).__name__,
_key_count,
type(self.data).__name__,
)
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
@ -1075,9 +1136,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 +1488,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 +1658,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 +1731,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 +1961,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,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

View file

@ -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
@ -81,7 +85,7 @@ async def test_async_anthropic_messages_handler_extra_headers():
extra_headers from kwargs with proper priority.
"""
handler = BaseLLMHTTPHandler()
# Mock the config
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
@ -90,7 +94,7 @@ async def test_async_anthropic_messages_handler_extra_headers():
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude-3-opus-20240229", "messages": []}
)
# Mock the client
mock_client = AsyncMock()
mock_response = Mock()
@ -104,13 +108,13 @@ async def test_async_anthropic_messages_handler_extra_headers():
"stop_reason": "end_turn",
}
mock_client.post = AsyncMock(return_value=mock_response)
# Mock logging object
mock_logging_obj = Mock()
mock_logging_obj.update_environment_variables = Mock()
mock_logging_obj.model_call_details = {}
mock_logging_obj.stream = False
# Test case 1: Only extra_headers in kwargs
kwargs = {
"extra_headers": {
@ -118,20 +122,21 @@ async def test_async_anthropic_messages_handler_extra_headers():
"X-Auth-Token": "token123",
}
}
with patch(
"litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers"
) as mock_provider_headers:
mock_provider_headers.return_value = None
# Capture what headers are passed to validate_anthropic_messages_environment
captured_headers = {}
def capture_validate(*args, **kwargs):
captured_headers.update(kwargs.get("headers", {}))
return ({"x-api-key": "test-key"}, "https://api.anthropic.com")
mock_config.validate_anthropic_messages_environment = capture_validate
try:
await handler.async_anthropic_messages_handler(
model="claude-3-opus-20240229",
@ -146,7 +151,7 @@ async def test_async_anthropic_messages_handler_extra_headers():
)
except Exception:
pass # We're testing header extraction, not the full flow
# Verify extra_headers were extracted and merged
assert "X-Custom-Header" in captured_headers
assert captured_headers["X-Custom-Header"] == "from-kwargs"
@ -219,9 +224,11 @@ async def test_async_anthropic_messages_handler_passes_litellm_metadata():
mock_logging_obj.update_from_kwargs.assert_called_once()
call_kwargs = mock_logging_obj.update_from_kwargs.call_args
kwargs_arg = call_kwargs.kwargs.get(
"kwargs", call_kwargs[1].get("kwargs", {})
) if call_kwargs.kwargs else call_kwargs[1].get("kwargs", {})
kwargs_arg = (
call_kwargs.kwargs.get("kwargs", call_kwargs[1].get("kwargs", {}))
if call_kwargs.kwargs
else call_kwargs[1].get("kwargs", {})
)
assert "litellm_metadata" in kwargs_arg
assert kwargs_arg["litellm_metadata"]["model_info"] == custom_model_info
@ -234,7 +241,7 @@ async def test_async_anthropic_messages_handler_header_priority():
forwarded < extra_headers < provider_specific
"""
handler = BaseLLMHTTPHandler()
# Mock the config
mock_config = Mock()
mock_client = AsyncMock()
@ -242,31 +249,32 @@ async def test_async_anthropic_messages_handler_header_priority():
mock_logging_obj.update_environment_variables = Mock()
mock_logging_obj.model_call_details = {}
mock_logging_obj.stream = False
# Test with all three header sources
kwargs = {
"headers": {"X-Priority": "forwarded", "X-Forwarded-Only": "keep"},
"extra_headers": {"X-Priority": "extra", "X-Extra-Only": "also-keep"},
}
with patch(
"litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers"
) as mock_provider_headers:
mock_provider_headers.return_value = {
"X-Priority": "provider",
"X-Provider-Only": "keep-this-too"
"X-Provider-Only": "keep-this-too",
}
captured_headers = {}
def capture_validate(*args, **kwargs):
captured_headers.update(kwargs.get("headers", {}))
return ({"x-api-key": "test-key"}, "https://api.anthropic.com")
mock_config.validate_anthropic_messages_environment = capture_validate
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude-3-opus-20240229", "messages": []}
)
try:
await handler.async_anthropic_messages_handler(
model="claude-3-opus-20240229",
@ -281,10 +289,36 @@ async def test_async_anthropic_messages_handler_header_priority():
)
except Exception:
pass
# Verify priority: provider_specific should win
assert captured_headers["X-Priority"] == "provider"
# Verify all unique headers from different sources are present
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"

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,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):
"""
@ -1009,13 +1146,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 +1644,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 +1871,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 +1889,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 +1901,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: