From beae5cfea97651bc6b63ffb78a663708f4ecf23d Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Sat, 10 May 2025 17:41:57 -0700 Subject: [PATCH] Litellm staging 05 10 2025 - openai pdf url support + sagemaker chat content length error fix (#10724) * Support pdf url's to openai (#10640) * fix(gpt_transformation.py): support pdf url input to openai pass as base64 as openai doesn't support image url's * fix(openai.py): support async message transformation allows async get request to convert url to base64 * fix(gpt_transformation.py): fix linting errrors and use common components across sync + async flows * fix: fix linting errors * fix(openai.py): pop correct var * Fix sagemaker chat calls - content length error (#10607) * fix(sagemaker_chat/): support passing dynamic aws params previously being ignored * refactor(sagemaker/chat): more refactoring * fix(sagemaker_chat/): make sure streaming is correctly handled post-refactor * refactor: more refactoring to support using signed json str * fix(sagemaker/chat): working sync streaming post refactor * fix(sagemaker/chat): support async streaming post refactor * fix(llm_http_handler.py): await async function * fix: remove print statements * test: update test * test: update test * fix(llm_http_handler.py): retain passing in data as json str * test: update test * fix(base_model_iterator.py): fix linting error * test: test auth * fix: fix linting error * test: update test * test: update translation test * fix(gpt_transformation.py): handle awaitable/non-awaitable object * fix: handle async flow for message transformation on openai compatible api's * test: cleanup testing * test: update test * test(test_router.py): use model with higher quota * test: simplify test * test: update test --- .../anthropic_messages/transformation.py | 6 +- litellm/llms/base_llm/base_model_iterator.py | 4 +- litellm/llms/base_llm/chat/transformation.py | 30 ++- litellm/llms/bedrock/base_aws_llm.py | 83 ++++++- litellm/llms/bedrock/chat/invoke_handler.py | 3 +- .../base_invoke_transformation.py | 66 ++---- .../anthropic_claude3_transformation.py | 4 +- litellm/llms/custom_httpx/llm_http_handler.py | 36 ++- .../llms/databricks/chat/transformation.py | 30 ++- litellm/llms/deepseek/chat/transformation.py | 28 ++- litellm/llms/groq/chat/transformation.py | 30 ++- .../llms/hosted_vllm/chat/transformation.py | 29 ++- .../mistral/mistral_chat_transformation.py | 24 +- .../llms/ollama/completion/transformation.py | 5 +- .../llms/openai/chat/gpt_transformation.py | 217 +++++++++++++++--- .../openai/chat/o_series_transformation.py | 36 ++- litellm/llms/openai/openai.py | 65 ++++-- litellm/llms/sagemaker/chat/transformation.py | 195 +++++++++++++++- litellm/llms/sagemaker/common_utils.py | 4 +- litellm/main.py | 14 +- litellm/utils.py | 4 +- .../base_image_generation_test.py | 2 + tests/litellm/log.txt | 48 ++-- tests/litellm/test_main.py | 7 +- .../test_bedrock_completion.py | 2 +- ..._bedrock_dynamic_auth_params_unit_tests.py | 4 +- tests/llm_translation/test_openai.py | 13 ++ tests/llm_translation/test_optional_params.py | 2 +- .../test_unit_test_bedrock_invoke.py | 2 +- tests/local_testing/test_router.py | 2 +- 30 files changed, 810 insertions(+), 185 deletions(-) diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index cd74efb202..710a107688 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional +from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional, Tuple import httpx @@ -87,7 +87,7 @@ class BaseAnthropicMessagesConfig(ABC): model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, - ) -> dict: + ) -> Tuple[dict, Optional[bytes]]: """ OPTIONAL @@ -95,7 +95,7 @@ class BaseAnthropicMessagesConfig(ABC): For all other providers, this is a no-op and we just return the headers """ - return headers + return headers, None def get_async_streaming_response_iterator( self, diff --git a/litellm/llms/base_llm/base_model_iterator.py b/litellm/llms/base_llm/base_model_iterator.py index 4cf757d6cd..9f293905d7 100644 --- a/litellm/llms/base_llm/base_model_iterator.py +++ b/litellm/llms/base_llm/base_model_iterator.py @@ -41,13 +41,13 @@ class BaseModelResponseIterator: self, str_line: str ) -> Union[GenericStreamingChunk, ModelResponseStream]: # chunk is a str at this point - + stripped_json_chunk: Optional[dict] = None stripped_chunk = litellm.CustomStreamWrapper._strip_sse_data_from_chunk( str_line ) try: if stripped_chunk is not None: - stripped_json_chunk: Optional[dict] = json.loads(stripped_chunk) + stripped_json_chunk = json.loads(stripped_chunk) else: stripped_json_chunk = None except json.JSONDecodeError: diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index fa278c805e..26faa4a5b8 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -11,6 +11,7 @@ from typing import ( Iterator, List, Optional, + Tuple, Type, Union, cast, @@ -277,7 +278,7 @@ class BaseConfig(ABC): model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, - ) -> dict: + ) -> Tuple[dict, Optional[bytes]]: """ Some providers like Bedrock require signing the request. The sign request funtion needs access to `request_data` and `complete_url` Args: @@ -290,7 +291,7 @@ class BaseConfig(ABC): Update the headers with the signed headers in this function. The return values will be sent as headers in the http request. """ - return headers + return headers, None def get_complete_url( self, @@ -323,6 +324,27 @@ class BaseConfig(ABC): ) -> dict: pass + async def async_transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + """ + Override to allow for http requests on async calls - e.g. converting url to base64 + + Currently only used by openai.py + """ + return self.transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers, + ) + @abstractmethod def transform_response( self, @@ -354,7 +376,7 @@ class BaseConfig(ABC): ) -> Any: pass - def get_async_custom_stream_wrapper( + async def get_async_custom_stream_wrapper( self, model: str, custom_llm_provider: str, @@ -365,6 +387,7 @@ class BaseConfig(ABC): messages: list, client: Optional[AsyncHTTPHandler] = None, json_mode: Optional[bool] = None, + signed_json_body: Optional[bytes] = None, ) -> CustomStreamWrapper: raise NotImplementedError @@ -379,6 +402,7 @@ class BaseConfig(ABC): messages: list, client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, json_mode: Optional[bool] = None, + signed_json_body: Optional[bytes] = None, ) -> CustomStreamWrapper: raise NotImplementedError diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 133ef6a952..a2832e69eb 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -2,7 +2,17 @@ import hashlib import json import os from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast, get_args +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + Optional, + Tuple, + cast, + get_args, +) import httpx from pydantic import BaseModel @@ -625,3 +635,74 @@ class BaseAWSLLM: prepped = request.prepare() return prepped + + def _sign_request( + self, + service_name: Literal["bedrock", "sagemaker"], + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + model: Optional[str] = None, + stream: Optional[bool] = None, + fake_stream: Optional[bool] = None, + ) -> Tuple[dict, Optional[bytes]]: + """ + Sign a request for Bedrock or Sagemaker + + Returns: + Tuple[dict, Optional[str]]: A tuple containing the headers and the json str body of the request + """ + try: + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + from botocore.credentials import Credentials + except ImportError: + raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") + + ## CREDENTIALS ## + # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them + aws_secret_access_key = optional_params.get("aws_secret_access_key", None) + aws_access_key_id = optional_params.get("aws_access_key_id", None) + aws_session_token = optional_params.get("aws_session_token", None) + aws_role_name = optional_params.get("aws_role_name", None) + aws_session_name = optional_params.get("aws_session_name", None) + aws_profile_name = optional_params.get("aws_profile_name", None) + aws_web_identity_token = optional_params.get("aws_web_identity_token", None) + aws_sts_endpoint = optional_params.get("aws_sts_endpoint", None) + aws_region_name = self._get_aws_region_name( + optional_params=optional_params, model=model + ) + + credentials: Credentials = self.get_credentials( + aws_access_key_id=aws_access_key_id, + aws_secret_access_key=aws_secret_access_key, + aws_session_token=aws_session_token, + aws_region_name=aws_region_name, + aws_session_name=aws_session_name, + aws_profile_name=aws_profile_name, + aws_role_name=aws_role_name, + aws_web_identity_token=aws_web_identity_token, + aws_sts_endpoint=aws_sts_endpoint, + ) + + sigv4 = SigV4Auth(credentials, service_name, aws_region_name) + if headers is not None: + headers = {"Content-Type": "application/json", **headers} + else: + headers = {"Content-Type": "application/json"} + + request = AWSRequest( + method="POST", + url=api_base, + data=json.dumps(request_data), + headers=headers, + ) + sigv4.add_auth(request) + + request_headers_dict = dict(request.headers) + if ( + headers is not None and "Authorization" in headers + ): # prevent sigv4 from overwriting the auth header + request_headers_dict["Authorization"] = headers["Authorization"] + return request_headers_dict, request.body diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index be80cb5881..2c3cf59585 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -272,6 +272,7 @@ def make_sync_call( api_base: str, headers: dict, data: str, + signed_json_body: Optional[bytes], model: str, messages: list, logging_obj: Logging, @@ -286,7 +287,7 @@ def make_sync_call( response = client.post( api_base, headers=headers, - data=data, + data=signed_json_body if signed_json_body is not None else data, stream=not fake_stream, logging_obj=logging_obj, ) diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index 67194e83e7..b5b13476c1 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -121,60 +121,17 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, - ) -> dict: - try: - from botocore.auth import SigV4Auth - from botocore.awsrequest import AWSRequest - from botocore.credentials import Credentials - except ImportError: - raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key = optional_params.get("aws_secret_access_key", None) - aws_access_key_id = optional_params.get("aws_access_key_id", None) - aws_session_token = optional_params.get("aws_session_token", None) - aws_role_name = optional_params.get("aws_role_name", None) - aws_session_name = optional_params.get("aws_session_name", None) - aws_profile_name = optional_params.get("aws_profile_name", None) - aws_web_identity_token = optional_params.get("aws_web_identity_token", None) - aws_sts_endpoint = optional_params.get("aws_sts_endpoint", None) - aws_region_name = self._get_aws_region_name( - optional_params=optional_params, model=model - ) - - credentials: Credentials = self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - ) - - sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name) - if headers is not None: - headers = {"Content-Type": "application/json", **headers} - else: - headers = {"Content-Type": "application/json"} - - request = AWSRequest( - method="POST", - url=api_base, - data=json.dumps(request_data), + ) -> Tuple[dict, Optional[bytes]]: + return self._sign_request( + service_name="bedrock", headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base, + model=model, + stream=stream, + fake_stream=fake_stream, ) - sigv4.add_auth(request) - - request_headers_dict = dict(request.headers) - if ( - headers is not None and "Authorization" in headers - ): # prevent sigv4 from overwriting the auth header - request_headers_dict["Authorization"] = headers["Authorization"] - return request_headers_dict def transform_request( self, @@ -454,7 +411,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): return BedrockError(status_code=status_code, message=error_message) @track_llm_api_timing() - def get_async_custom_stream_wrapper( + async def get_async_custom_stream_wrapper( self, model: str, custom_llm_provider: str, @@ -465,6 +422,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): messages: list, client: Optional[AsyncHTTPHandler] = None, json_mode: Optional[bool] = None, + signed_json_body: Optional[bytes] = None, ) -> CustomStreamWrapper: streaming_response = CustomStreamWrapper( completion_stream=None, @@ -499,6 +457,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): messages: list, client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, json_mode: Optional[bool] = None, + signed_json_body: Optional[bytes] = None, ) -> CustomStreamWrapper: if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -510,6 +469,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): api_base=api_base, headers=headers, data=json.dumps(data), + signed_json_body=signed_json_body, model=model, messages=messages, logging_obj=logging_obj, diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index f11a3b0b85..ff475a95db 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional, Tuple, Union import httpx @@ -47,7 +47,7 @@ class AmazonAnthropicClaude3MessagesConfig( model: Optional[str] = None, stream: Optional[bool] = None, fake_stream: Optional[bool] = None, - ) -> dict: + ) -> Tuple[dict, Optional[bytes]]: return AmazonInvokeConfig.sign_request( self=self, headers=headers, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index cd894a459a..921ce17f90 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -78,6 +78,7 @@ class BaseLLMHTTPHandler: litellm_params: dict, logging_obj: LiteLLMLoggingObj, stream: bool = False, + signed_json_body: Optional[bytes] = None, ) -> httpx.Response: """Common implementation across stream + non-stream calls. Meant to ensure consistent error-handling.""" max_retry_on_unprocessable_entity_error = ( @@ -90,7 +91,9 @@ class BaseLLMHTTPHandler: response = await async_httpx_client.post( url=api_base, headers=headers, - data=json.dumps(data), + data=signed_json_body + if signed_json_body is not None + else json.dumps(data), timeout=timeout, stream=stream, logging_obj=logging_obj, @@ -133,6 +136,7 @@ class BaseLLMHTTPHandler: litellm_params: dict, logging_obj: LiteLLMLoggingObj, stream: bool = False, + signed_json_body: Optional[bytes] = None, ) -> httpx.Response: max_retry_on_unprocessable_entity_error = ( provider_config.max_retry_on_unprocessable_entity_error @@ -145,7 +149,9 @@ class BaseLLMHTTPHandler: response = sync_httpx_client.post( url=api_base, headers=headers, - data=json.dumps(data), + data=signed_json_body + if signed_json_body is not None + else json.dumps(data), timeout=timeout, stream=stream, logging_obj=logging_obj, @@ -195,6 +201,7 @@ class BaseLLMHTTPHandler: api_key: Optional[str] = None, client: Optional[AsyncHTTPHandler] = None, json_mode: bool = False, + signed_json_body: Optional[bytes] = None, ): if client is None: async_httpx_client = get_async_httpx_client( @@ -214,6 +221,7 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, stream=False, logging_obj=logging_obj, + signed_json_body=signed_json_body, ) return provider_config.transform_response( model=model, @@ -295,7 +303,7 @@ class BaseLLMHTTPHandler: if extra_body is not None: data = {**data, **extra_body} - headers = provider_config.sign_request( + headers, signed_json_body = provider_config.sign_request( headers=headers, optional_params=optional_params, request_data=data, @@ -342,6 +350,7 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, json_mode=json_mode, optional_params=optional_params, + signed_json_body=signed_json_body, ) else: @@ -366,6 +375,7 @@ class BaseLLMHTTPHandler: else None ), json_mode=json_mode, + signed_json_body=signed_json_body, ) if stream is True: @@ -382,6 +392,7 @@ class BaseLLMHTTPHandler: api_base=api_base, headers=headers, data=data, + signed_json_body=signed_json_body, messages=messages, client=client, json_mode=json_mode, @@ -391,6 +402,8 @@ class BaseLLMHTTPHandler: api_base=api_base, headers=headers, # type: ignore data=data, + signed_json_body=signed_json_body, + original_data=data, model=model, messages=messages, logging_obj=logging_obj, @@ -425,6 +438,7 @@ class BaseLLMHTTPHandler: api_base=api_base, headers=headers, data=data, + signed_json_body=signed_json_body, timeout=timeout, litellm_params=litellm_params, logging_obj=logging_obj, @@ -449,6 +463,8 @@ class BaseLLMHTTPHandler: api_base: str, headers: dict, data: dict, + signed_json_body: Optional[bytes], + original_data: dict, model: str, messages: list, logging_obj, @@ -477,6 +493,7 @@ class BaseLLMHTTPHandler: api_base=api_base, headers=headers, data=data, + signed_json_body=signed_json_body, timeout=timeout, litellm_params=litellm_params, stream=stream, @@ -489,7 +506,7 @@ class BaseLLMHTTPHandler: raw_response=response, model_response=litellm.ModelResponse(), logging_obj=logging_obj, - request_data=data, + request_data=original_data, messages=messages, optional_params=optional_params, litellm_params=litellm_params, @@ -533,9 +550,10 @@ class BaseLLMHTTPHandler: fake_stream: bool = False, client: Optional[AsyncHTTPHandler] = None, json_mode: Optional[bool] = None, + signed_json_body: Optional[bytes] = None, ): if provider_config.has_custom_stream_wrapper is True: - return provider_config.get_async_custom_stream_wrapper( + return await provider_config.get_async_custom_stream_wrapper( model=model, custom_llm_provider=custom_llm_provider, logging_obj=logging_obj, @@ -545,6 +563,7 @@ class BaseLLMHTTPHandler: messages=messages, client=client, json_mode=json_mode, + signed_json_body=signed_json_body, ) completion_stream, _response_headers = await self.make_async_call_stream_helper( @@ -562,6 +581,7 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, optional_params=optional_params, json_mode=json_mode, + signed_json_body=signed_json_body, ) streamwrapper = CustomStreamWrapper( completion_stream=completion_stream, @@ -587,6 +607,7 @@ class BaseLLMHTTPHandler: fake_stream: bool = False, client: Optional[AsyncHTTPHandler] = None, json_mode: Optional[bool] = None, + signed_json_body: Optional[bytes] = None, ) -> Tuple[Any, httpx.Headers]: """ Helper function for making an async call with stream. @@ -610,6 +631,7 @@ class BaseLLMHTTPHandler: api_base=api_base, headers=headers, data=data, + signed_json_body=signed_json_body, timeout=timeout, litellm_params=litellm_params, stream=stream, @@ -1094,7 +1116,7 @@ class BaseLLMHTTPHandler: stream=stream, ) - headers = anthropic_messages_provider_config.sign_request( + headers, signed_json_body = anthropic_messages_provider_config.sign_request( headers=headers, optional_params=anthropic_messages_optional_request_params, request_data=request_body, @@ -1117,7 +1139,7 @@ class BaseLLMHTTPHandler: response = await async_httpx_client.post( url=request_url, headers=headers, - data=json.dumps(request_body), + data=signed_json_body or json.dumps(request_body), stream=stream or False, logging_obj=logging_obj, ) diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index fa41849add..ba22f7ac44 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -6,12 +6,15 @@ from typing import ( TYPE_CHECKING, Any, AsyncIterator, + Coroutine, Iterator, List, + Literal, Optional, Tuple, Union, cast, + overload, ) import httpx @@ -276,9 +279,24 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): return False + @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, + messages: List[AllMessageValues], + model: str, + is_async: Literal[False] = False, ) -> List[AllMessageValues]: + ... + + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: """ Databricks does not support: - content in list format. @@ -293,7 +311,15 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): new_messages.append(_message) new_messages = handle_messages_with_content_list_to_str_conversion(new_messages) new_messages = strip_name_from_messages(new_messages) - return super()._transform_messages(messages=new_messages, model=model) + + if is_async: + return super()._transform_messages( + messages=new_messages, model=model, is_async=cast(Literal[True], True) + ) + else: + return super()._transform_messages( + messages=new_messages, model=model, is_async=cast(Literal[False], False) + ) @staticmethod def extract_content_str( diff --git a/litellm/llms/deepseek/chat/transformation.py b/litellm/llms/deepseek/chat/transformation.py index f429f46331..a7defa886b 100644 --- a/litellm/llms/deepseek/chat/transformation.py +++ b/litellm/llms/deepseek/chat/transformation.py @@ -2,7 +2,7 @@ Translates from OpenAI's `/v1/chat/completions` to DeepSeek's `/v1/chat/completions` """ -from typing import List, Optional, Tuple +from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -14,14 +14,36 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig class DeepSeekChatConfig(OpenAIGPTConfig): + @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, + messages: List[AllMessageValues], + model: str, + is_async: Literal[False] = False, ) -> List[AllMessageValues]: + ... + + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: """ DeepSeek does not support content in list format. """ messages = handle_messages_with_content_list_to_str_conversion(messages) - return super()._transform_messages(messages=messages, model=model) + if is_async: + return super()._transform_messages( + messages=messages, model=model, is_async=True + ) + else: + return super()._transform_messages( + messages=messages, model=model, is_async=False + ) def _get_openai_compatible_provider_info( self, api_base: Optional[str], api_key: Optional[str] diff --git a/litellm/llms/groq/chat/transformation.py b/litellm/llms/groq/chat/transformation.py index 4befdc504e..877d9a6edb 100644 --- a/litellm/llms/groq/chat/transformation.py +++ b/litellm/llms/groq/chat/transformation.py @@ -2,7 +2,7 @@ Translate from OpenAI's `/v1/chat/completions` to Groq's `/v1/chat/completions` """ -from typing import List, Optional, Tuple, Union +from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload from pydantic import BaseModel @@ -65,7 +65,24 @@ class GroqChatConfig(OpenAILikeChatConfig): pass return base_params - def _transform_messages(self, messages: List[AllMessageValues], model: str) -> List: + @overload + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, + messages: List[AllMessageValues], + model: str, + is_async: Literal[False] = False, + ) -> List[AllMessageValues]: + ... + + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: for idx, message in enumerate(messages): """ 1. Don't pass 'null' function_call assistant message to groq - https://github.com/BerriAI/litellm/issues/5839 @@ -82,7 +99,14 @@ class GroqChatConfig(OpenAILikeChatConfig): new_message[k] = v # type: ignore messages[idx] = new_message - return messages + if is_async: + return super()._transform_messages( + messages=messages, model=model, is_async=True + ) + else: + return super()._transform_messages( + messages=messages, model=model, is_async=False + ) def _get_openai_compatible_provider_info( self, api_base: Optional[str], api_key: Optional[str] diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index e328bf2881..529354f80e 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -2,7 +2,7 @@ Translate from OpenAI's `/v1/chat/completions` to VLLM's `/v1/chat/completions` """ -from typing import List, Optional, Tuple, cast +from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, cast, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( _get_image_mime_type_from_url, @@ -92,9 +92,24 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): ) raise ValueError("file_id or file_data is required") + @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, + messages: List[AllMessageValues], + model: str, + is_async: Literal[False] = False, ) -> List[AllMessageValues]: + ... + + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: """ Support translating video files from file_id or file_data to video_url """ @@ -114,5 +129,11 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): message_content[idx] = self._convert_file_to_video_url( content_item ) - transformed_messages = super()._transform_messages(messages, model) - return transformed_messages + if is_async: + return super()._transform_messages( + messages, model, is_async=cast(Literal[True], True) + ) + else: + return super()._transform_messages( + messages, model, is_async=cast(Literal[False], False) + ) diff --git a/litellm/llms/mistral/mistral_chat_transformation.py b/litellm/llms/mistral/mistral_chat_transformation.py index 67d88868d3..f580414df8 100644 --- a/litellm/llms/mistral/mistral_chat_transformation.py +++ b/litellm/llms/mistral/mistral_chat_transformation.py @@ -6,7 +6,7 @@ Why separate file? Make it easy to see how transformation works Docs - https://docs.mistral.ai/api/ """ -from typing import List, Literal, Optional, Tuple, Union +from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -152,9 +152,24 @@ class MistralConfig(OpenAIGPTConfig): ) return api_base, dynamic_api_key + @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, + messages: List[AllMessageValues], + model: str, + is_async: Literal[False] = False, ) -> List[AllMessageValues]: + ... + + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: """ - handles scenario where content is list and not string - content list is just text, and no images @@ -182,7 +197,10 @@ class MistralConfig(OpenAIGPTConfig): m = strip_none_values_from_message(m) # prevents 'extra_forbidden' error new_messages.append(m) - return new_messages + if is_async: + return super()._transform_messages(new_messages, model, True) + else: + return super()._transform_messages(new_messages, model, False) @classmethod def _handle_name_in_message(cls, message: AllMessageValues) -> AllMessageValues: diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index b2557a53e9..133554befe 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -22,6 +22,7 @@ from litellm.types.utils import ( GenericStreamingChunk, ModelInfoBase, ModelResponse, + ModelResponseStream, ProviderField, ) @@ -415,7 +416,9 @@ class OllamaConfig(BaseConfig): class OllamaTextCompletionResponseIterator(BaseModelResponseIterator): - def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk: + def _handle_string_chunk( + self, str_line: str + ) -> Union[GenericStreamingChunk, ModelResponseStream]: return self.chunk_parser(json.loads(str_line)) def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 9e3d9e5fc9..70d4d58564 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -6,11 +6,14 @@ from typing import ( TYPE_CHECKING, Any, AsyncIterator, + Coroutine, Iterator, List, + Literal, Optional, Union, cast, + overload, ) import httpx @@ -22,6 +25,10 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo _should_convert_tool_call_to_json_mode, ) from litellm.litellm_core_utils.prompt_templates.common_utils import get_tool_call_names +from litellm.litellm_core_utils.prompt_templates.image_handling import ( + async_convert_url_to_base64, + convert_url_to_base64, +) from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.base_utils import BaseLLMModelInfo from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException @@ -33,6 +40,7 @@ from litellm.types.llms.openai import ( ChatCompletionImageObject, ChatCompletionImageUrlObject, OpenAIChatCompletionChoices, + OpenAIMessageContentListBlock, ) from litellm.types.utils import ( ChatCompletionMessageToolCall, @@ -196,42 +204,169 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): drop_params=drop_params, ) + @overload + def _handle_pdf_url( + self, content_item: ChatCompletionFileObjectFile, is_async: Literal[True] + ) -> Coroutine[Any, Any, ChatCompletionFileObjectFile]: + ... + + @overload + def _handle_pdf_url( + self, + content_item: ChatCompletionFileObjectFile, + is_async: Literal[False] = False, + ) -> ChatCompletionFileObjectFile: + ... + + def _handle_pdf_url( + self, content_item: ChatCompletionFileObjectFile, is_async: bool = False + ) -> Union[ + ChatCompletionFileObjectFile, Coroutine[Any, Any, ChatCompletionFileObjectFile] + ]: + potential_pdf_url_starts = ["https://", "http://", "www."] + content_copy = content_item.copy() + file_id = content_copy.get("file_id") + if file_id and any( + file_id.startswith(start) for start in potential_pdf_url_starts + ): + if is_async: + return self._async_handle_pdf_url_helper(content_item) + else: + base64_data = convert_url_to_base64(file_id) + content_copy["file_data"] = base64_data + content_copy["filename"] = "my_file.pdf" + content_copy.pop("file_id") + return content_copy + + async def _async_handle_pdf_url_helper( + self, content_item: ChatCompletionFileObjectFile + ) -> ChatCompletionFileObjectFile: + file_id = content_item.get("file_id") + if file_id is not None: # check for file id being url done in _handle_pdf_url + base64_data = await async_convert_url_to_base64(file_id) + content_item["file_data"] = base64_data + content_item["filename"] = "my_file.pdf" + content_item.pop("file_id") + return content_item + + @overload def _transform_messages( - self, messages: List[AllMessageValues], model: str + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, + messages: List[AllMessageValues], + model: str, + is_async: Literal[False] = False, ) -> List[AllMessageValues]: + ... + + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: """OpenAI no longer supports image_url as a string, so we need to convert it to a dict""" - for message in messages: - message_content = message.get("content") - if message_content and isinstance(message_content, list): - for content_item in message_content: - litellm_specific_params = {"format"} - if content_item.get("type") == "image_url": - content_item = cast(ChatCompletionImageObject, content_item) - if isinstance(content_item["image_url"], str): - content_item["image_url"] = { - "url": content_item["image_url"], - } - elif isinstance(content_item["image_url"], dict): - new_image_url_obj = ChatCompletionImageUrlObject( - **{ # type: ignore - k: v - for k, v in content_item["image_url"].items() - if k not in litellm_specific_params - } - ) - content_item["image_url"] = new_image_url_obj - elif content_item.get("type") == "file": - content_item = cast(ChatCompletionFileObject, content_item) - file_obj = content_item["file"] - new_file_obj = ChatCompletionFileObjectFile( - **{ # type: ignore - k: v - for k, v in file_obj.items() - if k not in litellm_specific_params - } + + def _apply_common_transform_content_item( + content_item: OpenAIMessageContentListBlock, + ) -> OpenAIMessageContentListBlock: + litellm_specific_params = {"format"} + if content_item.get("type") == "image_url": + content_item = cast(ChatCompletionImageObject, content_item) + if isinstance(content_item["image_url"], str): + content_item["image_url"] = { + "url": content_item["image_url"], + } + elif isinstance(content_item["image_url"], dict): + new_image_url_obj = ChatCompletionImageUrlObject( + **{ # type: ignore + k: v + for k, v in content_item["image_url"].items() + if k not in litellm_specific_params + } + ) + content_item["image_url"] = new_image_url_obj + elif content_item.get("type") == "file": + content_item = cast(ChatCompletionFileObject, content_item) + file_obj = content_item["file"] + new_file_obj = ChatCompletionFileObjectFile( + **{ # type: ignore + k: v + for k, v in file_obj.items() + if k not in litellm_specific_params + } + ) + content_item["file"] = new_file_obj + + return content_item + + def _transform_content_item( + content_item: OpenAIMessageContentListBlock, + ) -> OpenAIMessageContentListBlock: + content_item = _apply_common_transform_content_item(content_item) + content_item_type = content_item.get("type") + potential_file_obj = content_item.get("file") + if content_item_type == "file" and potential_file_obj: + file_obj = cast(ChatCompletionFileObjectFile, potential_file_obj) + content_item_typed = cast(ChatCompletionFileObject, content_item) + content_item_typed["file"] = self._handle_pdf_url(file_obj) + content_item = content_item_typed + return content_item + + async def _async_transform_content_item( + content_item: OpenAIMessageContentListBlock, is_async: bool = False + ) -> OpenAIMessageContentListBlock: + content_item = _apply_common_transform_content_item(content_item) + content_item_type = content_item.get("type") + potential_file_obj = content_item.get("file") + if content_item_type == "file" and potential_file_obj: + file_obj = cast(ChatCompletionFileObjectFile, potential_file_obj) + content_item_typed = cast(ChatCompletionFileObject, content_item) + content_item_typed["file"] = await self._handle_pdf_url( + file_obj, is_async=True + ) + content_item = content_item_typed + return content_item + + async def _async_transform(): + for message in messages: + message_content = message.get("content") + message_role = message.get("role") + if ( + message_role == "user" + and message_content + and isinstance(message_content, list) + ): + message_content_types = cast( + List[OpenAIMessageContentListBlock], message_content + ) + for i, content_item in enumerate(message_content_types): + message_content_types[i] = await _async_transform_content_item( + cast(OpenAIMessageContentListBlock, content_item), ) - content_item["file"] = new_file_obj - return messages + return messages + + if is_async: + return _async_transform() + else: + for message in messages: + message_content = message.get("content") + message_role = message.get("role") + if ( + message_role == "user" + and message_content + and isinstance(message_content, list) + ): + message_content_types = cast( + List[OpenAIMessageContentListBlock], message_content + ) + for i, content_item in enumerate(message_content): + message_content_types[i] = _transform_content_item( + cast(OpenAIMessageContentListBlock, content_item) + ) + return messages def transform_request( self, @@ -254,6 +389,24 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): **optional_params, } + async def async_transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + transformed_messages = await self._transform_messages( + messages=messages, model=model, is_async=True + ) + + return { + "model": model, + "messages": transformed_messages, + **optional_params, + } + def _passed_in_tools(self, optional_params: dict) -> bool: return optional_params.get("tools", None) is not None diff --git a/litellm/llms/openai/chat/o_series_transformation.py b/litellm/llms/openai/chat/o_series_transformation.py index af9b7ed460..30647f5868 100644 --- a/litellm/llms/openai/chat/o_series_transformation.py +++ b/litellm/llms/openai/chat/o_series_transformation.py @@ -11,7 +11,7 @@ Translations handled by LiteLLM: - Logprobs => drop param (if user opts in to dropping param) """ -from typing import List, Optional +from typing import Any, Coroutine, List, Literal, Optional, Union, cast, overload import litellm from litellm import verbose_logger @@ -130,13 +130,29 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): ) def is_model_o_series_model(self, model: str) -> bool: - model = model.split("/")[-1] # could be "openai/o3" or "o3" + model = model.split("/")[-1] # could be "openai/o3" or "o3" return model in litellm.open_ai_chat_completion_models and any( - model.startswith(pfx) for pfx in ("o1", "o3", "o4")) + model.startswith(pfx) for pfx in ("o1", "o3", "o4") + ) + + @overload + def _transform_messages( + self, messages: List[AllMessageValues], model: str, is_async: Literal[True] + ) -> Coroutine[Any, Any, List[AllMessageValues]]: + ... + + @overload + def _transform_messages( + self, + messages: List[AllMessageValues], + model: str, + is_async: Literal[False] = False, + ) -> List[AllMessageValues]: + ... def _transform_messages( - self, messages: List[AllMessageValues], model: str - ) -> List[AllMessageValues]: + self, messages: List[AllMessageValues], model: str, is_async: bool = False + ) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]: """ Handles limitations of O-1 model family. - modalities: image => drop param (if user opts in to dropping param) @@ -150,5 +166,11 @@ class OpenAIOSeriesConfig(OpenAIGPTConfig): ) messages[i] = new_message # Replace the old message with the new one - messages = super()._transform_messages(messages, model) - return messages + if is_async: + return super()._transform_messages( + messages, model, is_async=cast(Literal[True], True) + ) + else: + return super()._transform_messages( + messages, model, is_async=cast(Literal[False], False) + ) diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 13412ef96a..e9bed019a9 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -527,6 +527,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): model=model, provider=LlmProviders(custom_llm_provider) ) + if provider_config is None: + provider_config = OpenAIConfig() + if provider_config: fake_stream = provider_config.should_fake_stream( model=model, custom_llm_provider=custom_llm_provider, stream=stream @@ -551,30 +554,17 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): for _ in range( 2 ): # if call fails due to alternating messages, retry with reformatted message - if provider_config is not None: - data = provider_config.transform_request( - model=model, - messages=messages, - optional_params=inference_params, - litellm_params=litellm_params, - headers=headers or {}, - ) - else: - data = OpenAIConfig().transform_request( - model=model, - messages=messages, - optional_params=inference_params, - litellm_params=litellm_params, - headers=headers or {}, - ) try: - max_retries = data.pop("max_retries", 2) + max_retries = inference_params.pop("max_retries", 2) if acompletion is True: if stream is True and fake_stream is False: return self.async_streaming( logging_obj=logging_obj, headers=headers, - data=data, + messages=messages, + optional_params=inference_params, + litellm_params=litellm_params, + provider_config=provider_config, model=model, api_base=api_base, api_key=api_key, @@ -588,7 +578,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) else: return self.acompletion( - data=data, + messages=messages, + optional_params=inference_params, + litellm_params=litellm_params, + provider_config=provider_config, headers=headers, model=model, logging_obj=logging_obj, @@ -603,7 +596,15 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): drop_params=drop_params, fake_stream=fake_stream, ) - elif stream is True and fake_stream is False: + + data = provider_config.transform_request( + model=model, + messages=messages, + optional_params=inference_params, + litellm_params=litellm_params, + headers=headers or {}, + ) + if stream is True and fake_stream is False: return self.streaming( logging_obj=logging_obj, headers=headers, @@ -741,7 +742,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): async def acompletion( self, - data: dict, + messages: list, + optional_params: dict, + litellm_params: dict, + provider_config: BaseConfig, model: str, model_response: ModelResponse, logging_obj: LiteLLMLoggingObj, @@ -758,6 +762,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): fake_stream: bool = False, ): response = None + data = await provider_config.async_transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers or {}, + ) for _ in range( 2 ): # if call fails due to alternating messages, retry with reformatted message @@ -903,7 +914,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): async def async_streaming( self, timeout: Union[float, httpx.Timeout], - data: dict, + messages: list, + optional_params: dict, + litellm_params: dict, + provider_config: BaseConfig, model: str, logging_obj: LiteLLMLoggingObj, api_key: Optional[str] = None, @@ -917,6 +931,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): stream_options: Optional[dict] = None, ): response = None + data = provider_config.transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers or {}, + ) data["stream"] = True data.update( self.get_stream_options(stream_options=stream_options, api_base=api_base) diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py index 42c7e0d5fc..14dde144af 100644 --- a/litellm/llms/sagemaker/chat/transformation.py +++ b/litellm/llms/sagemaker/chat/transformation.py @@ -7,20 +7,209 @@ LiteLLM Docs: https://docs.litellm.ai/docs/providers/aws_sagemaker#sagemaker-mes Huggingface Docs: https://huggingface.co/docs/text-generation-inference/en/messages_api """ -from typing import Union +from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast +import httpx from httpx._models import Headers +from litellm.litellm_core_utils.logging_utils import track_llm_api_timing +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + HTTPHandler, + _get_httpx_client, + get_async_httpx_client, +) +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import LlmProviders from ...openai.chat.gpt_transformation import OpenAIGPTConfig -from ..common_utils import SagemakerError +from ..common_utils import AWSEventStreamDecoder, SagemakerError + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any -class SagemakerChatConfig(OpenAIGPTConfig): +class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): + def __init__(self, **kwargs): + OpenAIGPTConfig.__init__(self, **kwargs) + BaseAWSLLM.__init__(self, **kwargs) + def get_error_class( self, error_message: str, status_code: int, headers: Union[dict, Headers] ) -> BaseLLMException: return SagemakerError( status_code=status_code, message=error_message, headers=headers ) + + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + ) -> dict: + return headers + + def get_complete_url( + self, + api_base: Optional[str], + api_key: Optional[str], + model: str, + optional_params: dict, + litellm_params: dict, + stream: Optional[bool] = None, + ) -> str: + aws_region_name = self._get_aws_region_name( + optional_params=optional_params, + model=model, + model_id=None, + ) + if stream is True: + api_base = f"https://runtime.sagemaker.{aws_region_name}.amazonaws.com/endpoints/{model}/invocations-response-stream" + else: + api_base = f"https://runtime.sagemaker.{aws_region_name}.amazonaws.com/endpoints/{model}/invocations" + + sagemaker_base_url = cast( + Optional[str], optional_params.get("sagemaker_base_url") + ) + if sagemaker_base_url is not None: + api_base = sagemaker_base_url + + return api_base + + def sign_request( + self, + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + model: Optional[str] = None, + stream: Optional[bool] = None, + fake_stream: Optional[bool] = None, + ) -> Tuple[dict, Optional[bytes]]: + return self._sign_request( + service_name="sagemaker", + headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base, + model=model, + stream=stream, + fake_stream=fake_stream, + ) + + @property + def has_custom_stream_wrapper(self) -> bool: + return True + + @property + def supports_stream_param_in_request_body(self) -> bool: + return False + + @track_llm_api_timing() + def get_sync_custom_stream_wrapper( + self, + model: str, + custom_llm_provider: str, + logging_obj: LiteLLMLoggingObj, + api_base: str, + headers: dict, + data: dict, + messages: list, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + json_mode: Optional[bool] = None, + signed_json_body: Optional[bytes] = None, + ) -> CustomStreamWrapper: + if client is None or isinstance(client, AsyncHTTPHandler): + client = _get_httpx_client(params={}) + + try: + response = client.post( + api_base, + headers=headers, + data=signed_json_body if signed_json_body is not None else data, + stream=True, + logging_obj=logging_obj, + ) + except httpx.HTTPStatusError as e: + raise SagemakerError( + status_code=e.response.status_code, message=e.response.text + ) + + if response.status_code != 200: + raise SagemakerError( + status_code=response.status_code, message=response.text + ) + + custom_stream_decoder = AWSEventStreamDecoder(model="", is_messages_api=True) + completion_stream = custom_stream_decoder.iter_bytes( + response.iter_bytes(chunk_size=1024) + ) + + streaming_response = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="sagemaker_chat", + logging_obj=logging_obj, + ) + return streaming_response + + @track_llm_api_timing() + async def get_async_custom_stream_wrapper( + self, + model: str, + custom_llm_provider: str, + logging_obj: LiteLLMLoggingObj, + api_base: str, + headers: dict, + data: dict, + messages: list, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, + json_mode: Optional[bool] = None, + signed_json_body: Optional[bytes] = None, + ) -> CustomStreamWrapper: + if client is None or isinstance(client, HTTPHandler): + client = get_async_httpx_client( + llm_provider=LlmProviders.SAGEMAKER_CHAT, params={} + ) + + try: + response = await client.post( + api_base, + headers=headers, + data=signed_json_body if signed_json_body is not None else data, + stream=True, + logging_obj=logging_obj, + ) + except httpx.HTTPStatusError as e: + raise SagemakerError( + status_code=e.response.status_code, message=e.response.text + ) + + if response.status_code != 200: + raise SagemakerError( + status_code=response.status_code, message=response.text + ) + + custom_stream_decoder = AWSEventStreamDecoder(model="", is_messages_api=True) + completion_stream = custom_stream_decoder.aiter_bytes( + response.aiter_bytes(chunk_size=1024) + ) + + streaming_response = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="sagemaker_chat", + logging_obj=logging_obj, + ) + return streaming_response diff --git a/litellm/llms/sagemaker/common_utils.py b/litellm/llms/sagemaker/common_utils.py index 031a0c7f05..ad6b24d85a 100644 --- a/litellm/llms/sagemaker/common_utils.py +++ b/litellm/llms/sagemaker/common_utils.py @@ -34,7 +34,9 @@ class AWSEventStreamDecoder: def _chunk_parser_messages_api( self, chunk_data: dict ) -> StreamingChatCompletionChunk: - openai_chunk = StreamingChatCompletionChunk(**chunk_data) + openai_chunk = StreamingChatCompletionChunk( + **{"model": self.model, **chunk_data} + ) return openai_chunk diff --git a/litellm/main.py b/litellm/main.py index 5198cc7b0a..f8c3ec9b98 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -2664,19 +2664,21 @@ def completion( # type: ignore # noqa: PLR0915 response = _model_response elif custom_llm_provider == "sagemaker_chat": # boto3 reads keys from .env - model_response = sagemaker_chat_completion.completion( + model_response = base_llm_http_handler.completion( model=model, + stream=stream, messages=messages, + acompletion=acompletion, + api_base=api_base, model_response=model_response, - print_verbose=print_verbose, optional_params=optional_params, litellm_params=litellm_params, + custom_llm_provider="sagemaker_chat", timeout=timeout, - custom_prompt_dict=custom_prompt_dict, - logger_fn=logger_fn, + headers=headers, encoding=encoding, - logging_obj=logging, - acompletion=acompletion, + api_key=api_key, + logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements client=client, ) diff --git a/litellm/utils.py b/litellm/utils.py index 45da62631a..b2e9caf1c4 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2659,7 +2659,8 @@ def get_optional_params( # noqa: PLR0915 special_params = passed_params.pop("kwargs") for k, v in special_params.items(): if k.startswith("aws_") and ( - custom_llm_provider != "bedrock" and custom_llm_provider != "sagemaker" + custom_llm_provider != "bedrock" + and not custom_llm_provider.startswith("sagemaker") ): # allow dynamically setting boto3 init logic continue elif k == "hf_model_name" and custom_llm_provider != "sagemaker": @@ -6722,6 +6723,7 @@ def get_non_default_completion_params(kwargs: dict) -> dict: non_default_params = { k: v for k, v in kwargs.items() if k not in default_params } # model-specific params - pass them straight to the model/provider + return non_default_params diff --git a/tests/image_gen_tests/base_image_generation_test.py b/tests/image_gen_tests/base_image_generation_test.py index 76fa2f5540..df64d353c7 100644 --- a/tests/image_gen_tests/base_image_generation_test.py +++ b/tests/image_gen_tests/base_image_generation_test.py @@ -81,6 +81,8 @@ class BaseImageGenTest(ABC): pass except litellm.ContentPolicyViolationError: pass # Azure randomly raises these errors - skip when they occur + except litellm.InternalServerError: + pass except Exception as e: if "Your task failed as a result of our safety system." in str(e): pass diff --git a/tests/litellm/log.txt b/tests/litellm/log.txt index 8d75e0c800..8324baf549 100644 --- a/tests/litellm/log.txt +++ b/tests/litellm/log.txt @@ -1,28 +1,38 @@ ============================= test session starts ============================== -platform darwin -- Python 3.11.4, pytest-7.4.1, pluggy-1.2.0 -- /Library/Frameworks/Python.framework/Versions/3.11/bin/python3 +platform darwin -- Python 3.13.1, pytest-8.3.5, pluggy-1.5.0 -- /Users/krrishdholakia/Documents/litellm/myenv/bin/python3.13 cachedir: .pytest_cache -rootdir: /Users/krrishdholakia/Documents/litellm/tests/litellm -plugins: snapshot-0.9.0, cov-5.0.0, timeout-2.2.0, respx-0.21.1, asyncio-0.21.1, langsmith-0.3.4, anyio-4.8.0, mock-3.11.1, Faker-25.9.2 -asyncio: mode=Mode.STRICT -collecting ... collected 4 items +rootdir: /Users/krrishdholakia/Documents/litellm +configfile: pyproject.toml +plugins: respx-0.22.0, postgresql-7.0.1, anyio-4.4.0, asyncio-0.26.0, mock-3.14.0, ddtrace-2.19.0rc1, xdist-3.6.1 +asyncio: mode=Mode.STRICT, asyncio_default_fixture_loop_scope=None, asyncio_default_test_loop_scope=function +collecting ... collected 8 items -test_main.py::test_url_with_format_param[True-gemini/gemini-1.5-flash] PASSED [ 25%] -test_main.py::test_url_with_format_param[True-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 50%] -test_main.py::test_url_with_format_param[False-gemini/gemini-1.5-flash] PASSED [ 75%] -test_main.py::test_url_with_format_param[False-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [100%] +test_main.py::test_url_with_format_param[False-anthropic/claude-3-5-sonnet] PASSED [ 12%] +test_main.py::test_url_with_format_param[False-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 25%] +test_main.py::test_url_with_format_param[False-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 37%] +test_main.py::test_url_with_format_param[False-gemini/gemini-1.5-flash] PASSED [ 50%] +test_main.py::test_url_with_format_param[True-anthropic/claude-3-5-sonnet] PASSED [ 62%] +test_main.py::test_url_with_format_param[True-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 75%] +test_main.py::test_url_with_format_param[True-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0] PASSED [ 87%] +test_main.py::test_url_with_format_param[True-gemini/gemini-1.5-flash] PASSED [100%] =============================== warnings summary =============================== -../../../../../../Library/Frameworks/Python.framework/Versions/3.11/lib/python3.11/site-packages/pydantic/_internal/_config.py:295 - /Library/Frameworks/Python.framework/Versions/3.11/lib/python3.11/site-packages/pydantic/_internal/_config.py:295: PydanticDeprecatedSince20: Support for class-based `config` is deprecated, use ConfigDict instead. Deprecated in Pydantic V2.0 to be removed in V3.0. See Pydantic V2 Migration Guide at https://errors.pydantic.dev/2.10/migration/ - warnings.warn(DEPRECATION_MESSAGE, DeprecationWarning) +tests/litellm/test_main.py::test_url_with_format_param[False-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0] +tests/litellm/test_main.py::test_url_with_format_param[False-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0] +tests/litellm/test_main.py::test_url_with_format_param[True-bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0] +tests/litellm/test_main.py::test_url_with_format_param[True-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0] + /Users/krrishdholakia/Documents/litellm/myenv/lib/python3.13/site-packages/botocore/auth.py:425: DeprecationWarning: datetime.datetime.utcnow() is deprecated and scheduled for removal in a future version. Use timezone-aware objects to represent datetimes in UTC: datetime.datetime.now(datetime.UTC). + datetime_now = datetime.datetime.utcnow() -../../litellm/litellm_core_utils/get_model_cost_map.py:24 - /Users/krrishdholakia/Documents/litellm/litellm/litellm_core_utils/get_model_cost_map.py:24: DeprecationWarning: open_text is deprecated. Use files() instead. Refer to https://importlib-resources.readthedocs.io/en/latest/using.html#migrating-from-legacy for migration advice. - with importlib.resources.open_text( +tests/litellm/test_main.py::test_url_with_format_param[True-anthropic/claude-3-5-sonnet] + /Users/krrishdholakia/Documents/litellm/myenv/lib/python3.13/site-packages/pydantic/main.py:421: UserWarning: Pydantic serializer warnings: + Expected `str` but got `MagicMock` with value `` - serialized value may not be as expected + return self.__pydantic_serializer__.to_python( -../../litellm/utils.py:168 - /Users/krrishdholakia/Documents/litellm/litellm/utils.py:168: DeprecationWarning: open_text is deprecated. Use files() instead. Refer to https://importlib-resources.readthedocs.io/en/latest/using.html#migrating-from-legacy for migration advice. - with resources.open_text( +tests/litellm/test_main.py::test_url_with_format_param[True-bedrock/invoke/anthropic.claude-3-5-sonnet-20240620-v1:0] + /Users/krrishdholakia/Documents/litellm/myenv/lib/python3.13/site-packages/pydantic/main.py:421: UserWarning: Pydantic serializer warnings: + Expected `str` but got `MagicMock` with value `` - serialized value may not be as expected + return self.__pydantic_serializer__.to_python( -- Docs: https://docs.pytest.org/en/stable/how-to/capture-warnings.html -======================== 4 passed, 3 warnings in 2.80s ========================= +======================== 8 passed, 6 warnings in 2.33s ========================= diff --git a/tests/litellm/test_main.py b/tests/litellm/test_main.py index 4997898446..db8241ded2 100644 --- a/tests/litellm/test_main.py +++ b/tests/litellm/test_main.py @@ -176,7 +176,7 @@ async def test_url_with_format_param(model, sync_mode, monkeypatch): response = await acompletion(**args, client=client) print(response) except Exception as e: - print(e) + pass mock_client.assert_called() @@ -186,6 +186,11 @@ async def test_url_with_format_param(model, sync_mode, monkeypatch): json_str = mock_client.call_args.kwargs["data"] else: json_str = json.dumps(mock_client.call_args.kwargs["json"]) + + if isinstance(json_str, bytes): + json_str = json_str.decode("utf-8") + + print(f"type of json_str: {type(json_str)}") assert "png" in json_str assert "jpeg" not in json_str diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index e9eefe1f8e..2c43a9237c 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -1070,7 +1070,7 @@ async def test_bedrock_custom_prompt_template(): pass print(f"mock_client_post.call_args: {mock_client_post.call_args}") - assert "prompt" in mock_client_post.call_args.kwargs["data"] + assert "prompt" in json.loads(mock_client_post.call_args.kwargs["data"]) prompt = json.loads(mock_client_post.call_args.kwargs["data"])["prompt"] assert prompt == "<|im_start|>user\nWhat's AWS?<|im_end|>" diff --git a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py index f62a37515b..7220ffbb2c 100644 --- a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py +++ b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py @@ -68,7 +68,9 @@ def test_bedrock_completion_with_region_name(): ) assert ( mock_post.call_args.kwargs["data"] - == '{"message": "Hello, world!", "chat_history": []}' + == json.dumps({"message": "Hello, world!", "chat_history": []}).encode( + "utf-8" + ) ) # Print the URL and body of the HTTP request. diff --git a/tests/llm_translation/test_openai.py b/tests/llm_translation/test_openai.py index aad1cca567..885b5248df 100644 --- a/tests/llm_translation/test_openai.py +++ b/tests/llm_translation/test_openai.py @@ -472,3 +472,16 @@ class TestOpenAIGPT4OAudioTranscription(BaseLLMAudioTranscriptionTest): def get_custom_llm_provider(self) -> litellm.LlmProviders: return litellm.LlmProviders.OPENAI +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ["gpt-4o"]) +async def test_openai_pdf_url(model): + from litellm.utils import return_raw_request, CallTypes + + request = return_raw_request(CallTypes.completion, { + "model": model, + "messages": [{"role": "user", "content": [{"type": "text", "text": "What is the first page of the PDF?"}, {"type": "file", "file": {"file_id": "https://arxiv.org/pdf/2303.08774"}}]}], + }) + print("request: ", request) + + assert "file_data" in request["raw_request_body"]["messages"][0]["content"][1]["file"] + diff --git a/tests/llm_translation/test_optional_params.py b/tests/llm_translation/test_optional_params.py index e207e367e4..d1d10d5064 100644 --- a/tests/llm_translation/test_optional_params.py +++ b/tests/llm_translation/test_optional_params.py @@ -597,7 +597,7 @@ def test_get_optional_params_num_retries(): """ Relevant issue - https://github.com/BerriAI/litellm/issues/5124 """ - with patch("litellm.main.get_optional_params", new=MagicMock()) as mock_client: + with patch("litellm.main.get_optional_params", new=MagicMock(return_value={"max_retries": 0})) as mock_client: _ = litellm.completion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello world"}], diff --git a/tests/llm_translation/test_unit_test_bedrock_invoke.py b/tests/llm_translation/test_unit_test_bedrock_invoke.py index a156c8722d..5ed0b4aae3 100644 --- a/tests/llm_translation/test_unit_test_bedrock_invoke.py +++ b/tests/llm_translation/test_unit_test_bedrock_invoke.py @@ -99,7 +99,7 @@ def test_sign_request_basic(mock_aws_request, mock_sigv4_auth, bedrock_transform api_base = "https://bedrock-runtime.us-east-1.amazonaws.com" # Call the method - result = bedrock_transformer.sign_request( + result, _ = bedrock_transformer.sign_request( headers=headers, optional_params=optional_params, request_data=request_data, diff --git a/tests/local_testing/test_router.py b/tests/local_testing/test_router.py index 06cf75705a..43e9a17525 100644 --- a/tests/local_testing/test_router.py +++ b/tests/local_testing/test_router.py @@ -334,7 +334,7 @@ async def test_router_retries(sync_mode): { "model_name": "gpt-3.5-turbo", "litellm_params": { - "model": "azure/chatgpt-v-3", + "model": "azure/gpt-4o-new-test", "api_key": os.getenv("AZURE_API_KEY"), "api_base": os.getenv("AZURE_API_BASE"), "api_version": os.getenv("AZURE_API_VERSION"),