mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #19329 from BerriAI/litellm_vector_store_sync
Fix: vector store sync issues
This commit is contained in:
commit
daf70f7221
2 changed files with 307 additions and 28 deletions
|
|
@ -245,6 +245,7 @@ async def list_vector_stores(
|
|||
"""
|
||||
List all available vector stores with optional filtering and pagination.
|
||||
Combines both in-memory vector stores and those stored in the database.
|
||||
Database is the source of truth - deleted stores are removed from memory, updated stores sync to memory.
|
||||
|
||||
Parameters:
|
||||
- page: int - Page number for pagination (default: 1)
|
||||
|
|
@ -252,29 +253,65 @@ async def list_vector_stores(
|
|||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
seen_vector_store_ids = set()
|
||||
vector_store_map: Dict[str, LiteLLM_ManagedVectorStore] = {}
|
||||
db_vector_store_ids: set = set()
|
||||
|
||||
try:
|
||||
# Get in-memory vector stores
|
||||
in_memory_vector_stores: List[LiteLLM_ManagedVectorStore] = []
|
||||
# Get vector stores from database first (source of truth)
|
||||
vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
|
||||
# Build map from database vector stores
|
||||
for vector_store in vector_stores_from_db:
|
||||
vector_store_id = vector_store.get("vector_store_id", None)
|
||||
if vector_store_id:
|
||||
vector_store_map[vector_store_id] = vector_store
|
||||
db_vector_store_ids.add(vector_store_id)
|
||||
|
||||
# Process in-memory vector stores
|
||||
if litellm.vector_store_registry is not None:
|
||||
in_memory_vector_stores = copy.deepcopy(
|
||||
litellm.vector_store_registry.vector_stores
|
||||
)
|
||||
|
||||
vector_stores_to_delete_from_memory: List[str] = []
|
||||
|
||||
for vector_store in in_memory_vector_stores:
|
||||
vector_store_id = vector_store.get("vector_store_id", None)
|
||||
if not vector_store_id:
|
||||
continue
|
||||
|
||||
# If vector store is in memory but NOT in database, it was deleted
|
||||
if vector_store_id not in db_vector_store_ids:
|
||||
verbose_proxy_logger.info(
|
||||
f"Vector store {vector_store_id} exists in memory but not in database - marking for deletion from cache"
|
||||
)
|
||||
vector_stores_to_delete_from_memory.append(vector_store_id)
|
||||
# If not in our map yet, add it (only in-memory, not in DB)
|
||||
elif vector_store_id not in vector_store_map:
|
||||
vector_store_map[vector_store_id] = vector_store
|
||||
|
||||
# Synchronize in-memory registry with database
|
||||
# 1. Remove deleted vector stores from memory
|
||||
for vs_id in vector_stores_to_delete_from_memory:
|
||||
litellm.vector_store_registry.delete_vector_store_from_registry(
|
||||
vector_store_id=vs_id
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Removed deleted vector store {vs_id} from in-memory registry"
|
||||
)
|
||||
|
||||
# 2. Update in-memory registry with database versions (for updates)
|
||||
for vector_store in vector_stores_from_db:
|
||||
vector_store_id = vector_store.get("vector_store_id", None)
|
||||
if vector_store_id:
|
||||
litellm.vector_store_registry.update_vector_store_in_registry(
|
||||
vector_store_id=vector_store_id,
|
||||
updated_data=vector_store
|
||||
)
|
||||
|
||||
# Get vector stores from database
|
||||
vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
|
||||
# Combine in-memory and database vector stores
|
||||
combined_vector_stores: List[LiteLLM_ManagedVectorStore] = []
|
||||
for vector_store in in_memory_vector_stores + vector_stores_from_db:
|
||||
vector_store_id = vector_store.get("vector_store_id", None)
|
||||
if vector_store_id not in seen_vector_store_ids:
|
||||
combined_vector_stores.append(vector_store)
|
||||
seen_vector_store_ids.add(vector_store_id)
|
||||
|
||||
combined_vector_stores = list(vector_store_map.values())
|
||||
total_count = len(combined_vector_stores)
|
||||
total_pages = (total_count + page_size - 1) // page_size
|
||||
|
||||
|
|
@ -303,7 +340,7 @@ async def delete_vector_store(
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Delete a vector store.
|
||||
Delete a vector store from both database and in-memory registry.
|
||||
|
||||
Parameters:
|
||||
- vector_store_id: str - ID of the vector store to delete
|
||||
|
|
@ -314,31 +351,53 @@ async def delete_vector_store(
|
|||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Check if vector store exists
|
||||
# Check if vector store exists in database or in-memory registry
|
||||
db_vector_store_exists = False
|
||||
memory_vector_store_exists = False
|
||||
|
||||
existing_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
|
||||
where={"vector_store_id": data.vector_store_id}
|
||||
)
|
||||
)
|
||||
if existing_vector_store is None:
|
||||
if existing_vector_store is not None:
|
||||
db_vector_store_exists = True
|
||||
|
||||
# Check in-memory registry
|
||||
if litellm.vector_store_registry is not None:
|
||||
memory_vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
|
||||
vector_store_id=data.vector_store_id
|
||||
)
|
||||
if memory_vector_store is not None:
|
||||
memory_vector_store_exists = True
|
||||
|
||||
# If not found in either location, raise 404
|
||||
if not db_vector_store_exists and not memory_vector_store_exists:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Vector store with ID {data.vector_store_id} not found",
|
||||
)
|
||||
|
||||
# Delete vector store
|
||||
await prisma_client.db.litellm_managedvectorstorestable.delete(
|
||||
where={"vector_store_id": data.vector_store_id}
|
||||
)
|
||||
# Delete from database if exists
|
||||
if db_vector_store_exists:
|
||||
await prisma_client.db.litellm_managedvectorstorestable.delete(
|
||||
where={"vector_store_id": data.vector_store_id}
|
||||
)
|
||||
|
||||
# Delete vector store from registry
|
||||
if litellm.vector_store_registry is not None:
|
||||
# Delete from in-memory registry if exists
|
||||
if memory_vector_store_exists and litellm.vector_store_registry is not None:
|
||||
litellm.vector_store_registry.delete_vector_store_from_registry(
|
||||
vector_store_id=data.vector_store_id
|
||||
)
|
||||
|
||||
return {"message": f"Vector store {data.vector_store_id} deleted successfully"}
|
||||
return {
|
||||
"status": "success",
|
||||
"message": f"Vector store {data.vector_store_id} deleted successfully"
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error deleting vector store: {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
|
|
@ -415,8 +474,12 @@ async def update_vector_store(
|
|||
data: VectorStoreUpdateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Update vector store details"""
|
||||
"""
|
||||
Update vector store details in both database and in-memory registry.
|
||||
The updated data is immediately synchronized to the in-memory registry.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
|
@ -424,11 +487,36 @@ async def update_vector_store(
|
|||
try:
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
vector_store_id = update_data.pop("vector_store_id")
|
||||
|
||||
# Handle metadata serialization
|
||||
if update_data.get("vector_store_metadata") is not None:
|
||||
update_data["vector_store_metadata"] = safe_dumps(
|
||||
update_data["vector_store_metadata"]
|
||||
)
|
||||
|
||||
# Handle litellm_params if provided
|
||||
if "litellm_params" in update_data:
|
||||
_input_litellm_params: dict = update_data.get("litellm_params", {}) or {}
|
||||
|
||||
# Auto-resolve embedding config if embedding model is provided but config is not
|
||||
embedding_model = _input_litellm_params.get("litellm_embedding_model")
|
||||
if embedding_model and not _input_litellm_params.get("litellm_embedding_config"):
|
||||
resolved_config = await _resolve_embedding_config_from_db(
|
||||
embedding_model=embedding_model,
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
if resolved_config:
|
||||
_input_litellm_params["litellm_embedding_config"] = resolved_config
|
||||
verbose_proxy_logger.info(
|
||||
f"Auto-resolved embedding config for model {embedding_model}"
|
||||
)
|
||||
|
||||
litellm_params_dict = GenericLiteLLMParams(
|
||||
**_input_litellm_params
|
||||
).model_dump(exclude_none=True)
|
||||
update_data["litellm_params"] = safe_dumps(litellm_params_dict)
|
||||
|
||||
# Update in database
|
||||
updated = await prisma_client.db.litellm_managedvectorstorestable.update(
|
||||
where={"vector_store_id": vector_store_id},
|
||||
data=update_data,
|
||||
|
|
@ -436,13 +524,21 @@ async def update_vector_store(
|
|||
|
||||
updated_vs = LiteLLM_ManagedVectorStore(**updated.model_dump())
|
||||
|
||||
# Immediately update in-memory registry to keep it in sync
|
||||
if litellm.vector_store_registry is not None:
|
||||
litellm.vector_store_registry.update_vector_store_in_registry(
|
||||
vector_store_id=vector_store_id,
|
||||
updated_data=updated_vs,
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Updated vector store {vector_store_id} in both database and in-memory registry"
|
||||
)
|
||||
|
||||
return {"vector_store": updated_vs}
|
||||
return {
|
||||
"status": "success",
|
||||
"message": f"Vector store {vector_store_id} updated successfully",
|
||||
"vector_store": updated_vs
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error updating vector store: {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
|
|
|||
|
|
@ -1051,6 +1051,189 @@ async def test_vector_store_synchronization_across_instances():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_update_and_list_synchronization():
|
||||
"""
|
||||
Test that vector store updates are properly synchronized across multiple instances.
|
||||
|
||||
This test simulates the scenario where:
|
||||
1. Instance 1 creates a vector store
|
||||
2. Instance 2 caches it in memory
|
||||
3. Instance 1 updates the vector store in the database
|
||||
4. Instance 2 should see the updated data when listing (database is source of truth)
|
||||
|
||||
This is a regression test to prevent the bug where Instance 2 would show
|
||||
stale cached data instead of the updated database version.
|
||||
"""
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
|
||||
# Simulate two instances with separate in-memory registries
|
||||
instance_1_registry = VectorStoreRegistry(vector_stores=[])
|
||||
instance_2_registry = VectorStoreRegistry(vector_stores=[])
|
||||
|
||||
# Mock database that both instances share
|
||||
mock_db_vector_stores = []
|
||||
|
||||
async def mock_find_many(order=None):
|
||||
"""Mock find_many for listing vector stores"""
|
||||
result = []
|
||||
for vs in mock_db_vector_stores:
|
||||
class MockVectorStore:
|
||||
def __init__(self, data):
|
||||
for key, value in data.items():
|
||||
setattr(self, key, value)
|
||||
self._data = data
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._data.items())
|
||||
result.append(MockVectorStore(vs))
|
||||
return result
|
||||
|
||||
async def mock_create(data):
|
||||
"""Mock create for adding vector store to DB"""
|
||||
vector_store = data.copy()
|
||||
mock_db_vector_stores.append(vector_store)
|
||||
mock_obj = MagicMock()
|
||||
mock_obj.model_dump.return_value = vector_store
|
||||
return mock_obj
|
||||
|
||||
async def mock_update(where, data):
|
||||
"""Mock update for modifying vector store in DB"""
|
||||
vector_store_id = where.get("vector_store_id")
|
||||
for i, vs in enumerate(mock_db_vector_stores):
|
||||
if vs.get("vector_store_id") == vector_store_id:
|
||||
# Update the vector store
|
||||
mock_db_vector_stores[i].update(data)
|
||||
mock_obj = MagicMock()
|
||||
mock_obj.model_dump.return_value = mock_db_vector_stores[i]
|
||||
return mock_obj
|
||||
raise Exception(f"Vector store {vector_store_id} not found")
|
||||
|
||||
# Create mock prisma client
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(
|
||||
side_effect=mock_find_many
|
||||
)
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock(
|
||||
side_effect=mock_create
|
||||
)
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock(
|
||||
side_effect=mock_update
|
||||
)
|
||||
|
||||
# Test vector store data
|
||||
test_vector_store_id = "test-update-store-001"
|
||||
original_name = "Original Name"
|
||||
updated_name = "Updated Name"
|
||||
|
||||
test_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": test_vector_store_id,
|
||||
"custom_llm_provider": "bedrock",
|
||||
"vector_store_name": original_name,
|
||||
"vector_store_description": "Testing update synchronization",
|
||||
"litellm_params": {
|
||||
"vector_store_id": test_vector_store_id,
|
||||
"custom_llm_provider": "bedrock",
|
||||
"region_name": "us-east-1"
|
||||
},
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
}
|
||||
|
||||
# Step 1: Create vector store on Instance 1
|
||||
await mock_prisma_client.db.litellm_managedvectorstorestable.create(
|
||||
data=test_vector_store
|
||||
)
|
||||
instance_1_registry.add_vector_store_to_registry(vector_store=test_vector_store)
|
||||
|
||||
# Step 2: Instance 2 fetches and caches the vector store
|
||||
vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db(
|
||||
prisma_client=mock_prisma_client
|
||||
)
|
||||
for vs in vector_stores_from_db:
|
||||
if vs.get("vector_store_id") == test_vector_store_id:
|
||||
instance_2_registry.add_vector_store_to_registry(vector_store=vs)
|
||||
|
||||
# Verify both instances have the original data
|
||||
instance_1_vs = instance_1_registry.get_litellm_managed_vector_store_from_registry(
|
||||
test_vector_store_id
|
||||
)
|
||||
instance_2_vs = instance_2_registry.get_litellm_managed_vector_store_from_registry(
|
||||
test_vector_store_id
|
||||
)
|
||||
assert instance_1_vs.get("vector_store_name") == original_name
|
||||
assert instance_2_vs.get("vector_store_name") == original_name
|
||||
|
||||
# Step 3: Instance 1 updates the vector store in the database
|
||||
# (Simulating what happens in update_vector_store endpoint)
|
||||
update_data = {"vector_store_name": updated_name}
|
||||
await mock_prisma_client.db.litellm_managedvectorstorestable.update(
|
||||
where={"vector_store_id": test_vector_store_id},
|
||||
data=update_data
|
||||
)
|
||||
|
||||
# Instance 1 updates its own cache
|
||||
updated_vs_instance_1 = test_vector_store.copy()
|
||||
updated_vs_instance_1["vector_store_name"] = updated_name
|
||||
instance_1_registry.update_vector_store_in_registry(
|
||||
vector_store_id=test_vector_store_id,
|
||||
updated_data=updated_vs_instance_1
|
||||
)
|
||||
|
||||
# Verify Instance 1 has the updated data
|
||||
instance_1_vs_after_update = instance_1_registry.get_litellm_managed_vector_store_from_registry(
|
||||
test_vector_store_id
|
||||
)
|
||||
assert instance_1_vs_after_update.get("vector_store_name") == updated_name
|
||||
|
||||
# Verify Instance 2 still has stale data in cache
|
||||
instance_2_vs_before_list = instance_2_registry.get_litellm_managed_vector_store_from_registry(
|
||||
test_vector_store_id
|
||||
)
|
||||
assert instance_2_vs_before_list.get("vector_store_name") == original_name, (
|
||||
"Instance 2 should still have stale cached data before list operation"
|
||||
)
|
||||
|
||||
# Step 4: Instance 2 calls list endpoint (which should sync with database)
|
||||
# This simulates what list_vector_stores endpoint does
|
||||
vector_stores_from_db_after_update = await VectorStoreRegistry._get_vector_stores_from_db(
|
||||
prisma_client=mock_prisma_client
|
||||
)
|
||||
|
||||
# Build map from database vector stores (database is source of truth)
|
||||
vector_store_map = {}
|
||||
for vector_store in vector_stores_from_db_after_update:
|
||||
vector_store_id = vector_store.get("vector_store_id")
|
||||
if vector_store_id:
|
||||
vector_store_map[vector_store_id] = vector_store
|
||||
|
||||
# Update in-memory registry with database versions (this is the key fix)
|
||||
instance_2_registry.update_vector_store_in_registry(
|
||||
vector_store_id=vector_store_id,
|
||||
updated_data=vector_store
|
||||
)
|
||||
|
||||
# Step 5: Verify Instance 2 now has the updated data
|
||||
instance_2_vs_after_list = instance_2_registry.get_litellm_managed_vector_store_from_registry(
|
||||
test_vector_store_id
|
||||
)
|
||||
assert instance_2_vs_after_list.get("vector_store_name") == updated_name, (
|
||||
"Instance 2 should have updated data after list operation syncs with database"
|
||||
)
|
||||
|
||||
# Verify the list returned the correct data
|
||||
combined_vector_stores = list(vector_store_map.values())
|
||||
assert len(combined_vector_stores) == 1
|
||||
assert combined_vector_stores[0].get("vector_store_id") == test_vector_store_id
|
||||
assert combined_vector_stores[0].get("vector_store_name") == updated_name, (
|
||||
"List should return updated data from database"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_from_db():
|
||||
"""Test that _resolve_embedding_config_from_db correctly resolves embedding config from database."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue