mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
d1df4e838b
commit
3ae14bd9ff
3 changed files with 186 additions and 36 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue