fix: align Databricks docs with config and fix query mode selection (#4477)

Co-authored-by: utkarsh240799 <utkarsh240799@users.noreply.github.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Utkarsh
2026-03-23 19:20:11 +05:30
committed by GitHub
parent 316dc67a0a
commit d8a6960b4a
3 changed files with 588 additions and 64 deletions
+29 -20
View File
@@ -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.
+60 -42
View File
@@ -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 []
+499 -2
View File
@@ -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