From c549cddada10190a3a593c6c92850fceaa10e009 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 16:44:52 -0700 Subject: [PATCH] fix(vertex_ai): price passthrough calls on the URL's serving location --- .../litellm_core_utils/llm_cost_calc/utils.py | 5 - litellm/llms/vertex_ai/cost_calculator.py | 2 - .../vertex_passthrough_logging_handler.py | 11 +++ .../test_llm_pass_through_endpoints.py | 92 +++++++++++++++++++ tests/test_litellm/test_cost_calculator.py | 26 ++++-- 5 files changed, 119 insertions(+), 17 deletions(-) diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index dec35d16ea0..0793fe20b21 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -999,9 +999,6 @@ def generic_cost_per_token( prompt_cost *= uplift completion_cost *= uplift - ## VERTEX REGIONAL-ENDPOINT UPLIFT - # Applied as a flat multiplier across all token costs for the request - # when the Vertex AI endpoint serving it is non-global. vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location) if vertex_uplift != 1.0: prompt_cost *= vertex_uplift @@ -1109,8 +1106,6 @@ def get_token_type_cost_breakdown( cache_read_cost *= uplift cache_creation_cost *= uplift - # Same flat uplift for Vertex AI non-global endpoints, keeping per-type - # costs reconciled with the totals for regional Vertex deployments. vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location) if vertex_uplift != 1.0: reasoning_cost *= vertex_uplift diff --git a/litellm/llms/vertex_ai/cost_calculator.py b/litellm/llms/vertex_ai/cost_calculator.py index a9f5d77350c..23cb1e5b580 100644 --- a/litellm/llms/vertex_ai/cost_calculator.py +++ b/litellm/llms/vertex_ai/cost_calculator.py @@ -164,8 +164,6 @@ def cost_per_character( usage=usage, ) - # Applied once here; the cost_per_token fallbacks above are called without - # vertex_location so the uplift can never compound. vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location) return prompt_cost * vertex_uplift, completion_cost * vertex_uplift diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 621b3ff9c83..f8e521410bf 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -10,6 +10,7 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.constants import VERTEX_BATCH_PREDICTION_JOBS_ROUTE from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.vertex_ai.common_utils import get_vertex_location_from_url from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( ModelResponseIterator as VertexModelResponseIterator, ) @@ -60,6 +61,7 @@ class VertexPassthroughLoggingHandler: request_body: dict | None = None, **kwargs, ) -> PassThroughEndpointLoggingTypedDict: + vertex_location: Final = get_vertex_location_from_url(url_route) if "predictLongRunning" in url_route: model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) @@ -82,6 +84,7 @@ class VertexPassthroughLoggingHandler: model=model, custom_llm_provider="vertex_ai", call_type="create_video", + vertex_location=vertex_location, ) # Set response_cost in _hidden_params to prevent recalculation @@ -123,6 +126,7 @@ class VertexPassthroughLoggingHandler: end_time=end_time, logging_obj=logging_obj, custom_llm_provider=VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route), + vertex_location=vertex_location, ) return { @@ -190,6 +194,7 @@ class VertexPassthroughLoggingHandler: end_time=end_time, logging_obj=logging_obj, custom_llm_provider="vertex_ai", + vertex_location=vertex_location, ) return { @@ -206,6 +211,7 @@ class VertexPassthroughLoggingHandler: model="vertex_ai/search_api", custom_llm_provider="vertex_ai", call_type="vector_store_search", + vertex_location=vertex_location, ) standard_pass_through_response_object: Final[StandardPassThroughResponseObject] = { @@ -302,6 +308,7 @@ class VertexPassthroughLoggingHandler: completion_response=litellm_prediction_response, model=model, custom_llm_provider="vertex_ai", + vertex_location=get_vertex_location_from_url(url_route), ) kwargs["response_cost"] = response_cost @@ -381,6 +388,7 @@ class VertexPassthroughLoggingHandler: completion_response=litellm_embedding_response, model=model, custom_llm_provider=custom_llm_provider, + vertex_location=get_vertex_location_from_url(url_route), ) kwargs["response_cost"] = response_cost @@ -438,6 +446,7 @@ class VertexPassthroughLoggingHandler: end_time=end_time, logging_obj=litellm_logging_obj, custom_llm_provider=VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route), + vertex_location=get_vertex_location_from_url(url_route), ) return { @@ -591,6 +600,7 @@ class VertexPassthroughLoggingHandler: end_time: datetime, logging_obj: LiteLLMLoggingObj, custom_llm_provider: str, + vertex_location: str | None, ) -> dict: """ Create the standard logging object for Vertex passthrough generateContent (streaming and non-streaming) @@ -601,6 +611,7 @@ class VertexPassthroughLoggingHandler: completion_response=litellm_model_response, model=model, custom_llm_provider="vertex_ai", + vertex_location=vertex_location, ) kwargs["response_cost"] = response_cost diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index b56a8da7c66..8da6d0fe81d 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -1109,6 +1109,98 @@ class TestVertexAIPassThroughHandler: assert result["kwargs"].get("model") == "gemini-embedding-2-preview" mock_completion_cost.assert_called_once() + @pytest.mark.parametrize("streaming", [False, True]) + def test_vertex_passthrough_handler_prices_regional_endpoint_with_uplift(self, monkeypatch, streaming): + """ + Passthrough cost is computed inside the handler and stored as response_cost before the + logging cost resolver runs, so the handler itself must read the serving location out of + the passthrough URL; otherwise regional Vertex passthrough traffic bills at the global + rate (#34393). + """ + import datetime + + from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import ( + VertexPassthroughLoggingHandler, + ) + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr( + litellm, + "model_cost", + { + **litellm.get_model_cost_map(url=""), + "vertex_ai/gemini-fake-regional": { + "litellm_provider": "vertex_ai", + "mode": "chat", + "input_cost_per_token": 1e-06, + "output_cost_per_token": 2e-06, + "regional_endpoint_uplift_multiplier": 1.1, + }, + }, + ) + + response_body: Final = { + "candidates": [ + { + "content": {"parts": [{"text": "hello"}], "role": "model"}, + "finishReason": "STOP", + } + ], + "usageMetadata": { + "promptTokenCount": 10, + "candidatesTokenCount": 20, + "totalTokenCount": 30, + }, + } + + def cost_for(location: str) -> float: + url_route: Final = ( + f"https://{location}-aiplatform.googleapis.com/v1/projects/p/locations/{location}" + "/publishers/google/models/gemini-fake-regional:" + f"{'streamGenerateContent' if streaming else 'generateContent'}" + ) + mock_logging_obj: Final = Mock() + mock_logging_obj.litellm_call_id = "call-id" + mock_logging_obj.model_call_details = {} + mock_logging_obj.optional_params = {} + start_time: Final = datetime.datetime.now() + end_time: Final = datetime.datetime.now() + if streaming: + result = VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks( + litellm_logging_obj=mock_logging_obj, + passthrough_success_handler_obj=Mock(), + url_route=url_route, + request_body={}, + endpoint_type="vertex_ai", + start_time=start_time, + all_chunks=[json.dumps(response_body)], + model=None, + end_time=end_time, + ) + else: + mock_httpx_response: Final = Mock() + mock_httpx_response.json.return_value = response_body + mock_httpx_response.headers = {} + mock_httpx_response.status_code = 200 + result = VertexPassthroughLoggingHandler.vertex_passthrough_handler( + httpx_response=mock_httpx_response, + logging_obj=mock_logging_obj, + url_route=url_route, + result="test-result", + start_time=start_time, + end_time=end_time, + cache_hit=False, + ) + return result["kwargs"]["response_cost"] + + global_cost: Final = cost_for("global") + regional_cost: Final = cost_for("us-east5") + + assert global_cost == pytest.approx(10 * 1e-06 + 20 * 2e-06, rel=1e-9) + assert regional_cost == pytest.approx(global_cost * 1.10, rel=1e-9), ( + "regional Vertex passthrough traffic must bill at 1.1x the global rate" + ) + class TestVertexAIDiscoveryPassThroughHandler: """ diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 6deaf5479e0..98938dee62e 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -1781,16 +1781,22 @@ def test_vertex_uplift_composes_with_above_128k_pricing(monkeypatch): including the above-128k dynamic rates, so a synthetic model carrying both keys prices regional above-128k usage at 1.1x the above-128k rate.""" monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") - monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) - litellm.model_cost["vertex_ai/fake-regional-128k-model"] = { - "litellm_provider": "vertex_ai", - "mode": "chat", - "input_cost_per_token": 1e-06, - "output_cost_per_token": 2e-06, - "input_cost_per_token_above_128k_tokens": 2e-06, - "output_cost_per_token_above_128k_tokens": 4e-06, - "regional_endpoint_uplift_multiplier": 1.1, - } + monkeypatch.setattr( + litellm, + "model_cost", + { + **litellm.get_model_cost_map(url=""), + "vertex_ai/fake-regional-128k-model": { + "litellm_provider": "vertex_ai", + "mode": "chat", + "input_cost_per_token": 1e-06, + "output_cost_per_token": 2e-06, + "input_cost_per_token_above_128k_tokens": 2e-06, + "output_cost_per_token_above_128k_tokens": 4e-06, + "regional_endpoint_uplift_multiplier": 1.1, + }, + }, + ) usage = Usage(prompt_tokens=200_000, completion_tokens=10, total_tokens=200_010) global_prompt, global_completion = cost_per_token(