From d7efe22d6ef35c3821a7f94a58e6d71079da9db8 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 20 Feb 2025 14:09:03 -0800 Subject: [PATCH] feat(bedrock/rerank/transformation.py): include search units for bedrock rerank result Resolves https://github.com/BerriAI/litellm/issues/7258#issuecomment-2671557137 --- litellm/llms/bedrock/rerank/transformation.py | 4 +++- .../litellm/llms/bedrock/rerank/transformation.py | 14 ++++++++++++++ tests/llm_translation/base_rerank_unit_tests.py | 14 +++++++------- 3 files changed, 24 insertions(+), 8 deletions(-) create mode 100644 tests/litellm/llms/bedrock/rerank/transformation.py diff --git a/litellm/llms/bedrock/rerank/transformation.py b/litellm/llms/bedrock/rerank/transformation.py index 7dc9b0aab1f..a5380febe9c 100644 --- a/litellm/llms/bedrock/rerank/transformation.py +++ b/litellm/llms/bedrock/rerank/transformation.py @@ -91,7 +91,9 @@ class BedrockRerankConfig: example input: {"results":[{"index":0,"relevanceScore":0.6847912669181824},{"index":1,"relevanceScore":0.5980774760246277}]} """ - _billed_units = RerankBilledUnits(**response.get("usage", {})) + _billed_units = RerankBilledUnits( + **response.get("usage", {"search_units": 1}) + ) # by default 1 search unit _tokens = RerankTokens(**response.get("usage", {})) rerank_meta = RerankResponseMeta(billed_units=_billed_units, tokens=_tokens) diff --git a/tests/litellm/llms/bedrock/rerank/transformation.py b/tests/litellm/llms/bedrock/rerank/transformation.py new file mode 100644 index 00000000000..870a7cb1f1e --- /dev/null +++ b/tests/litellm/llms/bedrock/rerank/transformation.py @@ -0,0 +1,14 @@ +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 diff --git a/tests/llm_translation/base_rerank_unit_tests.py b/tests/llm_translation/base_rerank_unit_tests.py index 54f6009fc66..b5d0302b785 100644 --- a/tests/llm_translation/base_rerank_unit_tests.py +++ b/tests/llm_translation/base_rerank_unit_tests.py @@ -52,13 +52,13 @@ def assert_response_shape(response, custom_llm_provider): response.meta["api_version"]["version"], expected_api_version_shape["version"], ) - assert isinstance( - response.meta["billed_units"], expected_meta_shape["billed_units"] - ) - assert isinstance( - response.meta["billed_units"]["search_units"], - expected_billed_units_shape["search_units"], - ) + assert isinstance( + response.meta["billed_units"], expected_meta_shape["billed_units"] + ) + assert isinstance( + response.meta["billed_units"]["search_units"], + expected_billed_units_shape["search_units"], + ) class BaseLLMRerankTest(ABC):