mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
(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:
parent
038863a1fe
commit
3ab1c31e4e
2 changed files with 103 additions and 9 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue