Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py: 26%
69 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1from datetime import datetime
2from typing import Final
4import httpx
6import litellm
7from litellm import stream_chunk_builder
8from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
9from litellm.litellm_core_utils.litellm_logging import (
10 get_standard_logging_object_payload,
11)
12from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
13from litellm.llms.base_llm.chat.transformation import BaseConfig
14from litellm.llms.cohere.chat.v2_transformation import CohereV2ChatConfig
15from litellm.llms.cohere.common_utils import (
16 ModelResponseIterator as CohereModelResponseIterator,
17)
18from litellm.llms.cohere.embed.v1_transformation import CohereEmbeddingConfig
19from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
20from litellm.types.passthrough_endpoints.pass_through_endpoints import (
21 PassthroughStandardLoggingPayload,
22)
23from litellm.types.utils import (
24 LlmProviders,
25 ModelResponse,
26 TextCompletionResponse,
27)
29from .base_passthrough_logging_handler import BasePassthroughLoggingHandler
32class CoherePassthroughLoggingHandler(BasePassthroughLoggingHandler):
33 @property
34 def llm_provider_name(self) -> LlmProviders:
35 return LlmProviders.COHERE
37 def get_provider_config(self, model: str) -> BaseConfig:
38 return CohereV2ChatConfig()
40 def _build_complete_streaming_response(
41 self,
42 all_chunks: list[str],
43 litellm_logging_obj: LiteLLMLoggingObj,
44 model: str,
45 ) -> ModelResponse | TextCompletionResponse | None:
46 cohere_model_response_iterator: Final = CohereModelResponseIterator(
47 streaming_response=None,
48 sync_stream=False,
49 )
50 litellm_custom_stream_wrapper: Final = CustomStreamWrapper(
51 completion_stream=cohere_model_response_iterator,
52 model=model,
53 logging_obj=litellm_logging_obj,
54 custom_llm_provider="cohere",
55 )
56 all_openai_chunks: Final = []
57 for _chunk_str in all_chunks:
58 try:
59 generic_chunk = cohere_model_response_iterator.convert_str_chunk_to_generic_chunk(chunk=_chunk_str)
60 litellm_chunk = litellm_custom_stream_wrapper.chunk_creator(chunk=generic_chunk)
61 if litellm_chunk is not None:
62 all_openai_chunks.append(litellm_chunk)
63 except (StopIteration, StopAsyncIteration):
64 break
65 complete_streaming_response: Final = stream_chunk_builder(chunks=all_openai_chunks)
66 return complete_streaming_response
68 def cohere_passthrough_handler(
69 self,
70 httpx_response: httpx.Response,
71 response_body: dict,
72 logging_obj: LiteLLMLoggingObj,
73 url_route: str,
74 result: str,
75 start_time: datetime,
76 end_time: datetime,
77 cache_hit: bool,
78 request_body: dict,
79 **kwargs,
80 ) -> PassThroughEndpointLoggingTypedDict:
81 """
82 Handle Cohere passthrough logging with route detection and cost tracking.
83 """
84 # Check if this is an embed endpoint
85 if "/v1/embed" in url_route and "/v1/embeddings" not in url_route:
86 model: Final = request_body.get("model", response_body.get("model", ""))
87 try:
88 cohere_embed_config: Final = CohereEmbeddingConfig()
89 litellm_model_response = litellm.EmbeddingResponse()
90 handler_instance: Final = CoherePassthroughLoggingHandler()
92 input_texts = request_body.get("texts", [])
93 if not input_texts:
94 input_texts = request_body.get("input", [])
96 # Transform the response
97 litellm_model_response = cohere_embed_config._transform_response(
98 response=httpx_response,
99 api_key="",
100 logging_obj=logging_obj,
101 data=request_body,
102 model_response=litellm_model_response,
103 model=model,
104 encoding=litellm.encoding,
105 input=input_texts,
106 )
108 # Calculate cost using LiteLLM's cost calculator
109 response_cost: Final = litellm.completion_cost(
110 completion_response=litellm_model_response,
111 model=model,
112 custom_llm_provider="cohere",
113 call_type="aembedding",
114 )
116 # Set the calculated cost in _hidden_params to prevent recalculation
117 if not hasattr(litellm_model_response, "_hidden_params"):
118 litellm_model_response._hidden_params = {}
119 litellm_model_response._hidden_params["response_cost"] = response_cost
121 kwargs["response_cost"] = response_cost
122 kwargs["model"] = model
123 kwargs["custom_llm_provider"] = "cohere"
125 # Extract user information for tracking
126 passthrough_logging_payload: Final[PassthroughStandardLoggingPayload | None] = kwargs.get(
127 "passthrough_logging_payload"
128 )
129 if passthrough_logging_payload:
130 user: Final = handler_instance._get_user_from_metadata(
131 passthrough_logging_payload=passthrough_logging_payload,
132 )
133 if user:
134 kwargs.setdefault("litellm_params", {})
135 kwargs["litellm_params"].update({"proxy_server_request": {"body": {"user": user}}})
137 # Create standard logging object
138 if litellm_model_response is not None:
139 get_standard_logging_object_payload(
140 kwargs=kwargs,
141 init_response_obj=litellm_model_response,
142 start_time=start_time,
143 end_time=end_time,
144 logging_obj=logging_obj,
145 status="success",
146 )
148 # Update logging object with cost information
149 logging_obj.model_call_details["model"] = model
150 logging_obj.model_call_details["custom_llm_provider"] = "cohere"
151 logging_obj.model_call_details["response_cost"] = response_cost
153 return {
154 "result": litellm_model_response,
155 "kwargs": kwargs,
156 }
157 except Exception:
158 # For other routes (e.g., /v2/chat), fall back to chat handler
159 return super().passthrough_chat_handler(
160 httpx_response=httpx_response,
161 response_body=response_body,
162 logging_obj=logging_obj,
163 url_route=url_route,
164 result=result,
165 start_time=start_time,
166 end_time=end_time,
167 cache_hit=cache_hit,
168 request_body=request_body,
169 **kwargs,
170 )
172 # For non-embed routes (e.g., /v2/chat), fall back to chat handler
173 return super().passthrough_chat_handler(
174 httpx_response=httpx_response,
175 response_body=response_body,
176 logging_obj=logging_obj,
177 url_route=url_route,
178 result=result,
179 start_time=start_time,
180 end_time=end_time,
181 cache_hit=cache_hit,
182 request_body=request_body,
183 **kwargs,
184 )