diff --git a/litellm/__init__.py b/litellm/__init__.py index bda12e5ddd..207ea3ebd2 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -838,7 +838,7 @@ from .llms.databricks import DatabricksConfig, DatabricksEmbeddingConfig from .llms.predibase import PredibaseConfig from .llms.anthropic_text import AnthropicTextConfig from .llms.replicate import ReplicateConfig -from .llms.cohere import CohereConfig +from .llms.cohere.completion import CohereConfig from .llms.clarifai import ClarifaiConfig from .llms.ai21 import AI21Config from .llms.together_ai import TogetherAIConfig diff --git a/litellm/llms/cohere_chat.py b/litellm/llms/cohere/chat.py similarity index 99% rename from litellm/llms/cohere_chat.py rename to litellm/llms/cohere/chat.py index f13e74614b..b2569b4291 100644 --- a/litellm/llms/cohere_chat.py +++ b/litellm/llms/cohere/chat.py @@ -13,7 +13,7 @@ import litellm from litellm.types.llms.cohere import ToolResultObject from litellm.utils import Choices, Message, ModelResponse, Usage -from .prompt_templates.factory import cohere_message_pt, cohere_messages_pt_v2 +from ..prompt_templates.factory import cohere_message_pt, cohere_messages_pt_v2 class CohereError(Exception): diff --git a/litellm/llms/cohere.py b/litellm/llms/cohere/completion.py similarity index 66% rename from litellm/llms/cohere.py rename to litellm/llms/cohere/completion.py index 8bd1051e84..3e8bd4ded2 100644 --- a/litellm/llms/cohere.py +++ b/litellm/llms/cohere/completion.py @@ -1,6 +1,5 @@ -#################### OLD ######################## -##### See `cohere_chat.py` for `/chat` calls #### -################################################# +##### Calls /generate endpoint ####### + import json import os import time @@ -252,163 +251,3 @@ def completion( ) setattr(model_response, "usage", usage) return model_response - - -def _process_embedding_response( - embeddings: list, - model_response: litellm.EmbeddingResponse, - model: str, - encoding: Any, - input: list, -) -> litellm.EmbeddingResponse: - output_data = [] - for idx, embedding in enumerate(embeddings): - output_data.append( - {"object": "embedding", "index": idx, "embedding": embedding} - ) - model_response.object = "list" - model_response.data = output_data - model_response.model = model - input_tokens = 0 - for text in input: - input_tokens += len(encoding.encode(text)) - - setattr( - model_response, - "usage", - Usage( - prompt_tokens=input_tokens, completion_tokens=0, total_tokens=input_tokens - ), - ) - - return model_response - - -async def async_embedding( - model: str, - data: dict, - input: list, - model_response: litellm.utils.EmbeddingResponse, - timeout: Union[float, httpx.Timeout], - logging_obj: LiteLLMLoggingObj, - optional_params: dict, - api_base: str, - api_key: Optional[str], - headers: dict, - encoding: Callable, - client: Optional[AsyncHTTPHandler] = None, -): - - ## LOGGING - logging_obj.pre_call( - input=input, - api_key=api_key, - additional_args={ - "complete_input_dict": data, - "headers": headers, - "api_base": api_base, - }, - ) - ## COMPLETION CALL - if client is None: - client = AsyncHTTPHandler(concurrent_limit=1) - - response = await client.post(api_base, headers=headers, data=json.dumps(data)) - - ## LOGGING - logging_obj.post_call( - input=input, - api_key=api_key, - additional_args={"complete_input_dict": data}, - original_response=response, - ) - - embeddings = response.json()["embeddings"] - - ## PROCESS RESPONSE ## - return _process_embedding_response( - embeddings=embeddings, - model_response=model_response, - model=model, - encoding=encoding, - input=input, - ) - - -def embedding( - model: str, - input: list, - model_response: litellm.EmbeddingResponse, - logging_obj: LiteLLMLoggingObj, - optional_params: dict, - headers: dict, - encoding: Any, - api_key: Optional[str] = None, - aembedding: Optional[bool] = None, - timeout: Union[float, httpx.Timeout] = httpx.Timeout(None), - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, -): - headers = validate_environment(api_key, headers=headers) - embed_url = "https://api.cohere.ai/v1/embed" - model = model - data = {"model": model, "texts": input, **optional_params} - - if "3" in model and "input_type" not in data: - # cohere v3 embedding models require input_type, if no input_type is provided, default to "search_document" - data["input_type"] = "search_document" - - ## LOGGING - logging_obj.pre_call( - input=input, - api_key=api_key, - additional_args={"complete_input_dict": data}, - ) - - ## ROUTING - if aembedding is True: - return async_embedding( - model=model, - data=data, - input=input, - model_response=model_response, - timeout=timeout, - logging_obj=logging_obj, - optional_params=optional_params, - api_base=embed_url, - api_key=api_key, - headers=headers, - encoding=encoding, - ) - ## COMPLETION CALL - if client is None or not isinstance(client, HTTPHandler): - client = HTTPHandler(concurrent_limit=1) - response = client.post(embed_url, headers=headers, data=json.dumps(data)) - ## LOGGING - logging_obj.post_call( - input=input, - api_key=api_key, - additional_args={"complete_input_dict": data}, - original_response=response, - ) - """ - response - { - 'object': "list", - 'data': [ - - ] - 'model', - 'usage' - } - """ - if response.status_code != 200: - raise CohereError(message=response.text, status_code=response.status_code) - embeddings = response.json()["embeddings"] - - return _process_embedding_response( - embeddings=embeddings, - model_response=model_response, - model=model, - encoding=encoding, - input=input, - ) diff --git a/litellm/llms/cohere/embed.py b/litellm/llms/cohere/embed.py new file mode 100644 index 0000000000..81c84c4221 --- /dev/null +++ b/litellm/llms/cohere/embed.py @@ -0,0 +1,201 @@ +import json +import os +import time +import traceback +import types +from enum import Enum +from typing import Any, Callable, Optional, Union + +import httpx # type: ignore +import requests # type: ignore + +import litellm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.utils import Choices, Message, ModelResponse, Usage + + +def validate_environment(api_key, headers: dict): + headers.update( + { + "Request-Source": "unspecified:litellm", + "accept": "application/json", + "content-type": "application/json", + } + ) + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + return headers + + +class CohereError(Exception): + def __init__(self, status_code, message): + self.status_code = status_code + self.message = message + self.request = httpx.Request( + method="POST", url="https://api.cohere.ai/v1/generate" + ) + self.response = httpx.Response(status_code=status_code, request=self.request) + super().__init__( + self.message + ) # Call the base class constructor with the parameters it needs + + +def _process_embedding_response( + embeddings: list, + model_response: litellm.EmbeddingResponse, + model: str, + encoding: Any, + input: list, +) -> litellm.EmbeddingResponse: + output_data = [] + for idx, embedding in enumerate(embeddings): + output_data.append( + {"object": "embedding", "index": idx, "embedding": embedding} + ) + model_response.object = "list" + model_response.data = output_data + model_response.model = model + input_tokens = 0 + for text in input: + input_tokens += len(encoding.encode(text)) + + setattr( + model_response, + "usage", + Usage( + prompt_tokens=input_tokens, completion_tokens=0, total_tokens=input_tokens + ), + ) + + return model_response + + +async def async_embedding( + model: str, + data: dict, + input: list, + model_response: litellm.utils.EmbeddingResponse, + timeout: Union[float, httpx.Timeout], + logging_obj: LiteLLMLoggingObj, + optional_params: dict, + api_base: str, + api_key: Optional[str], + headers: dict, + encoding: Callable, + client: Optional[AsyncHTTPHandler] = None, +): + + ## LOGGING + logging_obj.pre_call( + input=input, + api_key=api_key, + additional_args={ + "complete_input_dict": data, + "headers": headers, + "api_base": api_base, + }, + ) + ## COMPLETION CALL + if client is None: + client = AsyncHTTPHandler(concurrent_limit=1) + + response = await client.post(api_base, headers=headers, data=json.dumps(data)) + + ## LOGGING + logging_obj.post_call( + input=input, + api_key=api_key, + additional_args={"complete_input_dict": data}, + original_response=response, + ) + + embeddings = response.json()["embeddings"] + + ## PROCESS RESPONSE ## + return _process_embedding_response( + embeddings=embeddings, + model_response=model_response, + model=model, + encoding=encoding, + input=input, + ) + + +def embedding( + model: str, + input: list, + model_response: litellm.EmbeddingResponse, + logging_obj: LiteLLMLoggingObj, + optional_params: dict, + headers: dict, + encoding: Any, + api_key: Optional[str] = None, + aembedding: Optional[bool] = None, + timeout: Union[float, httpx.Timeout] = httpx.Timeout(None), + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, +): + headers = validate_environment(api_key, headers=headers) + embed_url = "https://api.cohere.ai/v1/embed" + model = model + data = {"model": model, "texts": input, **optional_params} + + if "3" in model and "input_type" not in data: + # cohere v3 embedding models require input_type, if no input_type is provided, default to "search_document" + data["input_type"] = "search_document" + + ## LOGGING + logging_obj.pre_call( + input=input, + api_key=api_key, + additional_args={"complete_input_dict": data}, + ) + + ## ROUTING + if aembedding is True: + return async_embedding( + model=model, + data=data, + input=input, + model_response=model_response, + timeout=timeout, + logging_obj=logging_obj, + optional_params=optional_params, + api_base=embed_url, + api_key=api_key, + headers=headers, + encoding=encoding, + ) + ## COMPLETION CALL + if client is None or not isinstance(client, HTTPHandler): + client = HTTPHandler(concurrent_limit=1) + response = client.post(embed_url, headers=headers, data=json.dumps(data)) + ## LOGGING + logging_obj.post_call( + input=input, + api_key=api_key, + additional_args={"complete_input_dict": data}, + original_response=response, + ) + """ + response + { + 'object': "list", + 'data': [ + + ] + 'model', + 'usage' + } + """ + if response.status_code != 200: + raise CohereError(message=response.text, status_code=response.status_code) + embeddings = response.json()["embeddings"] + + return _process_embedding_response( + embeddings=embeddings, + model_response=model_response, + model=model, + encoding=encoding, + input=input, + ) diff --git a/litellm/main.py b/litellm/main.py index 2cf836890a..dd6a9e1b65 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -83,7 +83,6 @@ from .llms import ( clarifai, cloudflare, cohere, - cohere_chat, gemini, huggingface_restapi, maritalk, @@ -107,6 +106,9 @@ from .llms.anthropic_text import AnthropicTextCompletion from .llms.azure import AzureChatCompletion, _check_dynamic_azure_params from .llms.azure_text import AzureTextCompletion from .llms.bedrock_httpx import BedrockConverseLLM, BedrockLLM +from .llms.cohere import chat as cohere_chat +from .llms.cohere import completion as cohere_completion +from .llms.cohere import embed as cohere_embed from .llms.custom_llm import CustomLLM, custom_chat_llm_router from .llms.databricks import DatabricksChatCompletion from .llms.huggingface_restapi import Huggingface @@ -1645,7 +1647,7 @@ def completion( if extra_headers is not None: headers.update(extra_headers) - model_response = cohere.completion( + model_response = cohere_completion.completion( model=model, messages=messages, api_base=api_base, @@ -3457,7 +3459,7 @@ def embedding( headers = extra_headers else: headers = {} - response = cohere.embedding( + response = cohere_embed.embedding( model=model, input=input, optional_params=optional_params,