From 81e8a127b88037ca026906c10842130d467e4331 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 29 Jan 2026 16:31:30 -0800 Subject: [PATCH 1/2] Allow config embedding models --- .../management_endpoints.py | 143 ++++++++- .../test_vector_store_endpoints.py | 283 ++++++++++++++++++ .../VectorStoreForm.tsx | 9 +- 3 files changed, 431 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 6185f1541f..d34f26db9b 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -37,6 +37,88 @@ from litellm.vector_stores.vector_store_registry import VectorStoreRegistry router = APIRouter() +def _resolve_embedding_config_from_router( + embedding_model: str, llm_router +) -> Optional[Dict[str, Any]]: + """ + Resolve embedding config from router's config-defined models. + + Config-defined models (from proxy_config.yaml) are stored in the router's model_list, + not in the database. This function looks up the model in the router and extracts + api_key, api_base, and api_version from the deployment's litellm_params. + + Args: + embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large") + llm_router: The LiteLLM router instance + + Returns: + Dictionary with api_key, api_base, and api_version if model found, None otherwise + """ + if not embedding_model or llm_router is None: + return None + + # Extract model name candidates - could be "text-embedding-ada-002" or "azure/text-embedding-3-large" + # Try exact match first, then try without provider prefix + model_name_candidates = [embedding_model] + if "/" in embedding_model: + # If it has a provider prefix, also try without it + _, model_name = embedding_model.split("/", 1) + model_name_candidates.append(model_name) + + # Try to find model in router + for model_name in model_name_candidates: + try: + # Try to get deployment by model group name (model_name in config) + deployment = llm_router.get_deployment_by_model_group_name( + model_group_name=model_name + ) + + if deployment is not None and deployment.litellm_params is not None: + litellm_params = deployment.litellm_params + + # Build embedding config from model params + embedding_config: Dict[str, Any] = {} + + # Extract api_key + api_key = getattr(litellm_params, "api_key", None) + if api_key: + # Handle os.environ/ prefix + if isinstance(api_key, str) and api_key.startswith("os.environ/"): + api_key = get_secret(api_key) + embedding_config["api_key"] = api_key + + # Extract api_base + api_base = getattr(litellm_params, "api_base", None) + if api_base: + # Handle os.environ/ prefix + if isinstance(api_base, str) and api_base.startswith("os.environ/"): + api_base = get_secret(api_base) + embedding_config["api_base"] = api_base + + # Extract api_version + api_version = getattr(litellm_params, "api_version", None) + if api_version: + embedding_config["api_version"] = api_version + + project_id = getattr(litellm_params, "project_id", None) + if project_id: + embedding_config["project_id"] = project_id + + # Only return config if we have at least api_key or api_base + if embedding_config: + verbose_proxy_logger.debug( + f"Resolved embedding config from router model {model_name}: {list(embedding_config.keys())}" + ) + return embedding_config + except Exception as e: + verbose_proxy_logger.debug( + f"Error resolving embedding config from router for model {model_name}: {str(e)}" + ) + continue + + return None + + async def _resolve_embedding_config_from_db( embedding_model: str, prisma_client ) -> Optional[Dict[str, Any]]: @@ -133,6 +215,63 @@ async def _resolve_embedding_config_from_db( return None +async def _resolve_embedding_config( + embedding_model: str, prisma_client, llm_router=None +) -> Optional[Dict[str, Any]]: + """ + Resolve embedding config from either router (config-defined) or database models. + + This function first checks the router for config-defined models, then falls back + to the database. This allows users to use models defined in either location. + + Args: + embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large") + prisma_client: The Prisma client instance + llm_router: The LiteLLM router instance (optional, will be imported if not provided) + + Returns: + Dictionary with api_key, api_base, and api_version if model found, None otherwise + """ + if not embedding_model: + return None + + # Import llm_router if not provided + if llm_router is None: + try: + from litellm.proxy.proxy_server import llm_router + except ImportError: + llm_router = None + + # First try to resolve from router (config-defined models) + if llm_router is not None: + router_config = _resolve_embedding_config_from_router( + embedding_model=embedding_model, + llm_router=llm_router + ) + if router_config: + verbose_proxy_logger.debug( + f"Resolved embedding config from router for model {embedding_model}" + ) + return router_config + + # Fall back to database + if prisma_client is not None: + db_config = await _resolve_embedding_config_from_db( + embedding_model=embedding_model, + prisma_client=prisma_client + ) + if db_config: + verbose_proxy_logger.debug( + f"Resolved embedding config from database for model {embedding_model}" + ) + return db_config + + verbose_proxy_logger.debug( + f"Could not resolve embedding config for model {embedding_model} from router or database" + ) + return None + + ######################################################## # Helper Functions ######################################################## @@ -236,7 +375,7 @@ async def create_vector_store_in_db( # Auto-resolve embedding config if embedding model is provided but config is not embedding_model = litellm_params.get("litellm_embedding_model") if embedding_model and not litellm_params.get("litellm_embedding_config"): - resolved_config = await _resolve_embedding_config_from_db( + resolved_config = await _resolve_embedding_config( embedding_model=embedding_model, prisma_client=prisma_client ) @@ -648,7 +787,7 @@ async def update_vector_store( # 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( + resolved_config = await _resolve_embedding_config( embedding_model=embedding_model, prisma_client=prisma_client ) 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 558fe18ae3..703c7f2f24 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 @@ -21,7 +21,9 @@ from litellm.proxy.vector_store_endpoints.endpoints import ( _update_request_data_with_litellm_managed_vector_store_registry, ) from litellm.proxy.vector_store_endpoints.management_endpoints import ( + _resolve_embedding_config, _resolve_embedding_config_from_db, + _resolve_embedding_config_from_router, new_vector_store, ) from litellm.proxy.vector_store_endpoints.utils import ( @@ -1316,6 +1318,8 @@ async def test_new_vector_store_auto_resolves_embedding_config(): # Mock user API key mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) mock_user_api_key.user_role = None + mock_user_api_key.team_id = None + mock_user_api_key.user_id = None # Mock database operations mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( @@ -1345,9 +1349,16 @@ async def test_new_vector_store_auto_resolves_embedding_config(): mock_registry = MagicMock() mock_registry.add_vector_store_to_registry = MagicMock() + # Mock router to return None (so it falls back to DB resolution) + mock_router = MagicMock() + mock_router.get_deployment_by_model_group_name.return_value = None + with patch( "litellm.proxy.proxy_server.prisma_client", mock_prisma_client + ), patch( + "litellm.proxy.proxy_server.llm_router", + mock_router ), patch( "litellm.proxy.vector_store_endpoints.management_endpoints.decrypt_value_helper", side_effect=lambda value, key, return_original_value: value @@ -1368,3 +1379,275 @@ async def test_new_vector_store_auto_resolves_embedding_config(): assert litellm_params_dict["litellm_embedding_config"]["api_key"] == "resolved-api-key" assert litellm_params_dict["litellm_embedding_config"]["api_base"] == "https://api.openai.com" assert litellm_params_dict["litellm_embedding_config"]["api_version"] == "2024-01-01" + + +def test_resolve_embedding_config_from_router(): + """Test that _resolve_embedding_config_from_router correctly extracts credentials from config-defined models.""" + from litellm.types.router import Deployment, LiteLLM_Params + + # Create a mock router with a model + mock_router = MagicMock() + + # Create a mock deployment with litellm_params + mock_litellm_params = MagicMock(spec=LiteLLM_Params) + mock_litellm_params.api_key = "config-api-key" + mock_litellm_params.api_base = "https://config-api-base.com" + mock_litellm_params.api_version = "2024-02-01" + + mock_deployment = MagicMock(spec=Deployment) + mock_deployment.litellm_params = mock_litellm_params + + mock_router.get_deployment_by_model_group_name.return_value = mock_deployment + + # Test resolution + result = _resolve_embedding_config_from_router( + embedding_model="text-embedding-ada-002", + llm_router=mock_router + ) + + assert result is not None + assert result["api_key"] == "config-api-key" + assert result["api_base"] == "https://config-api-base.com" + assert result["api_version"] == "2024-02-01" + + mock_router.get_deployment_by_model_group_name.assert_called_once_with( + model_group_name="text-embedding-ada-002" + ) + + +def test_resolve_embedding_config_from_router_with_provider_prefix(): + """Test that _resolve_embedding_config_from_router handles provider prefixes like 'azure/model-name'.""" + from litellm.types.router import Deployment, LiteLLM_Params + + # Create a mock router + mock_router = MagicMock() + + # Create a mock deployment + mock_litellm_params = MagicMock(spec=LiteLLM_Params) + mock_litellm_params.api_key = "azure-api-key" + mock_litellm_params.api_base = "https://azure-endpoint.openai.azure.com" + mock_litellm_params.api_version = "2024-02-15" + + mock_deployment = MagicMock(spec=Deployment) + mock_deployment.litellm_params = mock_litellm_params + + # First call with full name returns None, second call with stripped name returns deployment + mock_router.get_deployment_by_model_group_name.side_effect = [None, mock_deployment] + + result = _resolve_embedding_config_from_router( + embedding_model="azure/text-embedding-3-large", + llm_router=mock_router + ) + + assert result is not None + assert result["api_key"] == "azure-api-key" + assert result["api_base"] == "https://azure-endpoint.openai.azure.com" + assert result["api_version"] == "2024-02-15" + + # Should have tried both the full name and stripped name + assert mock_router.get_deployment_by_model_group_name.call_count == 2 + + +def test_resolve_embedding_config_from_router_returns_none_when_not_found(): + """Test that _resolve_embedding_config_from_router returns None when model is not in router.""" + mock_router = MagicMock() + mock_router.get_deployment_by_model_group_name.return_value = None + + result = _resolve_embedding_config_from_router( + embedding_model="nonexistent-model", + llm_router=mock_router + ) + + assert result is None + + +def test_resolve_embedding_config_from_router_handles_os_environ(): + """Test that _resolve_embedding_config_from_router handles os.environ/ prefixed values.""" + from litellm.types.router import Deployment, LiteLLM_Params + + mock_router = MagicMock() + + mock_litellm_params = MagicMock(spec=LiteLLM_Params) + mock_litellm_params.api_key = "os.environ/OPENAI_API_KEY" + mock_litellm_params.api_base = "https://direct-url.com" + mock_litellm_params.api_version = None + + mock_deployment = MagicMock(spec=Deployment) + mock_deployment.litellm_params = mock_litellm_params + + mock_router.get_deployment_by_model_group_name.return_value = mock_deployment + + with patch( + "litellm.proxy.vector_store_endpoints.management_endpoints.get_secret", + return_value="resolved-from-env" + ) as mock_get_secret: + result = _resolve_embedding_config_from_router( + embedding_model="text-embedding-ada-002", + llm_router=mock_router + ) + + assert result is not None + assert result["api_key"] == "resolved-from-env" + assert result["api_base"] == "https://direct-url.com" + assert "api_version" not in result + + mock_get_secret.assert_called_once_with("os.environ/OPENAI_API_KEY") + + +@pytest.mark.asyncio +async def test_resolve_embedding_config_tries_router_then_db(): + """Test that _resolve_embedding_config tries router first, then falls back to DB.""" + from litellm.types.router import Deployment, LiteLLM_Params + + mock_prisma_client = MagicMock() + mock_router = MagicMock() + + # Router has the model + mock_litellm_params = MagicMock(spec=LiteLLM_Params) + mock_litellm_params.api_key = "router-api-key" + mock_litellm_params.api_base = "https://router-api-base.com" + mock_litellm_params.api_version = None + + mock_deployment = MagicMock(spec=Deployment) + mock_deployment.litellm_params = mock_litellm_params + + mock_router.get_deployment_by_model_group_name.return_value = mock_deployment + + # DB should NOT be called since router has the model + mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock() + + result = await _resolve_embedding_config( + embedding_model="text-embedding-ada-002", + prisma_client=mock_prisma_client, + llm_router=mock_router + ) + + assert result is not None + assert result["api_key"] == "router-api-key" + + # DB should NOT have been called since router found the model + mock_prisma_client.db.litellm_proxymodeltable.find_first.assert_not_called() + + +@pytest.mark.asyncio +async def test_resolve_embedding_config_falls_back_to_db(): + """Test that _resolve_embedding_config falls back to DB when router doesn't have the model.""" + mock_prisma_client = MagicMock() + mock_router = MagicMock() + + # Router doesn't have the model + mock_router.get_deployment_by_model_group_name.return_value = None + + # DB has the model + mock_db_model = MagicMock() + mock_db_model.litellm_params = { + "api_key": "db-api-key", + "api_base": "https://db-api-base.com", + } + mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock( + return_value=mock_db_model + ) + + with patch( + "litellm.proxy.vector_store_endpoints.management_endpoints.decrypt_value_helper", + side_effect=lambda value, key, return_original_value: value + ): + result = await _resolve_embedding_config( + embedding_model="text-embedding-ada-002", + prisma_client=mock_prisma_client, + llm_router=mock_router + ) + + assert result is not None + assert result["api_key"] == "db-api-key" + + # DB should have been called since router didn't find the model + mock_prisma_client.db.litellm_proxymodeltable.find_first.assert_called() + + +@pytest.mark.asyncio +async def test_new_vector_store_auto_resolves_from_router(): + """Test that new_vector_store auto-resolves embedding config from router when model is config-defined.""" + import json + + from litellm.types.router import Deployment, LiteLLM_Params + from litellm.types.vector_stores import LiteLLM_ManagedVectorStore + + mock_prisma_client = MagicMock() + + # Mock vector store request with embedding_model but no embedding_config + vector_store_data: LiteLLM_ManagedVectorStore = { + "vector_store_id": "test-store-router-001", + "custom_llm_provider": "openai", + "litellm_params": { + "litellm_embedding_model": "config-embedding-model", + # Note: litellm_embedding_config is not provided + } + } + + # Mock router with the model + mock_router = MagicMock() + mock_litellm_params = MagicMock(spec=LiteLLM_Params) + mock_litellm_params.api_key = "router-resolved-api-key" + mock_litellm_params.api_base = "https://router-resolved-base.com" + mock_litellm_params.api_version = "2024-03-01" + + mock_deployment = MagicMock(spec=Deployment) + mock_deployment.litellm_params = mock_litellm_params + + mock_router.get_deployment_by_model_group_name.return_value = mock_deployment + + # Mock user API key + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.user_role = None + mock_user_api_key.team_id = None + mock_user_api_key.user_id = None + + # Mock database operations + mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( + return_value=None # Vector store doesn't exist yet + ) + + # Track what was passed to create + captured_create_data = {} + + async def mock_create(*args, **kwargs): + captured_create_data.update(kwargs.get("data", {})) + mock_created_vector_store = MagicMock() + mock_created_vector_store.model_dump.return_value = { + "vector_store_id": "test-store-router-001", + "custom_llm_provider": "openai", + "litellm_params": kwargs.get("data", {}).get("litellm_params") + } + return mock_created_vector_store + + mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock( + side_effect=mock_create + ) + + mock_registry = MagicMock() + mock_registry.add_vector_store_to_registry = MagicMock() + + with patch( + "litellm.proxy.proxy_server.prisma_client", + mock_prisma_client + ), patch( + "litellm.proxy.proxy_server.llm_router", + mock_router + ), patch.object( + litellm, "vector_store_registry", mock_registry + ): + result = await new_vector_store( + vector_store=vector_store_data, + user_api_key_dict=mock_user_api_key + ) + + assert result["status"] == "success" + # Verify that embedding config was resolved from router and included in the create call + litellm_params_json = captured_create_data.get("litellm_params") + assert litellm_params_json is not None + litellm_params_dict = json.loads(litellm_params_json) + assert "litellm_embedding_config" in litellm_params_dict + assert litellm_params_dict["litellm_embedding_config"]["api_key"] == "router-resolved-api-key" + assert litellm_params_dict["litellm_embedding_config"]["api_base"] == "https://router-resolved-base.com" + assert litellm_params_dict["litellm_embedding_config"]["api_version"] == "2024-03-01" diff --git a/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreForm.tsx b/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreForm.tsx index 8be879b223..506543eb42 100644 --- a/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreForm.tsx +++ b/ui/litellm-dashboard/src/components/vector_store_management/VectorStoreForm.tsx @@ -76,7 +76,12 @@ const VectorStoreForm: React.FC = ({ const providerFields = getProviderSpecificFields(formValues.custom_llm_provider); const litellmParams = providerFields.reduce( (acc, field) => { - acc[field.name] = formValues[field.name]; + // Special handling for Milvus: rename embedding_model to litellm_embedding_model + if (formValues.custom_llm_provider === "milvus" && field.name === "embedding_model") { + acc["litellm_embedding_model"] = formValues[field.name]; + } else { + acc[field.name] = formValues[field.name]; + } return acc; }, {} as Record, @@ -229,7 +234,7 @@ const VectorStoreForm: React.FC = ({ {getProviderSpecificFields(selectedProvider).map((field: VectorStoreFieldConfig) => { if (field.type === "select") { const embeddingModels = modelInfo - .filter((option: ModelGroup) => option.mode === "embedding") + .filter((option: ModelGroup) => option.mode === "embedding" || option.mode === null) .map((option: ModelGroup) => ({ value: option.model_group, label: option.model_group, From 0b6bacb6d39697a51f7787442d64306f02ce7a6b Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 29 Jan 2026 16:34:21 -0800 Subject: [PATCH 2/2] adding tests --- .../test_vector_store_endpoints.py | 186 ++++++++++++++++++ 1 file changed, 186 insertions(+) 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 703c7f2f24..b24f0004f2 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 @@ -21,9 +21,11 @@ from litellm.proxy.vector_store_endpoints.endpoints import ( _update_request_data_with_litellm_managed_vector_store_registry, ) from litellm.proxy.vector_store_endpoints.management_endpoints import ( + _check_vector_store_access, _resolve_embedding_config, _resolve_embedding_config_from_db, _resolve_embedding_config_from_router, + create_vector_store_in_db, new_vector_store, ) from litellm.proxy.vector_store_endpoints.utils import ( @@ -1651,3 +1653,187 @@ async def test_new_vector_store_auto_resolves_from_router(): assert litellm_params_dict["litellm_embedding_config"]["api_key"] == "router-resolved-api-key" assert litellm_params_dict["litellm_embedding_config"]["api_base"] == "https://router-resolved-base.com" assert litellm_params_dict["litellm_embedding_config"]["api_version"] == "2024-03-01" + + +class TestCheckVectorStoreAccess: + """Test suite for _check_vector_store_access function.""" + + def test_access_granted_when_no_team_id(self): + """Test that access is granted when vector store has no team_id (legacy behavior).""" + vector_store: LiteLLM_ManagedVectorStore = { + "vector_store_id": "test-store", + "custom_llm_provider": "openai", + # No team_id field + } + + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.team_id = "team-123" + + result = _check_vector_store_access(vector_store, mock_user_api_key) + assert result is True + + def test_access_granted_when_team_ids_match(self): + """Test that access is granted when user's team_id matches vector store's team_id.""" + vector_store: LiteLLM_ManagedVectorStore = { + "vector_store_id": "test-store", + "custom_llm_provider": "openai", + "team_id": "team-123", + } + + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.team_id = "team-123" + + result = _check_vector_store_access(vector_store, mock_user_api_key) + assert result is True + + def test_access_denied_when_team_ids_dont_match(self): + """Test that access is denied when user's team_id doesn't match vector store's team_id.""" + vector_store: LiteLLM_ManagedVectorStore = { + "vector_store_id": "test-store", + "custom_llm_provider": "openai", + "team_id": "team-123", + } + + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.team_id = "team-456" + + result = _check_vector_store_access(vector_store, mock_user_api_key) + assert result is False + + def test_access_denied_when_vector_store_has_team_id_but_user_doesnt(self): + """Test that access is denied when vector store has team_id but user doesn't.""" + vector_store: LiteLLM_ManagedVectorStore = { + "vector_store_id": "test-store", + "custom_llm_provider": "openai", + "team_id": "team-123", + } + + mock_user_api_key = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key.team_id = None + + result = _check_vector_store_access(vector_store, mock_user_api_key) + assert result is False + + +@pytest.mark.asyncio +async def test_create_vector_store_in_db(): + """Test that create_vector_store_in_db correctly creates a vector store in the database.""" + from datetime import datetime, timezone + + mock_prisma_client = MagicMock() + + # Mock vector store data + vector_store_id = "test-create-store-001" + custom_llm_provider = "openai" + vector_store_name = "Test Store" + vector_store_description = "Test Description" + vector_store_metadata = {"key": "value"} + litellm_params = {"api_key": "test-key"} + team_id = "team-123" + user_id = "user-456" + + # Mock database operations + mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( + return_value=None # Vector store doesn't exist yet + ) + + created_vector_store_data = { + "vector_store_id": vector_store_id, + "custom_llm_provider": custom_llm_provider, + "vector_store_name": vector_store_name, + "vector_store_description": vector_store_description, + "vector_store_metadata": '{"key": "value"}', + "litellm_params": '{"api_key": "test-key"}', + "team_id": team_id, + "user_id": user_id, + "created_at": datetime.now(timezone.utc), + "updated_at": datetime.now(timezone.utc), + } + + mock_created_vector_store = MagicMock() + mock_created_vector_store.model_dump.return_value = created_vector_store_data + + mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock( + return_value=mock_created_vector_store + ) + + mock_registry = MagicMock() + mock_registry.add_vector_store_to_registry = MagicMock() + + with patch.object(litellm, "vector_store_registry", mock_registry): + result = await create_vector_store_in_db( + vector_store_id=vector_store_id, + custom_llm_provider=custom_llm_provider, + prisma_client=mock_prisma_client, + vector_store_name=vector_store_name, + vector_store_description=vector_store_description, + vector_store_metadata=vector_store_metadata, + litellm_params=litellm_params, + team_id=team_id, + user_id=user_id, + ) + + # Verify the result + assert result is not None + assert result["vector_store_id"] == vector_store_id + assert result["custom_llm_provider"] == custom_llm_provider + + # Verify database was called correctly + mock_prisma_client.db.litellm_managedvectorstorestable.find_unique.assert_called_once_with( + where={"vector_store_id": vector_store_id} + ) + mock_prisma_client.db.litellm_managedvectorstorestable.create.assert_called_once() + + # Verify registry was updated + mock_registry.add_vector_store_to_registry.assert_called_once() + + # Verify that create was called with correct data structure + create_call_args = mock_prisma_client.db.litellm_managedvectorstorestable.create.call_args + create_data = create_call_args.kwargs.get("data", {}) + assert create_data["vector_store_id"] == vector_store_id + assert create_data["custom_llm_provider"] == custom_llm_provider + assert create_data["vector_store_name"] == vector_store_name + assert create_data["vector_store_description"] == vector_store_description + assert create_data["team_id"] == team_id + assert create_data["user_id"] == user_id + + +@pytest.mark.asyncio +async def test_create_vector_store_in_db_raises_when_exists(): + """Test that create_vector_store_in_db raises HTTPException when vector store already exists.""" + mock_prisma_client = MagicMock() + + vector_store_id = "existing-store" + + # Mock that vector store already exists + existing_vector_store = MagicMock() + mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( + return_value=existing_vector_store + ) + + with pytest.raises(HTTPException) as exc_info: + await create_vector_store_in_db( + vector_store_id=vector_store_id, + custom_llm_provider="openai", + prisma_client=mock_prisma_client, + ) + + assert exc_info.value.status_code == 400 + assert "already exists" in exc_info.value.detail.lower() + + # Verify create was not called + mock_prisma_client.db.litellm_managedvectorstorestable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_create_vector_store_in_db_raises_when_no_db(): + """Test that create_vector_store_in_db raises HTTPException when database is not connected.""" + with pytest.raises(HTTPException) as exc_info: + await create_vector_store_in_db( + vector_store_id="test-store", + custom_llm_provider="openai", + prisma_client=None, + ) + + assert exc_info.value.status_code == 500 + assert "database not connected" in exc_info.value.detail.lower()