From 5bbf906c830376e0c9ad3dcdc7ace7b452383026 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Mon, 9 Dec 2024 15:58:25 -0800 Subject: [PATCH] Litellm code qa common config (#7113) * feat(base_llm): initial commit for common base config class Addresses code qa critique https://github.com/andrewyng/aisuite/issues/113#issuecomment-2512369132 * feat(base_llm/): add transform request/response abstract methods to base config class * feat(cohere-+-clarifai): refactor integrations to use common base config class * fix: fix linting errors * refactor(anthropic/): move anthropic + vertex anthropic to use base config * test: fix xai test * test: fix tests * fix: fix linting errors * test: comment out WIP test * fix(transformation.py): fix is pdf used check * fix: fix linting error --- litellm/__init__.py | 10 +- litellm/constants.py | 64 ++ litellm/integrations/mlflow.py | 2 +- .../get_supported_openai_params.py | 60 +- .../OpenAI/chat/gpt_audio_transformation.py | 16 +- .../llms/OpenAI/chat/gpt_transformation.py | 89 ++- litellm/llms/OpenAI/chat/o1_transformation.py | 16 +- litellm/llms/OpenAI/common_utils.py | 36 +- litellm/llms/OpenAI/completion/handler.py | 314 +++++++++ .../llms/OpenAI/completion/transformation.py | 178 +++++ litellm/llms/OpenAI/openai.py | 482 +------------- litellm/llms/anthropic/chat/handler.py | 64 +- litellm/llms/anthropic/chat/transformation.py | 181 ++++-- litellm/llms/anthropic/common_utils.py | 15 +- litellm/llms/azure_text.py | 3 +- litellm/llms/base_llm/transformation.py | 129 ++++ litellm/llms/clarifai.py | 378 ----------- litellm/llms/clarifai/chat/handler.py | 177 +++++ litellm/llms/clarifai/chat/transformation.py | 201 ++++++ litellm/llms/clarifai/common_utils.py | 8 + .../llms/cohere/{chat.py => chat/handler.py} | 127 +--- litellm/llms/cohere/chat/transformation.py | 184 ++++++ litellm/llms/cohere/common_utils.py | 19 + .../cohere/{ => completion}/completion.py | 100 +-- .../llms/cohere/completion/transformation.py | 183 ++++++ litellm/llms/databricks/chat/old_handler.py | 611 ------------------ litellm/llms/groq/chat/handler.py | 2 +- litellm/llms/groq/chat/transformation.py | 16 +- litellm/llms/openai_like/chat/handler.py | 2 +- .../llms/together_ai/completion/handler.py | 2 +- .../together_ai/completion/transformation.py | 2 +- .../anthropic/transformation.py | 20 + .../vertex_ai_partner_models/main.py | 2 + ...ai_transformation.py => transformation.py} | 2 - litellm/main.py | 13 +- .../anthropic_passthrough_logging_handler.py | 28 +- litellm/utils.py | 95 +-- .../test_max_completion_tokens.py | 6 +- tests/llm_translation/test_xai.py | 4 +- .../test_amazing_vertex_completion.py | 2 +- tests/local_testing/test_config.py | 32 + 41 files changed, 1877 insertions(+), 1998 deletions(-) create mode 100644 litellm/llms/OpenAI/completion/handler.py create mode 100644 litellm/llms/OpenAI/completion/transformation.py create mode 100644 litellm/llms/base_llm/transformation.py delete mode 100644 litellm/llms/clarifai.py create mode 100644 litellm/llms/clarifai/chat/handler.py create mode 100644 litellm/llms/clarifai/chat/transformation.py create mode 100644 litellm/llms/clarifai/common_utils.py rename litellm/llms/cohere/{chat.py => chat/handler.py} (74%) create mode 100644 litellm/llms/cohere/chat/transformation.py create mode 100644 litellm/llms/cohere/common_utils.py rename litellm/llms/cohere/{ => completion}/completion.py (52%) create mode 100644 litellm/llms/cohere/completion/transformation.py delete mode 100644 litellm/llms/databricks/chat/old_handler.py rename litellm/llms/xai/chat/{xai_transformation.py => transformation.py} (97%) diff --git a/litellm/__init__.py b/litellm/__init__.py index 4b872b01430..60b13e45a4d 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -22,6 +22,7 @@ from litellm.constants import ( DEFAULT_FLUSH_INTERVAL_SECONDS, ROUTER_MAX_FALLBACKS, DEFAULT_MAX_RETRIES, + LITELLM_CHAT_PROVIDERS, ) from litellm.types.guardrails import GuardrailItem from litellm.proxy._types import ( @@ -1064,8 +1065,8 @@ from .llms.databricks.chat.transformation import DatabricksConfig from .llms.databricks.embed.transformation import DatabricksEmbeddingConfig from .llms.predibase import PredibaseConfig from .llms.replicate import ReplicateConfig -from .llms.cohere.completion import CohereConfig -from .llms.clarifai import ClarifaiConfig +from .llms.cohere.completion.transformation import CohereTextConfig as CohereConfig +from .llms.clarifai.chat.transformation import ClarifaiConfig from .llms.ai21.completion import AI21Config from .llms.ai21.chat import AI21ChatConfig from .llms.together_ai.chat import TogetherAIConfig @@ -1125,13 +1126,14 @@ from .llms.bedrock.embed.amazon_titan_multimodal_transformation import ( from .llms.bedrock.embed.amazon_titan_v2_transformation import ( AmazonTitanV2Config, ) +from .llms.cohere.chat.transformation import CohereChatConfig from .llms.bedrock.embed.cohere_transformation import BedrockCohereEmbeddingConfig from .llms.OpenAI.openai import ( OpenAIConfig, - OpenAITextCompletionConfig, MistralEmbeddingConfig, DeepInfraConfig, ) +from litellm.llms.OpenAI.completion.transformation import OpenAITextCompletionConfig from .llms.groq.chat.transformation import GroqChatConfig from .llms.azure_ai.chat.transformation import AzureAIStudioConfig from .llms.mistral.mistral_chat_transformation import MistralConfig @@ -1165,7 +1167,7 @@ from .llms.fireworks_ai.embed.fireworks_ai_transformation import ( FireworksAIEmbeddingConfig, ) from .llms.jina_ai.embedding.transformation import JinaAIEmbeddingConfig -from .llms.xai.chat.xai_transformation import XAIChatConfig +from .llms.xai.chat.transformation import XAIChatConfig from .llms.volcengine import VolcEngineConfig from .llms.text_completion_codestral import MistralTextCompletionConfig from .llms.azure.azure import ( diff --git a/litellm/constants.py b/litellm/constants.py index 97dc6c7348b..c0aa2a36907 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2,3 +2,67 @@ ROUTER_MAX_FALLBACKS = 5 DEFAULT_BATCH_SIZE = 512 DEFAULT_FLUSH_INTERVAL_SECONDS = 5 DEFAULT_MAX_RETRIES = 2 +LITELLM_CHAT_PROVIDERS = [ + "openai", + "openai_like", + "xai", + "custom_openai", + "text-completion-openai", + "cohere", + "cohere_chat", + "clarifai", + "anthropic", + "replicate", + "huggingface", + "together_ai", + "openrouter", + "vertex_ai", + "vertex_ai_beta", + "palm", + "gemini", + "ai21", + "baseten", + "azure", + "azure_text", + "azure_ai", + "sagemaker", + "sagemaker_chat", + "bedrock", + "vllm", + "nlp_cloud", + "petals", + "oobabooga", + "ollama", + "ollama_chat", + "deepinfra", + "perplexity", + "anyscale", + "mistral", + "groq", + "nvidia_nim", + "cerebras", + "ai21_chat", + "volcengine", + "codestral", + "text-completion-codestral", + "deepseek", + "sambanova", + "maritalk", + "voyage", + "cloudflare", + "xinference", + "fireworks_ai", + "friendliai", + "watsonx", + "watsonx_text", + "triton", + "predibase", + "databricks", + "empower", + "github", + "custom", + "litellm_proxy", + "hosted_vllm", + "lm_studio", + "galadriel", +] diff --git a/litellm/integrations/mlflow.py b/litellm/integrations/mlflow.py index 7268350d19f..ad3392595aa 100644 --- a/litellm/integrations/mlflow.py +++ b/litellm/integrations/mlflow.py @@ -65,7 +65,7 @@ class MlflowLogger(CustomLogger): # Record exception info as event if exception := kwargs.get("exception"): - span.add_event(SpanEvent.from_exception(exception)) + span.add_event(SpanEvent.from_exception(exception)) # type: ignore self._end_span_or_trace( span=span, diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 554df8092b3..2efed0da365 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -33,7 +33,7 @@ def get_supported_openai_params( # noqa: PLR0915 elif custom_llm_provider == "ollama_chat": return litellm.OllamaChatConfig().get_supported_openai_params() elif custom_llm_provider == "anthropic": - return litellm.AnthropicConfig().get_supported_openai_params() + return litellm.AnthropicConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "fireworks_ai": if request_type == "embeddings": return litellm.FireworksAIEmbeddingConfig().get_supported_openai_params( @@ -75,33 +75,9 @@ def get_supported_openai_params( # noqa: PLR0915 "tool_choice", ] elif custom_llm_provider == "cohere": - return [ - "stream", - "temperature", - "max_tokens", - "logit_bias", - "top_p", - "frequency_penalty", - "presence_penalty", - "stop", - "n", - "extra_headers", - ] + return litellm.CohereConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "cohere_chat": - return [ - "stream", - "temperature", - "max_tokens", - "top_p", - "frequency_penalty", - "presence_penalty", - "stop", - "n", - "tools", - "tool_choice", - "seed", - "extra_headers", - ] + return litellm.CohereChatConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "maritalk": return [ "stream", @@ -194,7 +170,9 @@ def get_supported_openai_params( # noqa: PLR0915 litellm.MistralTextCompletionConfig().get_supported_openai_params() ) if model.startswith("claude"): - return litellm.VertexAIAnthropicConfig().get_supported_openai_params() + return litellm.VertexAIAnthropicConfig().get_supported_openai_params( + model=model + ) return litellm.VertexAIConfig().get_supported_openai_params() elif request_type == "embeddings": return litellm.VertexAITextEmbeddingConfig().get_supported_openai_params() @@ -255,27 +233,7 @@ def get_supported_openai_params( # noqa: PLR0915 elif custom_llm_provider == "watsonx": return litellm.IBMWatsonXChatConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "custom_openai" or "text-completion-openai": - return [ - "functions", - "function_call", - "temperature", - "top_p", - "n", - "stream", - "stream_options", - "stop", - "max_tokens", - "presence_penalty", - "frequency_penalty", - "logit_bias", - "user", - "response_format", - "seed", - "tools", - "tool_choice", - "max_retries", - "logprobs", - "top_logprobs", - "extra_headers", - ] + return litellm.OpenAITextCompletionConfig().get_supported_openai_params( + model=model + ) return None diff --git a/litellm/llms/OpenAI/chat/gpt_audio_transformation.py b/litellm/llms/OpenAI/chat/gpt_audio_transformation.py index 59f7dc01e54..867575e7962 100644 --- a/litellm/llms/OpenAI/chat/gpt_audio_transformation.py +++ b/litellm/llms/OpenAI/chat/gpt_audio_transformation.py @@ -20,21 +20,7 @@ class OpenAIGPTAudioConfig(OpenAIGPTConfig): @classmethod def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } + return super().get_config() def get_supported_openai_params(self, model: str) -> list: """ diff --git a/litellm/llms/OpenAI/chat/gpt_transformation.py b/litellm/llms/OpenAI/chat/gpt_transformation.py index c0c7e14dd8f..7c709e9541f 100644 --- a/litellm/llms/OpenAI/chat/gpt_transformation.py +++ b/litellm/llms/OpenAI/chat/gpt_transformation.py @@ -3,13 +3,26 @@ Support for gpt model family """ import types -from typing import List, Optional, Union +from typing import TYPE_CHECKING, Any, List, Optional, Union, cast + +import httpx import litellm +from litellm.llms.base_llm.transformation import BaseConfig, BaseLLMException from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage +from litellm.types.utils import ModelResponse + +from ..common_utils import OpenAIError + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + LoggingClass = LiteLLMLoggingObj +else: + LoggingClass = Any -class OpenAIGPTConfig: +class OpenAIGPTConfig(BaseConfig): """ Reference: https://platform.openai.com/docs/api-reference/chat/create @@ -69,21 +82,7 @@ class OpenAIGPTConfig: @classmethod def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } + return super().get_config() def get_supported_openai_params(self, model: str) -> list: base_params = [ @@ -168,3 +167,59 @@ class OpenAIGPTConfig: self, messages: List[AllMessageValues] ) -> List[AllMessageValues]: return messages + + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + """ + Transform the overall request to be sent to the API. + + Returns: + dict: The transformed request. Sent as the body of the API call. + """ + raise NotImplementedError + + def transform_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ModelResponse, + logging_obj: LoggingClass, + api_key: str, + request_data: dict, + messages: List[AllMessageValues], + optional_params: dict, + encoding: Any, + json_mode: Optional[bool] = None, + ) -> ModelResponse: + """ + Transform the response from the API. + + Returns: + dict: The transformed response. + """ + raise NotImplementedError + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> BaseLLMException: + return OpenAIError( + status_code=status_code, + message=error_message, + headers=cast(httpx.Headers, headers), + ) + + def validate_environment( + self, + api_key: str, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + ) -> dict: + raise NotImplementedError diff --git a/litellm/llms/OpenAI/chat/o1_transformation.py b/litellm/llms/OpenAI/chat/o1_transformation.py index 2dd70afbb4a..115bb29b1da 100644 --- a/litellm/llms/OpenAI/chat/o1_transformation.py +++ b/litellm/llms/OpenAI/chat/o1_transformation.py @@ -27,21 +27,7 @@ class OpenAIO1Config(OpenAIGPTConfig): @classmethod def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } + return super().get_config() def get_supported_openai_params(self, model: str) -> list: """ diff --git a/litellm/llms/OpenAI/common_utils.py b/litellm/llms/OpenAI/common_utils.py index 01c3ae9435e..5da8c4925f6 100644 --- a/litellm/llms/OpenAI/common_utils.py +++ b/litellm/llms/OpenAI/common_utils.py @@ -3,10 +3,44 @@ Common helpers / utils across al OpenAI endpoints """ import json -from typing import Any, Dict, List +from typing import Any, Dict, List, Optional +import httpx import openai +from litellm.llms.base_llm.transformation import BaseLLMException + + +class OpenAIError(BaseLLMException): + def __init__( + self, + status_code: int, + message: str, + request: Optional[httpx.Request] = None, + response: Optional[httpx.Response] = None, + headers: Optional[httpx.Headers] = None, + ): + self.status_code = status_code + self.message = message + self.headers = headers + if request: + self.request = request + else: + self.request = httpx.Request(method="POST", url="https://api.openai.com/v1") + if response: + self.response = response + else: + self.response = httpx.Response( + status_code=status_code, request=self.request + ) + super().__init__( + status_code=status_code, + message=self.message, + headers=self.headers, + request=self.request, + response=self.response, + ) + ####### Error Handling Utils for OpenAI API ####################### ################################################################### diff --git a/litellm/llms/OpenAI/completion/handler.py b/litellm/llms/OpenAI/completion/handler.py new file mode 100644 index 00000000000..43f1ed92329 --- /dev/null +++ b/litellm/llms/OpenAI/completion/handler.py @@ -0,0 +1,314 @@ +import json +from typing import Callable, List, Optional, Union + +from openai import AsyncOpenAI, OpenAI + +import litellm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.llms.base import BaseLLM +from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUserMessage +from litellm.types.utils import ModelResponse, TextCompletionResponse + +from ..common_utils import OpenAIError +from .transformation import OpenAITextCompletionConfig + + +class OpenAITextCompletion(BaseLLM): + openai_text_completion_global_config = OpenAITextCompletionConfig() + + def __init__(self) -> None: + super().__init__() + + def validate_environment(self, api_key): + headers = { + "content-type": "application/json", + } + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + return headers + + def completion( + self, + model_response: ModelResponse, + api_key: str, + model: str, + messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + timeout: float, + logging_obj: LiteLLMLoggingObj, + optional_params: dict, + print_verbose: Optional[Callable] = None, + api_base: Optional[str] = None, + acompletion: bool = False, + litellm_params=None, + logger_fn=None, + client=None, + organization: Optional[str] = None, + headers: Optional[dict] = None, + ): + try: + if headers is None: + headers = self.validate_environment(api_key=api_key) + if model is None or messages is None: + raise OpenAIError(status_code=422, message="Missing model or messages") + + # don't send max retries to the api, if set + + prompt = self.openai_text_completion_global_config._transform_prompt( + messages + ) + + data = {"model": model, "prompt": prompt, **optional_params} + max_retries = data.pop("max_retries", 2) + ## LOGGING + logging_obj.pre_call( + input=messages, + api_key=api_key, + additional_args={ + "headers": headers, + "api_base": api_base, + "complete_input_dict": data, + }, + ) + if acompletion is True: + if optional_params.get("stream", False): + return self.async_streaming( + logging_obj=logging_obj, + api_base=api_base, + api_key=api_key, + data=data, + headers=headers, + model_response=model_response, + model=model, + timeout=timeout, + max_retries=max_retries, + client=client, + organization=organization, + ) + else: + return self.acompletion(api_base=api_base, data=data, headers=headers, model_response=model_response, prompt=prompt, api_key=api_key, logging_obj=logging_obj, model=model, timeout=timeout, max_retries=max_retries, organization=organization, client=client) # type: ignore + elif optional_params.get("stream", False): + return self.streaming( + logging_obj=logging_obj, + api_base=api_base, + api_key=api_key, + data=data, + headers=headers, + model_response=model_response, + model=model, + timeout=timeout, + max_retries=max_retries, # type: ignore + client=client, + organization=organization, + ) + else: + if client is None: + openai_client = OpenAI( + api_key=api_key, + base_url=api_base, + http_client=litellm.client_session, + timeout=timeout, + max_retries=max_retries, # type: ignore + organization=organization, + ) + else: + openai_client = client + + raw_response = openai_client.completions.with_raw_response.create(**data) # type: ignore + response = raw_response.parse() + response_json = response.model_dump() + + ## LOGGING + logging_obj.post_call( + input=prompt, + api_key=api_key, + original_response=response_json, + additional_args={ + "headers": headers, + "api_base": api_base, + }, + ) + + ## RESPONSE OBJECT + return TextCompletionResponse(**response_json) + except Exception as e: + status_code = getattr(e, "status_code", 500) + error_headers = getattr(e, "headers", None) + error_text = getattr(e, "text", str(e)) + error_response = getattr(e, "response", None) + if error_headers is None and error_response: + error_headers = getattr(error_response, "headers", None) + raise OpenAIError( + status_code=status_code, message=error_text, headers=error_headers + ) + + async def acompletion( + self, + logging_obj, + api_base: str, + data: dict, + headers: dict, + model_response: ModelResponse, + prompt: str, + api_key: str, + model: str, + timeout: float, + max_retries: int, + organization: Optional[str] = None, + client=None, + ): + try: + if client is None: + openai_aclient = AsyncOpenAI( + api_key=api_key, + base_url=api_base, + http_client=litellm.aclient_session, + timeout=timeout, + max_retries=max_retries, + organization=organization, + ) + else: + openai_aclient = client + + raw_response = await openai_aclient.completions.with_raw_response.create( + **data + ) + response = raw_response.parse() + response_json = response.model_dump() + + ## LOGGING + logging_obj.post_call( + input=prompt, + api_key=api_key, + original_response=response, + additional_args={ + "headers": headers, + "api_base": api_base, + }, + ) + ## RESPONSE OBJECT + response_obj = TextCompletionResponse(**response_json) + response_obj._hidden_params.original_response = json.dumps(response_json) + return response_obj + except Exception as e: + status_code = getattr(e, "status_code", 500) + error_headers = getattr(e, "headers", None) + error_text = getattr(e, "text", str(e)) + error_response = getattr(e, "response", None) + if error_headers is None and error_response: + error_headers = getattr(error_response, "headers", None) + raise OpenAIError( + status_code=status_code, message=error_text, headers=error_headers + ) + + def streaming( + self, + logging_obj, + api_key: str, + data: dict, + headers: dict, + model_response: ModelResponse, + model: str, + timeout: float, + api_base: Optional[str] = None, + max_retries=None, + client=None, + organization=None, + ): + + if client is None: + openai_client = OpenAI( + api_key=api_key, + base_url=api_base, + http_client=litellm.client_session, + timeout=timeout, + max_retries=max_retries, # type: ignore + organization=organization, + ) + else: + openai_client = client + + try: + raw_response = openai_client.completions.with_raw_response.create(**data) + response = raw_response.parse() + except Exception as e: + status_code = getattr(e, "status_code", 500) + error_headers = getattr(e, "headers", None) + error_text = getattr(e, "text", str(e)) + error_response = getattr(e, "response", None) + if error_headers is None and error_response: + error_headers = getattr(error_response, "headers", None) + raise OpenAIError( + status_code=status_code, message=error_text, headers=error_headers + ) + streamwrapper = CustomStreamWrapper( + completion_stream=response, + model=model, + custom_llm_provider="text-completion-openai", + logging_obj=logging_obj, + stream_options=data.get("stream_options", None), + ) + + try: + for chunk in streamwrapper: + yield chunk + except Exception as e: + status_code = getattr(e, "status_code", 500) + error_headers = getattr(e, "headers", None) + error_text = getattr(e, "text", str(e)) + error_response = getattr(e, "response", None) + if error_headers is None and error_response: + error_headers = getattr(error_response, "headers", None) + raise OpenAIError( + status_code=status_code, message=error_text, headers=error_headers + ) + + async def async_streaming( + self, + logging_obj, + api_key: str, + data: dict, + headers: dict, + model_response: ModelResponse, + model: str, + timeout: float, + max_retries: int, + api_base: Optional[str] = None, + client=None, + organization=None, + ): + if client is None: + openai_client = AsyncOpenAI( + api_key=api_key, + base_url=api_base, + http_client=litellm.aclient_session, + timeout=timeout, + max_retries=max_retries, + organization=organization, + ) + else: + openai_client = client + + raw_response = await openai_client.completions.with_raw_response.create(**data) + response = raw_response.parse() + streamwrapper = CustomStreamWrapper( + completion_stream=response, + model=model, + custom_llm_provider="text-completion-openai", + logging_obj=logging_obj, + stream_options=data.get("stream_options", None), + ) + + try: + async for transformed_chunk in streamwrapper: + yield transformed_chunk + except Exception as e: + status_code = getattr(e, "status_code", 500) + error_headers = getattr(e, "headers", None) + error_text = getattr(e, "text", str(e)) + error_response = getattr(e, "response", None) + if error_headers is None and error_response: + error_headers = getattr(error_response, "headers", None) + raise OpenAIError( + status_code=status_code, message=error_text, headers=error_headers + ) diff --git a/litellm/llms/OpenAI/completion/transformation.py b/litellm/llms/OpenAI/completion/transformation.py new file mode 100644 index 00000000000..cb6baba3393 --- /dev/null +++ b/litellm/llms/OpenAI/completion/transformation.py @@ -0,0 +1,178 @@ +""" +Support for gpt model family +""" + +import types +from typing import List, Optional, Union, cast + +import litellm +from litellm.llms.base_llm.transformation import BaseConfig +from litellm.types.llms.openai import ( + AllMessageValues, + AllPromptValues, + OpenAITextCompletionUserMessage, +) +from litellm.types.utils import Choices, Message, ModelResponse, TextCompletionResponse + +from ...prompt_templates.common_utils import convert_content_list_to_str +from ..chat.gpt_transformation import OpenAIGPTConfig +from ..common_utils import OpenAIError +from .utils import is_tokens_or_list_of_tokens + + +class OpenAITextCompletionConfig(OpenAIGPTConfig): + """ + Reference: https://platform.openai.com/docs/api-reference/completions/create + + The class `OpenAITextCompletionConfig` provides configuration for the OpenAI's text completion API interface. Below are the parameters: + + - `best_of` (integer or null): This optional parameter generates server-side completions and returns the one with the highest log probability per token. + + - `echo` (boolean or null): This optional parameter will echo back the prompt in addition to the completion. + + - `frequency_penalty` (number or null): Defaults to 0. It is a numbers from -2.0 to 2.0, where positive values decrease the model's likelihood to repeat the same line. + + - `logit_bias` (map): This optional parameter modifies the likelihood of specified tokens appearing in the completion. + + - `logprobs` (integer or null): This optional parameter includes the log probabilities on the most likely tokens as well as the chosen tokens. + + - `max_tokens` (integer or null): This optional parameter sets the maximum number of tokens to generate in the completion. + + - `n` (integer or null): This optional parameter sets how many completions to generate for each prompt. + + - `presence_penalty` (number or null): Defaults to 0 and can be between -2.0 and 2.0. Positive values increase the model's likelihood to talk about new topics. + + - `stop` (string / array / null): Specifies up to 4 sequences where the API will stop generating further tokens. + + - `suffix` (string or null): Defines the suffix that comes after a completion of inserted text. + + - `temperature` (number or null): This optional parameter defines the sampling temperature to use. + + - `top_p` (number or null): An alternative to sampling with temperature, used for nucleus sampling. + """ + + best_of: Optional[int] = None + echo: Optional[bool] = None + frequency_penalty: Optional[int] = None + logit_bias: Optional[dict] = None + logprobs: Optional[int] = None + max_tokens: Optional[int] = None + n: Optional[int] = None + presence_penalty: Optional[int] = None + stop: Optional[Union[str, list]] = None + suffix: Optional[str] = None + + def __init__( + self, + best_of: Optional[int] = None, + echo: Optional[bool] = None, + frequency_penalty: Optional[int] = None, + logit_bias: Optional[dict] = None, + logprobs: Optional[int] = None, + max_tokens: Optional[int] = None, + n: Optional[int] = None, + presence_penalty: Optional[int] = None, + stop: Optional[Union[str, list]] = None, + suffix: Optional[str] = None, + temperature: Optional[float] = None, + top_p: Optional[float] = None, + ) -> None: + locals_ = locals().copy() + for key, value in locals_.items(): + if key != "self" and value is not None: + setattr(self.__class__, key, value) + + @classmethod + def get_config(cls): + return super().get_config() + + def _transform_prompt( + self, + messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], + ) -> AllPromptValues: + if len(messages) == 1: # base case + message_content = messages[0].get("content") + if ( + message_content + and isinstance(message_content, list) + and is_tokens_or_list_of_tokens(message_content) + ): + openai_prompt: AllPromptValues = cast(AllPromptValues, message_content) + else: + openai_prompt = "" + content = convert_content_list_to_str( + cast(AllMessageValues, messages[0]) + ) + openai_prompt += content + else: + prompt_str_list: List[str] = [] + for m in messages: + try: # expect list of int/list of list of int to be a 1 message array only. + content = convert_content_list_to_str(cast(AllMessageValues, m)) + prompt_str_list.append(content) + except Exception as e: + raise e + openai_prompt = prompt_str_list + return openai_prompt + + def convert_to_chat_model_response_object( + self, + response_object: Optional[TextCompletionResponse] = None, + model_response_object: Optional[ModelResponse] = None, + ): + try: + ## RESPONSE OBJECT + if response_object is None or model_response_object is None: + raise ValueError("Error in response object format") + choice_list = [] + for idx, choice in enumerate(response_object["choices"]): + message = Message( + content=choice["text"], + role="assistant", + ) + choice = Choices( + finish_reason=choice["finish_reason"], index=idx, message=message + ) + choice_list.append(choice) + model_response_object.choices = choice_list + + if "usage" in response_object: + setattr(model_response_object, "usage", response_object["usage"]) + + if "id" in response_object: + model_response_object.id = response_object["id"] + + if "model" in response_object: + model_response_object.model = response_object["model"] + + model_response_object._hidden_params["original_response"] = ( + response_object # track original response, if users make a litellm.text_completion() request, we can return the original response + ) + return model_response_object + except Exception as e: + raise e + + def get_supported_openai_params(self, model: str) -> List: + return [ + "functions", + "function_call", + "temperature", + "top_p", + "n", + "stream", + "stream_options", + "stop", + "max_tokens", + "presence_penalty", + "frequency_penalty", + "logit_bias", + "user", + "response_format", + "seed", + "tools", + "tool_choice", + "max_retries", + "logprobs", + "top_logprobs", + "extra_headers", + ] diff --git a/litellm/llms/OpenAI/openai.py b/litellm/llms/OpenAI/openai.py index 66ce7570188..108a31d19a9 100644 --- a/litellm/llms/OpenAI/openai.py +++ b/litellm/llms/OpenAI/openai.py @@ -34,37 +34,8 @@ from litellm.utils import ( from ...types.llms.openai import * from ..base import BaseLLM -from ..prompt_templates.common_utils import convert_content_list_to_str from ..prompt_templates.factory import custom_prompt, prompt_factory -from .common_utils import drop_params_from_unprocessable_entity_error -from .completion.utils import is_tokens_or_list_of_tokens - - -class OpenAIError(Exception): - def __init__( - self, - status_code, - message, - request: Optional[httpx.Request] = None, - response: Optional[httpx.Response] = None, - headers: Optional[httpx.Headers] = None, - ): - self.status_code = status_code - self.message = message - self.headers = headers - if request: - self.request = request - else: - self.request = httpx.Request(method="POST", url="https://api.openai.com/v1") - if response: - self.response = response - else: - 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 +from .common_utils import OpenAIError, drop_params_from_unprocessable_entity_error class MistralEmbeddingConfig: @@ -379,155 +350,6 @@ class OpenAIConfig: ) -class OpenAITextCompletionConfig: - """ - Reference: https://platform.openai.com/docs/api-reference/completions/create - - The class `OpenAITextCompletionConfig` provides configuration for the OpenAI's text completion API interface. Below are the parameters: - - - `best_of` (integer or null): This optional parameter generates server-side completions and returns the one with the highest log probability per token. - - - `echo` (boolean or null): This optional parameter will echo back the prompt in addition to the completion. - - - `frequency_penalty` (number or null): Defaults to 0. It is a numbers from -2.0 to 2.0, where positive values decrease the model's likelihood to repeat the same line. - - - `logit_bias` (map): This optional parameter modifies the likelihood of specified tokens appearing in the completion. - - - `logprobs` (integer or null): This optional parameter includes the log probabilities on the most likely tokens as well as the chosen tokens. - - - `max_tokens` (integer or null): This optional parameter sets the maximum number of tokens to generate in the completion. - - - `n` (integer or null): This optional parameter sets how many completions to generate for each prompt. - - - `presence_penalty` (number or null): Defaults to 0 and can be between -2.0 and 2.0. Positive values increase the model's likelihood to talk about new topics. - - - `stop` (string / array / null): Specifies up to 4 sequences where the API will stop generating further tokens. - - - `suffix` (string or null): Defines the suffix that comes after a completion of inserted text. - - - `temperature` (number or null): This optional parameter defines the sampling temperature to use. - - - `top_p` (number or null): An alternative to sampling with temperature, used for nucleus sampling. - """ - - best_of: Optional[int] = None - echo: Optional[bool] = None - frequency_penalty: Optional[int] = None - logit_bias: Optional[dict] = None - logprobs: Optional[int] = None - max_tokens: Optional[int] = None - n: Optional[int] = None - presence_penalty: Optional[int] = None - stop: Optional[Union[str, list]] = None - suffix: Optional[str] = None - temperature: Optional[float] = None - top_p: Optional[float] = None - - def __init__( - self, - best_of: Optional[int] = None, - echo: Optional[bool] = None, - frequency_penalty: Optional[int] = None, - logit_bias: Optional[dict] = None, - logprobs: Optional[int] = None, - max_tokens: Optional[int] = None, - n: Optional[int] = None, - presence_penalty: Optional[int] = None, - stop: Optional[Union[str, list]] = None, - suffix: Optional[str] = None, - temperature: Optional[float] = None, - top_p: Optional[float] = None, - ) -> None: - locals_ = locals().copy() - for key, value in locals_.items(): - if key != "self" and value is not None: - setattr(self.__class__, key, value) - - @classmethod - def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } - - def _transform_prompt( - self, - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], - ) -> AllPromptValues: - if len(messages) == 1: # base case - message_content = messages[0].get("content") - if ( - message_content - and isinstance(message_content, list) - and is_tokens_or_list_of_tokens(message_content) - ): - openai_prompt: AllPromptValues = cast(AllPromptValues, message_content) - else: - openai_prompt = "" - content = convert_content_list_to_str( - cast(AllMessageValues, messages[0]) - ) - openai_prompt += content - else: - prompt_str_list: List[str] = [] - for m in messages: - try: # expect list of int/list of list of int to be a 1 message array only. - content = convert_content_list_to_str(cast(AllMessageValues, m)) - prompt_str_list.append(content) - except Exception as e: - raise e - openai_prompt = prompt_str_list - return openai_prompt - - def convert_to_chat_model_response_object( - self, - response_object: Optional[TextCompletionResponse] = None, - model_response_object: Optional[ModelResponse] = None, - ): - try: - ## RESPONSE OBJECT - if response_object is None or model_response_object is None: - raise ValueError("Error in response object format") - choice_list = [] - for idx, choice in enumerate(response_object["choices"]): - message = Message( - content=choice["text"], - role="assistant", - ) - choice = Choices( - finish_reason=choice["finish_reason"], index=idx, message=message - ) - choice_list.append(choice) - model_response_object.choices = choice_list - - if "usage" in response_object: - setattr(model_response_object, "usage", response_object["usage"]) - - if "id" in response_object: - model_response_object.id = response_object["id"] - - if "model" in response_object: - model_response_object.model = response_object["model"] - - model_response_object._hidden_params["original_response"] = ( - response_object # track original response, if users make a litellm.text_completion() request, we can return the original response - ) - return model_response_object - except Exception as e: - raise e - - class OpenAIChatCompletion(BaseLLM): def __init__(self) -> None: @@ -710,7 +532,7 @@ class OpenAIChatCompletion(BaseLLM): custom_llm_provider=custom_llm_provider, ) if messages is not None and custom_llm_provider is not None: - provider_config = ProviderConfigManager.get_provider_config( + provider_config = ProviderConfigManager.get_provider_chat_config( model=model, provider=LlmProviders(custom_llm_provider) ) messages = provider_config._transform_messages(messages) @@ -1584,306 +1406,6 @@ class OpenAIChatCompletion(BaseLLM): return response -class OpenAITextCompletion(BaseLLM): - openai_text_completion_global_config = OpenAITextCompletionConfig() - - def __init__(self) -> None: - super().__init__() - - def validate_environment(self, api_key): - headers = { - "content-type": "application/json", - } - if api_key: - headers["Authorization"] = f"Bearer {api_key}" - return headers - - def completion( - self, - model_response: ModelResponse, - api_key: str, - model: str, - messages: Union[List[AllMessageValues], List[OpenAITextCompletionUserMessage]], - timeout: float, - logging_obj: LiteLLMLoggingObj, - optional_params: dict, - print_verbose: Optional[Callable] = None, - api_base: Optional[str] = None, - acompletion: bool = False, - litellm_params=None, - logger_fn=None, - client=None, - organization: Optional[str] = None, - headers: Optional[dict] = None, - ): - try: - if headers is None: - headers = self.validate_environment(api_key=api_key) - if model is None or messages is None: - raise OpenAIError(status_code=422, message="Missing model or messages") - - # don't send max retries to the api, if set - - prompt = self.openai_text_completion_global_config._transform_prompt( - messages - ) - - data = {"model": model, "prompt": prompt, **optional_params} - max_retries = data.pop("max_retries", 2) - ## LOGGING - logging_obj.pre_call( - input=messages, - api_key=api_key, - additional_args={ - "headers": headers, - "api_base": api_base, - "complete_input_dict": data, - }, - ) - if acompletion is True: - if optional_params.get("stream", False): - return self.async_streaming( - logging_obj=logging_obj, - api_base=api_base, - api_key=api_key, - data=data, - headers=headers, - model_response=model_response, - model=model, - timeout=timeout, - max_retries=max_retries, - client=client, - organization=organization, - ) - else: - return self.acompletion(api_base=api_base, data=data, headers=headers, model_response=model_response, prompt=prompt, api_key=api_key, logging_obj=logging_obj, model=model, timeout=timeout, max_retries=max_retries, organization=organization, client=client) # type: ignore - elif optional_params.get("stream", False): - return self.streaming( - logging_obj=logging_obj, - api_base=api_base, - api_key=api_key, - data=data, - headers=headers, - model_response=model_response, - model=model, - timeout=timeout, - max_retries=max_retries, # type: ignore - client=client, - organization=organization, - ) - else: - if client is None: - openai_client = OpenAI( - api_key=api_key, - base_url=api_base, - http_client=litellm.client_session, - timeout=timeout, - max_retries=max_retries, # type: ignore - organization=organization, - ) - else: - openai_client = client - - raw_response = openai_client.completions.with_raw_response.create(**data) # type: ignore - response = raw_response.parse() - response_json = response.model_dump() - - ## LOGGING - logging_obj.post_call( - input=prompt, - api_key=api_key, - original_response=response_json, - additional_args={ - "headers": headers, - "api_base": api_base, - }, - ) - - ## RESPONSE OBJECT - return TextCompletionResponse(**response_json) - except Exception as e: - status_code = getattr(e, "status_code", 500) - error_headers = getattr(e, "headers", None) - error_text = getattr(e, "text", str(e)) - error_response = getattr(e, "response", None) - if error_headers is None and error_response: - error_headers = getattr(error_response, "headers", None) - raise OpenAIError( - status_code=status_code, message=error_text, headers=error_headers - ) - - async def acompletion( - self, - logging_obj, - api_base: str, - data: dict, - headers: dict, - model_response: ModelResponse, - prompt: str, - api_key: str, - model: str, - timeout: float, - max_retries: int, - organization: Optional[str] = None, - client=None, - ): - try: - if client is None: - openai_aclient = AsyncOpenAI( - api_key=api_key, - base_url=api_base, - http_client=litellm.aclient_session, - timeout=timeout, - max_retries=max_retries, - organization=organization, - ) - else: - openai_aclient = client - - raw_response = await openai_aclient.completions.with_raw_response.create( - **data - ) - response = raw_response.parse() - response_json = response.model_dump() - - ## LOGGING - logging_obj.post_call( - input=prompt, - api_key=api_key, - original_response=response, - additional_args={ - "headers": headers, - "api_base": api_base, - }, - ) - ## RESPONSE OBJECT - response_obj = TextCompletionResponse(**response_json) - response_obj._hidden_params.original_response = json.dumps(response_json) - return response_obj - except Exception as e: - status_code = getattr(e, "status_code", 500) - error_headers = getattr(e, "headers", None) - error_text = getattr(e, "text", str(e)) - error_response = getattr(e, "response", None) - if error_headers is None and error_response: - error_headers = getattr(error_response, "headers", None) - raise OpenAIError( - status_code=status_code, message=error_text, headers=error_headers - ) - - def streaming( - self, - logging_obj, - api_key: str, - data: dict, - headers: dict, - model_response: ModelResponse, - model: str, - timeout: float, - api_base: Optional[str] = None, - max_retries=None, - client=None, - organization=None, - ): - - if client is None: - openai_client = OpenAI( - api_key=api_key, - base_url=api_base, - http_client=litellm.client_session, - timeout=timeout, - max_retries=max_retries, # type: ignore - organization=organization, - ) - else: - openai_client = client - - try: - raw_response = openai_client.completions.with_raw_response.create(**data) - response = raw_response.parse() - except Exception as e: - status_code = getattr(e, "status_code", 500) - error_headers = getattr(e, "headers", None) - error_text = getattr(e, "text", str(e)) - error_response = getattr(e, "response", None) - if error_headers is None and error_response: - error_headers = getattr(error_response, "headers", None) - raise OpenAIError( - status_code=status_code, message=error_text, headers=error_headers - ) - streamwrapper = CustomStreamWrapper( - completion_stream=response, - model=model, - custom_llm_provider="text-completion-openai", - logging_obj=logging_obj, - stream_options=data.get("stream_options", None), - ) - - try: - for chunk in streamwrapper: - yield chunk - except Exception as e: - status_code = getattr(e, "status_code", 500) - error_headers = getattr(e, "headers", None) - error_text = getattr(e, "text", str(e)) - error_response = getattr(e, "response", None) - if error_headers is None and error_response: - error_headers = getattr(error_response, "headers", None) - raise OpenAIError( - status_code=status_code, message=error_text, headers=error_headers - ) - - async def async_streaming( - self, - logging_obj, - api_key: str, - data: dict, - headers: dict, - model_response: ModelResponse, - model: str, - timeout: float, - max_retries: int, - api_base: Optional[str] = None, - client=None, - organization=None, - ): - if client is None: - openai_client = AsyncOpenAI( - api_key=api_key, - base_url=api_base, - http_client=litellm.aclient_session, - timeout=timeout, - max_retries=max_retries, - organization=organization, - ) - else: - openai_client = client - - raw_response = await openai_client.completions.with_raw_response.create(**data) - response = raw_response.parse() - streamwrapper = CustomStreamWrapper( - completion_stream=response, - model=model, - custom_llm_provider="text-completion-openai", - logging_obj=logging_obj, - stream_options=data.get("stream_options", None), - ) - - try: - async for transformed_chunk in streamwrapper: - yield transformed_chunk - except Exception as e: - status_code = getattr(e, "status_code", 500) - error_headers = getattr(e, "headers", None) - error_text = getattr(e, "text", str(e)) - error_response = getattr(e, "response", None) - if error_headers is None and error_response: - error_headers = getattr(error_response, "headers", None) - raise OpenAIError( - status_code=status_code, message=error_text, headers=error_headers - ) - - class OpenAIFilesAPI(BaseLLM): """ OpenAI methods to support for batches diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index be46051c68d..444082fac50 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -20,7 +20,7 @@ import litellm import litellm.litellm_core_utils import litellm.types import litellm.types.utils -from litellm import verbose_logger +from litellm import LlmProviders, verbose_logger from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -45,7 +45,7 @@ from litellm.types.llms.openai import ( ChatCompletionUsageBlock, ) from litellm.types.utils import GenericStreamingChunk -from litellm.utils import CustomStreamWrapper, ModelResponse +from litellm.utils import CustomStreamWrapper, ModelResponse, ProviderConfigManager from ...base import BaseLLM from ..common_utils import AnthropicError, process_anthropic_headers @@ -63,28 +63,7 @@ def validate_environment( anthropic_version: Optional[str] = None, ): - if api_key is None: - raise litellm.AuthenticationError( - message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` in your environment vars", - llm_provider="anthropic", - model=model, - ) - - prompt_caching_set = AnthropicConfig().is_cache_control_set(messages=messages) - computer_tool_used = AnthropicConfig().is_computer_tool_used(tools=tools) - pdf_used = AnthropicConfig().is_pdf_used(messages=messages) - headers = AnthropicConfig().get_anthropic_headers( - anthropic_version=anthropic_version, - computer_tool_used=computer_tool_used, - prompt_caching_set=prompt_caching_set, - pdf_used=pdf_used, - api_key=api_key, - is_vertex_request=is_vertex_request, - ) - - if user_headers is not None and isinstance(user_headers, dict): - headers = {**headers, **user_headers} - return headers + pass async def make_call( @@ -295,16 +274,14 @@ class AnthropicChatCompletion(BaseLLM): headers=error_headers, ) - return AnthropicConfig._process_response( + return AnthropicConfig().transform_response( model=model, - response=response, + raw_response=response, model_response=model_response, - stream=stream, logging_obj=logging_obj, api_key=api_key, - data=data, + request_data=data, messages=messages, - print_verbose=print_verbose, optional_params=optional_params, encoding=encoding, json_mode=json_mode, @@ -315,6 +292,7 @@ class AnthropicChatCompletion(BaseLLM): model: str, messages: list, api_base: str, + custom_llm_provider: str, custom_prompt_dict: dict, model_response: ModelResponse, print_verbose: Callable, @@ -335,23 +313,25 @@ class AnthropicChatCompletion(BaseLLM): is_vertex_request: bool = optional_params.pop("is_vertex_request", False) _is_function_call = False messages = copy.deepcopy(messages) - headers = validate_environment( - api_key, - headers, - model, + headers = AnthropicConfig().validate_environment( + api_key=api_key, + headers=headers, + model=model, messages=messages, - tools=optional_params.get("tools"), - is_vertex_request=is_vertex_request, + optional_params={**optional_params, "is_vertex_request": is_vertex_request}, ) - data = AnthropicConfig()._transform_request( + config = ProviderConfigManager.get_provider_chat_config( + model=model, + provider=LlmProviders(custom_llm_provider), + ) + + data = config.transform_request( model=model, messages=messages, optional_params=optional_params, litellm_params=litellm_params, headers=headers, - _is_function_call=_is_function_call, - is_vertex_request=is_vertex_request, ) ## LOGGING @@ -471,16 +451,14 @@ class AnthropicChatCompletion(BaseLLM): headers=error_headers, ) - return AnthropicConfig._process_response( + return AnthropicConfig().transform_response( model=model, - response=response, + raw_response=response, model_response=model_response, - stream=stream, logging_obj=logging_obj, api_key=api_key, - data=data, # type: ignore + request_data=data, messages=messages, - print_verbose=print_verbose, optional_params=optional_params, encoding=encoding, json_mode=json_mode, diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index feb5b8646a2..51eb6ad5c53 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -2,13 +2,25 @@ import json import time import types from re import A -from typing import Dict, List, Literal, Optional, Tuple, Union +from typing import ( + TYPE_CHECKING, + Any, + Callable, + Dict, + List, + Literal, + Optional, + Tuple, + Union, + cast, +) import httpx import requests import litellm from litellm.litellm_core_utils.core_helpers import map_finish_reason +from litellm.llms.base_llm.transformation import BaseConfig, BaseLLMException from litellm.llms.prompt_templates.factory import anthropic_messages_pt from litellm.types.llms.anthropic import ( AllAnthropicToolsValues, @@ -43,8 +55,15 @@ from litellm.utils import ( from ..common_utils import AnthropicError, process_anthropic_headers +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -class AnthropicConfig: + LoggingClass = LiteLLMLoggingObj +else: + LoggingClass = Any + + +class AnthropicConfig(BaseConfig): """ Reference: https://docs.anthropic.com/claude/reference/messages_post @@ -80,23 +99,9 @@ class AnthropicConfig: @classmethod def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } + return super().get_config() - def get_supported_openai_params(self): + def get_supported_openai_params(self, model: str): return [ "stream", "stop", @@ -297,8 +302,9 @@ class AnthropicConfig: self, non_default_params: dict, optional_params: dict, - messages: Optional[List[AllMessageValues]] = None, - ): + model: str, + drop_params: bool, + ) -> dict: for param, value in non_default_params.items(): if param == "max_tokens": optional_params["max_tokens"] = value @@ -347,25 +353,6 @@ class AnthropicConfig: optional_params["json_mode"] = True if param == "user": optional_params["metadata"] = {"user_id": value} - ## VALIDATE REQUEST - """ - Anthropic doesn't support tool calling without `tools=` param specified. - """ - if ( - "tools" not in non_default_params - and messages is not None - and has_tool_call_blocks(messages) - ): - if litellm.modify_params: - optional_params["tools"] = self._map_tools( - add_dummy_tool(custom_llm_provider="anthropic") - ) - else: - raise litellm.UnsupportedParamsError( - message="Anthropic doesn't support tool calling without `tools=` param specified. Pass `tools=` param OR set `litellm.modify_params = True` // `litellm_settings::modify_params: True` to add dummy tool to the request.", - model="", - llm_provider="anthropic", - ) return optional_params @@ -436,7 +423,7 @@ class AnthropicConfig: and isinstance(message["content"], list) ): for content in message["content"]: - if "type" in content: + if "type" in content and content["type"] != "text": return True return False @@ -493,19 +480,37 @@ class AnthropicConfig: return anthropic_system_message_list - def _transform_request( + def transform_request( self, model: str, messages: List[AllMessageValues], optional_params: dict, litellm_params: dict, headers: dict, - _is_function_call: bool, - is_vertex_request: bool, ) -> dict: """ Translate messages to anthropic format. """ + ## VALIDATE REQUEST + """ + Anthropic doesn't support tool calling without `tools=` param specified. + """ + if ( + "tools" not in optional_params + and messages is not None + and has_tool_call_blocks(messages) + ): + if litellm.modify_params: + optional_params["tools"] = self._map_tools( + add_dummy_tool(custom_llm_provider="anthropic") + ) + else: + raise litellm.UnsupportedParamsError( + message="Anthropic doesn't support tool calling without `tools=` param specified. Pass `tools=` param OR set `litellm.modify_params = True` // `litellm_settings::modify_params: True` to add dummy tool to the request.", + model="", + llm_provider="anthropic", + ) + # Separate system prompt from rest of message anthropic_system_message_list = self.translate_system_message(messages=messages) # Handling anthropic API Prompt Caching @@ -546,57 +551,55 @@ class AnthropicConfig: optional_params["metadata"] = {"user_id": _litellm_metadata["user_id"]} data = { + "model": model, "messages": anthropic_messages, **optional_params, } - if not is_vertex_request: - data["model"] = model + return data - @staticmethod - def _process_response( + def transform_response( + self, model: str, - response: Union[requests.Response, httpx.Response], + raw_response: httpx.Response, model_response: ModelResponse, - stream: bool, - logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, # type: ignore - optional_params: dict, + logging_obj: LoggingClass, api_key: str, - data: Union[dict, str], - messages: List, - print_verbose, - encoding, - json_mode: bool, + request_data: Dict, + messages: List[AllMessageValues], + optional_params: Dict, + encoding: Any, + json_mode: Optional[bool] = None, ) -> ModelResponse: _hidden_params: Dict = {} _hidden_params["additional_headers"] = process_anthropic_headers( - dict(response.headers) + dict(raw_response.headers) ) ## LOGGING logging_obj.post_call( input=messages, api_key=api_key, - original_response=response.text, - additional_args={"complete_input_dict": data}, + original_response=raw_response.text, + additional_args={"complete_input_dict": request_data}, ) - print_verbose(f"raw model_response: {response.text}") + ## RESPONSE OBJECT try: - completion_response = response.json() + completion_response = raw_response.json() except Exception as e: - response_headers = getattr(response, "headers", None) + response_headers = getattr(raw_response, "headers", None) raise AnthropicError( message="Unable to get json response - {}, Original Response: {}".format( - str(e), response.text + str(e), raw_response.text ), - status_code=response.status_code, + status_code=raw_response.status_code, headers=response_headers, ) if "error" in completion_response: - response_headers = getattr(response, "headers", None) + response_headers = getattr(raw_response, "headers", None) raise AnthropicError( message=str(completion_response["error"]), - status_code=response.status_code, + status_code=raw_response.status_code, headers=response_headers, ) else: @@ -625,7 +628,7 @@ class AnthropicConfig: ) ## HANDLE JSON MODE - anthropic returns single function call - if json_mode and len(tool_calls) == 1: + if json_mode is True and len(tool_calls) == 1: json_mode_content_str: Optional[str] = tool_calls[0]["function"].get( "arguments" ) @@ -711,3 +714,47 @@ class AnthropicConfig: # json decode error does occur, return the original tool response str return litellm.Message(content=json_mode_content_str) return None + + def _transform_messages( + self, messages: List[AllMessageValues] + ) -> List[AllMessageValues]: + return messages + + def get_error_class( + self, error_message: str, status_code: int, headers: Dict + ) -> BaseLLMException: + return AnthropicError( + status_code=status_code, + message=error_message, + headers=cast(httpx.Headers, headers), + ) + + def validate_environment( + self, + api_key: str, + headers: Dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + ) -> Dict: + if api_key is None: + raise litellm.AuthenticationError( + message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` in your environment vars", + llm_provider="anthropic", + model=model, + ) + + tools = optional_params.get("tools") + prompt_caching_set = self.is_cache_control_set(messages=messages) + computer_tool_used = self.is_computer_tool_used(tools=tools) + pdf_used = self.is_pdf_used(messages=messages) + anthropic_headers = self.get_anthropic_headers( + computer_tool_used=computer_tool_used, + prompt_caching_set=prompt_caching_set, + pdf_used=pdf_used, + api_key=api_key, + is_vertex_request=False, + ) + + headers = {**headers, **anthropic_headers} + return headers diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index cd268cb12ab..8ef79f95058 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -6,24 +6,17 @@ from typing import Optional, Union import httpx +from litellm.llms.base_llm.transformation import BaseLLMException -class AnthropicError(Exception): + +class AnthropicError(BaseLLMException): def __init__( self, status_code: int, message, headers: Optional[httpx.Headers] = None, ): - self.status_code = status_code - self.message: str = message - self.headers = headers - self.request = httpx.Request( - method="POST", url="https://api.anthropic.com/v1/messages" - ) - 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 + super().__init__(status_code=status_code, message=message, headers=headers) def process_anthropic_headers(headers: Union[httpx.Headers, dict]) -> dict: diff --git a/litellm/llms/azure_text.py b/litellm/llms/azure_text.py index c75965a8f50..9f52f214d70 100644 --- a/litellm/llms/azure_text.py +++ b/litellm/llms/azure_text.py @@ -20,7 +20,8 @@ from litellm.utils import ( ) from .base import BaseLLM -from .OpenAI.openai import OpenAITextCompletion, OpenAITextCompletionConfig +from .OpenAI.completion.handler import OpenAITextCompletion +from .OpenAI.completion.transformation import OpenAITextCompletionConfig from .prompt_templates.factory import custom_prompt, prompt_factory openai_text_completion_config = OpenAITextCompletionConfig() diff --git a/litellm/llms/base_llm/transformation.py b/litellm/llms/base_llm/transformation.py new file mode 100644 index 00000000000..cb600c33b9c --- /dev/null +++ b/litellm/llms/base_llm/transformation.py @@ -0,0 +1,129 @@ +""" +Common base config for all LLM providers +""" + +import types +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING, Any, Callable, List, Optional, Union + +import httpx + +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import ModelResponse + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + LoggingClass = LiteLLMLoggingObj +else: + LoggingClass = Any + + +class BaseLLMException(Exception): + def __init__( + self, + status_code: int, + message: str, + headers: Optional[httpx.Headers] = None, + request: Optional[httpx.Request] = None, + response: Optional[httpx.Response] = None, + ): + self.status_code = status_code + self.message: str = message + self.headers = headers + self.request = httpx.Request(method="POST", url="https://docs.litellm.ai/docs") + 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 + + +class BaseConfig(ABC): + def __init__(self): + pass + + @classmethod + def get_config(cls): + return { + k: v + for k, v in cls.__dict__.items() + if not k.startswith("__") + and not k.startswith("_abc") + and not isinstance( + v, + ( + types.FunctionType, + types.BuiltinFunctionType, + classmethod, + staticmethod, + ), + ) + and v is not None + } + + @abstractmethod + def get_supported_openai_params(self, model: str) -> list: + pass + + @abstractmethod + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + pass + + @abstractmethod + def validate_environment( + self, + api_key: str, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + ) -> dict: + pass + + @abstractmethod + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + pass + + @abstractmethod + def _transform_messages( + self, messages: List[AllMessageValues] + ) -> List[AllMessageValues]: + pass + + @abstractmethod + def transform_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ModelResponse, + logging_obj: LoggingClass, + api_key: str, + request_data: dict, + messages: List[AllMessageValues], + optional_params: dict, + encoding: Any, + json_mode: Optional[bool] = None, + ) -> ModelResponse: + pass + + @abstractmethod + def get_error_class( + self, + error_message: str, + status_code: int, + headers: dict, + ) -> BaseLLMException: + pass diff --git a/litellm/llms/clarifai.py b/litellm/llms/clarifai.py deleted file mode 100644 index 61d44542312..00000000000 --- a/litellm/llms/clarifai.py +++ /dev/null @@ -1,378 +0,0 @@ -import json -import os -import time -import traceback -import types -from typing import Callable, Optional - -import httpx -import requests - -import litellm -from litellm.llms.custom_httpx.http_handler import ( - AsyncHTTPHandler, - get_async_httpx_client, -) -from litellm.utils import Choices, CustomStreamWrapper, Message, ModelResponse, Usage - -from .prompt_templates.factory import custom_prompt, prompt_factory - - -class ClarifaiError(Exception): - def __init__(self, status_code, message, url): - self.status_code = status_code - self.message = message - self.request = httpx.Request(method="POST", url=url) - self.response = httpx.Response(status_code=status_code, request=self.request) - super().__init__(self.message) - - -class ClarifaiConfig: - """ - Reference: https://clarifai.com/meta/Llama-2/models/llama2-70b-chat - """ - - max_tokens: Optional[int] = None - temperature: Optional[int] = None - top_k: Optional[int] = None - - def __init__( - self, - max_tokens: Optional[int] = None, - temperature: Optional[int] = None, - top_k: Optional[int] = None, - ) -> None: - locals_ = locals() - for key, value in locals_.items(): - if key != "self" and value is not None: - setattr(self.__class__, key, value) - - @classmethod - def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } - - -def validate_environment(api_key): - headers = { - "accept": "application/json", - "content-type": "application/json", - } - if api_key: - headers["Authorization"] = f"Bearer {api_key}" - return headers - - -def completions_to_model(payload): - # if payload["n"] != 1: - # raise HTTPException( - # status_code=422, - # detail="Only one generation is supported. Please set candidate_count to 1.", - # ) - - params = {} - if temperature := payload.get("temperature"): - params["temperature"] = temperature - if max_tokens := payload.get("max_tokens"): - params["max_tokens"] = max_tokens - return { - "inputs": [{"data": {"text": {"raw": payload["prompt"]}}}], - "model": {"output_info": {"params": params}}, - } - - -def process_response( - model, - prompt, - response, - model_response: litellm.ModelResponse, - api_key, - data, - encoding, - logging_obj, -): - logging_obj.post_call( - input=prompt, - api_key=api_key, - original_response=response.text, - additional_args={"complete_input_dict": data}, - ) - ## RESPONSE OBJECT - try: - completion_response = response.json() - except Exception: - raise ClarifaiError( - message=response.text, status_code=response.status_code, url=model - ) - # print(completion_response) - try: - choices_list = [] - for idx, item in enumerate(completion_response["outputs"]): - if len(item["data"]["text"]["raw"]) > 0: - message_obj = Message(content=item["data"]["text"]["raw"]) - else: - message_obj = Message(content=None) - choice_obj = Choices( - finish_reason="stop", - index=idx + 1, # check - message=message_obj, - ) - choices_list.append(choice_obj) - model_response.choices = choices_list # type: ignore - - except Exception: - raise ClarifaiError( - message=traceback.format_exc(), status_code=response.status_code, url=model - ) - - # Calculate Usage - prompt_tokens = len(encoding.encode(prompt)) - completion_tokens = len( - encoding.encode(model_response["choices"][0]["message"].get("content")) - ) - model_response.model = model - setattr( - model_response, - "usage", - Usage( - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - total_tokens=prompt_tokens + completion_tokens, - ), - ) - return model_response - - -def convert_model_to_url(model: str, api_base: str): - user_id, app_id, model_id = model.split(".") - return f"{api_base}/users/{user_id}/apps/{app_id}/models/{model_id}/outputs" - - -def get_prompt_model_name(url: str): - clarifai_model_name = url.split("/")[-2] - if "claude" in clarifai_model_name: - return "anthropic", clarifai_model_name.replace("_", ".") - if ("llama" in clarifai_model_name) or ("mistral" in clarifai_model_name): - return "", "meta-llama/llama-2-chat" - else: - return "", clarifai_model_name - - -async def async_completion( - model: str, - prompt: str, - api_base: str, - custom_prompt_dict: dict, - model_response: ModelResponse, - print_verbose: Callable, - encoding, - api_key, - logging_obj, - data=None, - optional_params=None, - litellm_params=None, - logger_fn=None, - headers={}, -): - - async_handler = get_async_httpx_client( - llm_provider=litellm.LlmProviders.CLARIFAI, - params={"timeout": 600.0}, - ) - response = await async_handler.post( - url=model, headers=headers, data=json.dumps(data) - ) - - logging_obj.post_call( - input=prompt, - api_key=api_key, - original_response=response.text, - additional_args={"complete_input_dict": data}, - ) - ## RESPONSE OBJECT - try: - completion_response = response.json() - except Exception: - raise ClarifaiError( - message=response.text, status_code=response.status_code, url=model - ) - # print(completion_response) - try: - choices_list = [] - for idx, item in enumerate(completion_response["outputs"]): - if len(item["data"]["text"]["raw"]) > 0: - message_obj = Message(content=item["data"]["text"]["raw"]) - else: - message_obj = Message(content=None) - choice_obj = Choices( - finish_reason="stop", - index=idx + 1, # check - message=message_obj, - ) - choices_list.append(choice_obj) - model_response.choices = choices_list # type: ignore - - except Exception: - raise ClarifaiError( - message=traceback.format_exc(), status_code=response.status_code, url=model - ) - - # Calculate Usage - prompt_tokens = len(encoding.encode(prompt)) - completion_tokens = len( - encoding.encode(model_response["choices"][0]["message"].get("content")) - ) - model_response.model = model - setattr( - model_response, - "usage", - Usage( - prompt_tokens=prompt_tokens, - completion_tokens=completion_tokens, - total_tokens=prompt_tokens + completion_tokens, - ), - ) - return model_response - - -def completion( - model: str, - messages: list, - api_base: str, - model_response: ModelResponse, - print_verbose: Callable, - encoding, - api_key, - logging_obj, - optional_params: dict, - custom_prompt_dict={}, - acompletion=False, - litellm_params=None, - logger_fn=None, -): - headers = validate_environment(api_key) - model = convert_model_to_url(model, api_base) - prompt = " ".join(message["content"] for message in messages) # TODO - - ## Load Config - config = litellm.ClarifaiConfig.get_config() - for k, v in config.items(): - if k not in optional_params: - optional_params[k] = v - - custom_llm_provider, orig_model_name = get_prompt_model_name(model) - prompt: str = prompt_factory( # type: ignore - model=orig_model_name, - messages=messages, - api_key=api_key, - custom_llm_provider="clarifai", - ) - # print(prompt); exit(0) - - data = { - "prompt": prompt, - **optional_params, - } - data = completions_to_model(data) - - ## LOGGING - logging_obj.pre_call( - input=prompt, - api_key=api_key, - additional_args={ - "complete_input_dict": data, - "headers": headers, - "api_base": model, - }, - ) - if acompletion is True: - return async_completion( - model=model, - prompt=prompt, - api_base=api_base, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - api_key=api_key, - logging_obj=logging_obj, - data=data, - optional_params=optional_params, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=headers, - ) - else: - ## COMPLETION CALL - response = requests.post( - model, - headers=headers, - data=json.dumps(data), - ) - # print(response.content); exit() - - if response.status_code != 200: - raise ClarifaiError( - status_code=response.status_code, message=response.text, url=model - ) - - if "stream" in optional_params and optional_params["stream"] is True: - completion_stream = response.iter_lines() - stream_response = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="clarifai", - logging_obj=logging_obj, - ) - return stream_response - - else: - return process_response( - model=model, - prompt=prompt, - response=response, - model_response=model_response, - api_key=api_key, - data=data, - encoding=encoding, - logging_obj=logging_obj, - ) - - -class ModelResponseIterator: - def __init__(self, model_response): - self.model_response = model_response - self.is_done = False - - # Sync iterator - def __iter__(self): - return self - - def __next__(self): - if self.is_done: - raise StopIteration - self.is_done = True - return self.model_response - - # Async iterator - def __aiter__(self): - return self - - async def __anext__(self): - if self.is_done: - raise StopAsyncIteration - self.is_done = True - return self.model_response diff --git a/litellm/llms/clarifai/chat/handler.py b/litellm/llms/clarifai/chat/handler.py new file mode 100644 index 00000000000..cf6da51cfdd --- /dev/null +++ b/litellm/llms/clarifai/chat/handler.py @@ -0,0 +1,177 @@ +import json +import os +import time +import traceback +import types +from typing import Callable, List, Optional + +import httpx +import requests + +import litellm +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + _get_httpx_client, + get_async_httpx_client, +) +from litellm.types.llms.openai import AllMessageValues +from litellm.utils import Choices, CustomStreamWrapper, Message, ModelResponse, Usage + +from ...prompt_templates.factory import custom_prompt, prompt_factory +from ..common_utils import ClarifaiError + + +async def async_completion( + model: str, + messages: List[AllMessageValues], + model_response: ModelResponse, + encoding, + api_key, + api_base: str, + logging_obj, + data: dict, + optional_params: dict, + litellm_params=None, + logger_fn=None, + headers={}, +): + + async_handler = get_async_httpx_client( + llm_provider=litellm.LlmProviders.CLARIFAI, + params={"timeout": 600.0}, + ) + response = await async_handler.post( + url=api_base, headers=headers, data=json.dumps(data) + ) + + return litellm.ClarifaiConfig().transform_response( + model=model, + raw_response=response, + model_response=model_response, + logging_obj=logging_obj, + api_key=api_key, + request_data=data, + messages=messages, + optional_params=optional_params, + encoding=encoding, + ) + + +def completion( + model: str, + messages: list, + api_base: str, + model_response: ModelResponse, + print_verbose: Callable, + encoding, + api_key, + logging_obj, + optional_params: dict, + litellm_params: dict, + custom_prompt_dict={}, + acompletion=False, + logger_fn=None, + headers={}, +): + headers = litellm.ClarifaiConfig().validate_environment( + api_key=api_key, + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + ) + data = litellm.ClarifaiConfig().transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers, + ) + + ## LOGGING + logging_obj.pre_call( + input=data, + api_key=api_key, + additional_args={ + "complete_input_dict": data, + "headers": headers, + "api_base": model, + }, + ) + if acompletion is True: + return async_completion( + model=model, + messages=messages, + api_base=api_base, + model_response=model_response, + encoding=encoding, + api_key=api_key, + logging_obj=logging_obj, + data=data, + optional_params=optional_params, + litellm_params=litellm_params, + logger_fn=logger_fn, + headers=headers, + ) + else: + ## COMPLETION CALL + httpx_client = _get_httpx_client( + params={"timeout": 600.0}, + ) + response = httpx_client.post( + url=api_base, + headers=headers, + data=json.dumps(data), + ) + + if response.status_code != 200: + raise ClarifaiError(status_code=response.status_code, message=response.text) + + if "stream" in optional_params and optional_params["stream"] is True: + completion_stream = response.iter_lines() + stream_response = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="clarifai", + logging_obj=logging_obj, + ) + return stream_response + + else: + return litellm.ClarifaiConfig().transform_response( + model=model, + raw_response=response, + model_response=model_response, + logging_obj=logging_obj, + api_key=api_key, + request_data=data, + messages=messages, + optional_params=optional_params, + encoding=encoding, + ) + + +class ModelResponseIterator: + def __init__(self, model_response): + self.model_response = model_response + self.is_done = False + + # Sync iterator + def __iter__(self): + return self + + def __next__(self): + if self.is_done: + raise StopIteration + self.is_done = True + return self.model_response + + # Async iterator + def __aiter__(self): + return self + + async def __anext__(self): + if self.is_done: + raise StopAsyncIteration + self.is_done = True + return self.model_response diff --git a/litellm/llms/clarifai/chat/transformation.py b/litellm/llms/clarifai/chat/transformation.py new file mode 100644 index 00000000000..7c1dbaaa090 --- /dev/null +++ b/litellm/llms/clarifai/chat/transformation.py @@ -0,0 +1,201 @@ +import types +from typing import TYPE_CHECKING, Any, List, Optional + +import httpx + +import litellm +from litellm.llms.base_llm.transformation import BaseConfig, BaseLLMException +from litellm.llms.prompt_templates.common_utils import convert_content_list_to_str +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import Choices, Message, ModelResponse, Usage +from litellm.utils import token_counter + +from ..common_utils import ClarifaiError + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + LoggingClass = LiteLLMLoggingObj +else: + LoggingClass = Any + + +class ClarifaiConfig(BaseConfig): + """ + Reference: https://clarifai.com/meta/Llama-2/models/llama2-70b-chat + """ + + max_tokens: Optional[int] = None + temperature: Optional[int] = None + top_k: Optional[int] = None + + def __init__( + self, + max_tokens: Optional[int] = None, + temperature: Optional[int] = None, + top_k: Optional[int] = None, + ) -> None: + locals_ = locals() + for key, value in locals_.items(): + if key != "self" and value is not None: + setattr(self.__class__, key, value) + + @classmethod + def get_config(cls): + return super().get_config() + + def get_supported_openai_params(self, model: str) -> list: + return [ + "temperature", + "max_tokens", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + for param, value in non_default_params.items(): + if param == "temperature": + optional_params["temperature"] = value + elif param == "max_tokens": + optional_params["max_tokens"] = value + + return optional_params + + def _completions_to_model(self, prompt: str, optional_params: dict) -> dict: + params = {} + if temperature := optional_params.get("temperature"): + params["temperature"] = temperature + if max_tokens := optional_params.get("max_tokens"): + params["max_tokens"] = max_tokens + return { + "inputs": [{"data": {"text": {"raw": prompt}}}], + "model": {"output_info": {"params": params}}, + } + + def _convert_model_to_url(self, model: str, api_base: str): + user_id, app_id, model_id = model.split(".") + return f"{api_base}/users/{user_id}/apps/{app_id}/models/{model_id}/outputs" + + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + prompt = " ".join(convert_content_list_to_str(message) for message in messages) + + ## Load Config + config = self.get_config() + for k, v in config.items(): + if k not in optional_params: + optional_params[k] = v + + data = self._completions_to_model( + prompt=prompt, optional_params=optional_params + ) + + return data + + def validate_environment( + self, + api_key: str, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + ) -> dict: + headers = { + "accept": "application/json", + "content-type": "application/json", + } + + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + return headers + + def _transform_messages( + self, messages: List[AllMessageValues] + ) -> List[AllMessageValues]: + raise NotImplementedError + + def get_error_class( + self, error_message: str, status_code: int, headers: dict + ) -> BaseLLMException: + return ClarifaiError(message=error_message, status_code=status_code) + + def transform_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ModelResponse, + logging_obj: LoggingClass, + api_key: str, + request_data: dict, + messages: List[AllMessageValues], + optional_params: dict, + encoding: str, + json_mode: Optional[bool] = None, + ) -> litellm.ModelResponse: + logging_obj.post_call( + input=messages, + api_key=api_key, + original_response=raw_response.text, + additional_args={"complete_input_dict": request_data}, + ) + ## RESPONSE OBJECT + try: + completion_response = raw_response.json() + except httpx.HTTPStatusError as e: + raise ClarifaiError( + message=str(e), + status_code=raw_response.status_code, + ) + except Exception as e: + raise ClarifaiError( + message=str(e), + status_code=422, + ) + # print(completion_response) + try: + choices_list = [] + for idx, item in enumerate(completion_response["outputs"]): + if len(item["data"]["text"]["raw"]) > 0: + message_obj = Message(content=item["data"]["text"]["raw"]) + else: + message_obj = Message(content=None) + choice_obj = Choices( + finish_reason="stop", + index=idx + 1, # check + message=message_obj, + ) + choices_list.append(choice_obj) + model_response.choices = choices_list # type: ignore + + except Exception as e: + raise ClarifaiError( + message=str(e), + status_code=422, + ) + + # Calculate Usage + prompt_tokens = token_counter(model=model, messages=messages) + completion_tokens = len( + encoding.encode(model_response["choices"][0]["message"].get("content")) + ) + model_response.model = model + setattr( + model_response, + "usage", + Usage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + ), + ) + return model_response diff --git a/litellm/llms/clarifai/common_utils.py b/litellm/llms/clarifai/common_utils.py new file mode 100644 index 00000000000..0f249a07204 --- /dev/null +++ b/litellm/llms/clarifai/common_utils.py @@ -0,0 +1,8 @@ +import httpx + +from litellm.llms.base_llm.transformation import BaseLLMException + + +class ClarifaiError(BaseLLMException): + def __init__(self, status_code: int, message: str): + super().__init__(status_code=status_code, message=message) diff --git a/litellm/llms/cohere/chat.py b/litellm/llms/cohere/chat/handler.py similarity index 74% rename from litellm/llms/cohere/chat.py rename to litellm/llms/cohere/chat/handler.py index d7dfc3eaae1..c5ac4548345 100644 --- a/litellm/llms/cohere/chat.py +++ b/litellm/llms/cohere/chat/handler.py @@ -10,9 +10,12 @@ import httpx # type: ignore import requests # type: ignore import litellm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.llms.base_llm.transformation import BaseConfig, BaseLLMException from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.llms.cohere import ToolResultObject +from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ( ChatCompletionToolCallChunk, ChatCompletionUsageBlock, @@ -20,120 +23,8 @@ from litellm.types.utils import ( ) from litellm.utils import Choices, Message, ModelResponse, Usage -from ..prompt_templates.factory import cohere_message_pt, cohere_messages_pt_v2 - - -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/chat") - 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 - - -class CohereChatConfig: - """ - Configuration class for Cohere's API interface. - - Args: - preamble (str, optional): When specified, the default Cohere preamble will be replaced with the provided one. - chat_history (List[Dict[str, str]], optional): A list of previous messages between the user and the model. - generation_id (str, optional): Unique identifier for the generated reply. - response_id (str, optional): Unique identifier for the response. - conversation_id (str, optional): An alternative to chat_history, creates or resumes a persisted conversation. - prompt_truncation (str, optional): Dictates how the prompt will be constructed. Options: 'AUTO', 'AUTO_PRESERVE_ORDER', 'OFF'. - connectors (List[Dict[str, str]], optional): List of connectors (e.g., web-search) to enrich the model's reply. - search_queries_only (bool, optional): When true, the response will only contain a list of generated search queries. - documents (List[Dict[str, str]], optional): A list of relevant documents that the model can cite. - temperature (float, optional): A non-negative float that tunes the degree of randomness in generation. - max_tokens (int, optional): The maximum number of tokens the model will generate as part of the response. - k (int, optional): Ensures only the top k most likely tokens are considered for generation at each step. - p (float, optional): Ensures that only the most likely tokens, with total probability mass of p, are considered for generation. - frequency_penalty (float, optional): Used to reduce repetitiveness of generated tokens. - presence_penalty (float, optional): Used to reduce repetitiveness of generated tokens. - tools (List[Dict[str, str]], optional): A list of available tools (functions) that the model may suggest invoking. - tool_results (List[Dict[str, Any]], optional): A list of results from invoking tools. - seed (int, optional): A seed to assist reproducibility of the model's response. - """ - - preamble: Optional[str] = None - chat_history: Optional[list] = None - generation_id: Optional[str] = None - response_id: Optional[str] = None - conversation_id: Optional[str] = None - prompt_truncation: Optional[str] = None - connectors: Optional[list] = None - search_queries_only: Optional[bool] = None - documents: Optional[list] = None - temperature: Optional[int] = None - max_tokens: Optional[int] = None - k: Optional[int] = None - p: Optional[int] = None - frequency_penalty: Optional[int] = None - presence_penalty: Optional[int] = None - tools: Optional[list] = None - tool_results: Optional[list] = None - seed: Optional[int] = None - - def __init__( - self, - preamble: Optional[str] = None, - chat_history: Optional[list] = None, - generation_id: Optional[str] = None, - response_id: Optional[str] = None, - conversation_id: Optional[str] = None, - prompt_truncation: Optional[str] = None, - connectors: Optional[list] = None, - search_queries_only: Optional[bool] = None, - documents: Optional[list] = None, - temperature: Optional[int] = None, - max_tokens: Optional[int] = None, - k: Optional[int] = None, - p: Optional[int] = None, - frequency_penalty: Optional[int] = None, - presence_penalty: Optional[int] = None, - tools: Optional[list] = None, - tool_results: Optional[list] = None, - seed: Optional[int] = None, - ) -> None: - locals_ = locals() - for key, value in locals_.items(): - if key != "self" and value is not None: - setattr(self.__class__, key, value) - - @classmethod - def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } - - -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 +from ...prompt_templates.factory import cohere_message_pt, cohere_messages_pt_v2 +from .transformation import CohereChatConfig, CohereError def translate_openai_tool_to_cohere(openai_tool): @@ -321,7 +212,13 @@ def completion( # noqa: PLR0915 client=None, timeout=None, ): - headers = validate_environment(api_key, headers=headers) + headers = litellm.CohereChatConfig().validate_environment( + api_key=api_key, + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + ) completion_url = api_base model = model most_recent_message, chat_history = cohere_messages_pt_v2( diff --git a/litellm/llms/cohere/chat/transformation.py b/litellm/llms/cohere/chat/transformation.py new file mode 100644 index 00000000000..d74b2110f18 --- /dev/null +++ b/litellm/llms/cohere/chat/transformation.py @@ -0,0 +1,184 @@ +import types +from typing import TYPE_CHECKING, Any, List, Optional + +import httpx + +from litellm.llms.base_llm.transformation import BaseConfig, BaseLLMException +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import ModelResponse + +from ..common_utils import CohereError +from ..common_utils import validate_environment as cohere_validate_environment + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + LoggingObj = LiteLLMLoggingObj +else: + LoggingObj = Any + + +class CohereChatConfig(BaseConfig): + """ + Configuration class for Cohere's API interface. + + Args: + preamble (str, optional): When specified, the default Cohere preamble will be replaced with the provided one. + chat_history (List[Dict[str, str]], optional): A list of previous messages between the user and the model. + generation_id (str, optional): Unique identifier for the generated reply. + response_id (str, optional): Unique identifier for the response. + conversation_id (str, optional): An alternative to chat_history, creates or resumes a persisted conversation. + prompt_truncation (str, optional): Dictates how the prompt will be constructed. Options: 'AUTO', 'AUTO_PRESERVE_ORDER', 'OFF'. + connectors (List[Dict[str, str]], optional): List of connectors (e.g., web-search) to enrich the model's reply. + search_queries_only (bool, optional): When true, the response will only contain a list of generated search queries. + documents (List[Dict[str, str]], optional): A list of relevant documents that the model can cite. + temperature (float, optional): A non-negative float that tunes the degree of randomness in generation. + max_tokens (int, optional): The maximum number of tokens the model will generate as part of the response. + k (int, optional): Ensures only the top k most likely tokens are considered for generation at each step. + p (float, optional): Ensures that only the most likely tokens, with total probability mass of p, are considered for generation. + frequency_penalty (float, optional): Used to reduce repetitiveness of generated tokens. + presence_penalty (float, optional): Used to reduce repetitiveness of generated tokens. + tools (List[Dict[str, str]], optional): A list of available tools (functions) that the model may suggest invoking. + tool_results (List[Dict[str, Any]], optional): A list of results from invoking tools. + seed (int, optional): A seed to assist reproducibility of the model's response. + """ + + preamble: Optional[str] = None + chat_history: Optional[list] = None + generation_id: Optional[str] = None + response_id: Optional[str] = None + conversation_id: Optional[str] = None + prompt_truncation: Optional[str] = None + connectors: Optional[list] = None + search_queries_only: Optional[bool] = None + documents: Optional[list] = None + temperature: Optional[int] = None + max_tokens: Optional[int] = None + k: Optional[int] = None + p: Optional[int] = None + frequency_penalty: Optional[int] = None + presence_penalty: Optional[int] = None + tools: Optional[list] = None + tool_results: Optional[list] = None + seed: Optional[int] = None + + def __init__( + self, + preamble: Optional[str] = None, + chat_history: Optional[list] = None, + generation_id: Optional[str] = None, + response_id: Optional[str] = None, + conversation_id: Optional[str] = None, + prompt_truncation: Optional[str] = None, + connectors: Optional[list] = None, + search_queries_only: Optional[bool] = None, + documents: Optional[list] = None, + temperature: Optional[int] = None, + max_tokens: Optional[int] = None, + k: Optional[int] = None, + p: Optional[int] = None, + frequency_penalty: Optional[int] = None, + presence_penalty: Optional[int] = None, + tools: Optional[list] = None, + tool_results: Optional[list] = None, + seed: Optional[int] = None, + ) -> None: + locals_ = locals() + for key, value in locals_.items(): + if key != "self" and value is not None: + setattr(self.__class__, key, value) + + @classmethod + def get_config(cls): + return super().get_config() + + def _transform_messages( + self, messages: List[AllMessageValues] + ) -> List[AllMessageValues]: + raise NotImplementedError + + def get_error_class( + self, error_message: str, status_code: int, headers: dict + ) -> BaseLLMException: + return CohereError(status_code=status_code, message=error_message) + + def get_supported_openai_params(self, model: str) -> List[str]: + return [ + "stream", + "temperature", + "max_tokens", + "top_p", + "frequency_penalty", + "presence_penalty", + "stop", + "n", + "tools", + "tool_choice", + "seed", + "extra_headers", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + for param, value in non_default_params.items(): + if param == "stream": + optional_params["stream"] = value + if param == "temperature": + optional_params["temperature"] = value + if param == "max_tokens": + optional_params["max_tokens"] = value + if param == "n": + optional_params["num_generations"] = value + if param == "top_p": + optional_params["p"] = value + if param == "frequency_penalty": + optional_params["frequency_penalty"] = value + if param == "presence_penalty": + optional_params["presence_penalty"] = value + if param == "stop": + optional_params["stop_sequences"] = value + if param == "tools": + optional_params["tools"] = value + if param == "seed": + optional_params["seed"] = value + return optional_params + + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + raise NotImplementedError + + def transform_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ModelResponse, + logging_obj: LoggingObj, + api_key: str, + request_data: dict, + messages: List[AllMessageValues], + optional_params: dict, + encoding: Any, + json_mode: Optional[bool] = None, + ) -> ModelResponse: + raise NotImplementedError + + def validate_environment( + self, + api_key: str, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + ) -> dict: + return cohere_validate_environment(api_key=api_key, headers=headers) diff --git a/litellm/llms/cohere/common_utils.py b/litellm/llms/cohere/common_utils.py new file mode 100644 index 00000000000..12c14977eba --- /dev/null +++ b/litellm/llms/cohere/common_utils.py @@ -0,0 +1,19 @@ +from litellm.llms.base_llm.transformation import BaseLLMException + + +class CohereError(BaseLLMException): + def __init__(self, status_code, message): + super().__init__(status_code=status_code, message=message) + + +def validate_environment(*, api_key: str, headers: dict) -> 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 diff --git a/litellm/llms/cohere/completion.py b/litellm/llms/cohere/completion/completion.py similarity index 52% rename from litellm/llms/cohere/completion.py rename to litellm/llms/cohere/completion/completion.py index 4743996247c..77ed5cc83c5 100644 --- a/litellm/llms/cohere/completion.py +++ b/litellm/llms/cohere/completion/completion.py @@ -16,18 +16,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.utils import Choices, Message, ModelResponse, Usage - -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 +from ..common_utils import CohereError def construct_cohere_tool(tools=None): @@ -36,93 +25,6 @@ def construct_cohere_tool(tools=None): return {"tools": tools} -class CohereConfig: - """ - Reference: https://docs.cohere.com/reference/generate - - The class `CohereConfig` provides configuration for the Cohere's API interface. Below are the parameters: - - - `num_generations` (integer): Maximum number of generations returned. Default is 1, with a minimum value of 1 and a maximum value of 5. - - - `max_tokens` (integer): Maximum number of tokens the model will generate as part of the response. Default value is 20. - - - `truncate` (string): Specifies how the API handles inputs longer than maximum token length. Options include NONE, START, END. Default is END. - - - `temperature` (number): A non-negative float controlling the randomness in generation. Lower temperatures result in less random generations. Default is 0.75. - - - `preset` (string): Identifier of a custom preset, a combination of parameters such as prompt, temperature etc. - - - `end_sequences` (array of strings): The generated text gets cut at the beginning of the earliest occurrence of an end sequence, which will be excluded from the text. - - - `stop_sequences` (array of strings): The generated text gets cut at the end of the earliest occurrence of a stop sequence, which will be included in the text. - - - `k` (integer): Limits generation at each step to top `k` most likely tokens. Default is 0. - - - `p` (number): Limits generation at each step to most likely tokens with total probability mass of `p`. Default is 0. - - - `frequency_penalty` (number): Reduces repetitiveness of generated tokens. Higher values apply stronger penalties to previously occurred tokens. - - - `presence_penalty` (number): Reduces repetitiveness of generated tokens. Similar to frequency_penalty, but this penalty applies equally to all tokens that have already appeared. - - - `return_likelihoods` (string): Specifies how and if token likelihoods are returned with the response. Options include GENERATION, ALL and NONE. - - - `logit_bias` (object): Used to prevent the model from generating unwanted tokens or to incentivize it to include desired tokens. e.g. {"hello_world": 1233} - """ - - num_generations: Optional[int] = None - max_tokens: Optional[int] = None - truncate: Optional[str] = None - temperature: Optional[int] = None - preset: Optional[str] = None - end_sequences: Optional[list] = None - stop_sequences: Optional[list] = None - k: Optional[int] = None - p: Optional[int] = None - frequency_penalty: Optional[int] = None - presence_penalty: Optional[int] = None - return_likelihoods: Optional[str] = None - logit_bias: Optional[dict] = None - - def __init__( - self, - num_generations: Optional[int] = None, - max_tokens: Optional[int] = None, - truncate: Optional[str] = None, - temperature: Optional[int] = None, - preset: Optional[str] = None, - end_sequences: Optional[list] = None, - stop_sequences: Optional[list] = None, - k: Optional[int] = None, - p: Optional[int] = None, - frequency_penalty: Optional[int] = None, - presence_penalty: Optional[int] = None, - return_likelihoods: Optional[str] = None, - logit_bias: Optional[dict] = None, - ) -> None: - locals_ = locals() - for key, value in locals_.items(): - if key != "self" and value is not None: - setattr(self.__class__, key, value) - - @classmethod - def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } - - def validate_environment(api_key, headers: dict): headers.update( { diff --git a/litellm/llms/cohere/completion/transformation.py b/litellm/llms/cohere/completion/transformation.py new file mode 100644 index 00000000000..d5dafb69b41 --- /dev/null +++ b/litellm/llms/cohere/completion/transformation.py @@ -0,0 +1,183 @@ +import types +from typing import TYPE_CHECKING, Any, List, Optional + +import httpx + +from litellm.llms.base_llm.transformation import ( + BaseConfig, + BaseLLMException, + LoggingClass, +) +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import ModelResponse + +from ..common_utils import CohereError +from ..common_utils import validate_environment as cohere_validate_environment + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + LoggingObj = LiteLLMLoggingObj +else: + LoggingObj = Any + + +class CohereTextConfig(BaseConfig): + """ + Reference: https://docs.cohere.com/reference/generate + + The class `CohereConfig` provides configuration for the Cohere's API interface. Below are the parameters: + + - `num_generations` (integer): Maximum number of generations returned. Default is 1, with a minimum value of 1 and a maximum value of 5. + + - `max_tokens` (integer): Maximum number of tokens the model will generate as part of the response. Default value is 20. + + - `truncate` (string): Specifies how the API handles inputs longer than maximum token length. Options include NONE, START, END. Default is END. + + - `temperature` (number): A non-negative float controlling the randomness in generation. Lower temperatures result in less random generations. Default is 0.75. + + - `preset` (string): Identifier of a custom preset, a combination of parameters such as prompt, temperature etc. + + - `end_sequences` (array of strings): The generated text gets cut at the beginning of the earliest occurrence of an end sequence, which will be excluded from the text. + + - `stop_sequences` (array of strings): The generated text gets cut at the end of the earliest occurrence of a stop sequence, which will be included in the text. + + - `k` (integer): Limits generation at each step to top `k` most likely tokens. Default is 0. + + - `p` (number): Limits generation at each step to most likely tokens with total probability mass of `p`. Default is 0. + + - `frequency_penalty` (number): Reduces repetitiveness of generated tokens. Higher values apply stronger penalties to previously occurred tokens. + + - `presence_penalty` (number): Reduces repetitiveness of generated tokens. Similar to frequency_penalty, but this penalty applies equally to all tokens that have already appeared. + + - `return_likelihoods` (string): Specifies how and if token likelihoods are returned with the response. Options include GENERATION, ALL and NONE. + + - `logit_bias` (object): Used to prevent the model from generating unwanted tokens or to incentivize it to include desired tokens. e.g. {"hello_world": 1233} + """ + + num_generations: Optional[int] = None + max_tokens: Optional[int] = None + truncate: Optional[str] = None + temperature: Optional[int] = None + preset: Optional[str] = None + end_sequences: Optional[list] = None + stop_sequences: Optional[list] = None + k: Optional[int] = None + p: Optional[int] = None + frequency_penalty: Optional[int] = None + presence_penalty: Optional[int] = None + return_likelihoods: Optional[str] = None + logit_bias: Optional[dict] = None + + def __init__( + self, + num_generations: Optional[int] = None, + max_tokens: Optional[int] = None, + truncate: Optional[str] = None, + temperature: Optional[int] = None, + preset: Optional[str] = None, + end_sequences: Optional[list] = None, + stop_sequences: Optional[list] = None, + k: Optional[int] = None, + p: Optional[int] = None, + frequency_penalty: Optional[int] = None, + presence_penalty: Optional[int] = None, + return_likelihoods: Optional[str] = None, + logit_bias: Optional[dict] = None, + ) -> None: + locals_ = locals() + for key, value in locals_.items(): + if key != "self" and value is not None: + setattr(self.__class__, key, value) + + @classmethod + def get_config(cls): + return super().get_config() + + def validate_environment( + self, + api_key: str, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + ) -> dict: + return cohere_validate_environment(api_key=api_key, headers=headers) + + def _transform_messages( + self, + messages: List[AllMessageValues], + ) -> List[AllMessageValues]: + raise NotImplementedError + + def get_error_class( + self, error_message: str, status_code: int, headers: dict + ) -> BaseLLMException: + return CohereError(status_code=status_code, message=error_message) + + def get_supported_openai_params(self, model: str) -> List: + return [ + "stream", + "temperature", + "max_tokens", + "logit_bias", + "top_p", + "frequency_penalty", + "presence_penalty", + "stop", + "n", + "extra_headers", + ] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + for param, value in non_default_params.items(): + if param == "stream": + optional_params["stream"] = value + elif param == "temperature": + optional_params["temperature"] = value + elif param == "max_tokens": + optional_params["max_tokens"] = value + elif param == "n": + optional_params["num_generations"] = value + elif param == "logit_bias": + optional_params["logit_bias"] = value + elif param == "top_p": + optional_params["p"] = value + elif param == "frequency_penalty": + optional_params["frequency_penalty"] = value + elif param == "presence_penalty": + optional_params["presence_penalty"] = value + elif param == "stop": + optional_params["stop_sequences"] = value + return optional_params + + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + raise NotImplementedError + + def transform_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ModelResponse, + logging_obj: LoggingObj, + api_key: str, + request_data: dict, + messages: List[AllMessageValues], + optional_params: dict, + encoding: Any, + json_mode: Optional[bool] = None, + ) -> ModelResponse: + raise NotImplementedError diff --git a/litellm/llms/databricks/chat/old_handler.py b/litellm/llms/databricks/chat/old_handler.py deleted file mode 100644 index 95cc1cfc6d0..00000000000 --- a/litellm/llms/databricks/chat/old_handler.py +++ /dev/null @@ -1,611 +0,0 @@ -# What is this? -## Handler file for databricks API https://docs.databricks.com/en/machine-learning/foundation-models/api-reference.html#chat-request -import copy -import json -import os -import time -import types -from enum import Enum -from functools import partial -from typing import Any, Callable, List, Literal, Optional, Tuple, Union - -import httpx # type: ignore -import requests # type: ignore - -import litellm -from litellm import LlmProviders -from litellm.litellm_core_utils.core_helpers import map_finish_reason -from litellm.llms.custom_httpx.http_handler import ( - AsyncHTTPHandler, - HTTPHandler, - get_async_httpx_client, -) -from litellm.llms.databricks.exceptions import DatabricksError -from litellm.llms.databricks.streaming_utils import ModelResponseIterator -from litellm.types.llms.openai import ( - ChatCompletionDeltaChunk, - ChatCompletionResponseMessage, - ChatCompletionToolCallChunk, - ChatCompletionToolCallFunctionChunk, - ChatCompletionUsageBlock, -) -from litellm.types.utils import ( - CustomStreamingDecoder, - GenericStreamingChunk, - ProviderField, -) -from litellm.utils import ( - CustomStreamWrapper, - EmbeddingResponse, - ModelResponse, - ProviderConfigManager, - Usage, -) - -from ...base import BaseLLM -from ...prompt_templates.factory import custom_prompt, prompt_factory -from .transformation import DatabricksConfig - - -async def make_call( - client: Optional[AsyncHTTPHandler], - api_base: str, - headers: dict, - data: str, - model: str, - messages: list, - logging_obj, - streaming_decoder: Optional[CustomStreamingDecoder] = None, -): - if client is None: - client = get_async_httpx_client( - llm_provider=litellm.LlmProviders.DATABRICKS - ) # Create a new client if none provided - response = await client.post(api_base, headers=headers, data=data, stream=True) - - if response.status_code != 200: - raise DatabricksError(status_code=response.status_code, message=response.text) - - if streaming_decoder is not None: - completion_stream: Any = streaming_decoder.aiter_bytes( - response.aiter_bytes(chunk_size=1024) - ) - else: - completion_stream = ModelResponseIterator( - streaming_response=response.aiter_lines(), sync_stream=False - ) - # LOGGING - logging_obj.post_call( - input=messages, - api_key="", - original_response=completion_stream, # Pass the completion stream for logging - additional_args={"complete_input_dict": data}, - ) - - return completion_stream - - -def make_sync_call( - client: Optional[HTTPHandler], - api_base: str, - headers: dict, - data: str, - model: str, - messages: list, - logging_obj, - streaming_decoder: Optional[CustomStreamingDecoder] = None, -): - if client is None: - client = litellm.module_level_client # Create a new client if none provided - - response = client.post(api_base, headers=headers, data=data, stream=True) - - if response.status_code != 200: - raise DatabricksError(status_code=response.status_code, message=response.read()) - - if streaming_decoder is not None: - completion_stream = streaming_decoder.iter_bytes( - response.iter_bytes(chunk_size=1024) - ) - else: - completion_stream = ModelResponseIterator( - streaming_response=response.iter_lines(), sync_stream=True - ) - - # LOGGING - logging_obj.post_call( - input=messages, - api_key="", - original_response="first stream response received", - additional_args={"complete_input_dict": data}, - ) - - return completion_stream - - -class DatabricksChatCompletion(BaseLLM): - def __init__(self) -> None: - super().__init__() - - # makes headers for API call - def _get_databricks_credentials( - self, api_key: Optional[str], api_base: Optional[str], headers: Optional[dict] - ) -> Tuple[str, dict]: - headers = headers or {"Content-Type": "application/json"} - try: - from databricks.sdk import WorkspaceClient - - databricks_client = WorkspaceClient() - api_base = api_base or f"{databricks_client.config.host}/serving-endpoints" - - if api_key is None: - databricks_auth_headers: dict[str, str] = ( - databricks_client.config.authenticate() - ) - headers = {**databricks_auth_headers, **headers} - - return api_base, headers - except ImportError: - raise DatabricksError( - status_code=400, - message=( - "If the Databricks base URL and API key are not set, the databricks-sdk " - "Python library must be installed. Please install the databricks-sdk, set " - "{LLM_PROVIDER}_API_BASE and {LLM_PROVIDER}_API_KEY environment variables, " - "or provide the base URL and API key as arguments." - ), - ) - - def _validate_environment( - self, - api_key: Optional[str], - api_base: Optional[str], - endpoint_type: Literal["chat_completions", "embeddings"], - custom_endpoint: Optional[bool], - headers: Optional[dict], - ) -> Tuple[str, dict]: - if api_key is None and headers is None: - if custom_endpoint: - raise DatabricksError( - status_code=400, - message="Missing API Key - A call is being made to LLM Provider but no key is set either in the environment variables ({LLM_PROVIDER}_API_KEY) or via params", - ) - else: - api_base, headers = self._get_databricks_credentials( - api_base=api_base, api_key=api_key, headers=headers - ) - - if api_base is None: - if custom_endpoint: - raise DatabricksError( - status_code=400, - message="Missing API Base - A call is being made to LLM Provider but no api base is set either in the environment variables ({LLM_PROVIDER}_API_KEY) or via params", - ) - else: - api_base, headers = self._get_databricks_credentials( - api_base=api_base, api_key=api_key, headers=headers - ) - - if headers is None: - headers = { - "Authorization": "Bearer {}".format(api_key), - "Content-Type": "application/json", - } - else: - if api_key is not None: - headers.update({"Authorization": "Bearer {}".format(api_key)}) - - if api_key is not None: - headers["Authorization"] = f"Bearer {api_key}" - - if endpoint_type == "chat_completions" and custom_endpoint is not True: - api_base = "{}/chat/completions".format(api_base) - elif endpoint_type == "embeddings" and custom_endpoint is not True: - api_base = "{}/embeddings".format(api_base) - return api_base, headers - - async def acompletion_stream_function( - self, - model: str, - messages: list, - custom_llm_provider: str, - api_base: str, - custom_prompt_dict: dict, - model_response: ModelResponse, - print_verbose: Callable, - encoding, - api_key, - logging_obj, - stream, - data: dict, - optional_params=None, - litellm_params=None, - logger_fn=None, - headers={}, - client: Optional[AsyncHTTPHandler] = None, - streaming_decoder: Optional[CustomStreamingDecoder] = None, - ) -> CustomStreamWrapper: - - data["stream"] = True - completion_stream = await make_call( - client=client, - api_base=api_base, - headers=headers, - data=json.dumps(data), - model=model, - messages=messages, - logging_obj=logging_obj, - streaming_decoder=streaming_decoder, - ) - streamwrapper = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider=custom_llm_provider, - logging_obj=logging_obj, - ) - return streamwrapper - - async def acompletion_function( - self, - model: str, - messages: list, - api_base: str, - custom_prompt_dict: dict, - model_response: ModelResponse, - custom_llm_provider: str, - print_verbose: Callable, - encoding, - api_key, - logging_obj, - stream, - data: dict, - base_model: Optional[str], - optional_params: dict, - litellm_params=None, - logger_fn=None, - headers={}, - timeout: Optional[Union[float, httpx.Timeout]] = None, - ) -> ModelResponse: - if timeout is None: - timeout = httpx.Timeout(timeout=600.0, connect=5.0) - - self.async_handler = get_async_httpx_client( - llm_provider=litellm.LlmProviders.DATABRICKS, - params={"timeout": timeout}, - ) - - try: - response = await self.async_handler.post( - api_base, headers=headers, data=json.dumps(data) - ) - response.raise_for_status() - - response_json = response.json() - except httpx.HTTPStatusError as e: - raise DatabricksError( - status_code=e.response.status_code, - message=e.response.text, - ) - except httpx.TimeoutException: - raise DatabricksError(status_code=408, message="Timeout error occurred.") - except Exception as e: - raise DatabricksError(status_code=500, message=str(e)) - - logging_obj.post_call( - input=messages, - api_key="", - original_response=response_json, - additional_args={"complete_input_dict": data}, - ) - response = ModelResponse(**response_json) - - response.model = custom_llm_provider + "/" + (response.model or "") - - if base_model is not None: - response._hidden_params["model"] = base_model - return response - - def completion( - self, - model: str, - messages: list, - api_base: str, - custom_llm_provider: str, - custom_prompt_dict: dict, - model_response: ModelResponse, - print_verbose: Callable, - encoding, - api_key: Optional[str], - logging_obj, - optional_params: dict, - acompletion=None, - litellm_params=None, - logger_fn=None, - headers: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, - custom_endpoint: Optional[bool] = None, - streaming_decoder: Optional[ - CustomStreamingDecoder - ] = None, # if openai-compatible api needs custom stream decoder - e.g. sagemaker - ): - custom_endpoint = custom_endpoint or optional_params.pop( - "custom_endpoint", None - ) - base_model: Optional[str] = optional_params.pop("base_model", None) - api_base, headers = self._validate_environment( - api_base=api_base, - api_key=api_key, - endpoint_type="chat_completions", - custom_endpoint=custom_endpoint, - headers=headers, - ) - ## Load Config - config = litellm.DatabricksConfig().get_config() - for k, v in config.items(): - if ( - k not in optional_params - ): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in - optional_params[k] = v - - stream: bool = optional_params.get("stream", None) or False - optional_params.pop( - "max_retries", None - ) # [TODO] add max retry support at llm api call level - optional_params["stream"] = stream - - if messages is not None and custom_llm_provider is not None: - provider_config = ProviderConfigManager.get_provider_config( - model=model, provider=LlmProviders(custom_llm_provider) - ) - messages = provider_config._transform_messages(messages) - - data = { - "model": model, - "messages": messages, - **optional_params, - } - - ## LOGGING - logging_obj.pre_call( - input=messages, - api_key=api_key, - additional_args={ - "complete_input_dict": data, - "api_base": api_base, - "headers": headers, - }, - ) - if acompletion is True: - if client is not None and isinstance(client, HTTPHandler): - client = None - if ( - stream is not None and stream is True - ): # if function call - fake the streaming (need complete blocks for output parsing in openai format) - print_verbose("makes async anthropic streaming POST request") - data["stream"] = stream - return self.acompletion_stream_function( - model=model, - messages=messages, - data=data, - api_base=api_base, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - api_key=api_key, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=headers, - client=client, - custom_llm_provider=custom_llm_provider, - streaming_decoder=streaming_decoder, - ) - else: - return self.acompletion_function( - model=model, - messages=messages, - data=data, - api_base=api_base, - custom_prompt_dict=custom_prompt_dict, - custom_llm_provider=custom_llm_provider, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - api_key=api_key, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=headers, - timeout=timeout, - base_model=base_model, - ) - else: - ## COMPLETION CALL - if stream is True: - completion_stream = make_sync_call( - client=( - client - if client is not None and isinstance(client, HTTPHandler) - else None - ), - api_base=api_base, - headers=headers, - data=json.dumps(data), - model=model, - messages=messages, - logging_obj=logging_obj, - streaming_decoder=streaming_decoder, - ) - # completion_stream.__iter__() - return CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider=custom_llm_provider, - logging_obj=logging_obj, - ) - else: - if client is None or not isinstance(client, HTTPHandler): - client = HTTPHandler(timeout=timeout) # type: ignore - try: - response = client.post( - api_base, headers=headers, data=json.dumps(data) - ) - response.raise_for_status() - - response_json = response.json() - except httpx.HTTPStatusError as e: - raise DatabricksError( - status_code=e.response.status_code, - message=e.response.text, - ) - except httpx.TimeoutException: - raise DatabricksError( - status_code=408, message="Timeout error occurred." - ) - except Exception as e: - raise DatabricksError(status_code=500, message=str(e)) - - response = ModelResponse(**response_json) - - response.model = custom_llm_provider + "/" + (response.model or "") - - if base_model is not None: - response._hidden_params["model"] = base_model - - return response - - async def aembedding( - self, - input: list, - data: dict, - model_response: ModelResponse, - timeout: float, - api_key: str, - api_base: str, - logging_obj, - headers: dict, - client=None, - ) -> EmbeddingResponse: - response = None - try: - if client is None or isinstance(client, AsyncHTTPHandler): - self.async_client = get_async_httpx_client( - llm_provider=litellm.LlmProviders.DATABRICKS, - params={"timeout": timeout}, - ) - else: - self.async_client = client - - try: - response = await self.async_client.post( - api_base, - headers=headers, - data=json.dumps(data), - ) # type: ignore - - response.raise_for_status() - - response_json = response.json() - except httpx.HTTPStatusError as e: - raise DatabricksError( - status_code=e.response.status_code, - message=response.text if response else str(e), - ) - except httpx.TimeoutException: - raise DatabricksError( - status_code=408, message="Timeout error occurred." - ) - except Exception as e: - raise DatabricksError(status_code=500, message=str(e)) - - ## LOGGING - logging_obj.post_call( - input=input, - api_key=api_key, - additional_args={"complete_input_dict": data}, - original_response=response_json, - ) - return EmbeddingResponse(**response_json) - except Exception as e: - ## LOGGING - logging_obj.post_call( - input=input, - api_key=api_key, - original_response=str(e), - ) - raise e - - def embedding( - self, - model: str, - input: list, - timeout: float, - logging_obj, - api_key: Optional[str], - api_base: Optional[str], - optional_params: dict, - model_response: Optional[litellm.utils.EmbeddingResponse] = None, - client=None, - aembedding=None, - headers: Optional[dict] = None, - ) -> EmbeddingResponse: - api_base, headers = self._validate_environment( - api_base=api_base, - api_key=api_key, - endpoint_type="embeddings", - custom_endpoint=False, - headers=headers, - ) - model = model - data = {"model": model, "input": input, **optional_params} - - ## LOGGING - logging_obj.pre_call( - input=input, - api_key=api_key, - additional_args={"complete_input_dict": data, "api_base": api_base}, - ) - - if aembedding is True: - return self.aembedding(data=data, input=input, logging_obj=logging_obj, model_response=model_response, api_base=api_base, api_key=api_key, timeout=timeout, client=client, headers=headers) # type: ignore - if client is None or isinstance(client, AsyncHTTPHandler): - self.client = HTTPHandler(timeout=timeout) # type: ignore - else: - self.client = client - - ## EMBEDDING CALL - try: - response = self.client.post( - api_base, - headers=headers, - data=json.dumps(data), - ) # type: ignore - - response.raise_for_status() # type: ignore - - response_json = response.json() # type: ignore - except httpx.HTTPStatusError as e: - raise DatabricksError( - status_code=e.response.status_code, - message=e.response.text, - ) - except httpx.TimeoutException: - raise DatabricksError(status_code=408, message="Timeout error occurred.") - except Exception as e: - raise DatabricksError(status_code=500, message=str(e)) - - ## LOGGING - logging_obj.post_call( - input=input, - api_key=api_key, - additional_args={"complete_input_dict": data}, - original_response=response_json, - ) - - return litellm.EmbeddingResponse(**response_json) diff --git a/litellm/llms/groq/chat/handler.py b/litellm/llms/groq/chat/handler.py index 1fe87844c55..a6d6822a5e8 100644 --- a/litellm/llms/groq/chat/handler.py +++ b/litellm/llms/groq/chat/handler.py @@ -40,7 +40,7 @@ class GroqChatCompletion(OpenAILikeChatHandler): client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, custom_endpoint: Optional[bool] = None, streaming_decoder: Optional[CustomStreamingDecoder] = None, - fake_stream: bool = False + fake_stream: bool = False, ): messages = GroqChatConfig()._transform_messages(messages) # type: ignore diff --git a/litellm/llms/groq/chat/transformation.py b/litellm/llms/groq/chat/transformation.py index dddc56a2cbf..a728d6e96d9 100644 --- a/litellm/llms/groq/chat/transformation.py +++ b/litellm/llms/groq/chat/transformation.py @@ -59,21 +59,7 @@ class GroqChatConfig(OpenAIGPTConfig): @classmethod def get_config(cls): - return { - k: v - for k, v in cls.__dict__.items() - if not k.startswith("__") - and not isinstance( - v, - ( - types.FunctionType, - types.BuiltinFunctionType, - classmethod, - staticmethod, - ), - ) - and v is not None - } + return super().get_config() def _transform_messages(self, messages: List[AllMessageValues]) -> List: for idx, message in enumerate(messages): diff --git a/litellm/llms/openai_like/chat/handler.py b/litellm/llms/openai_like/chat/handler.py index baa9703049a..831051a2c28 100644 --- a/litellm/llms/openai_like/chat/handler.py +++ b/litellm/llms/openai_like/chat/handler.py @@ -277,7 +277,7 @@ class OpenAILikeChatHandler(OpenAILikeBase): optional_params["stream"] = stream if messages is not None and custom_llm_provider is not None: - provider_config = ProviderConfigManager.get_provider_config( + provider_config = ProviderConfigManager.get_provider_chat_config( model=model, provider=LlmProviders(custom_llm_provider) ) messages = provider_config._transform_messages(messages) diff --git a/litellm/llms/together_ai/completion/handler.py b/litellm/llms/together_ai/completion/handler.py index fab2a39c571..fac87944733 100644 --- a/litellm/llms/together_ai/completion/handler.py +++ b/litellm/llms/together_ai/completion/handler.py @@ -12,7 +12,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUserMessage from litellm.utils import ModelResponse -from ...OpenAI.openai import OpenAITextCompletion +from ...OpenAI.completion.handler import OpenAITextCompletion from .transformation import TogetherAITextCompletionConfig together_ai_text_completion_global_config = TogetherAITextCompletionConfig() diff --git a/litellm/llms/together_ai/completion/transformation.py b/litellm/llms/together_ai/completion/transformation.py index 65b9ad69bfa..6ec855de8ba 100644 --- a/litellm/llms/together_ai/completion/transformation.py +++ b/litellm/llms/together_ai/completion/transformation.py @@ -15,7 +15,7 @@ from litellm.types.llms.openai import ( OpenAITextCompletionUserMessage, ) -from ...OpenAI.openai import OpenAITextCompletionConfig +from ...OpenAI.completion.transformation import OpenAITextCompletionConfig class TogetherAITextCompletionConfig(OpenAITextCompletionConfig): diff --git a/litellm/llms/vertex_ai_and_google_ai_studio/vertex_ai_partner_models/anthropic/transformation.py b/litellm/llms/vertex_ai_and_google_ai_studio/vertex_ai_partner_models/anthropic/transformation.py index 0c3d3965d1b..882a1aed271 100644 --- a/litellm/llms/vertex_ai_and_google_ai_studio/vertex_ai_partner_models/anthropic/transformation.py +++ b/litellm/llms/vertex_ai_and_google_ai_studio/vertex_ai_partner_models/anthropic/transformation.py @@ -16,6 +16,7 @@ import litellm from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.llms.openai import ( + AllMessageValues, ChatCompletionToolParam, ChatCompletionToolParamFunctionChunk, ) @@ -69,6 +70,25 @@ class VertexAIAnthropicConfig(AnthropicConfig): Note: Please make sure to modify the default parameters as required for your use case. """ + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + data = super().transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers, + ) + + data.pop("model", None) # vertex anthropic doesn't accept 'model' parameter + return data + @classmethod def is_supported_model( cls, model: str, custom_llm_provider: Optional[str] = None diff --git a/litellm/llms/vertex_ai_and_google_ai_studio/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai_and_google_ai_studio/vertex_ai_partner_models/main.py index 5b2ba511a9a..fa2f9f7ff14 100644 --- a/litellm/llms/vertex_ai_and_google_ai_studio/vertex_ai_partner_models/main.py +++ b/litellm/llms/vertex_ai_and_google_ai_studio/vertex_ai_partner_models/main.py @@ -7,6 +7,7 @@ from typing import Callable, Literal, Optional, Union import httpx # type: ignore import litellm +from litellm import LlmProviders from litellm.utils import ModelResponse from ..vertex_llm_base import VertexBase @@ -211,6 +212,7 @@ class VertexAIPartnerModels(VertexBase): headers=headers, timeout=timeout, client=client, + custom_llm_provider=LlmProviders.VERTEX_AI.value, ) return openai_like_chat_completions.completion( diff --git a/litellm/llms/xai/chat/xai_transformation.py b/litellm/llms/xai/chat/transformation.py similarity index 97% rename from litellm/llms/xai/chat/xai_transformation.py rename to litellm/llms/xai/chat/transformation.py index 3bd41ed9073..ac3c4236511 100644 --- a/litellm/llms/xai/chat/xai_transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -22,8 +22,6 @@ class XAIChatConfig(OpenAIGPTConfig): "logit_bias", "logprobs", "max_tokens", - "messages", - "model", "n", "presence_penalty", "response_format", diff --git a/litellm/main.py b/litellm/main.py index f574b9339c9..44b123dd796 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -86,7 +86,6 @@ from .litellm_core_utils.streaming_chunk_builder_utils import ChunkProcessor from .llms import ( aleph_alpha, baseten, - clarifai, cloudflare, maritalk, nlp_cloud, @@ -111,8 +110,9 @@ from .llms.azure_text import AzureTextCompletion from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM from .llms.bedrock.embed.embedding import BedrockEmbedding from .llms.bedrock.image.image_handler import BedrockImageGeneration -from .llms.cohere import chat as cohere_chat -from .llms.cohere import completion as cohere_completion # type: ignore +from .llms.clarifai.chat import handler +from .llms.cohere.chat import handler as cohere_chat +from .llms.cohere.completion import completion as cohere_completion # type: ignore from .llms.cohere.embed import handler as cohere_embed from .llms.custom_llm import CustomLLM, custom_chat_llm_router from .llms.databricks.chat.handler import DatabricksChatCompletion @@ -121,7 +121,8 @@ from .llms.groq.chat.handler import GroqChatCompletion from .llms.huggingface_restapi import Huggingface from .llms.OpenAI.audio_transcriptions import OpenAIAudioTranscription from .llms.OpenAI.chat.o1_handler import OpenAIO1ChatCompletion -from .llms.OpenAI.openai import OpenAIChatCompletion, OpenAITextCompletion +from .llms.OpenAI.completion.handler import OpenAITextCompletion +from .llms.OpenAI.openai import OpenAIChatCompletion from .llms.openai_like.embedding.handler import OpenAILikeEmbeddingHandler from .llms.predibase import PredibaseChatCompletion from .llms.prompt_templates.common_utils import get_completion_messages @@ -1685,8 +1686,9 @@ def completion( # type: ignore # noqa: PLR0915 or get_secret("CLARIFAI_API_BASE") or "https://api.clarifai.com/v2" ) + api_base = litellm.ClarifaiConfig()._convert_model_to_url(model, api_base) custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict - model_response = clarifai.completion( + model_response = handler.completion( model=model, messages=messages, api_base=api_base, @@ -1789,6 +1791,7 @@ def completion( # type: ignore # noqa: PLR0915 headers=headers, timeout=timeout, client=client, + custom_llm_provider=custom_llm_provider, ) if optional_params.get("stream", False) or acompletion is True: ## LOGGING diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index fb244db67a1..696e864cb60 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -43,21 +43,19 @@ class AnthropicPassthroughLoggingHandler: Transforms Anthropic response to OpenAI response, generates a standard logging object so downstream logging can be handled """ model = response_body.get("model", "") - litellm_model_response: litellm.ModelResponse = ( - AnthropicConfig._process_response( - response=httpx_response, - model_response=litellm.ModelResponse(), - model=model, - stream=False, - messages=[], - logging_obj=logging_obj, - optional_params={}, - api_key="", - data={}, - print_verbose=litellm.print_verbose, - encoding=None, - json_mode=False, - ) + litellm_model_response: ( + litellm.ModelResponse + ) = AnthropicConfig().transform_response( + raw_response=httpx_response, + model_response=litellm.ModelResponse(), + model=model, + messages=[], + logging_obj=logging_obj, + optional_params={}, + api_key="", + request_data={}, + encoding=litellm.encoding, + json_mode=False, ) kwargs = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload( diff --git a/litellm/utils.py b/litellm/utils.py index bd36e211cff..d2c82d487f4 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2821,9 +2821,14 @@ def get_optional_params( # noqa: PLR0915 ) _check_valid_arg(supported_params=supported_params) optional_params = litellm.AnthropicConfig().map_openai_params( + model=model, non_default_params=non_default_params, optional_params=optional_params, - messages=messages, + drop_params=( + drop_params + if drop_params is not None and isinstance(drop_params, bool) + else False + ), ) elif custom_llm_provider == "cohere": ## check if unsupported param passed in @@ -2832,24 +2837,16 @@ def get_optional_params( # noqa: PLR0915 ) _check_valid_arg(supported_params=supported_params) # handle cohere params - if stream: - optional_params["stream"] = stream - if temperature is not None: - optional_params["temperature"] = temperature - if max_tokens is not None: - optional_params["max_tokens"] = max_tokens - if n is not None: - optional_params["num_generations"] = n - if logit_bias is not None: - optional_params["logit_bias"] = logit_bias - if top_p is not None: - optional_params["p"] = top_p - if frequency_penalty is not None: - optional_params["frequency_penalty"] = frequency_penalty - if presence_penalty is not None: - optional_params["presence_penalty"] = presence_penalty - if stop is not None: - optional_params["stop_sequences"] = stop + optional_params = litellm.CohereConfig().map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=model, + drop_params=( + drop_params + if drop_params is not None and isinstance(drop_params, bool) + else False + ), + ) elif custom_llm_provider == "cohere_chat": ## check if unsupported param passed in supported_params = get_supported_openai_params( @@ -2857,26 +2854,17 @@ def get_optional_params( # noqa: PLR0915 ) _check_valid_arg(supported_params=supported_params) # handle cohere params - if stream: - optional_params["stream"] = stream - if temperature is not None: - optional_params["temperature"] = temperature - if max_tokens is not None: - optional_params["max_tokens"] = max_tokens - if n is not None: - optional_params["num_generations"] = n - if top_p is not None: - optional_params["p"] = top_p - if frequency_penalty is not None: - optional_params["frequency_penalty"] = frequency_penalty - if presence_penalty is not None: - optional_params["presence_penalty"] = presence_penalty - if stop is not None: - optional_params["stop_sequences"] = stop - if tools is not None: - optional_params["tools"] = tools - if seed is not None: - optional_params["seed"] = seed + optional_params = litellm.CohereChatConfig().map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=model, + drop_params=( + drop_params + if drop_params is not None and isinstance(drop_params, bool) + else False + ), + ) + elif custom_llm_provider == "maritalk": ## check if unsupported param passed in supported_params = get_supported_openai_params( @@ -3071,8 +3059,14 @@ def get_optional_params( # noqa: PLR0915 ) _check_valid_arg(supported_params=supported_params) optional_params = litellm.VertexAIAnthropicConfig().map_openai_params( + model=model, non_default_params=non_default_params, optional_params=optional_params, + drop_params=( + drop_params + if drop_params is not None and isinstance(drop_params, bool) + else False + ), ) elif custom_llm_provider == "vertex_ai" and model in litellm.vertex_llama3_models: supported_params = get_supported_openai_params( @@ -6220,14 +6214,14 @@ def validate_chat_completion_user_messages(messages: List[AllMessageValues]): return messages -from litellm.llms.OpenAI.chat.gpt_transformation import OpenAIGPTConfig +from litellm.llms.base_llm.transformation import BaseConfig class ProviderConfigManager: @staticmethod - def get_provider_config( + def get_provider_chat_config( model: str, provider: litellm.LlmProviders - ) -> OpenAIGPTConfig: + ) -> BaseConfig: """ Returns the provider config for a given provider. """ @@ -6239,8 +6233,23 @@ class ProviderConfigManager: return litellm.GroqChatConfig() elif litellm.LlmProviders.DATABRICKS == provider: return litellm.DatabricksConfig() + elif litellm.LlmProviders.XAI == provider: + return litellm.XAIChatConfig() + elif litellm.LlmProviders.TEXT_COMPLETION_OPENAI == provider: + return litellm.OpenAITextCompletionConfig() + elif litellm.LlmProviders.COHERE_CHAT == provider: + return litellm.CohereChatConfig() + elif litellm.LlmProviders.COHERE == provider: + return litellm.CohereConfig() + elif litellm.LlmProviders.CLARIFAI == provider: + return litellm.ClarifaiConfig() + elif litellm.LlmProviders.ANTHROPIC == provider: + return litellm.AnthropicConfig() + elif litellm.LlmProviders.VERTEX_AI == provider: + if "claude" in model: + return litellm.VertexAIAnthropicConfig() - return OpenAIGPTConfig() + return litellm.OpenAIGPTConfig() def get_end_user_id_for_cost_tracking( diff --git a/tests/llm_translation/test_max_completion_tokens.py b/tests/llm_translation/test_max_completion_tokens.py index 363125a6038..0b1e9b71a05 100644 --- a/tests/llm_translation/test_max_completion_tokens.py +++ b/tests/llm_translation/test_max_completion_tokens.py @@ -309,12 +309,16 @@ def test_all_model_configs(): assert ( "max_completion_tokens" - in VertexAIAnthropicConfig().get_supported_openai_params() + in VertexAIAnthropicConfig().get_supported_openai_params( + model="claude-3-5-sonnet-20240620" + ) ) assert VertexAIAnthropicConfig().map_openai_params( non_default_params={"max_completion_tokens": 10}, optional_params={}, + model="claude-3-5-sonnet-20240620", + drop_params=False, ) == {"max_tokens": 10} from litellm.llms.vertex_ai_and_google_ai_studio.gemini.vertex_and_google_ai_studio_gemini import ( diff --git a/tests/llm_translation/test_xai.py b/tests/llm_translation/test_xai.py index 3701d39ce9a..2e336ba97b6 100644 --- a/tests/llm_translation/test_xai.py +++ b/tests/llm_translation/test_xai.py @@ -17,7 +17,7 @@ import litellm from litellm import Choices, Message, ModelResponse, EmbeddingResponse, Usage from litellm import completion from unittest.mock import patch -from litellm.llms.xai.chat.xai_transformation import XAIChatConfig, XAI_API_BASE +from litellm.llms.xai.chat.transformation import XAIChatConfig, XAI_API_BASE def test_xai_chat_config_get_openai_compatible_provider_info(): @@ -91,8 +91,6 @@ def test_xai_chat_config_map_openai_params(): assert result["frequency_penalty"] == 0.5 assert result["logit_bias"] == {"50256": -100} assert result["logprobs"] == 5 - assert result["messages"] == [{"role": "user", "content": "Hello"}] - assert result["model"] == "xai/grok-beta" assert result["n"] == 2 assert result["presence_penalty"] == 0.2 assert result["response_format"] == {"type": "json_object"} diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py index a41f1a6ec7e..2e5e71041f1 100644 --- a/tests/local_testing/test_amazing_vertex_completion.py +++ b/tests/local_testing/test_amazing_vertex_completion.py @@ -208,7 +208,7 @@ async def test_get_router_response(): # @pytest.mark.skip( # reason="Local test. Vertex AI Quota is low. Leads to rate limit errors on ci/cd." # ) -@pytest.mark.flaky(retries=3, delay=1) +# @pytest.mark.flaky(retries=3, delay=1) def test_vertex_ai_anthropic(): model = "claude-3-sonnet@20240229" diff --git a/tests/local_testing/test_config.py b/tests/local_testing/test_config.py index 28d144e4dc7..c5896793a7f 100644 --- a/tests/local_testing/test_config.py +++ b/tests/local_testing/test_config.py @@ -288,3 +288,35 @@ async def test_add_and_delete_deployments(llm_router, model_list_flag_value): assert len(llm_router.model_list) == len(model_list) else: assert len(llm_router.model_list) == len(model_list) + prev_llm_router_val + + +# def test_provider_config_manager(): +# from litellm import LITELLM_CHAT_PROVIDERS, LlmProviders +# from litellm.utils import ProviderConfigManager +# from litellm.llms.base_llm.transformation import BaseConfig +# from litellm.llms.OpenAI.chat.gpt_transformation import OpenAIGPTConfig + +# for provider in LITELLM_CHAT_PROVIDERS: +# assert isinstance( +# ProviderConfigManager.get_provider_chat_config( +# model="gpt-3.5-turbo", provider=LlmProviders(provider) +# ), +# BaseConfig, +# ), f"Provider {provider} is not a subclass of BaseConfig" + +# config = ProviderConfigManager.get_provider_chat_config( +# model="gpt-3.5-turbo", provider=LlmProviders(provider) +# ) + +# if ( +# provider != litellm.LlmProviders.OPENAI +# and provider != litellm.LlmProviders.OPENAI_LIKE +# and provider != litellm.LlmProviders.CUSTOM_OPENAI +# ): +# assert ( +# config.__class__.__name__ != "OpenAIGPTConfig" +# ), f"Provider {provider} is an instance of OpenAIGPTConfig" + +# assert ( +# "_abc_impl" not in config.get_config() +# ), f"Provider {provider} has _abc_impl"