mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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:
parent
97da33494a
commit
2bb8048864
19 changed files with 1252 additions and 7 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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 *
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
73
litellm/llms/base_llm/vector_store/transformation.py
Normal file
73
litellm/llms/base_llm/vector_store/transformation.py
Normal 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,
|
||||
)
|
||||
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
107
litellm/llms/openai/vector_stores/transformation.py
Normal file
107
litellm/llms/openai/vector_stores/transformation.py
Normal 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
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
@ -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]]
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
4
litellm/vector_stores/__init__.py
Normal file
4
litellm/vector_stores/__init__.py
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
from .main import asearch, search
|
||||
from .vector_store_registry import VectorStoreRegistry
|
||||
|
||||
__all__ = ["search", "asearch", "VectorStoreRegistry"]
|
||||
232
litellm/vector_stores/main.py
Normal file
232
litellm/vector_stores/main.py
Normal 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,
|
||||
)
|
||||
28
litellm/vector_stores/utils.py
Normal file
28
litellm/vector_stores/utils.py
Normal 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)
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
136
tests/vector_store_tests/base_vector_store_test.py
Normal file
136
tests/vector_store_tests/base_vector_store_test.py
Normal 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})")
|
||||
63
tests/vector_store_tests/conftest.py
Normal file
63
tests/vector_store_tests/conftest.py
Normal 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
|
||||
11
tests/vector_store_tests/test_openai_vector_store.py
Normal file
11
tests/vector_store_tests/test_openai_vector_store.py
Normal 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",
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue