From 1140366beeadf659ed4214e7004af63490e6e4dd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 17:24:44 -0700 Subject: [PATCH] fix(vertex_ai): resolve passthrough serving location in the logging cost recompute --- .../vertex_passthrough_logging_handler.py | 7 ++- .../test_llm_pass_through_endpoints.py | 56 +++++++++++++------ 2 files changed, 46 insertions(+), 17 deletions(-) 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 f8e521410bf..ddcca1d372b 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 @@ -62,6 +62,8 @@ class VertexPassthroughLoggingHandler: **kwargs, ) -> PassThroughEndpointLoggingTypedDict: vertex_location: Final = get_vertex_location_from_url(url_route) + if vertex_location is not None: + logging_obj.optional_params["vertex_location"] = vertex_location if "predictLongRunning" in url_route: model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route) @@ -421,6 +423,9 @@ class VertexPassthroughLoggingHandler: - Logs in litellm callbacks """ kwargs: dict[str, Any] = {} + vertex_location: Final = get_vertex_location_from_url(url_route) + if vertex_location is not None: + litellm_logging_obj.optional_params["vertex_location"] = vertex_location model = model or VertexPassthroughLoggingHandler.extract_model_from_url(url_route) complete_streaming_response: Final = VertexPassthroughLoggingHandler._build_complete_streaming_response( all_chunks=all_chunks, @@ -446,7 +451,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), + vertex_location=vertex_location, ) return { 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 8da6d0fe81d..a9454854948 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 @@ -719,6 +719,7 @@ class TestVertexAIPassThroughHandler: # Create mock logging object mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.optional_params = {} mock_logging_obj.litellm_call_id = "test-call-id-123" mock_logging_obj.model_call_details = {} @@ -895,6 +896,7 @@ class TestVertexAIPassThroughHandler: # Create mock logging object mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.optional_params = {} mock_logging_obj.litellm_call_id = "test-call-id-123" mock_logging_obj.model_call_details = {} @@ -965,6 +967,7 @@ class TestVertexAIPassThroughHandler: mock_httpx_response.status_code = 200 mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.optional_params = {} mock_logging_obj.litellm_call_id = "test-call-id-embed" mock_logging_obj.model_call_details = {} @@ -1023,6 +1026,7 @@ class TestVertexAIPassThroughHandler: mock_httpx_response.status_code = 200 mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.optional_params = {} mock_logging_obj.litellm_call_id = "test-call-id-batch" mock_logging_obj.model_call_details = {} @@ -1079,6 +1083,7 @@ class TestVertexAIPassThroughHandler: mock_httpx_response.status_code = 200 mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.optional_params = {} mock_logging_obj.litellm_call_id = "test-call-id-gemini-studio" mock_logging_obj.model_call_details = {} @@ -1112,13 +1117,14 @@ class TestVertexAIPassThroughHandler: @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). + Both cost computations for a passthrough call must price on the URL's serving location: + the handler-computed cost, and the async success recompute, which re-resolves the + location from the logging object and previously fell through empty optional_params to + the us-central1 default, billing the regional uplift on global traffic too (#34393). """ import datetime + from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import ( VertexPassthroughLoggingHandler, ) @@ -1153,21 +1159,33 @@ class TestVertexAIPassThroughHandler: }, } - def cost_for(location: str) -> float: + def costs_for(location: str) -> tuple[float, 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() + logging_obj: Final = Logging( + model="gemini-fake-regional", + messages=[{"role": "user", "content": "hi"}], + stream=streaming, + call_type="pass_through_endpoint", + start_time=start_time, + litellm_call_id="call-id", + function_id="fn-id", + ) + logging_obj.update_environment_variables( + model="gemini-fake-regional", + user="unknown", + optional_params={}, + litellm_params={}, + call_type="pass_through_endpoint", + ) if streaming: result = VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks( - litellm_logging_obj=mock_logging_obj, + litellm_logging_obj=logging_obj, passthrough_success_handler_obj=Mock(), url_route=url_route, request_body={}, @@ -1184,22 +1202,28 @@ class TestVertexAIPassThroughHandler: mock_httpx_response.status_code = 200 result = VertexPassthroughLoggingHandler.vertex_passthrough_handler( httpx_response=mock_httpx_response, - logging_obj=mock_logging_obj, + logging_obj=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"] + recomputed: Final = logging_obj._response_cost_calculator(result=result["result"]) + return result["kwargs"]["response_cost"], recomputed - global_cost: Final = cost_for("global") - regional_cost: Final = cost_for("us-east5") + global_handler_cost, global_recomputed_cost = costs_for("global") + regional_handler_cost, regional_recomputed_cost = costs_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), ( + plain_cost: Final = 10 * 1e-06 + 20 * 2e-06 + assert global_handler_cost == pytest.approx(plain_cost, rel=1e-9) + assert regional_handler_cost == pytest.approx(plain_cost * 1.10, rel=1e-9), ( "regional Vertex passthrough traffic must bill at 1.1x the global rate" ) + assert global_recomputed_cost == pytest.approx(plain_cost, rel=1e-9), ( + "the logging recompute must not price global passthrough traffic as regional" + ) + assert regional_recomputed_cost == pytest.approx(plain_cost * 1.10, rel=1e-9) class TestVertexAIDiscoveryPassThroughHandler: