(Feat) Add cost tracking for Vertex AI Passthrough /predict endpoint (#15019)

* Add cost tracking for passthrough for predict endpoint

* restore file
This commit is contained in:
Sameer Kankute 2025-09-30 01:49:04 +05:30 • committed by GitHub
parent 038863a1fe
commit 3ab1c31e4e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 103 additions and 9 deletions

View file

@ -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

View file

@ -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:
"""