From 514ebb0d96c13967ab040611cf426f92157992b0 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 19 Jan 2026 13:17:08 +0530 Subject: [PATCH] Fix: vector store sync issues --- .../management_endpoints.py | 152 ++++++++++++--- .../test_vector_store_endpoints.py | 183 ++++++++++++++++++ 2 files changed, 307 insertions(+), 28 deletions(-) diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 661f94e5f04..bc61a60fe5a 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -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)) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 352e84719f1..558fe18ae38 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -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."""