mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
aec6590486
commit
629404a100
3 changed files with 293 additions and 3 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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__])
|
||||
|
||||
Loading…
Add table
Reference in a new issue