mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
(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:
parent
b02be1ba70
commit
43aacf2dc0
26 changed files with 1823 additions and 280 deletions
|
|
@ -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)
|
||||
|
||||
```
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
},
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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' \
|
||||
|
|
|
|||
|
|
@ -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' \
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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..."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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("/")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
120
litellm/proxy/vector_store_endpoints/utils.py
Normal file
120
litellm/proxy/vector_store_endpoints/utils.py
Normal 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
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue