From 922471f43b08dd5bfc91918cbc0618dfe2354dee Mon Sep 17 00:00:00 2001 From: Josh Hayes <35790761+hayescode@users.noreply.github.com> Date: Thu, 9 Oct 2025 08:49:07 -0500 Subject: [PATCH] fix: Databricks Vector Store (#3546) --- mem0/vector_stores/databricks.py | 6 +- tests/vector_stores/test_databricks.py | 87 +++++++++++++++++--------- 2 files changed, 61 insertions(+), 32 deletions(-) diff --git a/mem0/vector_stores/databricks.py b/mem0/vector_stores/databricks.py index 6b5660e74..b77058c5f 100644 --- a/mem0/vector_stores/databricks.py +++ b/mem0/vector_stores/databricks.py @@ -589,7 +589,8 @@ class Databricks(VectorStoreBase): raise KeyError(f"Vector with ID {vector_id} not found") result = data_array[0] - row_data = result if isinstance(result, dict) else result.__dict__ + columns = 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 payload = { @@ -608,7 +609,7 @@ class Databricks(VectorStoreBase): payload[field] = row_data[field] # Add metadata - if "metadata" in row_data: + if "metadata" in row_data and row_data.get('metadata'): try: metadata = json.loads(extract_json(row_data["metadata"])) payload.update(metadata) @@ -707,6 +708,7 @@ class Databricks(VectorStoreBase): except Exception: pass memory_id = row_dict.get("memory_id") or row_dict.get("id") + payload['data'] = payload['memory'] memory_results.append(MemoryResult(id=memory_id, payload=payload)) return [memory_results] except Exception as e: diff --git a/tests/vector_stores/test_databricks.py b/tests/vector_stores/test_databricks.py index 9b2d1c560..7f0c82e7e 100644 --- a/tests/vector_stores/test_databricks.py +++ b/tests/vector_stores/test_databricks.py @@ -1,6 +1,6 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch -from databricks.sdk.service.vectorsearch import VectorIndexType +from databricks.sdk.service.vectorsearch import VectorIndexType, QueryVectorIndexResponse, ResultManifest, ResultData, ColumnInfo from mem0.vector_stores.databricks import Databricks import pytest @@ -241,21 +241,33 @@ def test_update_vector(db_instance_direct, mock_workspace_client): def test_get_vector(db_instance_delta, mock_workspace_client): - mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( - result=SimpleNamespace( + 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="score"), + ]), + result=ResultData( data_array=[ - { - "memory_id": "id-get", - "hash": "h", - "agent_id": "a", - "run_id": "r", - "user_id": "u", - "memory": "some memory", - "metadata": '{"tag":"x"}', - "created_at": "2024-01-01T00:00:00", - "updated_at": "2024-01-01T00:00:00", - "score": 0.99, - } + [ + "id-get", + "h", + "a", + "r", + "u", + "some memory", + '{"tag":"x"}', + "2024-01-01T00:00:00", + "2024-01-01T00:00:00", + "0.99", + ] ] ) ) @@ -284,25 +296,40 @@ def test_col_info(db_instance_delta): def test_list_memories(db_instance_delta, mock_workspace_client): - row = { - "memory_id": "id3", - "hash": "hash3", - "agent_id": "agent3", - "run_id": "run3", - "user_id": "user3", - "memory": "memory three", - "metadata": '{"topic":"misc"}', - "created_at": "2024-01-03T00:00:00", - "updated_at": "2024-01-03T00:00:00", - "score": 0.33, - } - mock_workspace_client.vector_search_indexes.query_index.return_value = SimpleNamespace( - result=SimpleNamespace(data_array=[row]) + 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="score"), + ]), + result=ResultData( + data_array=[ + [ + "id-get", + "h", + "a", + "r", + "u", + "some memory", + '{"tag":"x"}', + "2024-01-01T00:00:00", + "2024-01-01T00:00:00", + "0.99", + ] + ] + ) ) res = db_instance_delta.list(limit=1) assert isinstance(res, list) assert len(res[0]) == 1 - assert res[0][0].id == "id3" + assert res[0][0].id == "id-get" # ---------------------- Reset Tests ---------------------- #