mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #26685 from BerriAI/litellm_bedrock_retrievalconfig_passthrough2
feat(vector-stores): support Bedrock retrievalConfiguration passthrough
This commit is contained in:
commit
295a36aa69
16 changed files with 159 additions and 9 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue