Merge pull request #26685 from BerriAI/litellm_bedrock_retrievalconfig_passthrough2

feat(vector-stores): support Bedrock retrievalConfiguration passthrough
This commit is contained in:
Mateo Wang 2026-04-29 12:36:14 -07:00 • committed by GitHub
commit 295a36aa69
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 159 additions and 9 deletions

View file

@ -95,6 +95,7 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform search request for Azure AI Search API

View file

@ -59,6 +59,7 @@ class BaseVectorStoreConfig:
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
pass
@ -70,6 +71,7 @@ class BaseVectorStoreConfig:
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
"""
Optional async version of transform_search_vector_store_request.
@ -84,6 +86,7 @@ class BaseVectorStoreConfig:
api_base=api_base,
litellm_logging_obj=litellm_logging_obj,
litellm_params=litellm_params,
extra_body=extra_body,
)
@abstractmethod

View file

@ -1,14 +1,16 @@
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
from urllib.parse import urlparse
import httpx
from litellm._logging import verbose_logger
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.types.integrations.rag.bedrock_knowledgebase import (
BedrockKBContent,
BedrockKBResponse,
BedrockKBRetrievalConfiguration,
BedrockKBResponse,
BedrockKBRetrievalQuery,
)
from litellm.types.router import GenericLiteLLMParams
@ -202,6 +204,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
if isinstance(query, list):
query = " ".join(query)
@ -213,24 +216,46 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
}
retrieval_config: Dict[str, Any] = {}
if isinstance(extra_body, dict):
retrieval_config = deepcopy(
extra_body.get("retrievalConfiguration")
or extra_body.get("retrieval_configuration")
or {}
)
max_results = vector_store_search_optional_params.get("max_num_results")
if max_results is not None:
existing_number_of_results = retrieval_config.get(
"vectorSearchConfiguration", {}
).get("numberOfResults")
if (
existing_number_of_results is not None
and existing_number_of_results != max_results
):
verbose_logger.debug(
"Overriding extra_body retrievalConfiguration.vectorSearchConfiguration.numberOfResults (%s) with max_num_results=%s",
existing_number_of_results,
max_results,
)
retrieval_config.setdefault("vectorSearchConfiguration", {})[
"numberOfResults"
] = max_results
filters = vector_store_search_optional_params.get("filters")
if filters is not None:
existing_filter = retrieval_config.get("vectorSearchConfiguration", {}).get(
"filter"
)
if existing_filter is not None and existing_filter != filters:
verbose_logger.debug(
"Overriding extra_body retrievalConfiguration.vectorSearchConfiguration.filter with filters from vector_store_search_optional_params"
)
retrieval_config.setdefault("vectorSearchConfiguration", {})[
"filter"
] = filters
if retrieval_config:
# Create a properly typed retrieval configuration
typed_retrieval_config: BedrockKBRetrievalConfiguration = {}
if "vectorSearchConfiguration" in retrieval_config:
typed_retrieval_config["vectorSearchConfiguration"] = retrieval_config[
"vectorSearchConfiguration"
]
request_body["retrievalConfiguration"] = typed_retrieval_config
request_body["retrievalConfiguration"] = cast(
BedrockKBRetrievalConfiguration, retrieval_config
)
litellm_logging_obj.model_call_details["query"] = query
return url, request_body

View file

@ -8609,6 +8609,7 @@ class BaseLLMHTTPHandler:
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
)
else:
(
@ -8621,6 +8622,7 @@ class BaseLLMHTTPHandler:
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
)
all_optional_params: Dict[str, Any] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
@ -8721,6 +8723,7 @@ class BaseLLMHTTPHandler:
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
)
all_optional_params: Dict[str, Any] = dict(litellm_params)

View file

@ -118,6 +118,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
"""
Transform search request to Gemini's generateContent format.

View file

@ -130,6 +130,7 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform search request for Azure AI Search API

View file

@ -106,6 +106,7 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
url = f"{api_base}/{vector_store_id}/search"
typed_request_body = VectorStoreSearchRequest(

View file

@ -80,6 +80,7 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
url = f"{api_base}/{vector_store_id}/search"
_, request_body = super().transform_search_vector_store_request(
@ -89,5 +90,6 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig):
api_base=api_base,
litellm_logging_obj=litellm_logging_obj,
litellm_params=litellm_params,
extra_body=extra_body,
)
return url, request_body

View file

@ -102,6 +102,7 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
"""RAGFlow vector stores are management-only, search is not supported."""
raise NotImplementedError(

View file

@ -79,6 +79,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
"""Sync version - generates embedding synchronously."""
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name
@ -140,6 +141,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
"""Async version - generates embedding asynchronously."""
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name

View file

@ -100,6 +100,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform search request for Vertex AI RAG API

View file

@ -107,6 +107,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform search request for Vertex AI RAG API

View file

@ -354,6 +354,7 @@ async def test_bedrock_kb_request_body_has_transformed_filters(
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=litellm_params_dict,
extra_body=None,
)
)
captured_request_body["url"] = url

View file

@ -21,7 +21,112 @@ def test_transform_search_request():
api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases",
litellm_logging_obj=mock_log,
litellm_params={},
extra_body=None,
)
assert url.endswith("/kb123/retrieve")
assert body["retrievalQuery"].get("text") == "hello"
def test_transform_search_request_uses_only_retrieval_config_from_extra_body():
config = BedrockVectorStoreConfig()
mock_log = MagicMock()
mock_log.model_call_details = {}
url, body = config.transform_search_vector_store_request(
vector_store_id="kb123",
query="hello",
vector_store_search_optional_params={},
api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases",
litellm_logging_obj=mock_log,
litellm_params={},
extra_body={
"retrievalConfiguration": {
"vectorSearchConfiguration": {
"overrideSearchType": "HYBRID",
"numberOfResults": 8,
}
},
"unrelatedField": {"should_not": "be_forwarded"},
},
)
assert url.endswith("/kb123/retrieve")
assert body["retrievalQuery"].get("text") == "hello"
assert (
body["retrievalConfiguration"]["vectorSearchConfiguration"][
"overrideSearchType"
]
== "HYBRID"
)
assert "unrelatedField" not in body
def test_transform_search_request_does_not_mutate_extra_body_and_overrides_number_of_results():
config = BedrockVectorStoreConfig()
mock_log = MagicMock()
mock_log.model_call_details = {}
extra_body = {
"retrievalConfiguration": {
"vectorSearchConfiguration": {
"overrideSearchType": "HYBRID",
"numberOfResults": 8,
}
}
}
_, body = config.transform_search_vector_store_request(
vector_store_id="kb123",
query="hello",
vector_store_search_optional_params={"max_num_results": 10},
api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases",
litellm_logging_obj=mock_log,
litellm_params={},
extra_body=extra_body,
)
assert (
body["retrievalConfiguration"]["vectorSearchConfiguration"]["numberOfResults"]
== 10
)
assert (
extra_body["retrievalConfiguration"]["vectorSearchConfiguration"][
"numberOfResults"
]
== 8
)
def test_transform_search_request_overrides_filter_without_mutating_extra_body():
config = BedrockVectorStoreConfig()
mock_log = MagicMock()
mock_log.model_call_details = {}
extra_body = {
"retrievalConfiguration": {
"vectorSearchConfiguration": {
"filter": {"equals": {"key": "tenant", "value": "a"}}
}
}
}
new_filter = {"equals": {"key": "tenant", "value": "b"}}
_, body = config.transform_search_vector_store_request(
vector_store_id="kb123",
query="hello",
vector_store_search_optional_params={"filters": new_filter},
api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases",
litellm_logging_obj=mock_log,
litellm_params={},
extra_body=extra_body,
)
assert (
body["retrievalConfiguration"]["vectorSearchConfiguration"]["filter"]
== new_filter
)
assert (
extra_body["retrievalConfiguration"]["vectorSearchConfiguration"]["filter"][
"equals"
]["value"]
== "a"
)

View file

@ -59,6 +59,7 @@ class TestS3VectorsVectorStoreConfig:
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={},
extra_body=None,
)
def test_transform_search_response(self):

View file

@ -267,6 +267,7 @@ class TestRAGFlowVectorStore(BaseVectorStoreTest):
api_base="http://localhost:9380",
litellm_logging_obj=logging_obj,
litellm_params={},
extra_body=None,
)
def test_transform_search_vector_store_response_not_implemented(self):