From ad7617c74f586a0e1b8a923241933f03fa2d32e3 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 20 Feb 2025 13:53:49 -0800 Subject: [PATCH] feat(bedrock/rerank): infer model region if model given as arn --- docs/my-website/docs/providers/bedrock.md | 1 + litellm/llms/bedrock/base_aws_llm.py | 78 ++++++++++++++----- .../base_invoke_transformation.py | 55 ------------- litellm/llms/bedrock/image/image_handler.py | 2 +- litellm/llms/bedrock/rerank/handler.py | 27 +++++-- litellm/rerank_api/main.py | 1 + tests/litellm/rerank_api/test_main.py | 47 +++++++++++ 7 files changed, 129 insertions(+), 82 deletions(-) create mode 100644 tests/litellm/rerank_api/test_main.py diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index 744be74c093..00fe45e99ea 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -10,6 +10,7 @@ ALL Bedrock models (Anthropic, Meta, Deepseek, Mistral, Amazon, etc.) are Suppor | Provider Route on LiteLLM | `bedrock/`, [`bedrock/converse/`](#set-converse--invoke-route), [`bedrock/invoke/`](#set-invoke-route), [`bedrock/converse_like/`](#calling-via-internal-proxy), [`bedrock/llama/`](#deepseek-not-r1), [`bedrock/deepseek_r1/`](#deepseek-r1) | | Provider Doc | [Amazon Bedrock ↗](https://docs.aws.amazon.com/bedrock/latest/userguide/what-is-bedrock.html) | | Supported OpenAI Endpoints | `/chat/completions`, `/completions`, `/embeddings`, `/images/generations` | +| Rerank Endpoint | `/rerank` | | Pass-through Endpoint | [Supported](../pass_through/bedrock.md) | diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 7b04b2c02aa..681d551e43f 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -200,6 +200,61 @@ class BaseAWSLLM: self.iam_cache.set_cache(cache_key, credentials, ttl=_cache_ttl) return credentials + def _get_aws_region_from_model_arn(self, model: Optional[str]) -> Optional[str]: + try: + # First check if the string contains the expected prefix + if not isinstance(model, str) or "arn:aws:bedrock" not in model: + return None + + # Split the ARN and check if we have enough parts + parts = model.split(":") + if len(parts) < 4: + return None + + # Get the region from the correct position + region = parts[3] + if not region: # Check if region is empty + return None + + return region + except Exception: + # Catch any unexpected errors and return None + return None + + def _get_aws_region_name( + self, optional_params: dict, model: Optional[str] = None + ) -> str: + """ + Get the AWS region name from the environment variables + """ + aws_region_name = optional_params.get("aws_region_name", None) + ### SET REGION NAME ### + if aws_region_name is None: + # check model arn # + aws_region_name = self._get_aws_region_from_model_arn(model) + # check env # + litellm_aws_region_name = get_secret("AWS_REGION_NAME", None) + + if ( + aws_region_name is None + and litellm_aws_region_name is not None + and isinstance(litellm_aws_region_name, str) + ): + aws_region_name = litellm_aws_region_name + + standard_aws_region_name = get_secret("AWS_REGION", None) + if ( + aws_region_name is None + and standard_aws_region_name is not None + and isinstance(standard_aws_region_name, str) + ): + aws_region_name = standard_aws_region_name + + if aws_region_name is None: + aws_region_name = "us-west-2" + + return aws_region_name + def _auth_with_web_identity_token( self, aws_web_identity_token: str, @@ -408,7 +463,7 @@ class BaseAWSLLM: return endpoint_url, proxy_endpoint_url def _get_boto_credentials_from_optional_params( - self, optional_params: dict + self, optional_params: dict, model: Optional[str] = None ) -> Boto3CredentialsInfo: """ Get boto3 credentials from optional params @@ -428,7 +483,7 @@ class BaseAWSLLM: aws_secret_access_key = optional_params.pop("aws_secret_access_key", None) aws_access_key_id = optional_params.pop("aws_access_key_id", None) aws_session_token = optional_params.pop("aws_session_token", None) - aws_region_name = optional_params.pop("aws_region_name", None) + aws_region_name = self._get_aws_region_name(optional_params, model) aws_role_name = optional_params.pop("aws_role_name", None) aws_session_name = optional_params.pop("aws_session_name", None) aws_profile_name = optional_params.pop("aws_profile_name", None) @@ -438,25 +493,6 @@ class BaseAWSLLM: "aws_bedrock_runtime_endpoint", None ) # https://bedrock-runtime.{region_name}.amazonaws.com - ### SET REGION NAME ### - if aws_region_name is None: - # check env # - litellm_aws_region_name = get_secret_str("AWS_REGION_NAME", None) - - if litellm_aws_region_name is not None and isinstance( - litellm_aws_region_name, str - ): - aws_region_name = litellm_aws_region_name - - standard_aws_region_name = get_secret_str("AWS_REGION", None) - if standard_aws_region_name is not None and isinstance( - standard_aws_region_name, str - ): - aws_region_name = standard_aws_region_name - - if aws_region_name is None: - aws_region_name = "us-west-2" - credentials: Credentials = self.get_credentials( aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, 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 f3690577446..cb535d507bd 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -598,61 +598,6 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): ) return modelId - def get_aws_region_from_model_arn(self, model: Optional[str]) -> Optional[str]: - try: - # First check if the string contains the expected prefix - if not isinstance(model, str) or "arn:aws:bedrock" not in model: - return None - - # Split the ARN and check if we have enough parts - parts = model.split(":") - if len(parts) < 4: - return None - - # Get the region from the correct position - region = parts[3] - if not region: # Check if region is empty - return None - - return region - except Exception: - # Catch any unexpected errors and return None - return None - - def _get_aws_region_name( - self, optional_params: dict, model: Optional[str] = None - ) -> str: - """ - Get the AWS region name from the environment variables - """ - aws_region_name = optional_params.get("aws_region_name", None) - ### SET REGION NAME ### - if aws_region_name is None: - # check model arn # - aws_region_name = self.get_aws_region_from_model_arn(model) - # check env # - litellm_aws_region_name = get_secret("AWS_REGION_NAME", None) - - if ( - aws_region_name is None - and litellm_aws_region_name is not None - and isinstance(litellm_aws_region_name, str) - ): - aws_region_name = litellm_aws_region_name - - standard_aws_region_name = get_secret("AWS_REGION", None) - if ( - aws_region_name is None - and standard_aws_region_name is not None - and isinstance(standard_aws_region_name, str) - ): - aws_region_name = standard_aws_region_name - - if aws_region_name is None: - aws_region_name = "us-west-2" - - return aws_region_name - def _get_model_id_from_model_with_spec( self, model: str, diff --git a/litellm/llms/bedrock/image/image_handler.py b/litellm/llms/bedrock/image/image_handler.py index 5b14833f426..4bd63fd21b5 100644 --- a/litellm/llms/bedrock/image/image_handler.py +++ b/litellm/llms/bedrock/image/image_handler.py @@ -163,7 +163,7 @@ class BedrockImageGeneration(BaseAWSLLM): except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") boto3_credentials_info = self._get_boto_credentials_from_optional_params( - optional_params + optional_params, model ) ### SET RUNTIME ENDPOINT ### diff --git a/litellm/llms/bedrock/rerank/handler.py b/litellm/llms/bedrock/rerank/handler.py index 3683be06b6c..049dc8cc4fa 100644 --- a/litellm/llms/bedrock/rerank/handler.py +++ b/litellm/llms/bedrock/rerank/handler.py @@ -4,8 +4,11 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast import httpx import litellm +from litellm import verbose_logger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + HTTPHandler, _get_httpx_client, get_async_httpx_client, ) @@ -27,8 +30,10 @@ class BedrockRerankHandler(BaseAWSLLM): async def arerank( self, prepared_request: BedrockPreparedRequest, + client: Optional[AsyncHTTPHandler] = None, ): - client = get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK) + if client is None: + client = get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK) try: response = await client.post(url=prepared_request["endpoint_url"], headers=prepared_request["prepped"].headers, data=prepared_request["body"]) # type: ignore response.raise_for_status() @@ -54,7 +59,9 @@ class BedrockRerankHandler(BaseAWSLLM): _is_async: Optional[bool] = False, api_base: Optional[str] = None, extra_headers: Optional[dict] = None, + client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None, ) -> RerankResponse: + request_data = RerankRequest( model=model, query=query, @@ -66,6 +73,7 @@ class BedrockRerankHandler(BaseAWSLLM): data = BedrockRerankConfig()._transform_request(request_data) prepared_request = self._prepare_request( + model=model, optional_params=optional_params, api_base=api_base, extra_headers=extra_headers, @@ -83,9 +91,10 @@ class BedrockRerankHandler(BaseAWSLLM): ) if _is_async: - return self.arerank(prepared_request) # type: ignore + return self.arerank(prepared_request, client=client if client is not None and isinstance(client, AsyncHTTPHandler) else None) # type: ignore - client = _get_httpx_client() + if client is None or not isinstance(client, HTTPHandler): + client = _get_httpx_client() try: response = client.post(url=prepared_request["endpoint_url"], headers=prepared_request["prepped"].headers, data=prepared_request["body"]) # type: ignore response.raise_for_status() @@ -95,10 +104,18 @@ class BedrockRerankHandler(BaseAWSLLM): except httpx.TimeoutException: raise BedrockError(status_code=408, message="Timeout error occurred.") - return BedrockRerankConfig()._transform_response(response.json()) + logging_obj.post_call( + original_response=response.text, + api_key="", + ) + + response_json = response.json() + + return BedrockRerankConfig()._transform_response(response_json) def _prepare_request( self, + model: str, api_base: Optional[str], extra_headers: Optional[dict], data: dict, @@ -110,7 +127,7 @@ class BedrockRerankHandler(BaseAWSLLM): except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") boto3_credentials_info = self._get_boto_credentials_from_optional_params( - optional_params + optional_params, model ) ### SET RUNTIME ENDPOINT ### diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 8ec05ddadf4..5986c187a31 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -295,6 +295,7 @@ def rerank( # noqa: PLR0915 optional_params=optional_params.model_dump(exclude_unset=True), api_base=api_base, logging_obj=litellm_logging_obj, + client=client, ) else: raise ValueError(f"Unsupported provider: {_custom_llm_provider}") diff --git a/tests/litellm/rerank_api/test_main.py b/tests/litellm/rerank_api/test_main.py new file mode 100644 index 00000000000..f55e05ea0f6 --- /dev/null +++ b/tests/litellm/rerank_api/test_main.py @@ -0,0 +1,47 @@ +import json +import os +import sys + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path +from unittest.mock import MagicMock, patch + +from litellm import rerank +from litellm.llms.custom_httpx.http_handler import HTTPHandler + + +def test_rerank_infer_region_from_model_arn(): + mock_response = MagicMock() + args = { + "model": "bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0", + "query": "hello", + "documents": ["hello", "world"], + } + + def return_val(): + return { + "results": [ + {"index": 0, "relevanceScore": 0.6716859340667725}, + {"index": 1, "relevanceScore": 0.0004994205664843321}, + ] + } + + mock_response.json = return_val + mock_response.headers = {"key": "value"} + mock_response.status_code = 200 + + client = HTTPHandler() + + with patch.object(client, "post", return_value=mock_response) as mock_post: + rerank( + model=args["model"], + query=args["query"], + documents=args["documents"], + client=client, + ) + mock_post.assert_called_once() + print(f"mock_post.call_args: {mock_post.call_args.kwargs}")