From cfe4bc678e612bcc37831b69582e886bf05e03da Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 28 Apr 2026 15:37:23 +0530 Subject: [PATCH] feat(vector-stores): support provider-specific Bedrock retrieval config Route vector store search `extra_body` into provider transformers and handle Bedrock `retrievalConfiguration` explicitly so only intended provider-specific fields are forwarded. Made-with: Cursor --- .../azure_ai/vector_stores/transformation.py | 1 + .../base_llm/vector_store/transformation.py | 3 ++ .../bedrock/vector_stores/transformation.py | 23 +++++++----- litellm/llms/custom_httpx/llm_http_handler.py | 3 ++ .../gemini/vector_stores/transformation.py | 1 + .../milvus/vector_stores/transformation.py | 1 + .../openai/vector_stores/transformation.py | 1 + .../pg_vector/vector_stores/transformation.py | 2 ++ .../ragflow/vector_stores/transformation.py | 1 + .../vector_stores/transformation.py | 2 ++ .../vector_stores/rag_api/transformation.py | 1 + .../search_api/transformation.py | 1 + ...est_bedrock_vector_store_transformation.py | 35 +++++++++++++++++++ 13 files changed, 66 insertions(+), 9 deletions(-) diff --git a/litellm/llms/azure_ai/vector_stores/transformation.py b/litellm/llms/azure_ai/vector_stores/transformation.py index b62acb65166..d2c8206ca9a 100644 --- a/litellm/llms/azure_ai/vector_stores/transformation.py +++ b/litellm/llms/azure_ai/vector_stores/transformation.py @@ -92,6 +92,7 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, diff --git a/litellm/llms/base_llm/vector_store/transformation.py b/litellm/llms/base_llm/vector_store/transformation.py index 5fbf0a4b19f..49d2f72db7c 100644 --- a/litellm/llms/base_llm/vector_store/transformation.py +++ b/litellm/llms/base_llm/vector_store/transformation.py @@ -56,6 +56,7 @@ class BaseVectorStoreConfig: vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -67,6 +68,7 @@ class BaseVectorStoreConfig: vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -81,6 +83,7 @@ class BaseVectorStoreConfig: vector_store_id=vector_store_id, query=query, vector_store_search_optional_params=vector_store_search_optional_params, + extra_body=extra_body, api_base=api_base, litellm_logging_obj=litellm_logging_obj, litellm_params=litellm_params, diff --git a/litellm/llms/bedrock/vector_stores/transformation.py b/litellm/llms/bedrock/vector_stores/transformation.py index 4da0a7c7791..5b0b64f2429 100644 --- a/litellm/llms/bedrock/vector_stores/transformation.py +++ b/litellm/llms/bedrock/vector_stores/transformation.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast from urllib.parse import urlparse import httpx @@ -7,8 +7,8 @@ from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreCon 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 @@ -199,6 +199,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -213,6 +214,14 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): } retrieval_config: Dict[str, Any] = {} + from litellm import verbose_logger + + if isinstance(extra_body, dict): + retrieval_config = dict( + 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: retrieval_config.setdefault("vectorSearchConfiguration", {})[ @@ -224,13 +233,9 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): "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 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index b9ada079f6b..99f748c0c1a 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -8582,6 +8582,7 @@ class BaseLLMHTTPHandler: vector_store_id=vector_store_id, query=query, vector_store_search_optional_params=vector_store_search_optional_params, + extra_body=extra_body, api_base=api_base, litellm_logging_obj=logging_obj, litellm_params=dict(litellm_params), @@ -8594,6 +8595,7 @@ class BaseLLMHTTPHandler: vector_store_id=vector_store_id, query=query, vector_store_search_optional_params=vector_store_search_optional_params, + extra_body=extra_body, api_base=api_base, litellm_logging_obj=logging_obj, litellm_params=dict(litellm_params), @@ -8694,6 +8696,7 @@ class BaseLLMHTTPHandler: vector_store_id=vector_store_id, query=query, vector_store_search_optional_params=vector_store_search_optional_params, + extra_body=extra_body, api_base=api_base, litellm_logging_obj=logging_obj, litellm_params=dict(litellm_params), diff --git a/litellm/llms/gemini/vector_stores/transformation.py b/litellm/llms/gemini/vector_stores/transformation.py index e6e8369643e..b6cace066c8 100644 --- a/litellm/llms/gemini/vector_stores/transformation.py +++ b/litellm/llms/gemini/vector_stores/transformation.py @@ -115,6 +115,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, diff --git a/litellm/llms/milvus/vector_stores/transformation.py b/litellm/llms/milvus/vector_stores/transformation.py index fcf5d14db7c..8c08b783387 100644 --- a/litellm/llms/milvus/vector_stores/transformation.py +++ b/litellm/llms/milvus/vector_stores/transformation.py @@ -127,6 +127,7 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, diff --git a/litellm/llms/openai/vector_stores/transformation.py b/litellm/llms/openai/vector_stores/transformation.py index c763ed1c8da..b6eae390d93 100644 --- a/litellm/llms/openai/vector_stores/transformation.py +++ b/litellm/llms/openai/vector_stores/transformation.py @@ -103,6 +103,7 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, diff --git a/litellm/llms/pg_vector/vector_stores/transformation.py b/litellm/llms/pg_vector/vector_stores/transformation.py index ba87a8f2b01..8261036cae0 100644 --- a/litellm/llms/pg_vector/vector_stores/transformation.py +++ b/litellm/llms/pg_vector/vector_stores/transformation.py @@ -77,6 +77,7 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -86,6 +87,7 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig): vector_store_id=vector_store_id, query=query, vector_store_search_optional_params=vector_store_search_optional_params, + extra_body=extra_body, api_base=api_base, litellm_logging_obj=litellm_logging_obj, litellm_params=litellm_params, diff --git a/litellm/llms/ragflow/vector_stores/transformation.py b/litellm/llms/ragflow/vector_stores/transformation.py index ed5397eef0c..ae28222c3ce 100644 --- a/litellm/llms/ragflow/vector_stores/transformation.py +++ b/litellm/llms/ragflow/vector_stores/transformation.py @@ -99,6 +99,7 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, diff --git a/litellm/llms/s3_vectors/vector_stores/transformation.py b/litellm/llms/s3_vectors/vector_stores/transformation.py index 19b59769863..0cf86358873 100644 --- a/litellm/llms/s3_vectors/vector_stores/transformation.py +++ b/litellm/llms/s3_vectors/vector_stores/transformation.py @@ -76,6 +76,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, @@ -137,6 +138,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, diff --git a/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py index 4baa5774c48..b3fcf4b394c 100644 --- a/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py @@ -97,6 +97,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index 179bd7aeff1..47dac8dca32 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -104,6 +104,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): vector_store_id: str, query: Union[str, List[str]], vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams, + extra_body: Optional[Dict[str, Any]], api_base: str, litellm_logging_obj: LiteLLMLoggingObj, litellm_params: dict, diff --git a/tests/test_litellm/llms/bedrock/vector_stores/test_bedrock_vector_store_transformation.py b/tests/test_litellm/llms/bedrock/vector_stores/test_bedrock_vector_store_transformation.py index bcd4c620055..5fa48703b12 100644 --- a/tests/test_litellm/llms/bedrock/vector_stores/test_bedrock_vector_store_transformation.py +++ b/tests/test_litellm/llms/bedrock/vector_stores/test_bedrock_vector_store_transformation.py @@ -18,6 +18,7 @@ def test_transform_search_request(): vector_store_id="kb123", query="hello", vector_store_search_optional_params={}, + extra_body=None, api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases", litellm_logging_obj=mock_log, litellm_params={}, @@ -25,3 +26,37 @@ def test_transform_search_request(): 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={}, + extra_body={ + "retrievalConfiguration": { + "vectorSearchConfiguration": { + "overrideSearchType": "HYBRID", + "numberOfResults": 8, + } + }, + "unrelatedField": {"should_not": "be_forwarded"}, + }, + api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases", + litellm_logging_obj=mock_log, + litellm_params={}, + ) + + assert url.endswith("/kb123/retrieve") + assert body["retrievalQuery"].get("text") == "hello" + assert ( + body["retrievalConfiguration"]["vectorSearchConfiguration"][ + "overrideSearchType" + ] + == "HYBRID" + ) + assert "unrelatedField" not in body