Add cost tracking for cohere embed passthrough endpoint (#17029)

* Add cost tracking for cohere embed passthrough endpoint

* update passthrough code

* update passthrough code

* fixed lint and mypy errors
This commit is contained in:
Sameer Kankute 2025-11-25 07:09:26 +05:30 • committed by GitHub
parent aec6590486
commit 629404a100
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 293 additions and 3 deletions

View file

@ -1,14 +1,30 @@
from datetime import datetime
from typing import List, Optional, Union
import httpx
import litellm
from litellm import stream_chunk_builder
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import (
get_standard_logging_object_payload,
)
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.cohere.chat.v2_transformation import CohereV2ChatConfig
from litellm.llms.cohere.common_utils import (
ModelResponseIterator as CohereModelResponseIterator,
)
from litellm.types.utils import LlmProviders, ModelResponse, TextCompletionResponse
from litellm.llms.cohere.embed.v1_transformation import CohereEmbeddingConfig
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
PassthroughStandardLoggingPayload,
)
from litellm.types.utils import (
LlmProviders,
ModelResponse,
TextCompletionResponse,
)
from .base_passthrough_logging_handler import BasePassthroughLoggingHandler
@ -54,3 +70,123 @@ class CoherePassthroughLoggingHandler(BasePassthroughLoggingHandler):
break
complete_streaming_response = stream_chunk_builder(chunks=all_openai_chunks)
return complete_streaming_response
def cohere_passthrough_handler( # noqa: PLR0915
self,
httpx_response: httpx.Response,
response_body: dict,
logging_obj: LiteLLMLoggingObj,
url_route: str,
result: str,
start_time: datetime,
end_time: datetime,
cache_hit: bool,
request_body: dict,
**kwargs,
) -> PassThroughEndpointLoggingTypedDict:
"""
Handle Cohere passthrough logging with route detection and cost tracking.
"""
# Check if this is an embed endpoint
if "/v1/embed" in url_route:
model = request_body.get("model", response_body.get("model", ""))
try:
cohere_embed_config = CohereEmbeddingConfig()
litellm_model_response = litellm.EmbeddingResponse()
handler_instance = CoherePassthroughLoggingHandler()
input_texts = request_body.get("texts", [])
if not input_texts:
input_texts = request_body.get("input", [])
# Transform the response
litellm_model_response = cohere_embed_config._transform_response(
response=httpx_response,
api_key="",
logging_obj=logging_obj,
data=request_body,
model_response=litellm_model_response,
model=model,
encoding=litellm.encoding,
input=input_texts,
)
# Calculate cost using LiteLLM's cost calculator
response_cost = litellm.completion_cost(
completion_response=litellm_model_response,
model=model,
custom_llm_provider="cohere",
call_type="aembedding",
)
# Set the calculated cost in _hidden_params to prevent recalculation
if not hasattr(litellm_model_response, "_hidden_params"):
litellm_model_response._hidden_params = {}
litellm_model_response._hidden_params["response_cost"] = response_cost
kwargs["response_cost"] = response_cost
kwargs["model"] = model
kwargs["custom_llm_provider"] = "cohere"
# Extract user information for tracking
passthrough_logging_payload: Optional[
PassthroughStandardLoggingPayload
] = kwargs.get("passthrough_logging_payload")
if passthrough_logging_payload:
user = handler_instance._get_user_from_metadata(
passthrough_logging_payload=passthrough_logging_payload,
)
if user:
kwargs.setdefault("litellm_params", {})
kwargs["litellm_params"].update(
{"proxy_server_request": {"body": {"user": user}}}
)
# Create standard logging object
if litellm_model_response is not None:
get_standard_logging_object_payload(
kwargs=kwargs,
init_response_obj=litellm_model_response,
start_time=start_time,
end_time=end_time,
logging_obj=logging_obj,
status="success",
)
# Update logging object with cost information
logging_obj.model_call_details["model"] = model
logging_obj.model_call_details["custom_llm_provider"] = "cohere"
logging_obj.model_call_details["response_cost"] = response_cost
return {
"result": litellm_model_response,
"kwargs": kwargs,
}
except Exception:
# For other routes (e.g., /v2/chat), fall back to chat handler
return super().passthrough_chat_handler(
httpx_response=httpx_response,
response_body=response_body,
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
request_body=request_body,
**kwargs,
)
# For non-embed routes (e.g., /v2/chat), fall back to chat handler
return super().passthrough_chat_handler(
httpx_response=httpx_response,
response_body=response_body,
logging_obj=logging_obj,
url_route=url_route,
result=result,
start_time=start_time,
end_time=end_time,
cache_hit=cache_hit,
request_body=request_body,
**kwargs,
)

View file

@ -48,7 +48,7 @@ class PassThroughEndpointLogging:
self.TRACKED_ANTHROPIC_ROUTES = ["/messages"]
# Cohere
self.TRACKED_COHERE_ROUTES = ["/v2/chat"]
self.TRACKED_COHERE_ROUTES = ["/v2/chat", "/v1/embed"]
self.assemblyai_passthrough_logging_handler = (
AssemblyAIPassthroughLoggingHandler()
)
@ -177,7 +177,7 @@ class PassThroughEndpointLogging:
kwargs = anthropic_passthrough_logging_handler_result["kwargs"]
elif self.is_cohere_route(url_route):
cohere_passthrough_logging_handler_result = (
cohere_passthrough_logging_handler.passthrough_chat_handler(
cohere_passthrough_logging_handler.cohere_passthrough_handler(
httpx_response=httpx_response,
response_body=response_body or {},
logging_obj=logging_obj,

View file

@ -0,0 +1,154 @@
import json
import os
import sys
from datetime import datetime
from unittest.mock import MagicMock, patch
import httpx
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.cohere_passthrough_logging_handler import (
CoherePassthroughLoggingHandler,
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
PassthroughStandardLoggingPayload,
)
class TestCoherePassthroughLoggingHandler:
"""Test the Cohere passthrough logging handler for embed cost tracking."""
def setup_method(self):
"""Set up test fixtures"""
self.start_time = datetime.now()
self.end_time = datetime.now()
self.handler = CoherePassthroughLoggingHandler()
# Mock Cohere embed response
self.mock_cohere_embed_response = {
"embeddings": [
[0.1, 0.2, 0.3, 0.4, 0.5],
[0.6, 0.7, 0.8, 0.9, 1.0],
],
"meta": {
"billed_units": {
"input_tokens": 3,
}
},
}
def _create_mock_logging_obj(self) -> LiteLLMLoggingObj:
"""Create a mock logging object"""
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {}
return mock_logging_obj
def _create_mock_httpx_response(self, response_data: dict = None) -> httpx.Response:
"""Create a mock httpx response"""
if response_data is None:
response_data = self.mock_cohere_embed_response
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.text = json.dumps(response_data)
mock_response.json.return_value = response_data
mock_response.headers = {"content-type": "application/json"}
return mock_response
def _create_passthrough_logging_payload(self) -> PassthroughStandardLoggingPayload:
"""Create a mock passthrough logging payload"""
return PassthroughStandardLoggingPayload(
url="https://api.cohere.com/v1/embed",
request_body={"model": "embed-english-v3.0", "texts": ["test passthrough"]},
request_method="POST",
)
@patch("litellm.completion_cost")
@patch(
"litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload"
)
@patch("litellm.llms.cohere.embed.v1_transformation.CohereEmbeddingConfig._transform_response")
def test_cohere_embed_passthrough_cost_tracking(
self, mock_transform_response, mock_get_standard_logging, mock_completion_cost
):
"""Test successful cost tracking for Cohere embed passthrough"""
# Arrange
from litellm.types.utils import EmbeddingResponse
# Create a mock embedding response
mock_embedding_response = EmbeddingResponse()
mock_embedding_response.data = [
{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]},
{"object": "embedding", "index": 1, "embedding": [0.4, 0.5, 0.6]},
]
mock_embedding_response.model = "embed-english-v3.0"
mock_embedding_response.object = "list"
from litellm.types.utils import Usage
mock_embedding_response.usage = Usage(
prompt_tokens=3, completion_tokens=0, total_tokens=3
)
mock_transform_response.return_value = mock_embedding_response
mock_completion_cost.return_value = 3.6e-07 # Expected cost for embed-v4.0
mock_get_standard_logging.return_value = {"test": "logging_payload"}
mock_httpx_response = self._create_mock_httpx_response()
mock_logging_obj = self._create_mock_logging_obj()
passthrough_payload = self._create_passthrough_logging_payload()
kwargs = {
"passthrough_logging_payload": passthrough_payload,
}
request_body = {
"model": "embed-english-v3.0",
"texts": ["test passthrough"],
}
# Act
result = self.handler.cohere_passthrough_handler(
httpx_response=mock_httpx_response,
response_body=self.mock_cohere_embed_response,
logging_obj=mock_logging_obj,
url_route="https://api.cohere.com/v1/embed",
result="",
start_time=self.start_time,
end_time=self.end_time,
cache_hit=False,
request_body=request_body,
**kwargs,
)
# Assert
assert result is not None
assert "result" in result
assert "kwargs" in result
assert result["kwargs"]["model"] == "embed-english-v3.0"
assert result["kwargs"]["custom_llm_provider"] == "cohere"
# Verify cost calculation was called with correct parameters
mock_completion_cost.assert_called_once()
call_args = mock_completion_cost.call_args
assert call_args.kwargs["model"] == "embed-english-v3.0"
assert call_args.kwargs["custom_llm_provider"] == "cohere"
assert call_args.kwargs["call_type"] == "aembedding"
# Verify logging object was updated
assert mock_logging_obj.model_call_details["response_cost"] == 3.6e-07
assert mock_logging_obj.model_call_details["model"] == "embed-english-v3.0"
assert mock_logging_obj.model_call_details["custom_llm_provider"] == "cohere"
# Verify result is an EmbeddingResponse
assert hasattr(result["result"], "data")
assert hasattr(result["result"], "model")
assert result["result"].model == "embed-english-v3.0"
if __name__ == "__main__":
pytest.main([__file__])