diff --git a/docs/my-website/docs/completion/knowledgebase.md b/docs/my-website/docs/completion/knowledgebase.md index b3e6a06aa9a..a1c926274cd 100644 --- a/docs/my-website/docs/completion/knowledgebase.md +++ b/docs/my-website/docs/completion/knowledgebase.md @@ -18,6 +18,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) +- [Azure Vector Stores](https://learn.microsoft.com/en-us/azure/ai-services/openai/how-to/file-search?tabs=python#vector-stores ## Quick Start diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index ed62175289c..f2a8defe13f 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -1,12 +1,11 @@ import json import os -from typing import Any, Callable, Dict, Optional, Union +from typing import Any, Callable, Dict, Literal, Optional, Union, cast import httpx from openai import AsyncAzureOpenAI, AzureOpenAI import litellm -from litellm.types.router import GenericLiteLLMParams from litellm._logging import verbose_logger from litellm.caching.caching import DualCache from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -15,6 +14,8 @@ from litellm.secret_managers.get_azure_ad_token_provider import ( get_azure_ad_token_provider, ) from litellm.secret_managers.main import get_secret_str +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import _add_path_to_api_base azure_ad_cache = DualCache() @@ -613,3 +614,78 @@ class BaseAzureLLM(BaseOpenAILLM): else: client = AzureOpenAI(**azure_client_params) # type: ignore return client + + @staticmethod + def _base_validate_azure_environment( + 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.azure_key + or get_secret_str("AZURE_OPENAI_API_KEY") + or get_secret_str("AZURE_API_KEY") + ) + + if api_key: + headers["api-key"] = api_key + return headers + + ### Fallback to Azure AD token-based authentication if no API key is available + ### Retrieves Azure AD token and adds it to the Authorization header + azure_ad_token = get_azure_ad_token(litellm_params) + if azure_ad_token: + headers["Authorization"] = f"Bearer {azure_ad_token}" + + return headers + + @staticmethod + def _get_base_azure_url( + api_base: Optional[str], + litellm_params: Optional[Union[GenericLiteLLMParams, Dict[str, Any]]], + route: Literal["/openai/responses", "/openai/vector_stores"] + ) -> str: + api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") + if api_base is None: + raise ValueError( + f"api_base is required for Azure AI Studio. Please set the api_base parameter. Passed `api_base={api_base}`" + ) + original_url = httpx.URL(api_base) + + # Extract api_version or use default + litellm_params = litellm_params or {} + api_version = cast(Optional[str], litellm_params.get("api_version")) + + # Create a new dictionary with existing params + query_params = dict(original_url.params) + + # Add api_version if needed + if "api-version" not in query_params and api_version: + query_params["api-version"] = api_version + + # Add the path to the base URL + if route not in api_base: + new_url = _add_path_to_api_base( + api_base=api_base, ending_path=route + ) + else: + new_url = api_base + + if BaseAzureLLM._is_azure_v1_api_version(api_version): + # ensure the request go to /openai/v1 and not just /openai + if "/openai/v1" not in new_url: + parsed_url = httpx.URL(new_url) + new_url = str(parsed_url.copy_with(path=parsed_url.path.replace("/openai", "/openai/v1"))) + + + # Use the new query_params dictionary + final_url = httpx.URL(new_url).copy_with(params=query_params) + + return str(final_url) + + @staticmethod + def _is_azure_v1_api_version(api_version: Optional[str]) -> bool: + if api_version is None: + return False + return api_version == "preview" or api_version == "latest" diff --git a/litellm/llms/azure/responses/transformation.py b/litellm/llms/azure/responses/transformation.py index d5425ef4391..e6f48179e49 100644 --- a/litellm/llms/azure/responses/transformation.py +++ b/litellm/llms/azure/responses/transformation.py @@ -1,16 +1,11 @@ -from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, cast +from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple -import httpx - -import litellm from litellm._logging import verbose_logger -from litellm.llms.azure.common_utils import get_azure_ad_token +from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig -from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import * from litellm.types.responses.main import * from litellm.types.router import GenericLiteLLMParams -from litellm.utils import _add_path_to_api_base if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -24,27 +19,11 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): def validate_environment( self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams] ) -> dict: - litellm_params = litellm_params or GenericLiteLLMParams() - api_key = ( - litellm_params.api_key - or litellm.api_key - or litellm.azure_key - or get_secret_str("AZURE_OPENAI_API_KEY") - or get_secret_str("AZURE_API_KEY") + return BaseAzureLLM._base_validate_azure_environment( + headers=headers, + litellm_params=litellm_params ) - if api_key: - headers["api-key"] = api_key - return headers - - ### Fallback to Azure AD token-based authentication if no API key is available - ### Retrieves Azure AD token and adds it to the Authorization header - azure_ad_token = get_azure_ad_token(litellm_params) - if azure_ad_token: - headers["Authorization"] = f"Bearer {azure_ad_token}" - - return headers - def get_complete_url( self, api_base: Optional[str], @@ -66,47 +45,12 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): - A complete URL string, e.g., "https://litellm8397336933.openai.azure.com/openai/responses?api-version=2024-05-01-preview" """ - api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") - if api_base is None: - raise ValueError( - f"api_base is required for Azure AI Studio. Please set the api_base parameter. Passed `api_base={api_base}`" - ) - original_url = httpx.URL(api_base) - - # Extract api_version or use default - api_version = cast(Optional[str], litellm_params.get("api_version")) - - # Create a new dictionary with existing params - query_params = dict(original_url.params) - - # Add api_version if needed - if "api-version" not in query_params and api_version: - query_params["api-version"] = api_version - - # Add the path to the base URL - if "/openai/responses" not in api_base: - new_url = _add_path_to_api_base( - api_base=api_base, ending_path="/openai/responses" - ) - else: - new_url = api_base - - if self._is_azure_v1_api_version(api_version): - # ensure the request go to /openai/v1 and not just /openai - if "/openai/v1" not in new_url: - parsed_url = httpx.URL(new_url) - new_url = str(parsed_url.copy_with(path=parsed_url.path.replace("/openai", "/openai/v1"))) - - - # Use the new query_params dictionary - final_url = httpx.URL(new_url).copy_with(params=query_params) - - return str(final_url) + return BaseAzureLLM._get_base_azure_url( + api_base=api_base, + litellm_params=litellm_params, + route="/openai/responses" + ) - def _is_azure_v1_api_version(self, api_version: Optional[str]) -> bool: - if api_version is None: - return False - return api_version == "preview" or api_version == "latest" ######################################################### ########## DELETE RESPONSE API TRANSFORMATION ############## diff --git a/litellm/llms/azure/vector_stores/transformation.py b/litellm/llms/azure/vector_stores/transformation.py new file mode 100644 index 00000000000..f1cd81b2bf2 --- /dev/null +++ b/litellm/llms/azure/vector_stores/transformation.py @@ -0,0 +1,27 @@ +from typing import Optional + +from litellm.llms.azure.common_utils import BaseAzureLLM +from litellm.llms.openai.vector_stores.transformation import OpenAIVectorStoreConfig +from litellm.types.router import GenericLiteLLMParams + + +class AzureOpenAIVectorStoreConfig(OpenAIVectorStoreConfig): + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + return BaseAzureLLM._get_base_azure_url( + api_base=api_base, + litellm_params=litellm_params, + route="/openai/vector_stores" + ) + + + def validate_environment( + self, headers: dict, litellm_params: Optional[GenericLiteLLMParams] + ) -> dict: + return BaseAzureLLM._base_validate_azure_environment( + headers=headers, + litellm_params=litellm_params + ) \ No newline at end of file diff --git a/litellm/llms/base_llm/vector_store/transformation.py b/litellm/llms/base_llm/vector_store/transformation.py index 4df63bd2f96..6ec9d25ae59 100644 --- a/litellm/llms/base_llm/vector_store/transformation.py +++ b/litellm/llms/base_llm/vector_store/transformation.py @@ -5,6 +5,8 @@ import httpx from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + VectorStoreCreateOptionalRequestParams, + VectorStoreCreateResponse, VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, ) @@ -35,6 +37,17 @@ class BaseVectorStoreConfig: def transform_search_vector_store_response(self, response: httpx.Response) -> VectorStoreSearchResponse: pass + @abstractmethod + def transform_create_vector_store_request( + self, + vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, + api_base: str, + ) -> Tuple[str, Dict]: + pass + + @abstractmethod + def transform_create_vector_store_response(self, response: httpx.Response) -> VectorStoreCreateResponse: + pass @abstractmethod def validate_environment( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index dfc5ae5277f..92f9d6d958b 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -62,6 +62,8 @@ 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 ( + VectorStoreCreateOptionalRequestParams, + VectorStoreCreateResponse, VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, ) @@ -2758,3 +2760,131 @@ class BaseLLMHTTPHandler: response=response, ) + async def async_vector_store_create_handler( + self, + vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, + 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, + ) -> VectorStoreCreateResponse: + 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_create_vector_store_request( + vector_store_create_optional_params=vector_store_create_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_create_vector_store_response( + response=response, + ) + + def vector_store_create_handler( + self, + vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, + 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[VectorStoreCreateResponse, Coroutine[Any, Any, VectorStoreCreateResponse]]: + if _is_async: + return self.async_vector_store_create_handler( + vector_store_create_optional_params=vector_store_create_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_create_vector_store_request( + vector_store_create_optional_params=vector_store_create_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_create_vector_store_response( + response=response, + ) + diff --git a/litellm/llms/openai/vector_stores/transformation.py b/litellm/llms/openai/vector_stores/transformation.py index 98657825684..11d76937ab4 100644 --- a/litellm/llms/openai/vector_stores/transformation.py +++ b/litellm/llms/openai/vector_stores/transformation.py @@ -7,6 +7,9 @@ from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreCon from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + VectorStoreCreateOptionalRequestParams, + VectorStoreCreateRequest, + VectorStoreCreateResponse, VectorStoreSearchOptionalRequestParams, VectorStoreSearchRequest, VectorStoreSearchResponse, @@ -101,6 +104,36 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): headers=response.headers ) + def transform_create_vector_store_request( + self, + vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams, + api_base: str, + ) -> Tuple[str, Dict]: + url = api_base # Base URL for creating vector stores + typed_request_body = VectorStoreCreateRequest( + name=vector_store_create_optional_params.get("name", None), + file_ids=vector_store_create_optional_params.get("file_ids", None), + expires_after=vector_store_create_optional_params.get("expires_after", None), + chunking_strategy=vector_store_create_optional_params.get("chunking_strategy", None), + metadata=vector_store_create_optional_params.get("metadata", None), + ) + + dict_request_body = cast(dict, typed_request_body) + return url, dict_request_body + + def transform_create_vector_store_response(self, response: httpx.Response) -> VectorStoreCreateResponse: + try: + response_json = response.json() + return VectorStoreCreateResponse( + **response_json + ) + except Exception as e: + raise self.get_error_class( + error_message=str(e), + status_code=response.status_code, + headers=response.headers + ) + diff --git a/litellm/types/vector_stores.py b/litellm/types/vector_stores.py index 0fcf7a2ab00..3d7190d7f9b 100644 --- a/litellm/types/vector_stores.py +++ b/litellm/types/vector_stores.py @@ -95,4 +95,73 @@ class VectorStoreSearchOptionalRequestParams(TypedDict, total=False): class VectorStoreSearchRequest(VectorStoreSearchOptionalRequestParams, total=False): """Request body for searching a vector store""" - query: Union[str, List[str]] \ No newline at end of file + query: Union[str, List[str]] + + +# Vector Store Creation Types +class VectorStoreExpirationPolicy(TypedDict, total=False): + """The expiration policy for a vector store""" + anchor: Literal["last_active_at"] # Anchor timestamp after which the expiration policy applies + days: int # Number of days after anchor time that the vector store will expire + + +class VectorStoreAutoChunkingStrategy(TypedDict, total=False): + """Auto chunking strategy configuration""" + type: Literal["auto"] # Always "auto" + + +class VectorStoreStaticChunkingStrategyConfig(TypedDict, total=False): + """Static chunking strategy configuration""" + max_chunk_size_tokens: int # Maximum number of tokens per chunk + chunk_overlap_tokens: int # Number of tokens to overlap between chunks + + +class VectorStoreStaticChunkingStrategy(TypedDict, total=False): + """Static chunking strategy""" + type: Literal["static"] # Always "static" + static: VectorStoreStaticChunkingStrategyConfig + + +class VectorStoreChunkingStrategy(TypedDict, total=False): + """Union type for chunking strategies""" + # This can be either auto or static + type: Literal["auto", "static"] + static: Optional[VectorStoreStaticChunkingStrategyConfig] + + +class VectorStoreFileCounts(TypedDict, total=False): + """File counts for a vector store""" + in_progress: int + completed: int + failed: int + cancelled: int + total: int + + +class VectorStoreCreateOptionalRequestParams(TypedDict, total=False): + """TypedDict for Optional parameters supported by the vector store create API.""" + name: Optional[str] # Name of the vector store + file_ids: Optional[List[str]] # List of File IDs that the vector store should use + expires_after: Optional[VectorStoreExpirationPolicy] # Expiration policy for the vector store + chunking_strategy: Optional[VectorStoreChunkingStrategy] # Chunking strategy for the files + metadata: Optional[Dict[str, str]] # Set of key-value pairs for metadata + + +class VectorStoreCreateRequest(VectorStoreCreateOptionalRequestParams, total=False): + """Request body for creating a vector store""" + pass # All fields are optional for vector store creation + + +class VectorStoreCreateResponse(TypedDict, total=False): + """Response after creating a vector store""" + id: str # ID of the vector store + object: Literal["vector_store"] # Always "vector_store" + created_at: int # Unix timestamp of when the vector store was created + name: Optional[str] # Name of the vector store + bytes: int # Size of the vector store in bytes + file_counts: VectorStoreFileCounts # File counts for the vector store + status: Literal["expired", "in_progress", "completed"] # Status of the vector store + expires_after: Optional[VectorStoreExpirationPolicy] # Expiration policy + expires_at: Optional[int] # Unix timestamp of when the vector store expires + last_active_at: Optional[int] # Unix timestamp of when the vector store was last active + metadata: Optional[Dict[str, str]] # Metadata associated with the vector store \ No newline at end of file diff --git a/litellm/utils.py b/litellm/utils.py index 25f46d4bc5c..05328c3750b 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6926,6 +6926,11 @@ class ProviderConfigManager: OpenAIVectorStoreConfig, ) return OpenAIVectorStoreConfig() + elif litellm.LlmProviders.AZURE == provider: + from litellm.llms.azure.vector_stores.transformation import ( + AzureOpenAIVectorStoreConfig, + ) + return AzureOpenAIVectorStoreConfig() return None @staticmethod diff --git a/litellm/vector_stores/__init__.py b/litellm/vector_stores/__init__.py index 6546f9159e0..6bcc6540328 100644 --- a/litellm/vector_stores/__init__.py +++ b/litellm/vector_stores/__init__.py @@ -1,4 +1,4 @@ -from .main import asearch, search +from .main import acreate, asearch, create, search from .vector_store_registry import VectorStoreRegistry -__all__ = ["search", "asearch", "VectorStoreRegistry"] +__all__ = ["search", "asearch", "create", "acreate", "VectorStoreRegistry"] diff --git a/litellm/vector_stores/main.py b/litellm/vector_stores/main.py index 245a3f6fa3d..80d6341146a 100644 --- a/litellm/vector_stores/main.py +++ b/litellm/vector_stores/main.py @@ -14,6 +14,9 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + VectorStoreCreateOptionalRequestParams, + VectorStoreCreateResponse, + VectorStoreFileCounts, VectorStoreResultContent, VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, @@ -52,6 +55,205 @@ def mock_vector_store_search_response( ) +def mock_vector_store_create_response( + mock_response: Optional[VectorStoreCreateResponse] = None, +): + """Mock response for vector store create""" + if mock_response is None: + mock_response = VectorStoreCreateResponse( + id="vs_mock123", + object="vector_store", + created_at=1699061776, + name="Mock Vector Store", + bytes=0, + file_counts=VectorStoreFileCounts( + in_progress=0, + completed=0, + failed=0, + cancelled=0, + total=0, + ), + status="completed", + expires_after=None, + expires_at=None, + last_active_at=None, + metadata=None, + ) + + return mock_response + + +@client +async def acreate( + name: Optional[str] = None, + file_ids: Optional[List[str]] = None, + expires_after: Optional[Dict] = None, + chunking_strategy: Optional[Dict] = None, + metadata: Optional[Dict[str, str]] = 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, +) -> VectorStoreCreateResponse: + """ + Async: Create a vector store. + """ + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["acreate"] = 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( + create, + name=name, + file_ids=file_ids, + expires_after=expires_after, + chunking_strategy=chunking_strategy, + metadata=metadata, + 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 create( + name: Optional[str] = None, + file_ids: Optional[List[str]] = None, + expires_after: Optional[Dict] = None, + chunking_strategy: Optional[Dict] = None, + metadata: Optional[Dict[str, str]] = 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[VectorStoreCreateResponse, Coroutine[Any, Any, VectorStoreCreateResponse]]: + """ + Create a vector store. + + Args: + name: The name of the vector store. + file_ids: A list of File IDs that the vector store should use. + expires_after: The expiration policy for the vector store. + chunking_strategy: The chunking strategy used to chunk the file(s). + metadata: Set of 16 key-value pairs that can be attached to an object. + + Returns: + VectorStoreCreateResponse containing the created vector store details. + """ + 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("acreate", 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, dict + ): + return mock_vector_store_create_response( + mock_response=VectorStoreCreateResponse(**litellm_params.mock_response) + ) + + # 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 create is not supported for {custom_llm_provider}" + ) + + local_vars.update(kwargs) + + # Get VectorStoreCreateOptionalRequestParams with only valid parameters + vector_store_create_optional_params: VectorStoreCreateOptionalRequestParams = ( + VectorStoreRequestUtils.get_requested_vector_store_create_optional_param( + local_vars + ) + ) + + # Pre Call logging + litellm_logging_obj.update_environment_variables( + model=None, + optional_params={ + "name": name, + **vector_store_create_optional_params, + }, + litellm_params={ + "litellm_call_id": litellm_call_id, + }, + custom_llm_provider=custom_llm_provider, + ) + + response = base_llm_http_handler.vector_store_create_handler( + vector_store_create_optional_params=vector_store_create_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, + ) + + @client async def asearch( vector_store_id: str, diff --git a/litellm/vector_stores/utils.py b/litellm/vector_stores/utils.py index 7817526d00e..b7eb7790add 100644 --- a/litellm/vector_stores/utils.py +++ b/litellm/vector_stores/utils.py @@ -1,6 +1,9 @@ from typing import Any, Dict, cast, get_type_hints -from litellm.types.vector_stores import VectorStoreSearchOptionalRequestParams +from litellm.types.vector_stores import ( + VectorStoreCreateOptionalRequestParams, + VectorStoreSearchOptionalRequestParams, +) class VectorStoreRequestUtils: @@ -26,3 +29,23 @@ class VectorStoreRequestUtils: return cast(VectorStoreSearchOptionalRequestParams, filtered_params) + @staticmethod + def get_requested_vector_store_create_optional_param( + params: Dict[str, Any], + ) -> VectorStoreCreateOptionalRequestParams: + """ + Filter parameters to only include those defined in VectorStoreCreateOptionalRequestParams. + + Args: + params: Dictionary of parameters to filter + + Returns: + VectorStoreCreateOptionalRequestParams instance with only the valid parameters + """ + valid_keys = get_type_hints(VectorStoreCreateOptionalRequestParams).keys() + filtered_params = { + k: v for k, v in params.items() if k in valid_keys and v is not None + } + + return cast(VectorStoreCreateOptionalRequestParams, filtered_params) + diff --git a/tests/test_litellm/llms/azure/response/test_azure_transformation.py b/tests/test_litellm/llms/azure/response/test_azure_transformation.py index c210f91170e..8edd87a12fc 100644 --- a/tests/test_litellm/llms/azure/response/test_azure_transformation.py +++ b/tests/test_litellm/llms/azure/response/test_azure_transformation.py @@ -54,9 +54,9 @@ def test_validate_environment_azure_key_within_litellm(): def test_validate_environment_azure_openai_api_key_within_secret_str(): azure_openai_responses_apiconfig = AzureOpenAIResponsesAPIConfig() - with patch( - "litellm.llms.azure.responses.transformation.get_secret_str" - ) as mock_get_secret_str: + with patch("litellm.api_key", None), \ + patch("litellm.azure_key", None), \ + patch("litellm.llms.azure.common_utils.get_secret_str") as mock_get_secret_str: # Configure the mock to return "test-api-key" when called with "AZURE_OPENAI_API_KEY" mock_get_secret_str.side_effect = ( lambda key: "test-api-key" if key == "AZURE_OPENAI_API_KEY" else None @@ -74,13 +74,19 @@ def test_validate_environment_azure_openai_api_key_within_secret_str(): def test_validate_environment_azure_api_key_within_secret_str(): azure_openai_responses_apiconfig = AzureOpenAIResponsesAPIConfig() - with patch( - "litellm.llms.azure.responses.transformation.get_secret_str" - ) as mock_get_secret_str: - # Configure the mock to return "test-api-key" when called with "AZURE_API_KEY" - mock_get_secret_str.side_effect = ( - lambda key: "test-api-key" if key == "AZURE_API_KEY" else None - ) + with patch("litellm.api_key", None), \ + patch("litellm.azure_key", None), \ + patch("litellm.llms.azure.common_utils.get_secret_str") as mock_get_secret_str: + # Configure the mock to return None for "AZURE_OPENAI_API_KEY" and "test-api-key" for "AZURE_API_KEY" + def mock_side_effect(key): + if key == "AZURE_OPENAI_API_KEY": + return None + elif key == "AZURE_API_KEY": + return "test-api-key" + else: + return None + + mock_get_secret_str.side_effect = mock_side_effect litellm_params = GenericLiteLLMParams() result = azure_openai_responses_apiconfig.validate_environment( diff --git a/tests/vector_store_tests/base_vector_store_test.py b/tests/vector_store_tests/base_vector_store_test.py index d23012847bb..ac1ec704f36 100644 --- a/tests/vector_store_tests/base_vector_store_test.py +++ b/tests/vector_store_tests/base_vector_store_test.py @@ -26,6 +26,11 @@ class BaseVectorStoreTest(ABC): def get_base_request_args(self) -> dict: """Must return the base request args""" pass + + @abstractmethod + def get_base_create_vector_store_args(self) -> dict: + """Must return the base create vector store args""" + pass @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio @@ -52,6 +57,39 @@ class BaseVectorStoreTest(ABC): # Validate response structure self._validate_vector_store_response(response) + @pytest.mark.parametrize("sync_mode", [True, False]) + @pytest.mark.asyncio + async def test_basic_create_vector_store(self, sync_mode): + litellm._turn_on_debug() + litellm.set_verbose = True + base_request_args = self.get_base_create_vector_store_args() + + # Extract custom_llm_provider from base args if present + create_args = base_request_args + try: + if sync_mode: + response = litellm.vector_stores.create( + name="Test Vector Store", + **create_args + ) + else: + response = await litellm.vector_stores.acreate( + name="Test Vector Store", + **create_args + ) + except litellm.InternalServerError: + pytest.skip("Skipping test due to litellm.InternalServerError") + except Exception as e: + # If this is an authentication or permission error, skip the test + if "authentication" in str(e).lower() or "permission" in str(e).lower() or "unauthorized" in str(e).lower(): + pytest.skip(f"Skipping test due to authentication/permission error: {e}") + raise + + print("litellm create response=", json.dumps(response, indent=4, default=str)) + + # Validate response structure + self._validate_vector_store_create_response(response) + def _validate_vector_store_response(self, response): """Validate the structure and content of a vector store search response""" @@ -84,6 +122,87 @@ class BaseVectorStoreTest(ABC): print(f"✅ Response validation passed: Found {len(response['data'])} search results") + def _validate_vector_store_create_response(self, response): + """Validate the structure and content of a vector store create 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 for create response + required_fields = ['id', 'object', 'created_at'] + for field in required_fields: + assert field in response, f"Missing required field '{field}' in create response" + + # Validate object field + assert response['object'] == 'vector_store', \ + f"Expected object to be 'vector_store', got '{response['object']}'" + + # Validate id field + assert isinstance(response['id'], str), \ + f"id should be a string, got {type(response['id'])}" + assert len(response['id']) > 0, "id should not be empty" + assert response['id'].startswith('vs_'), \ + f"id should start with 'vs_', got '{response['id']}'" + + # Validate created_at field + assert isinstance(response['created_at'], int), \ + f"created_at should be an integer, got {type(response['created_at'])}" + assert response['created_at'] > 0, "created_at should be a positive timestamp" + + # Validate optional fields if present + if 'name' in response: + assert isinstance(response['name'], str), \ + f"name should be a string, got {type(response['name'])}" + + if 'bytes' in response: + assert isinstance(response['bytes'], int), \ + f"bytes should be an integer, got {type(response['bytes'])}" + assert response['bytes'] >= 0, "bytes should be non-negative" + + if 'file_counts' in response: + self._validate_file_counts(response['file_counts']) + + if 'status' in response: + valid_statuses = ['expired', 'in_progress', 'completed'] + assert response['status'] in valid_statuses, \ + f"status should be one of {valid_statuses}, got '{response['status']}'" + + if 'expires_at' in response and response['expires_at'] is not None: + assert isinstance(response['expires_at'], int), \ + f"expires_at should be an integer, got {type(response['expires_at'])}" + + if 'last_active_at' in response and response['last_active_at'] is not None: + assert isinstance(response['last_active_at'], int), \ + f"last_active_at should be an integer, got {type(response['last_active_at'])}" + + if 'metadata' in response and response['metadata'] is not None: + assert isinstance(response['metadata'], dict), \ + f"metadata should be a dict, got {type(response['metadata'])}" + + print(f"✅ Create response validation passed: Vector store '{response['id']}' created successfully") + + def _validate_file_counts(self, file_counts): + """Validate file_counts structure""" + assert isinstance(file_counts, dict), \ + f"file_counts should be a dict, got {type(file_counts)}" + + required_count_fields = ['in_progress', 'completed', 'failed', 'cancelled', 'total'] + for field in required_count_fields: + assert field in file_counts, f"Missing required field '{field}' in file_counts" + assert isinstance(file_counts[field], int), \ + f"{field} should be an integer, got {type(file_counts[field])}" + assert file_counts[field] >= 0, f"{field} should be non-negative" + + # Validate that total equals sum of other counts + calculated_total = ( + file_counts['in_progress'] + + file_counts['completed'] + + file_counts['failed'] + + file_counts['cancelled'] + ) + assert file_counts['total'] == calculated_total, \ + f"total should equal sum of other counts ({calculated_total}), got {file_counts['total']}" + def _validate_search_result(self, result, index): """Validate an individual search result""" diff --git a/tests/vector_store_tests/test_azure_vector_store.py b/tests/vector_store_tests/test_azure_vector_store.py new file mode 100644 index 00000000000..783417cb742 --- /dev/null +++ b/tests/vector_store_tests/test_azure_vector_store.py @@ -0,0 +1,25 @@ +from base_vector_store_test import BaseVectorStoreTest +import os +import pytest + +class TestAzureOpenAIVectorStore(BaseVectorStoreTest): + def get_base_request_args(self) -> dict: + """Must return the base request args""" + return {} + + @pytest.mark.parametrize("sync_mode", [True, False]) + @pytest.mark.asyncio + async def test_basic_search_vector_store(self, sync_mode): + pass + + + def get_base_create_vector_store_args(self) -> dict: + """ + This is a real vector store on Azure + """ + return { + "custom_llm_provider": "azure", + "api_base": os.getenv("AZURE_RESPONSES_OPENAI_ENDPOINT"), + "api_key": os.getenv("AZURE_RESPONSES_OPENAI_API_KEY"), + "api_version": "2025-04-01-preview", + } \ No newline at end of file diff --git a/tests/vector_store_tests/test_openai_vector_store.py b/tests/vector_store_tests/test_openai_vector_store.py index 20980f2ff56..3e27be2f64a 100644 --- a/tests/vector_store_tests/test_openai_vector_store.py +++ b/tests/vector_store_tests/test_openai_vector_store.py @@ -8,4 +8,13 @@ class TestOpenAIVectorStore(BaseVectorStoreTest): return { "vector_store_id": "vs_685b14b1a1b88191bc27e04f1917fddd", "custom_llm_provider": "openai", + } + + + def get_base_create_vector_store_args(self) -> dict: + """ + This is a real vector store on OpenAI + """ + return { + "custom_llm_provider": "openai", } \ No newline at end of file