[Feat] Add OpenAI Search Vector Store Operation (#12018)

* add BaseVectorStoreTransformation

* fix BaseVectorStoreTransformation

* add OpenAIVectorStoreTransformation

* fix transform

* add search, asearch vector stores

* add skeleton for vector store searching

* fix VectorStoreSearchOptionalRequestParams

* fix VectorStoreRequestUtils

* fix litellm.asearch/litellm.search

* fix BaseVectorStoreConfig

* add vector_store_search_handler to llm http handler

* use llm http handler for searching vector stores

* fix base vector store config

* fix vector_store_search_handler

* async_vector_store_search_handler

* add conftest

* add BaseVectorStoreTest

* move litellm.integrations.vector_store_integrations

* fix working OAI OpenAIVectorStoreConfig

* add Search vector store

* add OpenAI Vector Stores
This commit is contained in:
Ishaan Jaff 2025-06-24 15:52:43 -07:00 • committed by GitHub
parent 97da33494a
commit 2bb8048864
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 1252 additions and 7 deletions

View file

@ -17,6 +17,7 @@ LiteLLM integrates with vector stores, allowing your models to access your organ
## Supported Vector Stores
- [Bedrock Knowledge Bases](https://aws.amazon.com/bedrock/knowledge-bases/)
- [OpenAI Vector Stores](https://platform.openai.com/docs/api-reference/vector-stores/search)
## Quick Start

View file

@ -1141,6 +1141,7 @@ from .router import Router
from .assistants.main import *
from .batches.main import *
from .images.main import *
from .vector_stores import *
from .batch_completion.main import * # type: ignore
from .rerank_api.main import *
from .llms.anthropic.experimental_pass_through.messages.handler import *

View file

@ -0,0 +1,409 @@
# +-------------------------------------------------------------+
#
# Add Bedrock Knowledge Base Context to your LLM calls
#
# +-------------------------------------------------------------+
# Thank you users! We ❤️ you! - Krrish & Ishaan
import json
from datetime import datetime
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.vector_store_integrations.base_vector_store import (
BaseVectorStore,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.integrations.rag.bedrock_knowledgebase import (
BedrockKBContent,
BedrockKBGuardrailConfiguration,
BedrockKBRequest,
BedrockKBResponse,
BedrockKBRetrievalConfiguration,
BedrockKBRetrievalQuery,
BedrockKBRetrievalResult,
)
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
from litellm.types.utils import StandardLoggingVectorStoreRequest
from litellm.types.vector_stores import (
VectorStoreResultContent,
VectorStoreSearchResponse,
VectorStoreSearchResult,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import StandardCallbackDynamicParams
else:
StandardCallbackDynamicParams = Any
class BedrockVectorStore(BaseVectorStore, BaseAWSLLM):
CONTENT_PREFIX_STRING = "Context: \n\n"
CUSTOM_LLM_PROVIDER = "bedrock"
def __init__(
self,
**kwargs,
):
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
# store kwargs as optional_params
self.optional_params = kwargs
super().__init__(**kwargs)
BaseAWSLLM.__init__(self)
async def async_get_chat_completion_prompt(
self,
model: str,
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
litellm_logging_obj: LiteLLMLoggingObj,
tools: Optional[List[Dict]] = None,
prompt_label: Optional[str] = None,
) -> Tuple[str, List[AllMessageValues], dict]:
"""
Retrieves the context from the Bedrock Knowledge Base and appends it to the messages.
"""
if litellm.vector_store_registry is None:
return model, messages, non_default_params
vector_store_ids = litellm.vector_store_registry.pop_vector_store_ids_to_run(
non_default_params=non_default_params, tools=tools
)
vector_store_request_metadata: List[StandardLoggingVectorStoreRequest] = []
if vector_store_ids:
for vector_store_id in vector_store_ids:
start_time = datetime.now()
query = self._get_kb_query_from_messages(messages)
bedrock_kb_response = await self.make_bedrock_kb_retrieve_request(
knowledge_base_id=vector_store_id,
query=query,
non_default_params=non_default_params,
)
verbose_logger.debug(
f"Bedrock Knowledge Base Response: {bedrock_kb_response}"
)
(
context_message,
context_string,
) = self.get_chat_completion_message_from_bedrock_kb_response(
bedrock_kb_response
)
if context_message is not None:
messages.append(context_message)
#################################################################################################
########## LOGGING for Standard Logging Payload, Langfuse, s3, LiteLLM DB etc. ##################
#################################################################################################
vector_store_search_response: VectorStoreSearchResponse = (
self.transform_bedrock_kb_response_to_vector_store_search_response(
bedrock_kb_response=bedrock_kb_response, query=query
)
)
vector_store_request_metadata.append(
StandardLoggingVectorStoreRequest(
vector_store_id=vector_store_id,
query=query,
vector_store_search_response=vector_store_search_response,
custom_llm_provider=self.CUSTOM_LLM_PROVIDER,
start_time=start_time.timestamp(),
end_time=datetime.now().timestamp(),
)
)
litellm_logging_obj.model_call_details[
"vector_store_request_metadata"
] = vector_store_request_metadata
return model, messages, non_default_params
def transform_bedrock_kb_response_to_vector_store_search_response(
self,
bedrock_kb_response: BedrockKBResponse,
query: str,
) -> VectorStoreSearchResponse:
"""
Transform a BedrockKBResponse to a VectorStoreSearchResponse
"""
retrieval_results: Optional[
List[BedrockKBRetrievalResult]
] = bedrock_kb_response.get("retrievalResults", None)
vector_store_search_response: VectorStoreSearchResponse = (
VectorStoreSearchResponse(search_query=query, data=[])
)
if retrieval_results is None:
return vector_store_search_response
vector_search_response_data: List[VectorStoreSearchResult] = []
for retrieval_result in retrieval_results:
content: Optional[BedrockKBContent] = retrieval_result.get("content", None)
if content is None:
continue
content_text: Optional[str] = content.get("text", None)
if content_text is None:
continue
vector_store_search_result: VectorStoreSearchResult = (
VectorStoreSearchResult(
score=retrieval_result.get("score", None),
content=[VectorStoreResultContent(text=content_text, type="text")],
)
)
vector_search_response_data.append(vector_store_search_result)
vector_store_search_response["data"] = vector_search_response_data
return vector_store_search_response
def _get_kb_query_from_messages(self, messages: List[AllMessageValues]) -> str:
"""
Uses the text `content` field of the last message in the list of messages
"""
if len(messages) == 0:
return ""
last_message = messages[-1]
last_message_content = last_message.get("content", None)
if last_message_content is None:
return ""
if isinstance(last_message_content, str):
return last_message_content
elif isinstance(last_message_content, list):
return "\n".join([item.get("text", "") for item in last_message_content])
return ""
def _prepare_request(
self,
credentials: Any,
data: BedrockKBRequest,
optional_params: dict,
aws_region_name: str,
api_base: str,
extra_headers: Optional[dict] = None,
) -> Any:
"""
Prepare a signed AWS request.
Args:
credentials: AWS credentials
data: Request data
optional_params: Additional parameters
aws_region_name: AWS region name
api_base: Base API URL
extra_headers: Additional headers
Returns:
AWSRequest: A signed AWS request
"""
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
sigv4 = SigV4Auth(credentials, "bedrock", aws_region_name)
encoded_data = json.dumps(data).encode("utf-8")
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
request = AWSRequest(
method="POST", url=api_base, data=encoded_data, headers=headers
)
sigv4.add_auth(request)
if extra_headers is not None and "Authorization" in extra_headers:
# prevent sigv4 from overwriting the auth header
request.headers["Authorization"] = extra_headers["Authorization"]
return request.prepare()
async def make_bedrock_kb_retrieve_request(
self,
knowledge_base_id: str,
query: str,
guardrail_id: Optional[str] = None,
guardrail_version: Optional[str] = None,
next_token: Optional[str] = None,
retrieval_configuration: Optional[BedrockKBRetrievalConfiguration] = None,
non_default_params: Optional[dict] = None,
) -> BedrockKBResponse:
"""
Make a Bedrock Knowledge Base retrieve request.
Args:
knowledge_base_id (str): The unique identifier of the knowledge base to query
query (str): The query text to search for
guardrail_id (Optional[str]): The guardrail ID to apply
guardrail_version (Optional[str]): The version of the guardrail to apply
next_token (Optional[str]): Token for pagination
retrieval_configuration (Optional[BedrockKBRetrievalConfiguration]): Configuration for the retrieval process
Returns:
BedrockKBRetrievalResponse: A typed response object containing the retrieval results
"""
from fastapi import HTTPException
non_default_params = non_default_params or {}
credentials_dict: Dict[str, Any] = {}
if litellm.vector_store_registry is not None:
credentials_dict = (
litellm.vector_store_registry.get_credentials_for_vector_store(
knowledge_base_id
)
)
credentials = self.get_credentials(
aws_access_key_id=credentials_dict.get(
"aws_access_key_id", non_default_params.get("aws_access_key_id", None)
),
aws_secret_access_key=credentials_dict.get(
"aws_secret_access_key",
non_default_params.get("aws_secret_access_key", None),
),
aws_session_token=credentials_dict.get(
"aws_session_token", non_default_params.get("aws_session_token", None)
),
aws_region_name=credentials_dict.get(
"aws_region_name", non_default_params.get("aws_region_name", None)
),
aws_session_name=credentials_dict.get(
"aws_session_name", non_default_params.get("aws_session_name", None)
),
aws_profile_name=credentials_dict.get(
"aws_profile_name", non_default_params.get("aws_profile_name", None)
),
aws_role_name=credentials_dict.get(
"aws_role_name", non_default_params.get("aws_role_name", None)
),
aws_web_identity_token=credentials_dict.get(
"aws_web_identity_token",
non_default_params.get("aws_web_identity_token", None),
),
aws_sts_endpoint=credentials_dict.get(
"aws_sts_endpoint", non_default_params.get("aws_sts_endpoint", None)
),
)
aws_region_name = self.get_aws_region_name_for_non_llm_api_calls(
aws_region_name=credentials_dict.get(
"aws_region_name", non_default_params.get("aws_region_name", None)
),
)
# Prepare request data
request_data: BedrockKBRequest = BedrockKBRequest(
retrievalQuery=BedrockKBRetrievalQuery(text=query),
)
if next_token:
request_data["nextToken"] = next_token
if retrieval_configuration:
request_data["retrievalConfiguration"] = retrieval_configuration
if guardrail_id and guardrail_version:
request_data["guardrailConfiguration"] = BedrockKBGuardrailConfiguration(
guardrailId=guardrail_id, guardrailVersion=guardrail_version
)
verbose_logger.debug(
f"Request Data: {json.dumps(request_data, indent=4, default=str)}"
)
# Prepare the request
api_base = f"https://bedrock-agent-runtime.{aws_region_name}.amazonaws.com/knowledgebases/{knowledge_base_id}/retrieve"
prepared_request = self._prepare_request(
credentials=credentials,
data=request_data,
optional_params=self.optional_params,
aws_region_name=aws_region_name,
api_base=api_base,
)
verbose_proxy_logger.debug(
"Bedrock Knowledge Base request body: %s, url %s, headers: %s",
request_data,
prepared_request.url,
prepared_request.headers,
)
response = await self.async_handler.post(
url=prepared_request.url,
data=prepared_request.body, # type: ignore
headers=prepared_request.headers, # type: ignore
)
verbose_proxy_logger.debug("Bedrock Knowledge Base response: %s", response.text)
if response.status_code == 200:
response_data = response.json()
return BedrockKBResponse(**response_data)
else:
verbose_proxy_logger.error(
"Bedrock Knowledge Base: error in response. Status code: %s, response: %s",
response.status_code,
response.text,
)
raise HTTPException(
status_code=response.status_code,
detail={
"error": "Error calling Bedrock Knowledge Base",
"response": response.text,
},
)
@staticmethod
def get_initialized_custom_logger() -> Optional[CustomLogger]:
from litellm.litellm_core_utils.litellm_logging import (
_init_custom_logger_compatible_class,
)
return _init_custom_logger_compatible_class(
logging_integration="bedrock_vector_store",
internal_usage_cache=None,
llm_router=None,
)
@staticmethod
def get_chat_completion_message_from_bedrock_kb_response(
response: BedrockKBResponse,
) -> Tuple[Optional[ChatCompletionUserMessage], str]:
"""
Retrieves the context from the Bedrock Knowledge Base response and returns a ChatCompletionUserMessage object.
"""
retrieval_results: Optional[List[BedrockKBRetrievalResult]] = response.get(
"retrievalResults", None
)
if retrieval_results is None:
return None, ""
# string to combine the context from the knowledge base
context_string: str = BedrockVectorStore.CONTENT_PREFIX_STRING
for retrieval_result in retrieval_results:
retrieval_result_content: Optional[BedrockKBContent] = (
retrieval_result.get("content", None) or {}
)
if retrieval_result_content is None:
continue
retrieval_result_text: Optional[str] = retrieval_result_content.get(
"text", None
)
if retrieval_result_text is None:
continue
context_string += retrieval_result_text
message = ChatCompletionUserMessage(
role="user",
content=context_string,
)
return message, context_string

View file

@ -12,7 +12,9 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
import litellm
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.vector_stores.base_vector_store import BaseVectorStore
from litellm.integrations.vector_store_integrations.base_vector_store import (
BaseVectorStore,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,

View file

@ -32,7 +32,9 @@ from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.integrations.opik.opik import OpikLogger
from litellm.integrations.prometheus import PrometheusLogger
from litellm.integrations.s3_v2 import S3Logger
from litellm.integrations.vector_stores.bedrock_vector_store import BedrockVectorStore
from litellm.integrations.vector_store_integrations.bedrock_vector_store import (
BedrockVectorStore,
)
from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler

View file

@ -54,7 +54,9 @@ from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.deepeval.deepeval import DeepEvalLogger
from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.vector_stores.bedrock_vector_store import BedrockVectorStore
from litellm.integrations.vector_store_integrations.bedrock_vector_store import (
BedrockVectorStore,
)
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
StandardBuiltInToolCostTracking,

View file

@ -0,0 +1,73 @@
from abc import abstractmethod
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx
from litellm.types.router import GenericLiteLLMParams
from litellm.types.vector_stores import (
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from ..chat.transformation import BaseLLMException as _BaseLLMException
LiteLLMLoggingObj = _LiteLLMLoggingObj
BaseLLMException = _BaseLLMException
else:
LiteLLMLoggingObj = Any
BaseLLMException = Any
class BaseVectorStoreConfig:
@abstractmethod
def transform_search_vector_store_request(
self,
vector_store_id: str,
query: Union[str, List[str]],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
) -> Tuple[str, Dict]:
pass
@abstractmethod
def transform_search_vector_store_response(self, response: httpx.Response) -> VectorStoreSearchResponse:
pass
@abstractmethod
def validate_environment(
self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]
) -> dict:
return {}
@abstractmethod
def get_complete_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
OPTIONAL
Get the complete url for the request
Some providers need `model` in `api_base`
"""
if api_base is None:
raise ValueError("api_base is required")
return api_base
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
from ..chat.transformation import BaseLLMException
raise BaseLLMException(
status_code=status_code,
message=error_message,
headers=headers,
)

View file

@ -35,6 +35,7 @@ from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -60,6 +61,10 @@ from litellm.types.rerank import OptionalRerankParams, RerankResponse
from litellm.types.responses.main import DeleteResponseResult
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import EmbeddingResponse, FileTypes, TranscriptionResponse
from litellm.types.vector_stores import (
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
)
from litellm.utils import (
CustomStreamWrapper,
ImageResponse,
@ -2342,7 +2347,7 @@ class BaseLLMHTTPHandler:
self,
e: Exception,
provider_config: Union[
BaseConfig, BaseRerankConfig, BaseResponsesAPIConfig, BaseImageEditConfig
BaseConfig, BaseRerankConfig, BaseResponsesAPIConfig, BaseImageEditConfig, BaseVectorStoreConfig
],
):
status_code = getattr(e, "status_code", 500)
@ -2613,3 +2618,143 @@ class BaseLLMHTTPHandler:
raw_response=response,
logging_obj=logging_obj,
)
###### VECTOR STORE HANDLER ######
async def async_vector_store_search_handler(
self,
vector_store_id: str,
query: Union[str, List[str]],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
vector_store_provider_config: BaseVectorStoreConfig,
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
_is_async: bool = False,
) -> VectorStoreSearchResponse:
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders(custom_llm_provider),
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
)
else:
async_httpx_client = client
headers = vector_store_provider_config.validate_environment(
headers=extra_headers or {},
litellm_params=litellm_params
)
if extra_headers:
headers.update(extra_headers)
api_base = vector_store_provider_config.get_complete_url(
api_base=litellm_params.api_base,
litellm_params=dict(litellm_params),
)
url, request_body = vector_store_provider_config.transform_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
)
logging_obj.pre_call(
input="",
api_key="",
additional_args={
"complete_input_dict": request_body,
"api_base": api_base,
"headers": headers,
},
)
try:
response = await async_httpx_client.post(url=url, headers=headers, json=request_body, timeout=timeout)
except Exception as e:
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
return vector_store_provider_config.transform_search_vector_store_response(
response=response,
)
def vector_store_search_handler(
self,
vector_store_id: str,
query: Union[str, List[str]],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
vector_store_provider_config: BaseVectorStoreConfig,
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
_is_async: bool = False,
) -> Union[VectorStoreSearchResponse, Coroutine[Any, Any, VectorStoreSearchResponse]]:
if _is_async:
return self.async_vector_store_search_handler(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
vector_store_provider_config=vector_store_provider_config,
litellm_params=litellm_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
client=client,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client(
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
)
else:
sync_httpx_client = client
headers = vector_store_provider_config.validate_environment(
headers=extra_headers or {},
litellm_params=litellm_params
)
if extra_headers:
headers.update(extra_headers)
api_base = vector_store_provider_config.get_complete_url(
api_base=litellm_params.api_base,
litellm_params=dict(litellm_params),
)
url, request_body = vector_store_provider_config.transform_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
)
logging_obj.pre_call(
input="",
api_key="",
additional_args={
"complete_input_dict": request_body,
"api_base": api_base,
"headers": headers,
},
)
try:
response = sync_httpx_client.post(url=url, headers=headers, json=request_body)
except Exception as e:
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
return vector_store_provider_config.transform_search_vector_store_response(
response=response,
)

View file

@ -0,0 +1,107 @@
from typing import Dict, List, Optional, Tuple, Union, cast
import httpx
import litellm
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.router import GenericLiteLLMParams
from litellm.types.vector_stores import (
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchRequest,
VectorStoreSearchResponse,
)
class OpenAIVectorStoreConfig(BaseVectorStoreConfig):
ASSISTANTS_HEADER_KEY = "OpenAI-Beta"
ASSISTANTS_HEADER_VALUE = "assistants=v2"
def validate_environment(
self, headers: dict, litellm_params: Optional[GenericLiteLLMParams]
) -> dict:
litellm_params = litellm_params or GenericLiteLLMParams()
api_key = (
litellm_params.api_key
or litellm.api_key
or litellm.openai_key
or get_secret_str("OPENAI_API_KEY")
)
headers.update(
{
"Authorization": f"Bearer {api_key}",
}
)
#########################################################
# Ensure OpenAI Assistants header is includes
#########################################################
if self.ASSISTANTS_HEADER_KEY not in headers:
headers.update(
{
self.ASSISTANTS_HEADER_KEY: self.ASSISTANTS_HEADER_VALUE,
}
)
return headers
def get_complete_url(
self,
api_base: Optional[str],
litellm_params: dict,
) -> str:
"""
Get the Base endpoint for OpenAI Vector Stores API
"""
api_base = (
api_base
or litellm.api_base
or get_secret_str("OPENAI_BASE_URL")
or get_secret_str("OPENAI_API_BASE")
or "https://api.openai.com/v1"
)
# Remove trailing slashes
api_base = api_base.rstrip("/")
return f"{api_base}/vector_stores"
def transform_search_vector_store_request(
self,
vector_store_id: str,
query: Union[str, List[str]],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
) -> Tuple[str, Dict]:
url = f"{api_base}/{vector_store_id}/search"
typed_request_body = VectorStoreSearchRequest(
query=query,
filters=vector_store_search_optional_params.get("filters", None),
max_num_results=vector_store_search_optional_params.get("max_num_results", None),
ranking_options=vector_store_search_optional_params.get("ranking_options", None),
rewrite_query=vector_store_search_optional_params.get("rewrite_query", None),
)
dict_request_body = cast(dict, typed_request_body)
return url, dict_request_body
def transform_search_vector_store_response(self, response: httpx.Response) -> VectorStoreSearchResponse:
try:
response_json = response.json()
return VectorStoreSearchResponse(
**response_json
)
except Exception as e:
raise self.get_error_class(
error_message=str(e),
status_code=response.status_code,
headers=response.headers
)

View file

@ -85,3 +85,14 @@ class VectorStoreSearchResponse(TypedDict, total=False):
] # Always "vector_store.search_results.page"
search_query: Optional[str]
data: Optional[List[VectorStoreSearchResult]]
class VectorStoreSearchOptionalRequestParams(TypedDict, total=False):
"""TypedDict for Optional parameters supported by the vector store search API."""
filters: Optional[Dict]
max_num_results: Optional[int]
ranking_options: Optional[Dict]
rewrite_query: Optional[bool]
class VectorStoreSearchRequest(VectorStoreSearchOptionalRequestParams, total=False):
"""Request body for searching a vector store"""
query: Union[str, List[str]]

View file

@ -78,7 +78,9 @@ from litellm.constants import (
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.vector_stores.base_vector_store import BaseVectorStore
from litellm.integrations.vector_store_integrations.base_vector_store import (
BaseVectorStore,
)
from litellm.litellm_core_utils.core_helpers import (
map_finish_reason,
process_response_headers,
@ -242,6 +244,7 @@ from litellm.llms.base_llm.image_variations.transformation import (
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from ._logging import _is_debugging_on, verbose_logger
from .caching.caching import (
@ -6902,13 +6905,28 @@ class ProviderConfigManager:
def get_provider_vector_store_config(
provider: LlmProviders,
) -> Optional[CustomLogger]:
from litellm.integrations.vector_stores.bedrock_vector_store import (
from litellm.integrations.vector_store_integrations.bedrock_vector_store import (
BedrockVectorStore,
)
if LlmProviders.BEDROCK == provider:
return BedrockVectorStore.get_initialized_custom_logger()
return None
@staticmethod
def get_provider_vector_stores_config(
provider: LlmProviders,
) -> Optional[BaseVectorStoreConfig]:
"""
v2 vector store config, use this for new vector store integrations
"""
if litellm.LlmProviders.OPENAI == provider:
from litellm.llms.openai.vector_stores.transformation import (
OpenAIVectorStoreConfig,
)
return OpenAIVectorStoreConfig()
return None
@staticmethod
def get_provider_image_generation_config(

View file

@ -0,0 +1,4 @@
from .main import asearch, search
from .vector_store_registry import VectorStoreRegistry
__all__ = ["search", "asearch", "VectorStoreRegistry"]

View file

@ -0,0 +1,232 @@
"""
LiteLLM SDK Functions for Creating and Searching Vector Stores
"""
import asyncio
import contextvars
from functools import partial
from typing import Any, Coroutine, Dict, List, Optional, Union
import httpx
import litellm
from litellm.constants import request_timeout
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.types.router import GenericLiteLLMParams
from litellm.types.vector_stores import (
VectorStoreResultContent,
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
VectorStoreSearchResult,
)
from litellm.utils import ProviderConfigManager, client
from litellm.vector_stores.utils import VectorStoreRequestUtils
####### ENVIRONMENT VARIABLES ###################
# Initialize any necessary instances or variables here
base_llm_http_handler = BaseLLMHTTPHandler()
#################################################
def mock_vector_store_search_response(
mock_results: Optional[List[VectorStoreSearchResult]] = None,
):
"""Mock response for vector store search"""
if mock_results is None:
mock_results = [
VectorStoreSearchResult(
score=0.95,
content=[
VectorStoreResultContent(
text="This is a sample search result from the vector store.",
type="text"
)
]
)
]
return VectorStoreSearchResponse(
object="vector_store.search_results.page",
search_query="sample query",
data=mock_results,
)
@client
async def asearch(
vector_store_id: str,
query: Union[str, List[str]],
filters: Optional[Dict] = None,
max_num_results: Optional[int] = None,
ranking_options: Optional[Dict] = None,
rewrite_query: Optional[bool] = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
# LiteLLM specific params,
custom_llm_provider: Optional[str] = None,
**kwargs,
) -> VectorStoreSearchResponse:
"""
Async: Search a vector store for relevant chunks based on a query and file attributes filter.
"""
local_vars = locals()
try:
loop = asyncio.get_event_loop()
kwargs["asearch"] = True
# get custom llm provider so we can use this for mapping exceptions
if custom_llm_provider is None:
custom_llm_provider = "openai" # Default to OpenAI for vector stores
func = partial(
search,
vector_store_id=vector_store_id,
query=query,
filters=filters,
max_num_results=max_num_results,
ranking_options=ranking_options,
rewrite_query=rewrite_query,
extra_headers=extra_headers,
extra_query=extra_query,
extra_body=extra_body,
timeout=timeout,
custom_llm_provider=custom_llm_provider,
**kwargs,
)
ctx = contextvars.copy_context()
func_with_context = partial(ctx.run, func)
init_response = await loop.run_in_executor(None, func_with_context)
if asyncio.iscoroutine(init_response):
response = await init_response
else:
response = init_response
return response
except Exception as e:
raise litellm.exception_type(
model=None,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
@client
def search(
vector_store_id: str,
query: Union[str, List[str]],
filters: Optional[Dict] = None,
max_num_results: Optional[int] = None,
ranking_options: Optional[Dict] = None,
rewrite_query: Optional[bool] = None,
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
# The extra values given here take precedence over values defined on the client or passed to this method.
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
# LiteLLM specific params,
custom_llm_provider: Optional[str] = None,
**kwargs,
) -> Union[VectorStoreSearchResponse, Coroutine[Any, Any, VectorStoreSearchResponse]]:
"""
Search a vector store for relevant chunks based on a query and file attributes filter.
Args:
vector_store_id: The ID of the vector store to search.
query: A query string or array for the search.
filters: Optional filter to apply based on file attributes.
max_num_results: Maximum number of results to return (1-50, default 10).
ranking_options: Optional ranking options for search.
rewrite_query: Whether to rewrite the natural language query for vector search.
Returns:
VectorStoreSearchResponse containing the search results.
"""
local_vars = locals()
try:
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
_is_async = kwargs.pop("asearch", False) is True
# get llm provider logic
litellm_params = GenericLiteLLMParams(**kwargs)
## MOCK RESPONSE LOGIC
if litellm_params.mock_response and isinstance(
litellm_params.mock_response, (str, list)
):
mock_results = None
if isinstance(litellm_params.mock_response, list):
mock_results = litellm_params.mock_response
return mock_vector_store_search_response(mock_results=mock_results)
# Default to OpenAI for vector stores
if custom_llm_provider is None:
custom_llm_provider = "openai"
# get provider config - using vector store custom logger for now
vector_store_provider_config = ProviderConfigManager.get_provider_vector_stores_config(
provider=litellm.LlmProviders(custom_llm_provider),
)
if vector_store_provider_config is None:
raise ValueError(
f"Vector store search is not supported for {custom_llm_provider}"
)
local_vars.update(kwargs)
# Get VectorStoreSearchOptionalRequestParams with only valid parameters
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams = (
VectorStoreRequestUtils.get_requested_vector_store_search_optional_param(
local_vars
)
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
model=None,
optional_params={
"vector_store_id": vector_store_id,
"query": query,
**vector_store_search_optional_params,
},
litellm_params={
"litellm_call_id": litellm_call_id,
"vector_store_id": vector_store_id,
},
custom_llm_provider=custom_llm_provider,
)
response = base_llm_http_handler.vector_store_search_handler(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
vector_store_provider_config=vector_store_provider_config,
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
logging_obj=litellm_logging_obj,
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout or request_timeout,
_is_async=_is_async,
client=kwargs.get("client"),
)
return response
except Exception as e:
raise litellm.exception_type(
model=None,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)

View file

@ -0,0 +1,28 @@
from typing import Any, Dict, cast, get_type_hints
from litellm.types.vector_stores import VectorStoreSearchOptionalRequestParams
class VectorStoreRequestUtils:
"""Helper utils for constructing Vector Store search requests"""
@staticmethod
def get_requested_vector_store_search_optional_param(
params: Dict[str, Any],
) -> VectorStoreSearchOptionalRequestParams:
"""
Filter parameters to only include those defined in VectorStoreSearchOptionalRequestParams.
Args:
params: Dictionary of parameters to filter
Returns:
VectorStoreSearchOptionalRequestParams instance with only the valid parameters
"""
valid_keys = get_type_hints(VectorStoreSearchOptionalRequestParams).keys()
filtered_params = {
k: v for k, v in params.items() if k in valid_keys and v is not None
}
return cast(VectorStoreSearchOptionalRequestParams, filtered_params)

View file

@ -19,7 +19,7 @@ import pytest
import litellm
from litellm import completion
from litellm._logging import verbose_logger
from litellm.integrations.vector_stores.bedrock_vector_store import BedrockVectorStore
from litellm.integrations.vector_store_integrations.bedrock_vector_store import BedrockVectorStore
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import StandardLoggingPayload, StandardLoggingVectorStoreRequest

View file

@ -0,0 +1,136 @@
import httpx
import json
import pytest
import sys
from typing import Any, Dict, List
from unittest.mock import MagicMock, Mock, patch
import os
import uuid
import time
import base64
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
from abc import ABC, abstractmethod
from litellm.integrations.custom_logger import CustomLogger
import json
from litellm.types.utils import StandardLoggingPayload
class BaseVectorStoreTest(ABC):
"""
Abstract base test class that enforces a common test across all test classes.
"""
@abstractmethod
def get_base_request_args(self) -> dict:
"""Must return the base request args"""
pass
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_basic_search_vector_store(self, sync_mode):
litellm._turn_on_debug()
litellm.set_verbose = True
base_request_args = self.get_base_request_args()
try:
if sync_mode:
response = litellm.vector_stores.search(
query="Basic ping",
**base_request_args
)
else:
response = await litellm.vector_stores.asearch(
query="Basic ping",
**base_request_args
)
except litellm.InternalServerError:
pytest.skip("Skipping test due to litellm.InternalServerError")
print("litellm response=", json.dumps(response, indent=4, default=str))
# Validate response structure
self._validate_vector_store_response(response)
def _validate_vector_store_response(self, response):
"""Validate the structure and content of a vector store search response"""
# Check that response is a dictionary
assert isinstance(response, dict), f"Response should be a dict, got {type(response)}"
# Check required top-level fields
required_fields = ['object', 'search_query', 'data']
for field in required_fields:
assert field in response, f"Missing required field '{field}' in response"
# Validate object field
assert response['object'] == 'vector_store.search_results.page', \
f"Expected object to be 'vector_store.search_results.page', got '{response['object']}'"
# Validate search_query field
assert isinstance(response['search_query'], list), \
f"search_query should be a list, got {type(response['search_query'])}"
assert len(response['search_query']) > 0, "search_query should not be empty"
assert all(isinstance(query, str) for query in response['search_query']), \
"All items in search_query should be strings"
# Validate data field
assert isinstance(response['data'], list), \
f"data should be a list, got {type(response['data'])}"
# Validate each result in data
for i, result in enumerate(response['data']):
self._validate_search_result(result, i)
print(f"✅ Response validation passed: Found {len(response['data'])} search results")
def _validate_search_result(self, result, index):
"""Validate an individual search result"""
# Check that result is a dictionary
assert isinstance(result, dict), f"Result {index} should be a dict, got {type(result)}"
# Check required fields in each result
required_result_fields = ['file_id', 'filename', 'score', 'attributes', 'content']
for field in required_result_fields:
assert field in result, f"Missing required field '{field}' in result {index}"
# Validate file_id
assert isinstance(result['file_id'], str), \
f"file_id should be a string, got {type(result['file_id'])} in result {index}"
assert len(result['file_id']) > 0, f"file_id should not be empty in result {index}"
# Validate filename
assert isinstance(result['filename'], str), \
f"filename should be a string, got {type(result['filename'])} in result {index}"
assert len(result['filename']) > 0, f"filename should not be empty in result {index}"
# Validate score
assert isinstance(result['score'], (int, float)), \
f"score should be a number, got {type(result['score'])} in result {index}"
assert 0.0 <= result['score'] <= 1.0, \
f"score should be between 0.0 and 1.0, got {result['score']} in result {index}"
# Validate attributes
assert isinstance(result['attributes'], dict), \
f"attributes should be a dict, got {type(result['attributes'])} in result {index}"
# Validate content
assert isinstance(result['content'], list), \
f"content should be a list, got {type(result['content'])} in result {index}"
assert len(result['content']) > 0, f"content should not be empty in result {index}"
# Validate each content item
for j, content_item in enumerate(result['content']):
assert isinstance(content_item, dict), \
f"Content item {j} in result {index} should be a dict, got {type(content_item)}"
assert 'type' in content_item, \
f"Content item {j} in result {index} missing 'type' field"
assert 'text' in content_item, \
f"Content item {j} in result {index} missing 'text' field"
assert isinstance(content_item['text'], str), \
f"Content text should be a string in item {j} of result {index}"
assert len(content_item['text']) > 0, \
f"Content text should not be empty in item {j} of result {index}"
print(f"✅ Result {index} validation passed: {result['filename']} (score: {result['score']:.4f})")

View file

@ -0,0 +1,63 @@
# conftest.py
import importlib
import os
import sys
import pytest
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
@pytest.fixture(scope="function", autouse=True)
def setup_and_teardown():
"""
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
"""
curr_dir = os.getcwd() # Get the current working directory
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the project directory to the system path
import litellm
from litellm import Router
importlib.reload(litellm)
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
import asyncio
loop = asyncio.get_event_loop_policy().new_event_loop()
asyncio.set_event_loop(loop)
print(litellm)
# from litellm import Router, completion, aembedding, acompletion, embedding
yield
# Teardown code (executes after the yield point)
loop.close() # Close the loop created earlier
asyncio.set_event_loop(None) # Remove the reference to the loop
def pytest_collection_modifyitems(config, items):
# Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests
custom_logger_tests = [
item for item in items if "custom_logger" in item.parent.name
]
other_tests = [item for item in items if "custom_logger" not in item.parent.name]
# Sort tests based on their names
custom_logger_tests.sort(key=lambda x: x.name)
other_tests.sort(key=lambda x: x.name)
# Reorder the items list
items[:] = custom_logger_tests + other_tests

View file

@ -0,0 +1,11 @@
from base_vector_store_test import BaseVectorStoreTest
class TestOpenAIVectorStore(BaseVectorStoreTest):
def get_base_request_args(self) -> dict:
"""
This is a real vector store on OpenAI
"""
return {
"vector_store_id": "vs_685b14b1a1b88191bc27e04f1917fddd",
"custom_llm_provider": "openai",
}