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 8cbda18ee3b..5b22b2746c9 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 @@ -110,7 +110,7 @@ class VertexPassthroughLoggingHandler: PassthroughCallTypes.passthrough_image_generation.value ) elif VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - json_response=_json_response, + json_response=_json_response, ): # Use multimodal embedding transformation vertex_multimodal_config = VertexAIMultimodalEmbeddingConfig() @@ -137,6 +137,15 @@ class VertexPassthroughLoggingHandler: logging_obj.model = model logging_obj.model_call_details["model"] = logging_obj.model + response_cost = litellm.completion_cost( + completion_response=litellm_prediction_response, + model=model, + custom_llm_provider="vertex_ai", + ) + + kwargs["response_cost"] = response_cost + kwargs["model"] = model + logging_obj.model_call_details["response_cost"] = response_cost return { "result": litellm_prediction_response, @@ -221,7 +230,9 @@ class VertexPassthroughLoggingHandler: - Logs in litellm callbacks """ kwargs: Dict[str, Any] = {} - model = model or VertexPassthroughLoggingHandler.extract_model_from_url(url_route) + model = model or VertexPassthroughLoggingHandler.extract_model_from_url( + url_route + ) complete_streaming_response = ( VertexPassthroughLoggingHandler._build_complete_streaming_response( all_chunks=all_chunks, @@ -340,13 +351,13 @@ class VertexPassthroughLoggingHandler: """ Detect if the response is from a multimodal embedding request. - Check if the response contains multimodal embedding fields: - - Docs: https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-embeddings-api#response-body - - + Check if the response contains multimodal embedding fields: + - Docs: https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-embeddings-api#response-body + + Args: json_response: The JSON response from Vertex AI - + Returns: bool: True if this is a multimodal embedding response """ @@ -358,10 +369,14 @@ class VertexPassthroughLoggingHandler: # Check for multimodal embedding response fields if any( key in prediction - for key in ["textEmbedding", "imageEmbedding", "videoEmbeddings"] + for key in [ + "textEmbedding", + "imageEmbedding", + "videoEmbeddings", + ] ): return True - + return False @staticmethod 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 702ae4bd42f..239f83b21ad 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,85 @@ class TestVertexAIPassThroughHandler: empty_response = {} assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(empty_response) is False + def test_vertex_passthrough_handler_predict_cost_tracking(self): + """ + Test that vertex_passthrough_handler correctly tracks costs for /predict endpoint + """ + import datetime + from unittest.mock import Mock, patch + + from litellm.litellm_core_utils.litellm_logging import ( + Logging as LiteLLMLoggingObj, + ) + from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import ( + VertexPassthroughLoggingHandler, + ) + + # Create mock embedding response data + embedding_response_data = { + "predictions": [ + { + "embeddings": { + "values": [0.1, 0.2, 0.3, 0.4, 0.5], + "statistics": { + "token_count": 10 + } + } + } + ] + } + + # Create mock httpx.Response + mock_httpx_response = Mock() + mock_httpx_response.json.return_value = embedding_response_data + mock_httpx_response.status_code = 200 + + # Create mock logging object + mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.litellm_call_id = "test-call-id-123" + mock_logging_obj.model_call_details = {} + + # Test URL with /predict endpoint + url_route = "/v1/projects/test-project/locations/us-central1/publishers/google/models/textembedding-gecko@001:predict" + + start_time = datetime.datetime.now() + end_time = datetime.datetime.now() + + with patch("litellm.completion_cost") as mock_completion_cost: + # Mock the completion cost calculation + mock_completion_cost.return_value = 0.0001 + + # Call the handler + 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 + ) + + # Verify cost tracking was implemented + assert result is not None + assert "result" in result + assert "kwargs" in result + + # Verify cost calculation was called + mock_completion_cost.assert_called_once() + + # Verify cost is set in kwargs + assert "response_cost" in result["kwargs"] + assert result["kwargs"]["response_cost"] == 0.0001 + + # Verify cost is set in logging object + assert "response_cost" in mock_logging_obj.model_call_details + assert mock_logging_obj.model_call_details["response_cost"] == 0.0001 + + # Verify model is set in kwargs + assert "model" in result["kwargs"] + assert result["kwargs"]["model"] == "textembedding-gecko@001" + class TestVertexAIDiscoveryPassThroughHandler: """