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

1from datetime import datetime 

2from typing import Final 

3 

4import httpx 

5 

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) 

28 

29from .base_passthrough_logging_handler import BasePassthroughLoggingHandler 

30 

31 

32class CoherePassthroughLoggingHandler(BasePassthroughLoggingHandler): 

33 @property 

34 def llm_provider_name(self) -> LlmProviders: 

35 return LlmProviders.COHERE 

36 

37 def get_provider_config(self, model: str) -> BaseConfig: 

38 return CohereV2ChatConfig() 

39 

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 

67 

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() 

91 

92 input_texts = request_body.get("texts", []) 

93 if not input_texts: 

94 input_texts = request_body.get("input", []) 

95 

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 ) 

107 

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 ) 

115 

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 

120 

121 kwargs["response_cost"] = response_cost 

122 kwargs["model"] = model 

123 kwargs["custom_llm_provider"] = "cohere" 

124 

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}}}) 

136 

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 ) 

147 

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 

152 

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 ) 

171 

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 )