diff --git a/docs/components/vectordbs/dbs/databricks.mdx b/docs/components/vectordbs/dbs/databricks.mdx index 032e5f511..db0056e86 100644 --- a/docs/components/vectordbs/dbs/databricks.mdx +++ b/docs/components/vectordbs/dbs/databricks.mdx @@ -17,8 +17,10 @@ config = { "workspace_url": "https://your-workspace.databricks.com", "access_token": "your-access-token", "endpoint_name": "your-vector-search-endpoint", - "index_name": "catalog.schema.index_name", - "source_table_name": "catalog.schema.source_table", + "catalog": "your_catalog", + "schema": "your_schema", + "table_name": "your_table", + "collection_name": "your_index_name", "embedding_dimension": 1536 } } @@ -42,17 +44,22 @@ Here are the parameters available for configuring Databricks Vector Search: | --- | --- | --- | | `workspace_url` | The URL of your Databricks workspace | **Required** | | `access_token` | Personal Access Token for authentication | `None` | -| `service_principal_client_id` | Service principal client ID (alternative to access_token) | `None` | -| `service_principal_client_secret` | Service principal client secret (required with client_id) | `None` | +| `client_id` | Service principal client ID (alternative to access_token) | `None` | +| `client_secret` | Service principal client secret (required with client_id) | `None` | +| `azure_client_id` | Azure AD application client ID (for Azure Databricks) | `None` | +| `azure_client_secret` | Azure AD application client secret (for Azure Databricks) | `None` | | `endpoint_name` | Name of the Vector Search endpoint | **Required** | -| `index_name` | Name of the vector index (Unity Catalog format: catalog.schema.index) | **Required** | -| `source_table_name` | Name of the source Delta table (Unity Catalog format: catalog.schema.table) | **Required** | -| `embedding_dimension` | Dimension of self-managed embeddings | `1536` | -| `embedding_source_column` | Column name for text when using Databricks-computed embeddings | `None` | +| `catalog` | Unity Catalog catalog name | **Required** | +| `schema` | Unity Catalog schema name | **Required** | +| `table_name` | Source Delta table name | **Required** | +| `collection_name` | Vector search index name | `mem0` | +| `index_type` | Index type: `DELTA_SYNC` or `DIRECT_ACCESS` | `DELTA_SYNC` | | `embedding_model_endpoint_name` | Databricks serving endpoint for embeddings | `None` | -| `embedding_vector_column` | Column name for self-managed embedding vectors | `embedding` | +| `embedding_dimension` | Dimension of self-managed embeddings | `1536` | | `endpoint_type` | Type of endpoint (`STANDARD` or `STORAGE_OPTIMIZED`) | `STANDARD` | -| `sync_computed_embeddings` | Whether to sync computed embeddings automatically | `True` | +| `pipeline_type` | Sync pipeline type: `TRIGGERED` or `CONTINUOUS` | `TRIGGERED` | +| `warehouse_name` | Databricks SQL warehouse name (if using SQL warehouse) | `None` | +| `query_type` | Query type: `ANN` or `HYBRID` | `ANN` | ### Authentication @@ -65,11 +72,13 @@ config = { "provider": "databricks", "config": { "workspace_url": "https://your-workspace.databricks.com", - "service_principal_client_id": "your-service-principal-id", - "service_principal_client_secret": "your-service-principal-secret", + "client_id": "your-service-principal-id", + "client_secret": "your-service-principal-secret", "endpoint_name": "your-endpoint", - "index_name": "catalog.schema.index_name", - "source_table_name": "catalog.schema.source_table" + "catalog": "your_catalog", + "schema": "your_schema", + "table_name": "your_table", + "collection_name": "your_index_name", } } } @@ -84,8 +93,10 @@ config = { "workspace_url": "https://your-workspace.databricks.com", "access_token": "your-personal-access-token", "endpoint_name": "your-endpoint", - "index_name": "catalog.schema.index_name", - "source_table_name": "catalog.schema.source_table" + "catalog": "your_catalog", + "schema": "your_schema", + "table_name": "your_table", + "collection_name": "your_index_name", } } } @@ -103,7 +114,6 @@ config = { "config": { # ... authentication config ... "embedding_dimension": 768, # Match your embedding model - "embedding_vector_column": "embedding" } } } @@ -118,7 +128,6 @@ config = { "provider": "databricks", "config": { # ... authentication config ... - "embedding_source_column": "text", "embedding_model_endpoint_name": "e5-small-v2" } } @@ -127,8 +136,8 @@ config = { ### Important Notes -- **Delta Sync Index**: This implementation uses Delta Sync Index, which automatically syncs with your source Delta table. Direct vector insertion/deletion/update operations will log warnings as they're not supported with Delta Sync. -- **Unity Catalog**: Both the source table and index must be in Unity Catalog format (`catalog.schema.table_name`). +- **Index Types**: This implementation supports both `DELTA_SYNC` (auto-syncs with source Delta table) and `DIRECT_ACCESS` (manage vectors directly) index types. +- **Unity Catalog**: The source table and index are created under the specified `catalog.schema` namespace. - **Endpoint Auto-Creation**: If the specified endpoint doesn't exist, it will be created automatically. - **Index Auto-Creation**: If the specified index doesn't exist, it will be created automatically with the provided configuration. - **Filter Support**: Supports filtering by metadata fields, with different syntax for STANDARD vs STORAGE_OPTIMIZED endpoints. diff --git a/mem0/vector_stores/databricks.py b/mem0/vector_stores/databricks.py index b77058c5f..b8ec82efd 100644 --- a/mem0/vector_stores/databricks.py +++ b/mem0/vector_stores/databricks.py @@ -65,7 +65,7 @@ class Databricks(VectorStoreBase): catalog (str): Unity Catalog catalog name. schema (str): Unity Catalog schema name. table_name (str): Source Delta table name. - index_name (str, optional): Vector search index name (default: "mem0"). + collection_name (str, optional): Vector search index name (default: "mem0"). index_type (str, optional): Index type, either "DELTA_SYNC" or "DIRECT_ACCESS" (default: "DELTA_SYNC"). embedding_model_endpoint_name (str, optional): Embedding model endpoint for Databricks-computed embeddings. embedding_dimension (int, optional): Vector embedding dimensions (default: 1536). @@ -85,7 +85,7 @@ class Databricks(VectorStoreBase): self.fully_qualified_index_name = f"{self.catalog}.{self.schema}.{self.index_name}" # Configuration - self.index_type = index_type + self.index_type = VectorIndexType(index_type) if isinstance(index_type, str) else index_type self.embedding_model_endpoint_name = embedding_model_endpoint_name self.embedding_dimension = embedding_dimension self.endpoint_type = endpoint_type @@ -261,11 +261,11 @@ class Databricks(VectorStoreBase): ) logger.info(f"Successfully created source table '{self.fully_qualified_table_name}'") self.client.table_constraints.create( - full_name_arg="logistics_dev.ai.dev_memory", + full_name_arg=self.fully_qualified_table_name, constraint=TableConstraint( primary_key_constraint=PrimaryKeyConstraint( - name="pk_dev_memory", # Name of the primary key constraint - child_columns=["memory_id"], # Columns that make up the primary key + name=f"pk_{self.table_name}", + child_columns=["memory_id"], ) ), ) @@ -439,29 +439,29 @@ class Databricks(VectorStoreBase): try: filters_json = json.dumps(filters) if filters else None - # Choose query type - if self.index_type == VectorIndexType.DELTA_SYNC and query: - # Text-based search - sdk_results = self.client.vector_search_indexes.query_index( - index_name=self.fully_qualified_index_name, - columns=self.column_names, - query_text=query, - num_results=limit, - query_type=self.query_type, - filters_json=filters_json, - ) - elif self.index_type == VectorIndexType.DIRECT_ACCESS and vectors: - # Vector-based search - sdk_results = self.client.vector_search_indexes.query_index( - index_name=self.fully_qualified_index_name, - columns=self.column_names, - query_vector=vectors, - num_results=limit, - query_type=self.query_type, - filters_json=filters_json, - ) + # Choose query mode per Databricks SDK contract: + # - query_text: for Delta Sync Index with model endpoint + # - query_vector: for Direct Access Index and Delta Sync Index with self-managed vectors + query_kwargs = { + "index_name": self.fully_qualified_index_name, + "columns": self.column_names, + "num_results": limit, + "query_type": self.query_type, + "filters_json": filters_json, + } + uses_model_endpoint = ( + self.index_type == VectorIndexType.DELTA_SYNC and self.embedding_model_endpoint_name + ) + if uses_model_endpoint: + if not query: + raise ValueError("Query text is required for Delta Sync Index with model endpoint.") + query_kwargs["query_text"] = query + elif vectors: + query_kwargs["query_vector"] = vectors else: - raise ValueError("Must provide query text for DELTA_SYNC or vectors for DIRECT_ACCESS.") + raise ValueError("Must provide vectors for search.") + + sdk_results = self.client.vector_search_indexes.query_index(**query_kwargs) # Parse results result_data = sdk_results.result if hasattr(sdk_results, "result") else sdk_results @@ -572,14 +572,23 @@ class Databricks(VectorStoreBase): filters = {"memory_id": vector_id} filters_json = json.dumps(filters) - results = self.client.vector_search_indexes.query_index( - index_name=self.fully_qualified_index_name, - columns=self.column_names, - query_text=" ", # Empty query, rely on filters - num_results=1, - query_type=self.query_type, - filters_json=filters_json, + # Use query_text for Delta Sync with model endpoint, query_vector otherwise + query_kwargs = { + "index_name": self.fully_qualified_index_name, + "columns": self.column_names, + "num_results": 1, + "query_type": self.query_type, + "filters_json": filters_json, + } + uses_model_endpoint = ( + self.index_type == VectorIndexType.DELTA_SYNC and self.embedding_model_endpoint_name ) + if uses_model_endpoint: + query_kwargs["query_text"] = " " + else: + query_kwargs["query_vector"] = [0.0] * self.embedding_dimension + + results = self.client.vector_search_indexes.query_index(**query_kwargs) # Process results result_data = results.result if hasattr(results, "result") else results @@ -589,7 +598,7 @@ class Databricks(VectorStoreBase): raise KeyError(f"Vector with ID {vector_id} not found") result = data_array[0] - columns = columns = [col.name for col in results.manifest.columns] if results.manifest and results.manifest.columns else [] + columns = [col.name for col in results.manifest.columns] if results.manifest and results.manifest.columns else [] row_data = dict(zip(columns, result)) # Build payload following the standard schema @@ -686,14 +695,23 @@ class Databricks(VectorStoreBase): filters_json = json.dumps(filters) if filters else None num_results = limit or 100 columns = self.column_names - sdk_results = self.client.vector_search_indexes.query_index( - index_name=self.fully_qualified_index_name, - columns=columns, - query_text=" ", - num_results=num_results, - query_type=self.query_type, - filters_json=filters_json, + # Use query_text for Delta Sync with model endpoint, query_vector otherwise + query_kwargs = { + "index_name": self.fully_qualified_index_name, + "columns": columns, + "num_results": num_results, + "query_type": self.query_type, + "filters_json": filters_json, + } + uses_model_endpoint = ( + self.index_type == VectorIndexType.DELTA_SYNC and self.embedding_model_endpoint_name ) + if uses_model_endpoint: + query_kwargs["query_text"] = " " + else: + query_kwargs["query_vector"] = [0.0] * self.embedding_dimension + + sdk_results = self.client.vector_search_indexes.query_index(**query_kwargs) result_data = sdk_results.result if hasattr(sdk_results, "result") else sdk_results data_array = result_data.data_array if hasattr(result_data, "data_array") else [] diff --git a/tests/vector_stores/test_databricks.py b/tests/vector_stores/test_databricks.py index 7f0c82e7e..b1ea2c94a 100644 --- a/tests/vector_stores/test_databricks.py +++ b/tests/vector_stores/test_databricks.py @@ -1,8 +1,12 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch + +import pytest + +pytest.importorskip("databricks", reason="databricks-sdk package not installed") + from databricks.sdk.service.vectorsearch import VectorIndexType, QueryVectorIndexResponse, ResultManifest, ResultData, ColumnInfo from mem0.vector_stores.databricks import Databricks -import pytest # ---------------------- Fixtures ---------------------- # @@ -205,9 +209,34 @@ def test_search_direct_access_vector(db_instance_direct, mock_workspace_client): assert results[0].score == 0.77 +def test_search_delta_sync_self_managed_vectors(mock_workspace_client): + """DELTA_SYNC without embedding model endpoint should use query_vector, not query_text.""" + mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True) + # DELTA_SYNC without embedding_model_endpoint_name = self-managed vectors + inst = Databricks( + workspace_url="https://test", + access_token="tok", + endpoint_name="vs-endpoint", + catalog="catalog", + schema="schema", + table_name="table", + warehouse_name="test-warehouse", + index_type=VectorIndexType.DELTA_SYNC, + embedding_dimension=4, + # NOTE: no embedding_model_endpoint_name + ) + mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( + result=SimpleNamespace(data_array=[]) + ) + inst.search(query="ignored", vectors=[0.1, 0.2, 0.3, 0.4], limit=5) + call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_vector" in call_kwargs + assert "query_text" not in call_kwargs + + def test_search_missing_params_raises(db_instance_delta): with pytest.raises(ValueError): - db_instance_delta.search(query="", vectors=[0.1, 0.2]) # DELTA_SYNC requires query text + db_instance_delta.search(query="", vectors=[0.1, 0.2]) # DELTA_SYNC with model endpoint requires query text # ---------------------- Delete Tests ---------------------- # @@ -275,6 +304,54 @@ def test_get_vector(db_instance_delta, mock_workspace_client): assert res.id == "id-get" assert res.payload["data"] == "some memory" assert res.payload["tag"] == "x" + # DELTA_SYNC should use query_text, not query_vector + call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_text" in call_kwargs + assert "query_vector" not in call_kwargs + + +def test_get_vector_direct_access(db_instance_direct, mock_workspace_client): + """get() on a DIRECT_ACCESS index must use query_vector instead of query_text.""" + mock_workspace_client.vector_search_indexes.query_index.return_value = QueryVectorIndexResponse( + manifest=ResultManifest(columns=[ + ColumnInfo(name="memory_id"), + ColumnInfo(name="hash"), + ColumnInfo(name="agent_id"), + ColumnInfo(name="run_id"), + ColumnInfo(name="user_id"), + ColumnInfo(name="memory"), + ColumnInfo(name="metadata"), + ColumnInfo(name="created_at"), + ColumnInfo(name="updated_at"), + ColumnInfo(name="embedding"), + ColumnInfo(name="score"), + ]), + result=ResultData( + data_array=[ + [ + "id-get-da", + "h", + "a", + "r", + "u", + "direct access memory", + '{"tag":"da"}', + "2024-01-01T00:00:00", + "2024-01-01T00:00:00", + [0.1, 0.2, 0.3, 0.4], + "0.88", + ] + ] + ) + ) + res = db_instance_direct.get("id-get-da") + assert res.id == "id-get-da" + assert res.payload["data"] == "direct access memory" + # DIRECT_ACCESS should use query_vector, not query_text + call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_vector" in call_kwargs + assert "query_text" not in call_kwargs + assert call_kwargs["query_vector"] == [0.0] * 4 # embedding_dimension=4 # ---------------------- Collection Info / Listing Tests ---------------------- # @@ -330,6 +407,185 @@ def test_list_memories(db_instance_delta, mock_workspace_client): assert isinstance(res, list) assert len(res[0]) == 1 assert res[0][0].id == "id-get" + # DELTA_SYNC should use query_text + call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_text" in call_kwargs + assert "query_vector" not in call_kwargs + + +def test_list_memories_direct_access(db_instance_direct, mock_workspace_client): + """list() on a DIRECT_ACCESS index must use query_vector instead of query_text.""" + mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( + result=SimpleNamespace( + data_array=[ + [ + "id-da-list", + "h", + "a", + "r", + "u", + "direct memory", + None, + "2024-01-01T00:00:00", + "2024-01-01T00:00:00", + [0.1, 0.2, 0.3, 0.4], + ] + ] + ) + ) + res = db_instance_direct.list(limit=5) + assert isinstance(res, list) + assert len(res[0]) == 1 + assert res[0][0].id == "id-da-list" + assert res[0][0].payload["data"] == "direct memory" + # DIRECT_ACCESS should use query_vector + call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_vector" in call_kwargs + assert "query_text" not in call_kwargs + assert call_kwargs["query_vector"] == [0.0] * 4 + + +def test_get_vector_delta_sync_self_managed(mock_workspace_client): + """get() on DELTA_SYNC without model endpoint should use query_vector.""" + mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True) + inst = Databricks( + workspace_url="https://test", + access_token="tok", + endpoint_name="vs-endpoint", + catalog="catalog", + schema="schema", + table_name="table", + warehouse_name="test-warehouse", + index_type=VectorIndexType.DELTA_SYNC, + embedding_dimension=4, + # NOTE: no embedding_model_endpoint_name + ) + mock_workspace_client.vector_search_indexes.query_index.return_value = QueryVectorIndexResponse( + manifest=ResultManifest(columns=[ + ColumnInfo(name="memory_id"), ColumnInfo(name="hash"), + ColumnInfo(name="agent_id"), ColumnInfo(name="run_id"), + ColumnInfo(name="user_id"), ColumnInfo(name="memory"), + ColumnInfo(name="metadata"), ColumnInfo(name="created_at"), + ColumnInfo(name="updated_at"), + ]), + result=ResultData(data_array=[["id-sm", "h", None, None, None, "self-managed mem", None, None, None]]), + ) + res = inst.get("id-sm") + assert res.id == "id-sm" + call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_vector" in call_kwargs + assert "query_text" not in call_kwargs + assert call_kwargs["query_vector"] == [0.0] * 4 + + +def test_list_memories_delta_sync_self_managed(mock_workspace_client): + """list() on DELTA_SYNC without model endpoint should use query_vector.""" + mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True) + inst = Databricks( + workspace_url="https://test", + access_token="tok", + endpoint_name="vs-endpoint", + catalog="catalog", + schema="schema", + table_name="table", + warehouse_name="test-warehouse", + index_type=VectorIndexType.DELTA_SYNC, + embedding_dimension=4, + # NOTE: no embedding_model_endpoint_name + ) + mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( + result=SimpleNamespace(data_array=[]) + ) + inst.list(limit=5) + call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_vector" in call_kwargs + assert "query_text" not in call_kwargs + + +def test_list_memories_default_limit(db_instance_delta, mock_workspace_client): + """list() with no limit should default to 100.""" + mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( + result=SimpleNamespace(data_array=[]) + ) + db_instance_delta.list(limit=None) + call_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert call_kwargs["num_results"] == 100 + + +# ---------------------- Table Creation Tests ---------------------- # + + +def test_ensure_source_table_uses_dynamic_names(mock_workspace_client): + """Verify _ensure_source_table_exists uses self.fully_qualified_table_name and + self.table_name for the PK constraint, not hardcoded values.""" + mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=False) + Databricks( + workspace_url="https://test", + access_token="tok", + endpoint_name="vs-endpoint", + catalog="my_catalog", + schema="my_schema", + table_name="my_memories", + collection_name="my_index", + warehouse_name="test-warehouse", + index_type=VectorIndexType.DELTA_SYNC, + embedding_model_endpoint_name="embedding-endpoint", + ) + # _ensure_source_table_exists was called during __init__ via create_col + constraint_call = mock_workspace_client.table_constraints.create.call_args + assert constraint_call.kwargs["full_name_arg"] == "my_catalog.my_schema.my_memories" + pk_name = constraint_call.kwargs["constraint"].primary_key_constraint.name + assert pk_name == "pk_my_memories" + + +# ---------------------- Config Validation Tests ---------------------- # + + +def test_config_rejects_old_doc_params(): + """Config should reject the old documentation parameter names like index_name and source_table_name.""" + from mem0.configs.vector_stores.databricks import DatabricksConfig + with pytest.raises(ValueError, match="Extra fields not allowed"): + DatabricksConfig( + workspace_url="https://test", + access_token="tok", + endpoint_name="ep", + catalog="cat", + schema="sch", + table_name="tbl", + index_name="catalog.schema.index", # old param from docs + ) + + +def test_config_rejects_source_table_name(): + """Config should reject source_table_name which was in old docs.""" + from mem0.configs.vector_stores.databricks import DatabricksConfig + with pytest.raises(ValueError, match="Extra fields not allowed"): + DatabricksConfig( + workspace_url="https://test", + access_token="tok", + endpoint_name="ep", + catalog="cat", + schema="sch", + table_name="tbl", + source_table_name="catalog.schema.table", # old param from docs + ) + + +def test_config_accepts_correct_params(): + """Config should accept all the correct parameter names.""" + from mem0.configs.vector_stores.databricks import DatabricksConfig + config = DatabricksConfig( + workspace_url="https://test", + access_token="tok", + endpoint_name="ep", + catalog="cat", + schema="sch", + table_name="tbl", + collection_name="my_index", + embedding_dimension=768, + ) + assert config.collection_name == "my_index" + assert config.embedding_dimension == 768 # ---------------------- Reset Tests ---------------------- # @@ -341,3 +597,244 @@ def test_reset(db_instance_delta, mock_workspace_client): with patch.object(db_instance_delta, "create_col", wraps=db_instance_delta.create_col) as create_spy: db_instance_delta.reset() assert create_spy.called + + +# ---------------------- End-to-End Config → Factory → CRUD Tests ---------------------- # + + +def test_e2e_config_to_factory_delta_sync(mock_workspace_client): + """End-to-end: VectorStoreConfig validates docs-correct params, factory creates Databricks instance.""" + from mem0.vector_stores.configs import VectorStoreConfig + from mem0.utils.factory import VectorStoreFactory + + # Step 1: Config validation (simulates what Memory.from_config does) + vs_config = VectorStoreConfig( + provider="databricks", + config={ + "workspace_url": "https://my-workspace.databricks.com", + "access_token": "my-token", + "endpoint_name": "my-endpoint", + "catalog": "prod_catalog", + "schema": "ai_schema", + "table_name": "memories_table", + "collection_name": "my_index", + "embedding_dimension": 768, + "warehouse_name": "test-warehouse", + }, + ) + assert vs_config.config.collection_name == "my_index" + assert vs_config.config.catalog == "prod_catalog" + + # Step 2: Factory instantiation (same as MemoryBase.__init__) + instance = VectorStoreFactory.create("databricks", vs_config.config) + assert isinstance(instance, Databricks) + assert instance.fully_qualified_table_name == "prod_catalog.ai_schema.memories_table" + assert instance.fully_qualified_index_name == "prod_catalog.ai_schema.my_index" + assert instance.embedding_dimension == 768 + + +def test_e2e_config_to_factory_direct_access(mock_workspace_client): + """End-to-end: DIRECT_ACCESS via config → factory creates correct instance.""" + from mem0.vector_stores.configs import VectorStoreConfig + from mem0.utils.factory import VectorStoreFactory + + mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True) + + vs_config = VectorStoreConfig( + provider="databricks", + config={ + "workspace_url": "https://my-workspace.databricks.com", + "access_token": "my-token", + "endpoint_name": "my-endpoint", + "catalog": "cat", + "schema": "sch", + "table_name": "tbl", + "index_type": "DIRECT_ACCESS", + "embedding_dimension": 4, + "warehouse_name": "test-warehouse", + }, + ) + instance = VectorStoreFactory.create("databricks", vs_config.config) + assert isinstance(instance, Databricks) + assert "embedding" in instance.column_names + + +def test_e2e_old_docs_config_rejected(): + """End-to-end: Config from old docs (with index_name, source_table_name) is rejected at validation.""" + from mem0.vector_stores.configs import VectorStoreConfig + + with pytest.raises(ValueError, match="Extra fields not allowed"): + VectorStoreConfig( + provider="databricks", + config={ + "workspace_url": "https://my-workspace.databricks.com", + "access_token": "my-token", + "endpoint_name": "my-endpoint", + "index_name": "catalog.schema.index_name", + "source_table_name": "catalog.schema.source_table", + "embedding_dimension": 1536, + }, + ) + + +def test_e2e_crud_lifecycle_delta_sync(mock_workspace_client): + """End-to-end CRUD lifecycle: insert → search → get → list → update → delete.""" + from mem0.vector_stores.configs import VectorStoreConfig + from mem0.utils.factory import VectorStoreFactory + + vs_config = VectorStoreConfig( + provider="databricks", + config={ + "workspace_url": "https://test", + "access_token": "tok", + "endpoint_name": "ep", + "catalog": "cat", + "schema": "sch", + "table_name": "tbl", + "warehouse_name": "test-warehouse", + "embedding_model_endpoint_name": "emb-ep", + }, + ) + db = VectorStoreFactory.create("databricks", vs_config.config) + + # INSERT + db.insert( + vectors=[[0.1, 0.2]], + payloads=[{"data": "test memory", "user_id": "u1", "hash": "h1"}], + ids=["mem-001"], + ) + insert_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"] + assert "INSERT INTO cat.sch.tbl" in insert_sql + assert "mem-001" in insert_sql + + # SEARCH + mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( + result=SimpleNamespace( + data_array=[["mem-001", "h1", None, None, "u1", "test memory", None, None, None, 0.95]] + ) + ) + results = db.search(query="test", vectors=None, limit=5) + assert len(results) == 1 + assert results[0].id == "mem-001" + assert results[0].payload["data"] == "test memory" + search_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert search_kwargs["query_text"] == "test" + + # GET + mock_workspace_client.vector_search_indexes.query_index.return_value = QueryVectorIndexResponse( + manifest=ResultManifest(columns=[ + ColumnInfo(name="memory_id"), ColumnInfo(name="hash"), + ColumnInfo(name="agent_id"), ColumnInfo(name="run_id"), + ColumnInfo(name="user_id"), ColumnInfo(name="memory"), + ColumnInfo(name="metadata"), ColumnInfo(name="created_at"), + ColumnInfo(name="updated_at"), + ]), + result=ResultData(data_array=[["mem-001", "h1", None, None, "u1", "test memory", None, None, None]]), + ) + got = db.get("mem-001") + assert got.id == "mem-001" + assert got.payload["data"] == "test memory" + get_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_text" in get_kwargs + assert "query_vector" not in get_kwargs + + # LIST + mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( + result=SimpleNamespace( + data_array=[["mem-001", "h1", None, None, "u1", "test memory", None, None, None]] + ) + ) + listed = db.list(filters={"user_id": "u1"}, limit=10) + assert len(listed[0]) == 1 + assert listed[0][0].id == "mem-001" + list_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_text" in list_kwargs + assert list_kwargs["num_results"] == 10 + + # UPDATE + db.update(vector_id="mem-001", payload={"memory": "updated memory"}) + update_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"] + assert "UPDATE cat.sch.tbl" in update_sql + assert "mem-001" in update_sql + assert "updated memory" in update_sql + + # DELETE + db.delete("mem-001") + delete_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"] + assert "DELETE FROM cat.sch.tbl" in delete_sql + assert "mem-001" in delete_sql + + +def test_e2e_crud_lifecycle_direct_access(mock_workspace_client): + """End-to-end CRUD lifecycle for DIRECT_ACCESS: insert → search → get → list.""" + from mem0.vector_stores.configs import VectorStoreConfig + from mem0.utils.factory import VectorStoreFactory + + mock_workspace_client.tables.exists.return_value = SimpleNamespace(table_exists=True) + + vs_config = VectorStoreConfig( + provider="databricks", + config={ + "workspace_url": "https://test", + "access_token": "tok", + "endpoint_name": "ep", + "catalog": "cat", + "schema": "sch", + "table_name": "tbl", + "index_type": "DIRECT_ACCESS", + "embedding_dimension": 4, + "warehouse_name": "test-warehouse", + "embedding_model_endpoint_name": "emb-ep", + }, + ) + db = VectorStoreFactory.create("databricks", vs_config.config) + assert "embedding" in db.column_names + + # INSERT with vector + db.insert( + vectors=[[0.1, 0.2, 0.3, 0.4]], + payloads=[{"data": "direct memory", "user_id": "u1", "hash": "h1"}], + ids=["mem-da-001"], + ) + insert_sql = mock_workspace_client.statement_execution.execute_statement.call_args.kwargs["statement"] + assert "array(0.1, 0.2, 0.3, 0.4)" in insert_sql + + # SEARCH with vector + mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( + result=SimpleNamespace( + data_array=[["mem-da-001", "h1", None, None, "u1", "direct memory", None, None, None, [0.1, 0.2, 0.3, 0.4], 0.9]] + ) + ) + results = db.search(query="", vectors=[0.1, 0.2, 0.3, 0.4], limit=5) + assert len(results) == 1 + search_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_vector" in search_kwargs + assert "query_text" not in search_kwargs + + # GET — must use query_vector for DIRECT_ACCESS + mock_workspace_client.vector_search_indexes.query_index.return_value = QueryVectorIndexResponse( + manifest=ResultManifest(columns=[ + ColumnInfo(name="memory_id"), ColumnInfo(name="hash"), + ColumnInfo(name="agent_id"), ColumnInfo(name="run_id"), + ColumnInfo(name="user_id"), ColumnInfo(name="memory"), + ColumnInfo(name="metadata"), ColumnInfo(name="created_at"), + ColumnInfo(name="updated_at"), ColumnInfo(name="embedding"), + ]), + result=ResultData(data_array=[["mem-da-001", "h1", None, None, "u1", "direct memory", None, None, None, [0.1, 0.2, 0.3, 0.4]]]), + ) + got = db.get("mem-da-001") + assert got.id == "mem-da-001" + get_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_vector" in get_kwargs + assert get_kwargs["query_vector"] == [0.0] * 4 + + # LIST — must use query_vector for DIRECT_ACCESS + mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( + result=SimpleNamespace( + data_array=[["mem-da-001", "h1", None, None, "u1", "direct memory", None, None, None, [0.1, 0.2, 0.3, 0.4]]] + ) + ) + db.list(limit=5) + list_kwargs = mock_workspace_client.vector_search_indexes.query_index.call_args.kwargs + assert "query_vector" in list_kwargs + assert "query_text" not in list_kwargs