From 2bb8048864285c1a62fe5c89fca3bc31865d1b59 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 24 Jun 2025 15:52:43 -0700 Subject: [PATCH] [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 --- .../docs/completion/knowledgebase.md | 1 + litellm/__init__.py | 1 + .../base_vector_store.py | 0 .../bedrock_vector_store.py | 409 ++++++++++++++++++ .../vector_stores/bedrock_vector_store.py | 4 +- .../custom_logger_registry.py | 4 +- litellm/litellm_core_utils/litellm_logging.py | 4 +- .../base_llm/vector_store/transformation.py | 73 ++++ litellm/llms/custom_httpx/llm_http_handler.py | 147 ++++++- .../openai/vector_stores/transformation.py | 107 +++++ litellm/types/vector_stores.py | 11 + litellm/utils.py | 22 +- litellm/vector_stores/__init__.py | 4 + litellm/vector_stores/main.py | 232 ++++++++++ litellm/vector_stores/utils.py | 28 ++ .../test_bedrock_knowledgebase_hook.py | 2 +- .../base_vector_store_test.py | 136 ++++++ tests/vector_store_tests/conftest.py | 63 +++ .../test_openai_vector_store.py | 11 + 19 files changed, 1252 insertions(+), 7 deletions(-) rename litellm/integrations/{vector_stores => vector_store_integrations}/base_vector_store.py (100%) create mode 100644 litellm/integrations/vector_store_integrations/bedrock_vector_store.py create mode 100644 litellm/llms/base_llm/vector_store/transformation.py create mode 100644 litellm/llms/openai/vector_stores/transformation.py create mode 100644 litellm/vector_stores/__init__.py create mode 100644 litellm/vector_stores/main.py create mode 100644 litellm/vector_stores/utils.py create mode 100644 tests/vector_store_tests/base_vector_store_test.py create mode 100644 tests/vector_store_tests/conftest.py create mode 100644 tests/vector_store_tests/test_openai_vector_store.py diff --git a/docs/my-website/docs/completion/knowledgebase.md b/docs/my-website/docs/completion/knowledgebase.md index 033dccea200..b3e6a06aa9a 100644 --- a/docs/my-website/docs/completion/knowledgebase.md +++ b/docs/my-website/docs/completion/knowledgebase.md @@ -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 diff --git a/litellm/__init__.py b/litellm/__init__.py index 5407dd85d58..2921c0600d9 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 * diff --git a/litellm/integrations/vector_stores/base_vector_store.py b/litellm/integrations/vector_store_integrations/base_vector_store.py similarity index 100% rename from litellm/integrations/vector_stores/base_vector_store.py rename to litellm/integrations/vector_store_integrations/base_vector_store.py diff --git a/litellm/integrations/vector_store_integrations/bedrock_vector_store.py b/litellm/integrations/vector_store_integrations/bedrock_vector_store.py new file mode 100644 index 00000000000..a00acefb6a3 --- /dev/null +++ b/litellm/integrations/vector_store_integrations/bedrock_vector_store.py @@ -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 diff --git a/litellm/integrations/vector_stores/bedrock_vector_store.py b/litellm/integrations/vector_stores/bedrock_vector_store.py index 0523dac8edd..a00acefb6a3 100644 --- a/litellm/integrations/vector_stores/bedrock_vector_store.py +++ b/litellm/integrations/vector_stores/bedrock_vector_store.py @@ -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, diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index 20dc2b0c903..1b75cc3e3df 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -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 diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index c608dfc8d88..aa0280210d4 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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, diff --git a/litellm/llms/base_llm/vector_store/transformation.py b/litellm/llms/base_llm/vector_store/transformation.py new file mode 100644 index 00000000000..4df63bd2f96 --- /dev/null +++ b/litellm/llms/base_llm/vector_store/transformation.py @@ -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, + ) + diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 4ed0016d8b2..dfc5ae5277f 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, + ) + diff --git a/litellm/llms/openai/vector_stores/transformation.py b/litellm/llms/openai/vector_stores/transformation.py new file mode 100644 index 00000000000..98657825684 --- /dev/null +++ b/litellm/llms/openai/vector_stores/transformation.py @@ -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 + ) + + + + + \ No newline at end of file diff --git a/litellm/types/vector_stores.py b/litellm/types/vector_stores.py index cd8280ac960..0fcf7a2ab00 100644 --- a/litellm/types/vector_stores.py +++ b/litellm/types/vector_stores.py @@ -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]] \ No newline at end of file diff --git a/litellm/utils.py b/litellm/utils.py index dacc1680458..25f46d4bc5c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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( diff --git a/litellm/vector_stores/__init__.py b/litellm/vector_stores/__init__.py new file mode 100644 index 00000000000..6546f9159e0 --- /dev/null +++ b/litellm/vector_stores/__init__.py @@ -0,0 +1,4 @@ +from .main import asearch, search +from .vector_store_registry import VectorStoreRegistry + +__all__ = ["search", "asearch", "VectorStoreRegistry"] diff --git a/litellm/vector_stores/main.py b/litellm/vector_stores/main.py new file mode 100644 index 00000000000..245a3f6fa3d --- /dev/null +++ b/litellm/vector_stores/main.py @@ -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, + ) \ No newline at end of file diff --git a/litellm/vector_stores/utils.py b/litellm/vector_stores/utils.py new file mode 100644 index 00000000000..7817526d00e --- /dev/null +++ b/litellm/vector_stores/utils.py @@ -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) + diff --git a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py index 532a5c11c52..4c75d485277 100644 --- a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py +++ b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py @@ -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 diff --git a/tests/vector_store_tests/base_vector_store_test.py b/tests/vector_store_tests/base_vector_store_test.py new file mode 100644 index 00000000000..d23012847bb --- /dev/null +++ b/tests/vector_store_tests/base_vector_store_test.py @@ -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})") diff --git a/tests/vector_store_tests/conftest.py b/tests/vector_store_tests/conftest.py new file mode 100644 index 00000000000..b3561d8a626 --- /dev/null +++ b/tests/vector_store_tests/conftest.py @@ -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 diff --git a/tests/vector_store_tests/test_openai_vector_store.py b/tests/vector_store_tests/test_openai_vector_store.py new file mode 100644 index 00000000000..20980f2ff56 --- /dev/null +++ b/tests/vector_store_tests/test_openai_vector_store.py @@ -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", + } \ No newline at end of file