This commit is contained in:
Ankit Jha 2026-10-01 15:58:21 +00:00 • committed by GitHub
commit 189aedd056
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 82 additions and 0 deletions

View file

@ -2,6 +2,7 @@ import asyncio
import contextvars
from collections.abc import Iterator
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, ClassVar, Final
import httpx
@ -192,12 +193,28 @@ class GenerateContentHelper:
if litellm_logging_obj is None:
raise ValueError("litellm_logging_obj is required, but got None")
# Record the configured location so a `global` model is not priced as
# us-central1 (+10%), matching the completion path.
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
explicit_vertex_location: Final = VertexBase.explicit_vertex_ai_location(
MappingProxyType(
{key: getattr(litellm_params, key, None) for key in ("vertex_location", "vertex_ai_location")}
)
)
vertex_location_params: Final = (
MappingProxyType({"vertex_location": explicit_vertex_location})
if explicit_vertex_location
else MappingProxyType({})
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=dict(generate_content_config_dict),
litellm_params={
"litellm_call_id": litellm_call_id,
**vertex_location_params,
},
custom_llm_provider=custom_llm_provider,
)

View file

@ -259,3 +259,68 @@ async def test_native_fields_forwarded_on_async_stream():
body = mock_post.call_args.kwargs["json"]
assert body["safetySettings"] == safety_settings
assert "safetySettings" not in body.get("generationConfig", {})
def _generate_content_logging_obj(call_id: str):
from datetime import datetime, timezone
from litellm.litellm_core_utils.litellm_logging import Logging
return Logging(
model="gemini-2.0-flash",
messages=[],
stream=False,
call_type="generate_content",
start_time=datetime(2025, 1, 1, tzinfo=timezone.utc),
litellm_call_id=call_id,
function_id=call_id,
)
@pytest.mark.parametrize("configured_location", ["global", "us-east5"])
def test_vertex_location_recorded_for_cost_calculation(configured_location):
"""
Regression for https://github.com/BerriAI/litellm/issues/40692
generate_content must record vertex_location so `global` is not priced as us-central1.
"""
from litellm.google_genai.main import GenerateContentHelper
logging_obj = _generate_content_logging_obj("vertex-location-test")
GenerateContentHelper.setup_generate_content_call(
model="vertex_ai/gemini-2.0-flash",
contents=[{"role": "user", "parts": [{"text": "say ok"}]}],
custom_llm_provider="vertex_ai",
litellm_logging_obj=logging_obj,
litellm_call_id="vertex-location-test",
vertex_project="test-project",
vertex_location=configured_location,
)
recorded = logging_obj.model_call_details["litellm_params"]
assert recorded.get("vertex_location") == configured_location, (
f"the configured vertex_location must reach the cost calculator; got {recorded.get('vertex_location')!r}"
)
def test_vertex_location_absent_when_not_configured():
"""
Nothing configured means nothing recorded, so the cost calculator keeps its
own fallback rather than being handed an empty value here.
"""
from litellm.google_genai.main import GenerateContentHelper
logging_obj = _generate_content_logging_obj("vertex-location-unset")
GenerateContentHelper.setup_generate_content_call(
model="gemini/gemini-2.0-flash",
contents=[{"role": "user", "parts": [{"text": "say ok"}]}],
custom_llm_provider="gemini",
litellm_logging_obj=logging_obj,
litellm_call_id="vertex-location-unset",
api_key="test-key",
)
recorded = logging_obj.model_call_details["litellm_params"]
assert "vertex_location" not in recorded