diff --git a/docs/my-website/docs/providers/azure_ai/azure_ai_vector_stores_passthrough.md b/docs/my-website/docs/providers/azure_ai/azure_ai_vector_stores_passthrough.md new file mode 100644 index 00000000000..a528b1ccfcf --- /dev/null +++ b/docs/my-website/docs/providers/azure_ai/azure_ai_vector_stores_passthrough.md @@ -0,0 +1,391 @@ +# Azure AI Search - Vector Store (Passthrough API) + +Use this to allow developers to **create** and **search** vector stores using the Azure AI Search API in the **native** Azure AI Search API format, without giving them the Azure AI credentials. + +This is for the proxy only. + +## Admin Flow + +### 1. Add the vector store to LiteLLM + +```yaml +model_list: + - model_name: embedding-model + litellm_params: + model: openai/text-embedding-3-large + + +vector_store_registry: + - vector_store_name: "azure-ai-search" + litellm_params: + vector_store_id: "can-be-anything" # vector store id can be anything for the purpose of passthrough api + custom_llm_provider: "azure_ai" + api_key: os.environ/AZURE_SEARCH_API_KEY + api_base: https://azure-kb-search.search.windows.net + litellm_embedding_model: "azure/text-embedding-3-large" + litellm_embedding_config: + api_base: https://krris-mh44uf7y-eastus2.cognitiveservices.azure.com/ + api_key: os.environ/AZURE_API_KEY + api_version: "2025-09-01" + +general_settings: + database_url: "postgresql://user:password@host:port/database" + master_key: "sk-1234" +``` + +Add your vector store credentials to LiteLLM. + +### 2. Start the proxy. + +```bash +litellm --config /path/to/config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +### 3. Create a virtual index. + +```bash +curl -L -X POST 'http://0.0.0.0:4000/v1/indexes' \ +-H 'Content-Type: application/json' \ +-H 'Authorization: Bearer sk-1234' \ +-d '{ + "index_name": "dall-e-4", + "litellm_params": { + "vector_store_index": "real-index-name-2", + "vector_store_name": "azure-ai-search" + } + +}' +``` + +This is a virtual index, which the developer can use to create and search vector stores. + +### 4. Create a key with the vector store permissions. + +```bash +curl -L -X POST 'http://0.0.0.0:4000/key/generate' \ +-H 'Content-Type: application/json' \ +-H 'Authorization: Bearer sk-1234' \ +-d '{ + "allowed_vector_store_indexes": [{"index_name": "dall-e-4", "index_permissions": ["write", "read"]}], + "models": ["embedding-model"] +}' +``` + +Give the key access to the virtual index and the embedding model. + +**Expected response** + +```json +{ + "key": "sk-my-virtual-key" +} +``` + +## Developer Flow + +### 1. Create a vector store with some documents. + +Note: Use the '/azure_ai' endpoint for the passthrough api that uses the `azure_ai` provider in your `_new_secret_config.yaml` file. + +```python +import requests +import json + +# ---------------------------- +# ๐Ÿ” CONFIGURATION +# ---------------------------- +# Azure OpenAI (for embeddings) +AZURE_OPENAI_ENDPOINT = "http://0.0.0.0:4000" +AZURE_OPENAI_KEY = "sk-my-virtual-key" +EMBEDDING_DEPLOYMENT_NAME = "embedding-model" + +# Azure AI Search +AZURE_AI_SEARCH_ENDPOINT = "http://0.0.0.0:4000/azure_ai" # IMPORTANT: Use the '/azure_ai' endpoint for the passthrough api to Azure +SEARCH_API_KEY = "sk-my-virtual-key" +INDEX_NAME = "dall-e-4" + + + +# Vector dimensions (text-embedding-3-large uses 3072 dimensions) +VECTOR_DIMENSIONS = 3072 + +# Example docs (replace with your own) +documents = [ + {"id": "1", "content": "Refunds must be requested within 30 days."}, + {"id": "2", "content": "We offer 24/7 support for all enterprise customers."}, +] + + +# ---------------------------- +# ๐Ÿ“‹ STEP 0 โ€” Create Index Schema +# ---------------------------- +def delete_index_if_exists(): + """Delete the index if it exists""" + index_url = f"{AZURE_AI_SEARCH_ENDPOINT}/indexes/{INDEX_NAME}?api-version=2024-07-01" + headers = {"api-key": SEARCH_API_KEY} + + response = requests.delete(index_url, headers=headers) + + if response.status_code == 204: + print(f"๐Ÿ—‘๏ธ Deleted existing index '{INDEX_NAME}'") + return True + elif response.status_code == 404: + print(f"โ„น๏ธ Index '{INDEX_NAME}' does not exist yet") + return False + else: + print(f"โš ๏ธ Delete response: {response.status_code}") + print(f" Message: {response.text}") + return False + + +def create_index(): + """Create the Azure AI Search index with proper schema""" + index_url = f"{AZURE_AI_SEARCH_ENDPOINT}/indexes/{INDEX_NAME}?api-version=2024-07-01" + headers = {"Content-Type": "application/json", "api-key": SEARCH_API_KEY} + + index_schema = { + "name": INDEX_NAME, + "fields": [ + {"name": "id", "type": "Edm.String", "key": True, "filterable": True}, + { + "name": "content", + "type": "Edm.String", + "searchable": True, + "filterable": False, + }, + { + "name": "contentVector", + "type": "Collection(Edm.Single)", + "searchable": True, + "dimensions": VECTOR_DIMENSIONS, + "vectorSearchProfile": "my-vector-profile", + }, + ], + "vectorSearch": { + "algorithms": [ + { + "name": "my-hnsw-algorithm", + "kind": "hnsw", + "hnswParameters": { + "metric": "cosine", + "m": 4, + "efConstruction": 400, + "efSearch": 500, + }, + } + ], + "profiles": [ + {"name": "my-vector-profile", "algorithm": "my-hnsw-algorithm"} + ], + }, + } + + # Create the index + response = requests.put(index_url, headers=headers, json=index_schema) + + if response.status_code == 201: + print(f"โœ… Index '{INDEX_NAME}' created successfully.") + return True + elif response.status_code == 204: + print(f"โœ… Index '{INDEX_NAME}' updated successfully.") + return True + else: + print(f"โŒ Failed to create index: {response.status_code}") + print(f" Message: {response.text}") + return False + + +# Delete and recreate the index with correct schema +print("๐Ÿ”„ Setting up Azure AI Search index...") +delete_index_if_exists() +if not create_index(): + print("โŒ Could not create index. Exiting.") + exit(1) + + +# ---------------------------- +# ๐Ÿง  STEP 1 โ€” Generate Embeddings +# ---------------------------- +def get_embedding(text: str): + url = f"{AZURE_OPENAI_ENDPOINT}/openai/deployments/{EMBEDDING_DEPLOYMENT_NAME}/embeddings?api-version=2024-10-21" + headers = {"Content-Type": "application/json", "api-key": AZURE_OPENAI_KEY} + payload = {"input": text} + response = requests.post(url, headers=headers, json=payload) + + if response.status_code != 200: + raise Exception(f"Embedding failed: {response.status_code}\n{response.text}") + return response.json()["data"][0]["embedding"] + + +# Generate embeddings for each document +for doc in documents: + doc["contentVector"] = get_embedding(doc["content"]) + print(f"โœ… Embedded doc {doc['id']} (vector length: {len(doc['contentVector'])})") + +# ---------------------------- +# ๐Ÿ“ค STEP 2 โ€” Upload to Azure AI Search +# ---------------------------- +upload_url = f"{AZURE_AI_SEARCH_ENDPOINT}/indexes/{INDEX_NAME}/docs/index?api-version=2024-07-01" +headers = {"Content-Type": "application/json", "api-key": SEARCH_API_KEY} + +payload = { + "value": [ + { + "@search.action": "upload", + "id": doc["id"], + "content": doc["content"], + "contentVector": doc["contentVector"], + } + for doc in documents + ] +} + +response = requests.post(upload_url, headers=headers, data=json.dumps(payload)) + +# ---------------------------- +# ๐Ÿงพ RESULT +# ---------------------------- +if response.status_code == 200: + print("โœ… Documents uploaded successfully.") +else: + print(f"โŒ Upload failed: {response.status_code}") + print(response.text) + +``` + + +### 2. Search the vector store. + + +```python +import requests +import json + +# ---------------------------- +# ๐Ÿ” CONFIGURATION +# ---------------------------- +# Azure OpenAI (for embeddings) +AZURE_OPENAI_ENDPOINT = "http://0.0.0.0:4000" +AZURE_OPENAI_KEY = "sk-my-virtual-key" +EMBEDDING_DEPLOYMENT_NAME = "embedding-model" + +# Azure AI Search +AZURE_AI_SEARCH_ENDPOINT = "http://0.0.0.0:4000/azure_ai" +SEARCH_API_KEY = "sk-my-virtual-key" +INDEX_NAME = "dall-e-4" + + +# ---------------------------- +# ๐Ÿง  Generate Query Embedding +# ---------------------------- +def get_embedding(text: str): + """Generate embedding for the query text""" + url = f"{AZURE_OPENAI_ENDPOINT}/openai/deployments/{EMBEDDING_DEPLOYMENT_NAME}/embeddings?api-version=2024-10-21" + headers = {"Content-Type": "application/json", "api-key": AZURE_OPENAI_KEY} + payload = {"input": text} + response = requests.post(url, headers=headers, json=payload) + + if response.status_code != 200: + raise Exception(f"Embedding failed: {response.status_code}\n{response.text}") + return response.json()["data"][0]["embedding"] + + +# ---------------------------- +# ๐Ÿ” Vector Search Function +# ---------------------------- +def search_knowledge_base(query: str, top_k: int = 3): + """ + Search the knowledge base using vector similarity + + Args: + query: The search query string + top_k: Number of top results to return (default: 3) + + Returns: + List of search results with content and scores + """ + print(f"๐Ÿ” Searching for: '{query}'") + + # Step 1: Generate embedding for the query + print(" Generating query embedding...") + query_vector = get_embedding(query) + + # Step 2: Perform vector search + search_url = f"{AZURE_AI_SEARCH_ENDPOINT}/indexes/{INDEX_NAME}/docs/search?api-version=2024-07-01" + headers = {"Content-Type": "application/json", "api-key": SEARCH_API_KEY} + + # Build the search request with vector search + search_payload = { + "search": "*", # Get all documents + "vectorQueries": [ + { + "vector": query_vector, + "fields": "contentVector", + "kind": "vector", + "k": top_k, # Number of nearest neighbors to return + } + ], + "select": "id,content", # Fields to return + "top": top_k, + } + + # Execute the search + response = requests.post(search_url, headers=headers, json=search_payload) + + if response.status_code != 200: + raise Exception(f"Search failed: {response.status_code}\n{response.text}") + + # Parse and return results + results = response.json() + return results.get("value", []) + + +# ---------------------------- +# ๐Ÿ“Š Display Results +# ---------------------------- +def display_results(results): + """Pretty print the search results""" + if not results: + print("\nโŒ No results found.") + return + + print(f"\nโœ… Found {len(results)} results:\n") + print("=" * 80) + + for i, result in enumerate(results, 1): + print(f"\n๐Ÿ“„ Result #{i}") + print(f" ID: {result.get('id', 'N/A')}") + print(f" Score: {result.get('@search.score', 'N/A')}") + print(f" Content: {result.get('content', 'N/A')}") + print("-" * 80) + + +# ---------------------------- +# ๐ŸŽฏ MAIN - Example Queries +# ---------------------------- +if __name__ == "__main__": + # Example 1: Search for refund policy + print("\n" + "=" * 80) + print("EXAMPLE 1: Refund Policy Query") + print("=" * 80) + results = search_knowledge_base("How do I get a refund?", top_k=2) + display_results(results) + + # Example 2: Search for customer support + print("\n\n" + "=" * 80) + print("EXAMPLE 2: Customer Support Query") + print("=" * 80) + results = search_knowledge_base("When can I contact support?", top_k=2) + display_results(results) + + # Example 3: Custom query - uncomment to use + # print("\n\n" + "=" * 80) + # print("CUSTOM QUERY") + # print("=" * 80) + # custom_query = input("Enter your query: ") + # results = search_knowledge_base(custom_query, top_k=3) + # display_results(results) + +``` \ No newline at end of file diff --git a/docs/my-website/docs/providers/azure_ai_vector_stores.md b/docs/my-website/docs/providers/azure_ai_vector_stores.md index d3abb78bbe4..b9dfa3bdc9c 100644 --- a/docs/my-website/docs/providers/azure_ai_vector_stores.md +++ b/docs/my-website/docs/providers/azure_ai_vector_stores.md @@ -1,9 +1,9 @@ import Tabs from '@theme/Tabs'; import TabItem from '@theme/TabItem'; -# Azure AI Search - Vector Store +# Azure AI Search - Vector Store (Unified API) -Use Azure AI Search as a vector store for RAG. +Use this to **search** Azure AI Search Vector Stores, with LiteLLM's unified `/chat/completions` API. ## Quick Start diff --git a/docs/my-website/docs/vector_stores/create.md b/docs/my-website/docs/vector_stores/create.md index c97f88a2543..19b4f39cd9e 100644 --- a/docs/my-website/docs/vector_stores/create.md +++ b/docs/my-website/docs/vector_stores/create.md @@ -12,7 +12,8 @@ Create a vector store which can be used to store and search document chunks for | Cost Tracking | โœ… | Tracked per vector store operation | | Logging | โœ… | Works across all integrations | | End-user Tracking | โœ… | | -| Support LLM Providers | **OpenAI** | Full vector stores API support across providers | +| Support LLM Providers (OpenAI `/vector_stores` API) | **OpenAI** | Full vector stores API support across providers | +| Support LLM Providers (Passthrough API) | [**Azure AI**](/docs/providers/azure_ai/azure_ai_vector_stores_passthrough) | Full vector stores API support across providers | ## Usage diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 1ed2840d29e..d8ba2b1c0a5 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -457,6 +457,7 @@ const sidebars = { "providers/azure_ai_speech", "providers/azure_ai_img", "providers/azure_ai_vector_stores", + "providers/azure_ai/azure_ai_vector_stores_passthrough", ] }, { diff --git a/litellm/__init__.py b/litellm/__init__.py index 04463266947..d141a3fb28f 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -90,7 +90,11 @@ from litellm.types.proxy.management_endpoints.ui_sso import ( DefaultTeamSSOParams, LiteLLM_UpperboundKeyGenerateParams, ) -from litellm.types.utils import StandardKeyGenerationConfig, LlmProviders, SearchProviders +from litellm.types.utils import ( + StandardKeyGenerationConfig, + LlmProviders, + SearchProviders, +) from litellm.types.utils import PriorityReservationSettings from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.logging_callback_manager import LoggingCallbackManager @@ -157,9 +161,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "cloudzero", "posthog", ] -cold_storage_custom_logger: Optional[ - _custom_logger_compatible_callbacks_literal -] = None +cold_storage_custom_logger: Optional[_custom_logger_compatible_callbacks_literal] = None logged_real_time_event_types: Optional[Union[List[str], Literal["*"]]] = None _known_custom_logger_compatible_callbacks: List = list( get_args(_custom_logger_compatible_callbacks_literal) @@ -174,22 +176,22 @@ prometheus_initialize_budget_metrics: Optional[bool] = False require_auth_for_metrics_endpoint: Optional[bool] = False argilla_batch_size: Optional[int] = None datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload. -gcs_pub_sub_use_v1: Optional[ - bool -] = False # if you want to use v1 gcs pubsub logged payload -generic_api_use_v1: Optional[ - bool -] = False # if you want to use v1 generic api logged payload +gcs_pub_sub_use_v1: Optional[bool] = ( + False # if you want to use v1 gcs pubsub logged payload +) +generic_api_use_v1: Optional[bool] = ( + False # if you want to use v1 generic api logged payload +) argilla_transformation_object: Optional[Dict[str, Any]] = None -_async_input_callback: List[ - Union[str, Callable, CustomLogger] -] = [] # internal variable - async custom callbacks are routed here. -_async_success_callback: List[ - Union[str, Callable, CustomLogger] -] = [] # internal variable - async custom callbacks are routed here. -_async_failure_callback: List[ - Union[str, Callable, CustomLogger] -] = [] # internal variable - async custom callbacks are routed here. +_async_input_callback: List[Union[str, Callable, CustomLogger]] = ( + [] +) # internal variable - async custom callbacks are routed here. +_async_success_callback: List[Union[str, Callable, CustomLogger]] = ( + [] +) # internal variable - async custom callbacks are routed here. +_async_failure_callback: List[Union[str, Callable, CustomLogger]] = ( + [] +) # internal variable - async custom callbacks are routed here. pre_call_rules: List[Callable] = [] post_call_rules: List[Callable] = [] turn_off_message_logging: Optional[bool] = False @@ -197,18 +199,18 @@ log_raw_request_response: bool = False redact_messages_in_exceptions: Optional[bool] = False redact_user_api_key_info: Optional[bool] = False filter_invalid_headers: Optional[bool] = False -add_user_information_to_llm_headers: Optional[ - bool -] = None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers +add_user_information_to_llm_headers: Optional[bool] = ( + None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers +) store_audit_logs = False # Enterprise feature, allow users to see audit logs ### end of callbacks ############# -email: Optional[ - str -] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -token: Optional[ - str -] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +email: Optional[str] = ( + None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +token: Optional[str] = ( + None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) telemetry = True max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False)) @@ -264,7 +266,9 @@ use_client: bool = False ssl_verify: Union[str, bool] = True ssl_security_level: Optional[str] = None ssl_certificate: Optional[str] = None -ssl_ecdh_curve: Optional[str] = None # Set to 'X25519' to disable PQC and improve performance +ssl_ecdh_curve: Optional[str] = ( + None # Set to 'X25519' to disable PQC and improve performance +) disable_streaming_logging: bool = False disable_token_counter: bool = False disable_add_transform_inline_image_block: bool = False @@ -310,20 +314,24 @@ enable_loadbalancing_on_batch_endpoints: Optional[bool] = None enable_caching_on_provider_specific_optional_params: bool = ( False # feature-flag for caching on optional params - e.g. 'top_k' ) -caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -caching_with_models: bool = False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -cache: Optional[ - Cache -] = None # cache object <- use this - https://docs.litellm.ai/docs/caching +caching: bool = ( + False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +caching_with_models: bool = ( + False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +cache: Optional[Cache] = ( + None # cache object <- use this - https://docs.litellm.ai/docs/caching +) default_in_memory_ttl: Optional[float] = None default_redis_ttl: Optional[float] = None default_redis_batch_cache_expiry: Optional[float] = None model_alias_map: Dict[str, str] = {} model_group_settings: Optional["ModelGroupSettings"] = None max_budget: float = 0.0 # set the max budget across all providers -budget_duration: Optional[ - str -] = None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). +budget_duration: Optional[str] = ( + None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). +) default_soft_budget: float = ( DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0 ) @@ -332,11 +340,15 @@ forward_traceparent_to_llm_provider: bool = False _current_cost = 0.0 # private variable, used if max budget is set error_logs: Dict = {} -add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt +add_function_to_prompt: bool = ( + False # if function calling not supported by api, append function call details to system prompt +) client_session: Optional[httpx.Client] = None aclient_session: Optional[httpx.AsyncClient] = None model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks' -model_cost_map_url: str = "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json" +model_cost_map_url: str = ( + "https://raw.githubusercontent.com/BerriAI/litellm/main/model_prices_and_context_window.json" +) suppress_debug_info = False dynamodb_table_name: Optional[str] = None s3_callback_params: Optional[Dict] = None @@ -366,7 +378,9 @@ prometheus_metrics_config: Optional[List] = None disable_add_prefix_to_prompt: bool = ( False # used by anthropic, to disable adding prefix to prompt ) -disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. +disable_copilot_system_to_assistant: bool = ( + False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. +) public_model_groups: Optional[List[str]] = None public_model_groups_links: Dict[str, str] = {} #### REQUEST PRIORITIZATION ####### @@ -377,13 +391,17 @@ priority_reservation_settings: "PriorityReservationSettings" = ( ######## Networking Settings ######## -use_aiohttp_transport: bool = True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead. +use_aiohttp_transport: bool = ( + True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead. +) aiohttp_trust_env: bool = False # set to true to use HTTP_ Proxy settings disable_aiohttp_transport: bool = False # Set this to true to use httpx instead disable_aiohttp_trust_env: bool = ( False # When False, aiohttp will respect HTTP(S)_PROXY env vars ) -force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. +force_ipv4: bool = ( + False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. +) module_level_aclient = AsyncHTTPHandler( timeout=request_timeout, client_alias="module level aclient" ) @@ -397,13 +415,13 @@ fallbacks: Optional[List] = None context_window_fallbacks: Optional[List] = None content_policy_fallbacks: Optional[List] = None allowed_fails: int = 3 -num_retries_per_request: Optional[ - int -] = None # for the request overall (incl. fallbacks + model retries) +num_retries_per_request: Optional[int] = ( + None # for the request overall (incl. fallbacks + model retries) +) ####### SECRET MANAGERS ##################### -secret_manager_client: Optional[ - Any -] = None # list of instantiated key management clients - e.g. azure kv, infisical, etc. +secret_manager_client: Optional[Any] = ( + None # list of instantiated key management clients - e.g. azure kv, infisical, etc. +) _google_kms_resource_name: Optional[str] = None _key_management_system: Optional[KeyManagementSystem] = None _key_management_settings: KeyManagementSettings = KeyManagementSettings() @@ -413,7 +431,9 @@ output_parse_pii: bool = False from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map model_cost = get_model_cost_map(url=model_cost_map_url) -cost_discount_config: Dict[str, float] = {} # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount +cost_discount_config: Dict[str, float] = ( + {} +) # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount custom_prompt_dict: Dict[str, dict] = {} check_provider_endpoint = False @@ -771,9 +791,9 @@ azure_llms = { "gpt-35-turbo": "azure/gpt-35-turbo", "gpt-35-turbo-16k": "azure/gpt-35-turbo-16k", "gpt-35-turbo-instruct": "azure/gpt-35-turbo-instruct", - "azure/gpt-41":"gpt-4.1", - "azure/gpt-41-mini":"gpt-4.1-mini", - "azure/gpt-41-nano":"gpt-4.1-nano" + "azure/gpt-41": "gpt-4.1", + "azure/gpt-41-mini": "gpt-4.1-mini", + "azure/gpt-41-nano": "gpt-4.1-nano", } azure_embedding_models = { @@ -1188,7 +1208,9 @@ from .llms.bedrock.embed.amazon_titan_v2_transformation import ( from .llms.cohere.chat.transformation import CohereChatConfig from .llms.cohere.chat.v2_transformation import CohereV2ChatConfig from .llms.bedrock.embed.cohere_transformation import BedrockCohereEmbeddingConfig -from .llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig +from .llms.bedrock.embed.twelvelabs_marengo_transformation import ( + TwelveLabsMarengoEmbeddingConfig, +) from .llms.openai.openai import OpenAIConfig, MistralEmbeddingConfig from .llms.openai.image_variations.transformation import OpenAIImageVariationConfig from .llms.deepinfra.chat.transformation import DeepInfraConfig @@ -1360,21 +1382,25 @@ import litellm.anthropic_interface as anthropic adapters: List[AdapterItem] = [] ### Vector Store Registry ### -from .vector_stores.vector_store_registry import VectorStoreRegistry +from .vector_stores.vector_store_registry import ( + VectorStoreRegistry, + VectorStoreIndexRegistry, +) vector_store_registry: Optional[VectorStoreRegistry] = None +vector_store_index_registry: Optional[VectorStoreIndexRegistry] = None ### CUSTOM LLMs ### from .types.llms.custom_llm import CustomLLMItem from .types.utils import GenericStreamingChunk custom_provider_map: List[CustomLLMItem] = [] -_custom_providers: List[ - str -] = [] # internal helper util, used to track names of custom providers -disable_hf_tokenizer_download: Optional[ - bool -] = None # disable huggingface tokenizer download. Defaults to openai clk100 +_custom_providers: List[str] = ( + [] +) # internal helper util, used to track names of custom providers +disable_hf_tokenizer_download: Optional[bool] = ( + None # disable huggingface tokenizer download. Defaults to openai clk100 +) global_disable_no_log_param: bool = False ### CLI UTILITIES ### @@ -1393,9 +1419,11 @@ def set_global_bitbucket_config(config: Dict[str, Any]) -> None: global global_bitbucket_config global_bitbucket_config = config + ### GLOBAL CONFIG ### global_gitlab_config: Optional[Dict[str, Any]] = None + def set_global_gitlab_config(config: Dict[str, Any]) -> None: """Set global BitBucket configuration for prompt management.""" global global_gitlab_config diff --git a/litellm/llms/azure_ai/vector_stores/transformation.py b/litellm/llms/azure_ai/vector_stores/transformation.py index f99d2c4c4b2..96cea064ce1 100644 --- a/litellm/llms/azure_ai/vector_stores/transformation.py +++ b/litellm/llms/azure_ai/vector_stores/transformation.py @@ -7,8 +7,10 @@ from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + BaseVectorStoreAuthCredentials, VectorStoreCreateOptionalRequestParams, VectorStoreCreateResponse, + VectorStoreIndexEndpoints, VectorStoreResultContent, VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, @@ -34,6 +36,25 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): def __init__(self): super().__init__() + def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints: + return { + "read": [("GET", "/docs/search"), ("POST", "/docs/search")], + "write": [("PUT", "/docs")], + } + + def get_auth_credentials( + self, litellm_params: dict + ) -> BaseVectorStoreAuthCredentials: + api_key = litellm_params.get("api_key") + if api_key is None: + raise ValueError("api_key is required") + + return { + "headers": { + "api-key": api_key, + } + } + def validate_environment( self, headers: dict, litellm_params: Optional[GenericLiteLLMParams] ) -> dict: diff --git a/litellm/llms/base_llm/vector_store/transformation.py b/litellm/llms/base_llm/vector_store/transformation.py index daca48fc323..89f2094d5df 100644 --- a/litellm/llms/base_llm/vector_store/transformation.py +++ b/litellm/llms/base_llm/vector_store/transformation.py @@ -5,9 +5,11 @@ import httpx from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + BaseVectorStoreAuthCredentials, VECTOR_STORE_OPENAI_PARAMS, VectorStoreCreateOptionalRequestParams, VectorStoreCreateResponse, + VectorStoreIndexEndpoints, VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, ) @@ -39,6 +41,16 @@ class BaseVectorStoreConfig: ) -> dict: return optional_params + @abstractmethod + def get_auth_credentials( + self, litellm_params: dict + ) -> BaseVectorStoreAuthCredentials: + pass + + @abstractmethod + def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints: + pass + @abstractmethod def transform_search_vector_store_request( self, diff --git a/litellm/llms/bedrock/vector_stores/transformation.py b/litellm/llms/bedrock/vector_stores/transformation.py index c485defe8a8..72e1e1470d3 100644 --- a/litellm/llms/bedrock/vector_stores/transformation.py +++ b/litellm/llms/bedrock/vector_stores/transformation.py @@ -13,6 +13,8 @@ from litellm.types.integrations.rag.bedrock_knowledgebase import ( ) from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + BaseVectorStoreAuthCredentials, + VectorStoreIndexEndpoints, VECTOR_STORE_OPENAI_PARAMS, VectorStoreResultContent, VectorStoreSearchOptionalRequestParams, @@ -33,6 +35,17 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM): BaseVectorStoreConfig.__init__(self) BaseAWSLLM.__init__(self) + def get_auth_credentials( + self, litellm_params: dict + ) -> BaseVectorStoreAuthCredentials: + return {} + + def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints: + return { + "read": [("POST", "/knowledgebases/{knowledge_base_id}/retrieve")], + "write": [], + } + def get_supported_openai_params( self, model: str ) -> List[VECTOR_STORE_OPENAI_PARAMS]: diff --git a/litellm/llms/openai/vector_stores/transformation.py b/litellm/llms/openai/vector_stores/transformation.py index 76cd12be8ee..5689e249bc0 100644 --- a/litellm/llms/openai/vector_stores/transformation.py +++ b/litellm/llms/openai/vector_stores/transformation.py @@ -7,9 +7,11 @@ 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 ( + BaseVectorStoreAuthCredentials, VectorStoreCreateOptionalRequestParams, VectorStoreCreateRequest, VectorStoreCreateResponse, + VectorStoreIndexEndpoints, VectorStoreSearchOptionalRequestParams, VectorStoreSearchRequest, VectorStoreSearchResponse, @@ -23,10 +25,29 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any + class OpenAIVectorStoreConfig(BaseVectorStoreConfig): ASSISTANTS_HEADER_KEY = "OpenAI-Beta" ASSISTANTS_HEADER_VALUE = "assistants=v2" + def get_auth_credentials( + self, litellm_params: dict + ) -> BaseVectorStoreAuthCredentials: + api_key = litellm_params.get("api_key") + if api_key is None: + raise ValueError("api_key is required") + return { + "headers": { + "Authorization": f"Bearer {api_key}", + }, + } + + def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints: + return { + "read": [("GET", "/vector_stores/{index_name}/search")], + "write": [("POST", "/vector_stores")], + } + def validate_environment( self, headers: dict, litellm_params: Optional[GenericLiteLLMParams] ) -> dict: @@ -51,8 +72,8 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): headers.update( { self.ASSISTANTS_HEADER_KEY: self.ASSISTANTS_HEADER_VALUE, - } - ) + } + ) return headers @@ -76,7 +97,6 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): api_base = api_base.rstrip("/") return f"{api_base}/vector_stores" - def transform_search_vector_store_request( self, @@ -91,27 +111,31 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): 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), + 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, litellm_logging_obj: LiteLLMLoggingObj) -> VectorStoreSearchResponse: + def transform_search_vector_store_response( + self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj + ) -> VectorStoreSearchResponse: try: response_json = response.json() - return VectorStoreSearchResponse( - **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 + error_message=str(e), + status_code=response.status_code, + headers=response.headers, ) def transform_create_vector_store_request( @@ -124,28 +148,27 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): 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), + expires_after=vector_store_create_optional_params.get( + "expires_after", None + ), + chunking_strategy=vector_store_create_optional_params.get( + "chunking_strategy", None + ), metadata=add_openai_metadata(metadata) if metadata is not None else 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: + def transform_create_vector_store_response( + self, response: httpx.Response + ) -> VectorStoreCreateResponse: try: response_json = response.json() - return VectorStoreCreateResponse( - **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 + error_message=str(e), + status_code=response.status_code, + headers=response.headers, ) - - - - - \ No newline at end of file diff --git a/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py index 5296b11e883..6f258bc04a6 100644 --- a/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/rag_api/transformation.py @@ -6,8 +6,10 @@ from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreCon from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + BaseVectorStoreAuthCredentials, VectorStoreCreateOptionalRequestParams, VectorStoreCreateResponse, + VectorStoreIndexEndpoints, VectorStoreResultContent, VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, @@ -25,13 +27,40 @@ else: class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): """ Configuration for Vertex AI Vector Store RAG API - + This implementation uses the Vertex AI RAG Engine API for vector store operations. """ def __init__(self): super().__init__() + def get_auth_credentials( + self, litellm_params: dict + ) -> BaseVectorStoreAuthCredentials: + # Get credentials and project info + vertex_credentials = self.get_vertex_ai_credentials(dict(litellm_params)) + vertex_project = self.get_vertex_ai_project(dict(litellm_params)) + + # Get access token using the base class method + access_token, project_id = self._ensure_access_token( + credentials=vertex_credentials, + project_id=vertex_project, + custom_llm_provider="vertex_ai", + ) + + return { + "headers": { + "Authorization": f"Bearer {access_token}", + "Content-Type": "application/json", + }, + } + + def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints: + return { + "read": [("POST", ":retrieveContexts")], + "write": [("POST", "/ragCorpora")], + } + def validate_environment( self, headers: dict, litellm_params: Optional[GenericLiteLLMParams] ) -> dict: @@ -39,23 +68,9 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): Validate and set up authentication for Vertex AI RAG API """ litellm_params = litellm_params or GenericLiteLLMParams() - - # Get credentials and project info - vertex_credentials = self.get_vertex_ai_credentials(dict(litellm_params)) - vertex_project = self.get_vertex_ai_project(dict(litellm_params)) - - # Get access token using the base class method - access_token, project_id = self._ensure_access_token( - credentials=vertex_credentials, - project_id=vertex_project, - custom_llm_provider="vertex_ai", - ) - - headers.update({ - "Authorization": f"Bearer {access_token}", - "Content-Type": "application/json", - }) - + + auth_headers = self.get_auth_credentials(litellm_params.model_dump()) + headers.update(auth_headers.get("headers", {})) return headers def get_complete_url( @@ -68,10 +83,10 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): """ vertex_location = self.get_vertex_ai_location(litellm_params) vertex_project = self.get_vertex_ai_project(litellm_params) - + if api_base: return api_base.rstrip("/") - + # Vertex AI RAG API endpoint for retrieveContexts return f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}" @@ -90,60 +105,52 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): # Convert query to string if it's a list if isinstance(query, list): query = " ".join(query) - + # Vertex AI RAG API endpoint for retrieving contexts url = f"{api_base}:retrieveContexts" - + # Use helper methods to get project and location, then construct full rag corpus path vertex_project = self.get_vertex_ai_project(litellm_params) vertex_location = self.get_vertex_ai_location(litellm_params) - + # Construct full rag corpus path full_rag_corpus = f"projects/{vertex_project}/locations/{vertex_location}/ragCorpora/{vector_store_id}" - + # Build the request body for Vertex AI RAG API request_body: Dict[str, Any] = { - "vertex_rag_store": { - "rag_resources": [ - { - "rag_corpus": full_rag_corpus - } - ] - }, - "query": { - "text": query - } + "vertex_rag_store": {"rag_resources": [{"rag_corpus": full_rag_corpus}]}, + "query": {"text": query}, } ######################################################### # Update logging object with details of the request ######################################################### litellm_logging_obj.model_call_details["query"] = query - + # Add optional parameters max_num_results = vector_store_search_optional_params.get("max_num_results") if max_num_results is not None: - request_body["query"]["rag_retrieval_config"] = { - "top_k": max_num_results - } - + request_body["query"]["rag_retrieval_config"] = {"top_k": max_num_results} + # Add filters if provided filters = vector_store_search_optional_params.get("filters") if filters is not None: if "rag_retrieval_config" not in request_body["query"]: request_body["query"]["rag_retrieval_config"] = {} request_body["query"]["rag_retrieval_config"]["filter"] = filters - + # Add ranking options if provided ranking_options = vector_store_search_optional_params.get("ranking_options") if ranking_options is not None: if "rag_retrieval_config" not in request_body["query"]: request_body["query"]["rag_retrieval_config"] = {} request_body["query"]["rag_retrieval_config"]["ranking"] = ranking_options - + return url, request_body - def transform_search_vector_store_response(self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj) -> VectorStoreSearchResponse: + def transform_search_vector_store_response( + self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj + ) -> VectorStoreSearchResponse: """ Transform Vertex AI RAG API response to standard vector store search response """ @@ -152,7 +159,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): response_json = response.json() # Extract contexts from Vertex AI response - handle nested structure contexts = response_json.get("contexts", {}).get("contexts", []) - + # Transform contexts to standard format search_results = [] for context in contexts: @@ -162,27 +169,29 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): type="text", ) ] - + # Extract file information source_uri = context.get("sourceUri", "") source_display_name = context.get("sourceDisplayName", "") - + # Generate file_id from source URI or use display name as fallback file_id = source_uri if source_uri else source_display_name - filename = source_display_name if source_display_name else "Unknown Document" - + filename = ( + source_display_name if source_display_name else "Unknown Document" + ) + # Build attributes with available metadata attributes = {} if source_uri: attributes["sourceUri"] = source_uri if source_display_name: attributes["sourceDisplayName"] = source_display_name - + # Add page span information if available page_span = context.get("pageSpan", {}) if page_span: attributes["pageSpan"] = page_span - + result = VectorStoreSearchResult( score=context.get("score", 0.0), content=content, @@ -191,18 +200,18 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): attributes=attributes, ) search_results.append(result) - + return VectorStoreSearchResponse( object="vector_store.search_results.page", search_query=litellm_logging_obj.model_call_details.get("query", ""), - data=search_results + data=search_results, ) - + except Exception as e: raise self.get_error_class( error_message=str(e), status_code=response.status_code, - headers=response.headers + headers=response.headers, ) def transform_create_vector_store_request( @@ -214,48 +223,55 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): Transform create request for Vertex AI RAG Corpus """ url = f"{api_base}/ragCorpora" # Base URL for creating RAG corpus - + # Build the request body for Vertex AI RAG Corpus creation request_body: Dict[str, Any] = { - "display_name": vector_store_create_optional_params.get("name", "litellm-vector-store"), - "description": "Vector store created via LiteLLM" + "display_name": vector_store_create_optional_params.get( + "name", "litellm-vector-store" + ), + "description": "Vector store created via LiteLLM", } - + # Add metadata if provided metadata = vector_store_create_optional_params.get("metadata") if metadata is not None: request_body["labels"] = metadata - + return url, request_body - def transform_create_vector_store_response(self, response: httpx.Response) -> VectorStoreCreateResponse: + def transform_create_vector_store_response( + self, response: httpx.Response + ) -> VectorStoreCreateResponse: """ Transform Vertex AI RAG Corpus creation response to standard vector store response """ try: response_json = response.json() - + # Extract the corpus ID from the response name corpus_name = response_json.get("name", "") - corpus_id = corpus_name.split("/")[-1] if "/" in corpus_name else corpus_name - + corpus_id = ( + corpus_name.split("/")[-1] if "/" in corpus_name else corpus_name + ) + # Handle createTime conversion create_time = response_json.get("createTime", 0) if isinstance(create_time, str): # Convert ISO timestamp to Unix timestamp from datetime import datetime + try: - dt = datetime.fromisoformat(create_time.replace('Z', '+00:00')) + dt = datetime.fromisoformat(create_time.replace("Z", "+00:00")) create_time = int(dt.timestamp()) except ValueError: create_time = 0 elif not isinstance(create_time, int): create_time = 0 - + # Handle labels safely labels = response_json.get("labels", {}) metadata = labels if isinstance(labels, dict) else {} - + return VectorStoreCreateResponse( id=corpus_id, object="vector_store", @@ -267,18 +283,18 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase): "completed": 0, "failed": 0, "cancelled": 0, - "total": 0 + "total": 0, }, status="completed", # Vertex AI corpus creation is typically synchronous expires_after=None, expires_at=None, last_active_at=None, - metadata=metadata + metadata=metadata, ) - + except Exception as e: raise self.get_error_class( error_message=str(e), status_code=response.status_code, - headers=response.headers - ) \ No newline at end of file + headers=response.headers, + ) diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index 31837146485..8fa65002d53 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -7,8 +7,10 @@ from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreCon from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.types.router import GenericLiteLLMParams from litellm.types.vector_stores import ( + BaseVectorStoreAuthCredentials, VectorStoreCreateOptionalRequestParams, VectorStoreCreateResponse, + VectorStoreIndexEndpoints, VectorStoreResultContent, VectorStoreSearchOptionalRequestParams, VectorStoreSearchResponse, @@ -33,14 +35,9 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): def __init__(self): super().__init__() - def validate_environment( - self, headers: dict, litellm_params: Optional[GenericLiteLLMParams] - ) -> dict: - """ - Validate and set up authentication for Vertex AI RAG API - """ - litellm_params = litellm_params or GenericLiteLLMParams() - + def get_auth_credentials( + self, litellm_params: dict + ) -> BaseVectorStoreAuthCredentials: # Get credentials and project info vertex_credentials = self.get_vertex_ai_credentials(dict(litellm_params)) vertex_project = self.get_vertex_ai_project(dict(litellm_params)) @@ -52,13 +49,28 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase): custom_llm_provider="vertex_ai", ) - headers.update( - { + return { + "headers": { "Authorization": f"Bearer {access_token}", "Content-Type": "application/json", - } - ) + }, + } + def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints: + return { + "read": [("POST", ":search")], + "write": [], + } + + def validate_environment( + self, headers: dict, litellm_params: Optional[GenericLiteLLMParams] + ) -> dict: + """ + Validate and set up authentication for Vertex AI RAG API + """ + litellm_params = litellm_params or GenericLiteLLMParams() + auth_headers = self.get_auth_credentials(litellm_params.model_dump()) + headers.update(auth_headers.get("headers", {})) return headers def get_complete_url( diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 8544956db92..6d0936d1980 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -4,10 +4,7 @@ model_list: model: bedrock/global.anthropic.claude-sonnet-4-5-20250929-v1:0 - model_name: embedding-model litellm_params: - model: azure/text-embedding-3-large - api_base: https://krris-mh44uf7y-eastus2.cognitiveservices.azure.com/ - api_key: os.environ/AZURE_API_KEY - api_version: "2025-09-01" + model: openai/text-embedding-3-large general_settings: master_key: sk-1234 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b1082c4cc49..20a78d839a3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -340,6 +340,7 @@ class LiteLLMRoutes(enum.Enum): "/anthropic", "/langfuse", "/azure", + "/azure_ai", "/openai", "/assemblyai", "/eu.assemblyai", @@ -777,6 +778,11 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): object_permission: Optional[LiteLLM_ObjectPermissionBase] = None +class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase): + index_name: str + index_permissions: List[Literal["read", "write"]] + + class KeyRequestBase(GenerateRequestBase): key: Optional[str] = None budget_id: Optional[str] = None @@ -784,6 +790,7 @@ class KeyRequestBase(GenerateRequestBase): enforced_params: Optional[List[str]] = None allowed_routes: Optional[list] = [] allowed_passthrough_routes: Optional[list] = None + allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None rpm_limit_type: Optional[ Literal["guaranteed_throughput", "best_effort_throughput", "dynamic"] ] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating rpm @@ -1311,6 +1318,7 @@ class NewTeamRequest(TeamBase): None # allow user to set TPM limit for all team members ) team_member_key_duration: Optional[str] = None # e.g. "1d", "1w", "1m" + allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None model_config = ConfigDict(protected_namespaces=()) @@ -1359,6 +1367,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): allowed_passthrough_routes: Optional[list] = None model_rpm_limit: Optional[Dict[str, int]] = None model_tpm_limit: Optional[Dict[str, int]] = None + allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase): @@ -3177,11 +3186,10 @@ LiteLLM_ManagementEndpoint_MetadataFields = [ "model_tpm_limit", "rpm_limit_type", "tpm_limit_type", - "guardrails", - "tags", "enforced_params", "temp_budget_increase", "temp_budget_expiry", + "allowed_vector_store_indexes", ] LiteLLM_ManagementEndpoint_MetadataFields_Premium = [ diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index eeae54055f8..472cb81f93d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -551,6 +551,7 @@ async def _common_key_generation_helper( # noqa: PLR0915 field_name=field, value=getattr(data, field), ) + delattr(data, field) for field in LiteLLM_ManagementEndpoint_MetadataFields: if getattr(data, field, None) is not None: @@ -1019,6 +1020,7 @@ async def generate_key_fn( - prompts: Optional[List[str]] - List of allowed prompts for the key. If specified, the key will only be able to use these specific prompts. - auto_rotate: Optional[bool] - Whether this key should be automatically rotated (regenerated) - rotation_interval: Optional[str] - How often to auto-rotate this key (e.g., '30s', '30m', '30h', '30d'). Required if auto_rotate=True. + - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. Examples: @@ -1167,6 +1169,8 @@ async def generate_service_account_key_fn( - allowed_routes: Optional[list] - List of allowed routes for the key. Store the actual route or store a wildcard pattern for a set of routes. Example - ["/chat/completions", "/embeddings", "/keys/*"] - object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission. Examples: + - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. + 1. Allow users to turn on/off pii masking @@ -1457,6 +1461,8 @@ async def update_key_fn( - object_permission: Optional[LiteLLM_ObjectPermissionBase] - key-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "mcp_tool_permissions": {"server_id_1": ["tool1", "tool2"]}}. IF null or {} then no object permission. - auto_rotate: Optional[bool] - Whether this key should be automatically rotated - rotation_interval: Optional[str] - How often to rotate this key (e.g., '30d', '90d'). Required if auto_rotate=True + - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. + Example: ```bash curl --location 'http://0.0.0.0:4000/key/update' \ diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 582f5f8654b..d4bdde02bbd 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -556,6 +556,8 @@ async def new_team( # noqa: PLR0915 - team_member_key_duration: Optional[str] - The duration for a team member's key. e.g. "1d", "1w", "1mo" - prompts: Optional[List[str]] - List of allowed prompts for the team. If specified, the team will only be able to use these specific prompts. - allowed_passthrough_routes: Optional[List[str]] - List of allowed pass through routes for the team. + - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. + Returns: @@ -1114,6 +1116,8 @@ async def update_team( - model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit per model for this team. Example: {"gpt-4": 100, "gpt-3.5-turbo": 200} - model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit per model for this team. Example: {"gpt-4": 10000, "gpt-3.5-turbo": 20000} Example - update team TPM Limit + - allowed_vector_store_indexes: Optional[List[dict]] - List of allowed vector store indexes for the key. Example - [{"index_name": "my-index", "index_permissions": ["write", "read"]}]. If specified, the key will only be able to use these specific vector store indexes. Create index, using `/v1/indexes` endpoint. + ``` curl --location 'http://0.0.0.0:4000/team/update' \ diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 749b421c427..5d1c2febad8 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -35,7 +35,11 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( websocket_passthrough_request, ) from litellm.proxy.utils import is_known_model +from litellm.proxy.vector_store_endpoints.utils import ( + is_allowed_to_call_vector_store_endpoint, +) from litellm.secret_managers.main import get_secret_str +from litellm.utils import ProviderConfigManager from .passthrough_endpoint_router import PassthroughEndpointRouter @@ -988,6 +992,11 @@ async def assemblyai_proxy_route( return received_value +@router.api_route( + "/azure_ai/{endpoint:path}", + methods=["GET", "POST", "PUT", "DELETE", "PATCH"], + tags=["Azure AI Pass-through", "pass-through"], +) @router.api_route( "/azure/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], @@ -1014,9 +1023,21 @@ async def azure_proxy_route( if len(parts) > 1 and llm_router: for part in parts: + # check if LLM MODEL is_router_model = is_passthrough_request_using_router_model( request_body={"model": part}, llm_router=llm_router ) + # check if vector store index + is_vector_store_index = ( + ( + litellm.vector_store_index_registry.is_vector_store_index( + vector_store_index_name=part + ) + ) + if litellm.vector_store_index_registry is not None + else False + ) + if is_router_model: request_body = await get_request_body(request) is_streaming_request = is_passthrough_request_streaming(request_body) @@ -1074,6 +1095,66 @@ async def azure_proxy_route( custom_headers=None, ), ) + elif is_vector_store_index: + # get the api key from the provider config + provider_config = ( + ProviderConfigManager.get_provider_vector_stores_config( + provider=litellm.LlmProviders.AZURE_AI + ) + ) + if provider_config is None: + raise Exception("Provider config not found for Azure AI") + # get the index from registry + if litellm.vector_store_registry is None: + raise Exception("Vector store registry not found") + + is_allowed_to_call_vector_store_endpoint( + index_name=part, + provider=litellm.LlmProviders.AZURE_AI, + request=request, + user_api_key_dict=user_api_key_dict, + ) + # get the vector store name from index registry + index_object = ( + ( + litellm.vector_store_index_registry.get_vector_store_index_by_name( + vector_store_index_name=part + ) + ) + if litellm.vector_store_index_registry is not None + else None + ) + if index_object is None: + raise Exception(f"Vector store index not found for {part}") + + vector_store_name = index_object.litellm_params.vector_store_name + + vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry_by_name( + vector_store_name=vector_store_name + ) + if vector_store is None: + raise Exception(f"Vector store not found for {vector_store_name}") + litellm_params = vector_store.get("litellm_params") or {} + auth_credentials = provider_config.get_auth_credentials( + litellm_params=litellm_params + ) + + extra_headers = auth_credentials.get("headers") or {} + + base_target_url = litellm_params.get("api_base") + if base_target_url is None: + raise Exception(f"API base not found for {part}") + return await BaseOpenAIPassThroughHandler._base_openai_pass_through_handler( + endpoint=endpoint, + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + base_target_url=base_target_url, + api_key=None, + custom_llm_provider=litellm.LlmProviders.AZURE_AI, + extra_headers=cast(dict, extra_headers), + ) + base_target_url = get_secret_str(secret_name="AZURE_API_BASE") if base_target_url is None: raise Exception( @@ -1577,8 +1658,9 @@ class BaseOpenAIPassThroughHandler: fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth, base_target_url: str, - api_key: str, + api_key: Optional[str], custom_llm_provider: litellm.LlmProviders, + extra_headers: Optional[dict] = None, ): encoded_endpoint = httpx.URL(endpoint).path # Ensure endpoint starts with '/' for proper URL construction @@ -1603,7 +1685,7 @@ class BaseOpenAIPassThroughHandler: endpoint=endpoint, target=str(updated_url), custom_headers=BaseOpenAIPassThroughHandler._assemble_headers( - api_key=api_key, request=request + api_key=api_key, request=request, extra_headers=extra_headers ), ) # dynamically construct pass-through endpoint based on incoming path received_value = await endpoint_func( @@ -1629,11 +1711,17 @@ class BaseOpenAIPassThroughHandler: return headers @staticmethod - def _assemble_headers(api_key: str, request: Request) -> dict: - base_headers = { - "authorization": "Bearer {}".format(api_key), - "api-key": "{}".format(api_key), - } + def _assemble_headers( + api_key: Optional[str], request: Request, extra_headers: Optional[dict] = None + ) -> dict: + base_headers = {} + if api_key is not None: + base_headers = { + "authorization": "Bearer {}".format(api_key), + "api-key": "{}".format(api_key), + } + if extra_headers is not None: + base_headers.update(extra_headers) return BaseOpenAIPassThroughHandler._append_openai_beta_header( headers=base_headers, request=request, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 72e139dea52..39187edb523 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -771,14 +771,6 @@ async def pass_through_request( # noqa: PLR0915 status_code=response.status_code, ) - verbose_proxy_logger.debug("request method: {}".format(request.method)) - verbose_proxy_logger.debug("request url: {}".format(url)) - verbose_proxy_logger.debug("request headers: {}".format(headers)) - verbose_proxy_logger.debug( - "requested_query_params={}".format(requested_query_params) - ) - verbose_proxy_logger.debug("request body: {}".format(_parsed_body)) - response = ( await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler( request=request, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 14cf848f70e..12e6164f424 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -269,9 +269,7 @@ from litellm.proxy.management_endpoints.customer_endpoints import ( from litellm.proxy.management_endpoints.internal_user_endpoints import ( router as internal_user_router, ) -from litellm.proxy.management_endpoints.internal_user_endpoints import ( - user_update, -) +from litellm.proxy.management_endpoints.internal_user_endpoints import user_update from litellm.proxy.management_endpoints.key_management_endpoints import ( delete_verification_tokens, duration_in_seconds, @@ -322,9 +320,7 @@ from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, ) -from litellm.proxy.openai_files_endpoints.files_endpoints import ( - set_files_config, -) +from litellm.proxy.openai_files_endpoints.files_endpoints import set_files_config from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( passthrough_endpoint_router, ) @@ -411,9 +407,7 @@ from litellm.types.proxy.management_endpoints.ui_sso import ( LiteLLM_UpperboundKeyGenerateParams, ) from litellm.types.realtime import RealtimeQueryParams -from litellm.types.router import ( - DeploymentTypedDict, -) +from litellm.types.router import DeploymentTypedDict from litellm.types.router import ModelInfo as RouterModelInfo from litellm.types.router import ( RouterGeneralSettings, @@ -1928,29 +1922,31 @@ class ProxyConfig: """ Parse and validate search tools from config. Loads environment variables and casts to SearchToolTypedDict. - + Args: config: Config dictionary containing search_tools - + Returns: List of validated SearchToolTypedDict or None if not configured """ search_tools_raw = config.get("search_tools", None) if not search_tools_raw: return None - + search_tools_parsed: List[SearchToolTypedDict] = [] - + print( # noqa "\033[32mLiteLLM: Proxy initialized with Search Tools:\033[0m" ) # noqa - + for search_tool in search_tools_raw: # Display loaded search tool search_tool_name = search_tool.get("search_tool_name", "") - search_provider = search_tool.get("litellm_params", {}).get("search_provider", "") + search_provider = search_tool.get("litellm_params", {}).get( + "search_provider", "" + ) print(f"\033[32m {search_tool_name} ({search_provider})\033[0m") # noqa - + # Cast to SearchToolTypedDict for type safety try: search_tool_typed: SearchToolTypedDict = SearchToolTypedDict(**search_tool) # type: ignore @@ -1960,7 +1956,7 @@ class ProxyConfig: f"Error parsing search tool {search_tool_name}: {str(e)}" ) continue - + return search_tools_parsed if search_tools_parsed else None def _load_environment_variables(self, config: dict): @@ -2466,7 +2462,9 @@ class ProxyConfig: assistants_config = AssistantsTypedDict(**assistant_settings) # type: ignore ## SEARCH TOOLS SETTINGS - search_tools: Optional[List[SearchToolTypedDict]] = self.parse_search_tools(config) + search_tools: Optional[List[SearchToolTypedDict]] = self.parse_search_tools( + config + ) ## /fine_tuning/jobs endpoints config finetuning_config = config.get("finetune_settings", None) @@ -3366,6 +3364,10 @@ class ProxyConfig: if self._should_load_db_object(object_type="vector_stores"): await self._init_vector_stores_in_db(prisma_client=prisma_client) + if self._should_load_db_object(object_type="vector_store_indexes"): + + await self._init_vector_store_indexes_in_db(prisma_client=prisma_client) + if self._should_load_db_object(object_type="mcp"): await self._init_mcp_servers_in_db() @@ -3394,21 +3396,26 @@ class ProxyConfig: """ Initialize SSO settings from database into the router on startup. """ - + try: sso_settings = await prisma_client.db.litellm_ssoconfig.find_unique( where={"id": "sso_config"} ) if sso_settings is not None: # Capitalize all keys in sso_settings dictionary - uppercase_sso_settings = {key.upper(): value for key, value in sso_settings.sso_settings.items()} - self._decrypt_and_set_db_env_variables(environment_variables=uppercase_sso_settings) + uppercase_sso_settings = { + key.upper(): value + for key, value in sso_settings.sso_settings.items() + } + self._decrypt_and_set_db_env_variables( + environment_variables=uppercase_sso_settings + ) except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.py::ProxyConfig:_init_sso_settings_in_db - {}".format( str(e) ) - ) + ) async def _check_and_reload_model_cost_map(self, prisma_client: PrismaClient): """ @@ -3579,6 +3586,36 @@ class ProxyConfig: ) ) + async def _init_vector_store_indexes_in_db(self, prisma_client: PrismaClient): + from litellm.vector_stores.vector_store_registry import VectorStoreIndexRegistry + + try: + # read vector stores from db table + vector_store_indexes = ( + await VectorStoreIndexRegistry._get_vector_store_indexes_from_db( + prisma_client=prisma_client + ) + ) + + if len(vector_store_indexes) <= 0: + return + + if litellm.vector_store_index_registry is None: + litellm.vector_store_index_registry = VectorStoreIndexRegistry( + vector_store_indexes=vector_store_indexes + ) + else: + for vector_store_index in vector_store_indexes: + litellm.vector_store_index_registry.upsert_vector_store_index( + vector_store_index=vector_store_index + ) + except Exception as e: + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.py::ProxyConfig:_init_vector_stores_in_db - {}".format( + str(e) + ) + ) + async def _init_mcp_servers_in_db(self): from litellm.proxy._experimental.mcp_server.utils import is_mcp_available @@ -3600,7 +3637,7 @@ class ProxyConfig: str(e) ) ) - + async def _init_search_tools_in_db(self, prisma_client: PrismaClient): """ Initialize search tools from database into the router on startup. @@ -3611,19 +3648,20 @@ class ProxyConfig: SearchToolRegistry, ) from litellm.router_utils.search_api_router import SearchAPIRouter - + try: - search_tools = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client) - + search_tools = await SearchToolRegistry.get_all_search_tools_from_db( + prisma_client=prisma_client + ) + verbose_proxy_logger.info( f"Loading {len(search_tools)} search tool(s) from database into router" ) - + if llm_router is not None: # Add search tools to the router await SearchAPIRouter.update_router_search_tools( - router_instance=llm_router, - search_tools=search_tools + router_instance=llm_router, search_tools=search_tools ) verbose_proxy_logger.info( f"Successfully loaded {len(search_tools)} search tool(s) into router" @@ -3632,14 +3670,14 @@ class ProxyConfig: verbose_proxy_logger.debug( "Router not initialized yet, search tools will be added when router is created" ) - + except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.py::ProxyConfig:_init_search_tools_in_db - {}".format( str(e) ) ) - + async def _init_pass_through_endpoints_in_db(self): from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( initialize_pass_through_endpoints_in_db, @@ -4138,21 +4176,24 @@ class ProxyStartupEvent: # It must be passed individually when calling add_job() }, # Limit job store size to prevent memory growth - jobstores={ - 'default': MemoryJobStore() # explicitly use memory job store - }, + jobstores={"default": MemoryJobStore()}, # explicitly use memory job store # Use simple executor to minimize overhead executors={ - 'default': AsyncIOExecutor(), + "default": AsyncIOExecutor(), }, # Disable timezone awareness to reduce computation - timezone=None + timezone=None, ) # Use fixed intervals with small random offset instead of jitter # This avoids the expensive jitter calculations in APScheduler - budget_interval = proxy_budget_rescheduler_min_time + random.randint(0, - min(30, proxy_budget_rescheduler_max_time - proxy_budget_rescheduler_min_time)) + budget_interval = proxy_budget_rescheduler_min_time + random.randint( + 0, + min( + 30, + proxy_budget_rescheduler_max_time - proxy_budget_rescheduler_min_time, + ), + ) # Ensure minimum interval of 30 seconds for batch writing to prevent memory issues batch_writing_interval = proxy_batch_write_at + random.randint(0, 5) @@ -4248,7 +4289,9 @@ class ProxyStartupEvent: # REMOVED jitter parameter - major cause of memory leak # Use random start time instead for distribution next_run_time=datetime.now() - + timedelta(seconds=10 + random.randint(0, 300)), # Random 0-5 min offset + + timedelta( + seconds=10 + random.randint(0, 300) + ), # Random 0-5 min offset args=[spend_report_frequency], id="weekly_spend_report_job", replace_existing=True, @@ -4292,7 +4335,8 @@ class ProxyStartupEvent: scheduler.add_job( spend_log_cleanup.cleanup_old_spend_logs, "interval", - seconds=interval_seconds + random.randint(0, 60), # Add small random offset + seconds=interval_seconds + + random.randint(0, 60), # Add small random offset # REMOVED jitter parameter - major cause of memory leak args=[prisma_client], id="spend_log_cleanup_job", @@ -4318,7 +4362,8 @@ class ProxyStartupEvent: scheduler.add_job( check_batch_cost_job.check_batch_cost, "interval", - seconds=proxy_batch_polling_interval + random.randint(0, 30), # Add small random offset + seconds=proxy_batch_polling_interval + + random.randint(0, 30), # Add small random offset # REMOVED jitter parameter - major cause of memory leak id="check_batch_cost_job", replace_existing=True, @@ -4327,9 +4372,7 @@ class ProxyStartupEvent: verbose_proxy_logger.info("Batch cost check job scheduled successfully") except Exception as e: - verbose_proxy_logger.error( - f"Failed to setup batch cost checking: {e}" - ) + verbose_proxy_logger.error(f"Failed to setup batch cost checking: {e}") verbose_proxy_logger.debug( "Checking batch cost for LiteLLM Managed Files is an Enterprise Feature. Skipping..." ) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index ac1acee908a..605a312db4d 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -590,6 +590,17 @@ model LiteLLM_SSOConfig { updated_at DateTime @updatedAt } +model LiteLLM_ManagedVectorStoreIndexTable { + id String @id @default(uuid()) + index_name String @unique + litellm_params Json + index_info Json? + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @updatedAt + updated_by String? +} + // Cache configuration table model LiteLLM_CacheConfig { id String @id @default("cache_config") diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index dbd795e1b31..7ad39bff6b1 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -985,7 +985,11 @@ class ProxyLogging: ) else: _callback = callback # type: ignore - if _callback is not None and isinstance(_callback, CustomGuardrail) and data is not None: + if ( + _callback is not None + and isinstance(_callback, CustomGuardrail) + and data is not None + ): result = await self._process_guardrail_callback( callback=_callback, data=data, @@ -3691,6 +3695,16 @@ def is_known_model(model: Optional[str], llm_router: Optional[Router]) -> bool: return is_in_list +def is_known_vector_store_index(index_name: str) -> bool: + """ + Returns True if the vector store index is in the llm_router vector store indexes + """ + + if litellm.vector_store_index_registry is None: + return False + return index_name in litellm.vector_store_index_registry.get_vector_store_indexes() + + def join_paths(base_path: str, route: str) -> str: # Remove trailing slashes from base_path and leading slashes from route base_path = base_path.rstrip("/") diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 7859c1c4b83..509869d9bb0 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -1,20 +1,23 @@ from typing import Dict, Optional -from fastapi import APIRouter, Depends, Request, Response +from fastapi import APIRouter, Depends, HTTPException, Request, Response import litellm from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( LiteLLM_ManagedVectorStore, ) -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.utils import jsonify_object +from litellm.types.vector_stores import IndexCreateRequest router = APIRouter() ######################################################## # OpenAI Compatible Endpoints ######################################################## + def _update_request_data_with_litellm_managed_vector_store_registry( data: Dict, vector_store_id: str, @@ -24,23 +27,35 @@ def _update_request_data_with_litellm_managed_vector_store_registry( """ if litellm.vector_store_registry is not None: - vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( - vector_store_id=vector_store_id + vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = ( + litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( + vector_store_id=vector_store_id + ) ) if vector_store_to_run is not None: if "custom_llm_provider" in vector_store_to_run: - data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider") - + data["custom_llm_provider"] = vector_store_to_run.get( + "custom_llm_provider" + ) + if "litellm_credential_name" in vector_store_to_run: - data["litellm_credential_name"] = vector_store_to_run.get("litellm_credential_name") + data["litellm_credential_name"] = vector_store_to_run.get( + "litellm_credential_name" + ) if "litellm_params" in vector_store_to_run: litellm_params = vector_store_to_run.get("litellm_params", {}) or {} data.update(litellm_params) return data -@router.post("/v1/vector_stores/{vector_store_id}/search", dependencies=[Depends(user_api_key_auth)]) -@router.post("/vector_stores/{vector_store_id}/search", dependencies=[Depends(user_api_key_auth)]) + +@router.post( + "/v1/vector_stores/{vector_store_id}/search", + dependencies=[Depends(user_api_key_auth)], +) +@router.post( + "/vector_stores/{vector_store_id}/search", dependencies=[Depends(user_api_key_auth)] +) async def vector_store_search( request: Request, vector_store_id: str, @@ -71,12 +86,11 @@ async def vector_store_search( data = await _read_request_body(request=request) if "vector_store_id" not in data: data["vector_store_id"] = vector_store_id - + data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, - vector_store_id=vector_store_id + data=data, vector_store_id=vector_store_id ) - + processor = ProxyBaseLLMRequestProcessing(data=data) try: return await processor.base_process_llm_request( @@ -106,7 +120,6 @@ async def vector_store_search( ) - @router.post("/v1/vector_stores", dependencies=[Depends(user_api_key_auth)]) @router.post("/vector_stores", dependencies=[Depends(user_api_key_auth)]) async def vector_store_create( @@ -163,3 +176,61 @@ async def vector_store_create( proxy_logging_obj=proxy_logging_obj, version=version, ) + + +@router.post( + "/v1/indexes", + dependencies=[Depends(user_api_key_auth)], +) +async def index_create( + request: Request, + index_create_request: IndexCreateRequest, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create an index. Just writes the index to the database. + + ```bash + curl -L -X POST 'http://0.0.0.0:4000/indexes/create' \ + -H 'Content-Type: application/json' \ + -H 'Authorization: Bearer sk-1234' \ + -H 'LiteLLM-Beta: indexes_beta=v1' \ + -d '{ + "index_name": "dall-e-3", + "vector_store_index": "real-index-name", + "vector_store_name": "azure-ai-search" + }' + ``` + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail=CommonProxyErrors.db_not_connected_error.value, + ) + ## 1. check if index already exists + existing_index = ( + await prisma_client.db.litellm_managedvectorstoreindextable.find_unique( + where={"index_name": index_create_request.index_name} + ) + ) + + ## 2. set created_by and updated_by + + if existing_index is not None: + raise HTTPException( + status_code=400, + detail=f"Index {index_create_request.index_name} already exists", + ) + + ## 2. create index + index_data = index_create_request.model_dump(exclude_none=True) + index_data["created_by"] = user_api_key_dict.user_id + index_data["updated_by"] = user_api_key_dict.user_id + new_index = await prisma_client.db.litellm_managedvectorstoreindextable.create( + data=jsonify_object(index_data) + ) + + return new_index.model_dump() diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py new file mode 100644 index 00000000000..1d88ad0edf4 --- /dev/null +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -0,0 +1,120 @@ +from typing import Any, Dict, Literal, Optional + +from fastapi import HTTPException, Request + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + + +def check_vector_store_permission( + index_name: str, + permission: str, + key_metadata: Optional[Dict[str, Any]], + team_metadata: Optional[Dict[str, Any]], +) -> bool: + """ + Check if a specific permission is allowed for a given vector store index. + + Args: + index_name: The name of the vector store index + permission: The permission to check (e.g., "read", "write") + key_metadata: Metadata from the API key + team_metadata: Metadata from the team + + Returns: + True if the permission is allowed, False otherwise + + Example metadata format: + "metadata": { + "allowed_vector_store_indexes": [ + { + "index_name": "dall-e-3", + "index_permissions": ["write"] + } + ] + } + """ + # Check both key_metadata and team_metadata + for metadata in [key_metadata, team_metadata]: + if metadata is None: + continue + + allowed_indexes = metadata.get("allowed_vector_store_indexes") + if not allowed_indexes or not isinstance(allowed_indexes, list): + continue + + # Look for matching index + for index_config in allowed_indexes: + if not isinstance(index_config, dict): + continue + + if index_config.get("index_name") == index_name: + index_permissions = index_config.get("index_permissions", []) + if ( + isinstance(index_permissions, list) + and permission in index_permissions + ): + return True + + return False + + +def is_allowed_to_call_vector_store_endpoint( + provider: LlmProviders, + index_name: str, + request: Request, + user_api_key_dict: UserAPIKeyAuth, +) -> Optional[Literal[True]]: + """ + Check if the user is allowed to call the vector store endpoint. + + Cover: + 1. Creating a vector store index + 2. Reading a vector store index (Search / List / Get) + """ + # check what allowed permissions are for the key + key_metadata = user_api_key_dict.metadata + team_metadata = user_api_key_dict.team_metadata + + provider_config = ProviderConfigManager.get_provider_vector_stores_config( + provider=provider + ) + if provider_config is None: + return None + + provider_vector_store_endpoints = ( + provider_config.get_vector_store_endpoints_by_type() + ) + + # Determine the permission type based on the request + permission_type = None + for endpoint in provider_vector_store_endpoints["read"]: + if request.method == endpoint[0] and endpoint[1] in request.url.path: + permission_type = "read" + break + + if permission_type is None: + for endpoint in provider_vector_store_endpoints["write"]: + if request.method == endpoint[0] and endpoint[1] in request.url.path: + permission_type = "write" + break + + if permission_type is None: + return None + + # Check if key has specific permission for allowed_vector_store_indexes + has_permission = check_vector_store_permission( + index_name=index_name, + permission=permission_type, + key_metadata=key_metadata, + team_metadata=team_metadata, + ) + + if not has_permission: + raise HTTPException( + status_code=403, + detail=f"User does not have permission to call vector store endpoint {index_name}. Ask your administrator to add the necessary permissions to your API key/Team.", + ) + + return has_permission diff --git a/litellm/types/vector_stores.py b/litellm/types/vector_stores.py index 144b1853382..5456eb90e30 100644 --- a/litellm/types/vector_stores.py +++ b/litellm/types/vector_stores.py @@ -1,6 +1,6 @@ from datetime import datetime from enum import Enum -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Dict, List, Literal, Optional, Tuple, Union from annotated_types import Ge from pydantic import BaseModel @@ -195,6 +195,51 @@ class VectorStoreCreateResponse(TypedDict, total=False): metadata: Optional[Dict[str, str]] # Metadata associated with the vector store +class IndexCreateLiteLLMParams(BaseModel): + vector_store_index: str + vector_store_name: str + + +class IndexCreateRequest(BaseModel): + index_name: str + litellm_params: IndexCreateLiteLLMParams + index_info: Optional[Dict[str, Any]] = None + + +class BaseVectorStoreAuthCredentials(TypedDict, total=False): + headers: dict + query_params: dict + + +class LiteLLM_ManagedVectorStoreIndex(BaseModel): + """LiteLLM managed vector store index object - this is is the object stored in the database""" + + id: str + index_name: str + litellm_params: IndexCreateLiteLLMParams + index_info: Optional[Dict[str, Any]] = None + created_at: Optional[datetime] = None + created_by: Optional[str] = None + updated_at: Optional[datetime] = None + updated_by: Optional[str] = None + + +class VectorStoreIndexType(str, Enum): + """Type of vector store index""" + + READ = "read" + WRITE = "write" + + +class VectorStoreIndexEndpoints(TypedDict): + """Endpoints for vector store index""" + + read: List[ + Tuple[Literal["GET", "POST", "PUT", "DELETE", "PATCH"], str] + ] # endpoints for reading a vector store index + write: List[ + Tuple[Literal["GET", "POST", "PUT", "DELETE", "PATCH"], str] + ] # endpoints for writing a vector store index VECTOR_STORE_OPENAI_PARAMS = Literal[ "filters", "max_num_results", diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index bdd14cb2d99..f7a2ddaec8d 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -7,6 +7,7 @@ from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import remove_items_at_indices from litellm.types.vector_stores import ( LiteLLM_ManagedVectorStore, + LiteLLM_ManagedVectorStoreIndex, LiteLLM_ManagedVectorStoreListResponse, LiteLLM_VectorStoreConfig, ) @@ -17,6 +18,91 @@ else: PrismaClient = Any +class VectorStoreIndexRegistry: + def __init__( + self, vector_store_indexes: List[LiteLLM_ManagedVectorStoreIndex] = [] + ): + self.vector_store_indexes: List[LiteLLM_ManagedVectorStoreIndex] = ( + vector_store_indexes + ) + + def get_vector_store_indexes(self) -> List[LiteLLM_ManagedVectorStoreIndex]: + """ + Returns the vector store indexes + """ + return self.vector_store_indexes + + def get_vector_store_index_by_name( + self, vector_store_index_name: str + ) -> Optional[LiteLLM_ManagedVectorStoreIndex]: + """ + Returns the vector store index by name + """ + for vector_store_index in self.vector_store_indexes: + if vector_store_index.index_name == vector_store_index_name: + return vector_store_index + return None + + def upsert_vector_store_index( + self, vector_store_index: LiteLLM_ManagedVectorStoreIndex + ): + """ + Adds a vector store index to the registry. + + If it already exists, it will be updated. + """ + for i, _vector_store_index in enumerate[LiteLLM_ManagedVectorStoreIndex]( + self.vector_store_indexes + ): + if _vector_store_index.index_name == vector_store_index.index_name: + self.vector_store_indexes[i] = vector_store_index + return + self.vector_store_indexes.append(vector_store_index) + + def delete_vector_store_index(self, vector_store_index: str): + """ + Deletes a vector store index from the registry + """ + self.vector_store_indexes = [ + index for index in self.vector_store_indexes if index != vector_store_index + ] + + def is_vector_store_index(self, vector_store_index_name: str) -> bool: + """ + Returns True if the vector store index is in the registry + """ + for vector_store_index in self.vector_store_indexes: + if vector_store_index.index_name == vector_store_index_name: + return True + return False + + ######################################################### + ########### DB management helpers for vector stores ########### + ######################################################### + + @staticmethod + async def _get_vector_store_indexes_from_db( + prisma_client: Optional[PrismaClient], + ) -> List[LiteLLM_ManagedVectorStoreIndex]: + """ + Get vector stores from the database + """ + vector_stores_from_db: List[LiteLLM_ManagedVectorStoreIndex] = [] + if prisma_client is not None: + _vector_stores_from_db = ( + await prisma_client.db.litellm_managedvectorstoreindextable.find_many( + order={"created_at": "desc"}, + ) + ) + for vector_store in _vector_stores_from_db: + _dict_vector_store = dict(vector_store) + _litellm_managed_vector_store = LiteLLM_ManagedVectorStoreIndex( + **_dict_vector_store + ) + vector_stores_from_db.append(_litellm_managed_vector_store) + return vector_stores_from_db + + class VectorStoreRegistry: def __init__(self, vector_stores: List[LiteLLM_ManagedVectorStore] = []): self.vector_stores: List[LiteLLM_ManagedVectorStore] = vector_stores @@ -139,6 +225,17 @@ class VectorStoreRegistry: return vector_store return None + def get_litellm_managed_vector_store_from_registry_by_name( + self, vector_store_name: str + ) -> Optional[LiteLLM_ManagedVectorStore]: + """ + Returns the vector store from the registry by name + """ + for vector_store in self.vector_stores: + if vector_store.get("vector_store_name") == vector_store_name: + return vector_store + return None + def pop_vector_stores_to_run( self, non_default_params: Dict, tools: Optional[List[Dict]] = None ) -> List[LiteLLM_ManagedVectorStore]: diff --git a/schema.prisma b/schema.prisma index ac1acee908a..a142b388b24 100644 --- a/schema.prisma +++ b/schema.prisma @@ -590,6 +590,17 @@ model LiteLLM_SSOConfig { updated_at DateTime @updatedAt } +model LiteLLM_IndexTable { + id String @id @default(uuid()) + index_name String @unique + litellm_params Json + index_info Json? + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @updatedAt + updated_by String? +} + // Cache configuration table model LiteLLM_CacheConfig { id String @id @default("cache_config") diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 384c37ec1e6..1a287e95746 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -3,48 +3,57 @@ import sys from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import Request sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path +from fastapi import HTTPException + import litellm from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( LiteLLM_ManagedVectorStore, ) +from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.vector_store_endpoints.endpoints import ( _update_request_data_with_litellm_managed_vector_store_registry, ) +from litellm.proxy.vector_store_endpoints.utils import ( + check_vector_store_permission, + is_allowed_to_call_vector_store_endpoint, +) +from litellm.types.utils import LlmProviders @pytest.mark.asyncio async def test_router_avector_store_search_passes_correct_args(): """ - Test that router.avector_store_search() passes the correct arguments + Test that router.avector_store_search() passes the correct arguments to downstream litellm.vector_stores.asearch() with custom_llm_provider and query. """ # Create a router router = litellm.Router(model_list=[]) - + # Mock the router's _init_vector_store_api_endpoints method to avoid real API calls - with patch.object(router, '_init_vector_store_api_endpoints') as mock_init: + with patch.object(router, "_init_vector_store_api_endpoints") as mock_init: mock_init.return_value = { "object": "vector_store.search_results.page", "search_query": "test query", - "data": [] + "data": [], } - + # Call router's avector_store_search result = await router.avector_store_search( vector_store_id="test_store_id", query="test query", - custom_llm_provider="bedrock" + custom_llm_provider="bedrock", ) - + # Verify the internal method was called with correct args mock_init.assert_called_once() call_args = mock_init.call_args - + # Check that the original function is passed correctly assert call_args[1]["vector_store_id"] == "test_store_id" assert call_args[1]["query"] == "test query" @@ -59,44 +68,553 @@ def test_update_request_data_with_litellm_managed_vector_store_registry(): # Setup test data data = {"existing_key": "existing_value"} vector_store_id = "test_store_id" - + # Mock vector store registry mock_vector_store: LiteLLM_ManagedVectorStore = { - "vector_store_id": "test_store_id", + "vector_store_id": "test_store_id", "custom_llm_provider": "bedrock", "litellm_credential_name": "test_credential", - "litellm_params": {"api_key": "test_key", "aws_region_name": "us-east-1"} + "litellm_params": {"api_key": "test_key", "aws_region_name": "us-east-1"}, } - + mock_registry = MagicMock() - mock_registry.get_litellm_managed_vector_store_from_registry.return_value = mock_vector_store - + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = ( + mock_vector_store + ) + # Test with vector store registry - with patch.object(litellm, 'vector_store_registry', mock_registry): + with patch.object(litellm, "vector_store_registry", mock_registry): result = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, - vector_store_id=vector_store_id + data=data, vector_store_id=vector_store_id ) - + # Verify the data was updated correctly assert result["existing_key"] == "existing_value" # Original data preserved assert result["custom_llm_provider"] == "bedrock" assert result["litellm_credential_name"] == "test_credential" assert result["api_key"] == "test_key" assert result["aws_region_name"] == "us-east-1" - + # Verify registry was called correctly mock_registry.get_litellm_managed_vector_store_from_registry.assert_called_once_with( vector_store_id="test_store_id" ) - + # Test with no vector store registry - with patch.object(litellm, 'vector_store_registry', None): + with patch.object(litellm, "vector_store_registry", None): original_data = {"existing_key": "existing_value"} result = _update_request_data_with_litellm_managed_vector_store_registry( - data=original_data, - vector_store_id=vector_store_id + data=original_data, vector_store_id=vector_store_id ) - + # Verify data remains unchanged when no registry - assert result == original_data \ No newline at end of file + assert result == original_data + + +class TestCheckVectorStorePermission: + """Test suite for check_vector_store_permission function.""" + + def test_permission_allowed_in_key_metadata(self): + """Test that permission is allowed when found in key metadata.""" + key_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["read", "write"]} + ] + } + + result = check_vector_store_permission( + index_name="my-index", + permission="read", + key_metadata=key_metadata, + team_metadata=None, + ) + + assert result is True + + def test_permission_allowed_in_team_metadata(self): + """Test that permission is allowed when found in team metadata.""" + team_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "team-index", "index_permissions": ["write"]} + ] + } + + result = check_vector_store_permission( + index_name="team-index", + permission="write", + key_metadata=None, + team_metadata=team_metadata, + ) + + assert result is True + + def test_permission_denied_wrong_permission(self): + """Test that permission is denied when index exists but wrong permission.""" + key_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["read"]} + ] + } + + result = check_vector_store_permission( + index_name="my-index", + permission="write", + key_metadata=key_metadata, + team_metadata=None, + ) + + assert result is False + + def test_permission_denied_index_not_found(self): + """Test that permission is denied when index doesn't exist.""" + key_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "other-index", "index_permissions": ["read", "write"]} + ] + } + + result = check_vector_store_permission( + index_name="my-index", + permission="read", + key_metadata=key_metadata, + team_metadata=None, + ) + + assert result is False + + def test_permission_denied_no_metadata(self): + """Test that permission is denied when no metadata provided.""" + result = check_vector_store_permission( + index_name="my-index", + permission="read", + key_metadata=None, + team_metadata=None, + ) + + assert result is False + + def test_permission_denied_no_allowed_indexes_field(self): + """Test that permission is denied when metadata has no allowed_vector_store_indexes.""" + key_metadata = {"some_other_field": "value"} + + result = check_vector_store_permission( + index_name="my-index", + permission="read", + key_metadata=key_metadata, + team_metadata=None, + ) + + assert result is False + + def test_key_metadata_takes_precedence(self): + """Test that key metadata is checked and returns permission successfully.""" + key_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["read"]} + ] + } + team_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["write"]} + ] + } + + # Should find permission in key_metadata (checked first) + result = check_vector_store_permission( + index_name="my-index", + permission="read", + key_metadata=key_metadata, + team_metadata=team_metadata, + ) + + assert result is True + + def test_team_metadata_as_fallback(self): + """Test that team metadata is checked when key metadata doesn't have permission.""" + key_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "other-index", "index_permissions": ["read"]} + ] + } + team_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["write"]} + ] + } + + # Should find permission in team_metadata + result = check_vector_store_permission( + index_name="my-index", + permission="write", + key_metadata=key_metadata, + team_metadata=team_metadata, + ) + + assert result is True + + def test_multiple_indexes_in_metadata(self): + """Test handling multiple indexes in metadata.""" + key_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "index-1", "index_permissions": ["read"]}, + {"index_name": "index-2", "index_permissions": ["write"]}, + {"index_name": "index-3", "index_permissions": ["read", "write"]}, + ] + } + + # Test each index + assert ( + check_vector_store_permission("index-1", "read", key_metadata, None) is True + ) + assert ( + check_vector_store_permission("index-1", "write", key_metadata, None) + is False + ) + assert ( + check_vector_store_permission("index-2", "write", key_metadata, None) + is True + ) + assert ( + check_vector_store_permission("index-2", "read", key_metadata, None) + is False + ) + assert ( + check_vector_store_permission("index-3", "read", key_metadata, None) is True + ) + assert ( + check_vector_store_permission("index-3", "write", key_metadata, None) + is True + ) + + def test_invalid_metadata_structure(self): + """Test handling of invalid metadata structures.""" + # Test when allowed_vector_store_indexes is not a list + key_metadata = {"allowed_vector_store_indexes": "not-a-list"} + + result = check_vector_store_permission( + index_name="my-index", + permission="read", + key_metadata=key_metadata, + team_metadata=None, + ) + + assert result is False + + # Test when index config is not a dict + key_metadata = { + "allowed_vector_store_indexes": [ + "not-a-dict", + {"index_name": "my-index", "index_permissions": ["read"]}, + ] + } + + result = check_vector_store_permission( + index_name="my-index", + permission="read", + key_metadata=key_metadata, + team_metadata=None, + ) + + # Should still work because it skips invalid entries + assert result is True + + def test_missing_index_permissions_field(self): + """Test when index_permissions field is missing.""" + key_metadata = { + "allowed_vector_store_indexes": [ + { + "index_name": "my-index" + # Missing index_permissions field + } + ] + } + + result = check_vector_store_permission( + index_name="my-index", + permission="read", + key_metadata=key_metadata, + team_metadata=None, + ) + + assert result is False + + +class TestIsAllowedToCallVectorStoreEndpoint: + """Test suite for is_allowed_to_call_vector_store_endpoint function.""" + + def test_read_permission_allowed(self): + """Test read permission is checked correctly.""" + # Mock request + mock_request = MagicMock(spec=Request) + mock_request.method = "GET" + mock_request.url.path = "/v1/vector_stores/my-index/search" + + # Mock user API key with permissions + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["read"]} + ] + } + mock_user_api_key.team_metadata = None + + # Mock provider config + mock_provider_config = MagicMock() + mock_provider_config.get_vector_store_endpoints_by_type.return_value = { + "read": [("GET", "/search")], + "write": [("POST", "/create")], + } + + with patch( + "litellm.proxy.vector_store_endpoints.utils.ProviderConfigManager.get_provider_vector_stores_config", + return_value=mock_provider_config, + ): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.OPENAI, + index_name="my-index", + request=mock_request, + user_api_key_dict=mock_user_api_key, + ) + + assert result is True + + def test_write_permission_allowed(self): + """Test write permission is checked correctly.""" + # Mock request + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url.path = "/v1/vector_stores/my-index/create" + + # Mock user API key with permissions + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["write"]} + ] + } + mock_user_api_key.team_metadata = None + + # Mock provider config + mock_provider_config = MagicMock() + mock_provider_config.get_vector_store_endpoints_by_type.return_value = { + "read": [("GET", "/search")], + "write": [("POST", "/create")], + } + + with patch( + "litellm.proxy.vector_store_endpoints.utils.ProviderConfigManager.get_provider_vector_stores_config", + return_value=mock_provider_config, + ): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.OPENAI, + index_name="my-index", + request=mock_request, + user_api_key_dict=mock_user_api_key, + ) + + assert result is True + + def test_permission_denied_wrong_permission(self): + """Test permission denied when user has read but tries write.""" + # Mock request for write operation + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url.path = "/v1/vector_stores/my-index/create" + + # Mock user API key with only read permissions + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["read"]} + ] + } + mock_user_api_key.team_metadata = None + + # Mock provider config + mock_provider_config = MagicMock() + mock_provider_config.get_vector_store_endpoints_by_type.return_value = { + "read": [("GET", "/search")], + "write": [("POST", "/create")], + } + + with patch( + "litellm.proxy.vector_store_endpoints.utils.ProviderConfigManager.get_provider_vector_stores_config", + return_value=mock_provider_config, + ): + with pytest.raises(HTTPException) as exc_info: + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.OPENAI, + index_name="my-index", + request=mock_request, + user_api_key_dict=mock_user_api_key, + ) + + assert exc_info.value.status_code == 403 + + def test_provider_config_not_found(self): + """Test when provider config is not found.""" + mock_request = MagicMock(spec=Request) + mock_request.method = "GET" + mock_request.url.path = "/v1/vector_stores/my-index/search" + + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.metadata = {} + mock_user_api_key.team_metadata = None + + with patch( + "litellm.proxy.vector_store_endpoints.utils.ProviderConfigManager.get_provider_vector_stores_config", + return_value=None, + ): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.OPENAI, + index_name="my-index", + request=mock_request, + user_api_key_dict=mock_user_api_key, + ) + + assert result is None + + def test_endpoint_not_recognized(self): + """Test when endpoint doesn't match any read or write patterns.""" + # Mock request with unrecognized path + mock_request = MagicMock(spec=Request) + mock_request.method = "DELETE" + mock_request.url.path = "/v1/vector_stores/my-index/unknown" + + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "my-index", "index_permissions": ["read", "write"]} + ] + } + mock_user_api_key.team_metadata = None + + # Mock provider config + mock_provider_config = MagicMock() + mock_provider_config.get_vector_store_endpoints_by_type.return_value = { + "read": [("GET", "/search")], + "write": [("POST", "/create")], + } + + with patch( + "litellm.proxy.vector_store_endpoints.utils.ProviderConfigManager.get_provider_vector_stores_config", + return_value=mock_provider_config, + ): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.OPENAI, + index_name="my-index", + request=mock_request, + user_api_key_dict=mock_user_api_key, + ) + + assert result is None + + def test_team_metadata_permissions(self): + """Test that team metadata permissions work.""" + # Mock request + mock_request = MagicMock(spec=Request) + mock_request.method = "GET" + mock_request.url.path = "/v1/vector_stores/team-index/search" + + # Mock user API key with no key metadata but team metadata + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.metadata = None + mock_user_api_key.team_metadata = { + "allowed_vector_store_indexes": [ + {"index_name": "team-index", "index_permissions": ["read"]} + ] + } + + # Mock provider config + mock_provider_config = MagicMock() + mock_provider_config.get_vector_store_endpoints_by_type.return_value = { + "read": [("GET", "/search")], + "write": [("POST", "/create")], + } + + with patch( + "litellm.proxy.vector_store_endpoints.utils.ProviderConfigManager.get_provider_vector_stores_config", + return_value=mock_provider_config, + ): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.OPENAI, + index_name="team-index", + request=mock_request, + user_api_key_dict=mock_user_api_key, + ) + + assert result is True + + def test_no_permissions_configured(self): + """Test when user has no vector store permissions configured.""" + + mock_request = MagicMock(spec=Request) + mock_request.method = "GET" + mock_request.url.path = "/v1/vector_stores/my-index/search" + + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.metadata = {} + mock_user_api_key.team_metadata = {} + + # Mock provider config + mock_provider_config = MagicMock() + mock_provider_config.get_vector_store_endpoints_by_type.return_value = { + "read": [("GET", "/search")], + "write": [("POST", "/create")], + } + + with patch( + "litellm.proxy.vector_store_endpoints.utils.ProviderConfigManager.get_provider_vector_stores_config", + return_value=mock_provider_config, + ): + with pytest.raises(HTTPException) as exc_info: + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.OPENAI, + index_name="my-index", + request=mock_request, + user_api_key_dict=mock_user_api_key, + ) + + assert exc_info.value.status_code == 403 + + def test_azure_ai_permission_allowed(self): + from datetime import datetime, timezone + from unittest.mock import MagicMock + + from fastapi import HTTPException + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.utils import LlmProviders + + mock_request = MagicMock() + mock_request.method = "GET" + mock_request.url.path = "/azure_ai/indexes/dall-e-4/docs/search" + mock_user_api_key = UserAPIKeyAuth( + token="b637312ebffb9745321224644430ba9e4916a291c8281f293d21182c5e80bc5a", + key_name="sk-...plNQ", + metadata={ + "allowed_vector_store_indexes": [ + {"index_name": "dall-e-4", "index_permissions": ["write"]} + ] + }, + spend=0.015, + ) + mock_provider_config = MagicMock() + mock_provider_config.get_vector_store_endpoints_by_type.return_value = { + "read": [("GET", "/docs/search")], + "write": [("POST", "/create")], + } + + with patch( + "litellm.proxy.vector_store_endpoints.utils.ProviderConfigManager.get_provider_vector_stores_config", + return_value=mock_provider_config, + ): + with pytest.raises(HTTPException) as exc_info: + is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name="dall-e-4", + request=mock_request, + user_api_key_dict=mock_user_api_key, + ) + + assert exc_info.value.status_code == 403