(feat) Azure AI Vector Stores - support "virtual" indexes + create vector store on passthrough API (#16160)

* feat(vector_store_endpoints/endpoints.py): add new index_create endpoint

allows admin to create a virtual index, to do permission management for

* feat(key_management_endpoints.py): enable setting allowed_vector_store_indexes on keys

proxy admin can enable dev to create an index on a vector stor

* feat: initial commit adding vector store index passthrough logic to litellm

* feat: add vector store table

* fix(azure_ai/transformation.py): fix headers

* feat: track read/write endpoints by vector store integration

enables permissions by index to work

* fix: azure_ai/vector_stores/search

document the vector store endpoints correctly

 ensures permission management works as expected

* fix(proxy/utils.py): improve error message

* docs(azure_ai_vector_stores_passthrough.md): document azure ai passthrough vector store support

* docs(create.md): document azure ai support via passthrough for vector store create

* fix: fix code qa errors

* fix: document new allowed_vector_store_indexes endpoint
This commit is contained in:
Krish Dholakia 2025-11-01 12:01:32 -07:00 • committed by GitHub
parent b02be1ba70
commit 43aacf2dc0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
26 changed files with 1823 additions and 280 deletions

View file

@ -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)
```

View file

@ -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

View file

@ -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

View file

@ -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",
]
},
{

View file

@ -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

View file

@ -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:

View file

@ -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,

View file

@ -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]:

View file

@ -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,
)

View file

@ -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
)
headers=response.headers,
)

View file

@ -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(

View file

@ -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

View file

@ -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 = [

View file

@ -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' \

View file

@ -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' \

View file

@ -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,

View file

@ -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,

View file

@ -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..."
)

View file

@ -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")

View file

@ -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("/")

View file

@ -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()

View file

@ -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

View file

@ -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",

View file

@ -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]:

View file

@ -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")

View file

@ -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
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