mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
9af320187c
Signed-off-by: Drew Cain <groksrc@gmail.com>
443 lines
15 KiB
Python
443 lines
15 KiB
Python
"""Semantic search service regression tests for local SQLite search."""
|
|
|
|
from datetime import datetime
|
|
from types import SimpleNamespace
|
|
from unittest.mock import AsyncMock
|
|
|
|
import pytest
|
|
from sqlalchemy import text
|
|
|
|
from basic_memory import db
|
|
from basic_memory.repository import EntityRepository
|
|
from basic_memory.repository.search_repository_base import VectorSyncBatchResult
|
|
from basic_memory.repository.semantic_errors import (
|
|
SemanticDependenciesMissingError,
|
|
SemanticSearchDisabledError,
|
|
)
|
|
from basic_memory.repository.sqlite_search_repository import SQLiteSearchRepository
|
|
from basic_memory.schemas.search import SearchItemType, SearchQuery, SearchRetrievalMode
|
|
|
|
|
|
def _sqlite_repo(search_service) -> SQLiteSearchRepository:
|
|
repository = search_service.repository
|
|
if not isinstance(repository, SQLiteSearchRepository):
|
|
pytest.skip("Semantic retrieval behavior is local SQLite-only in this phase.")
|
|
return repository
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_vector_search_fails_when_disabled(search_service, test_graph):
|
|
"""Vector mode should fail fast when semantic search is disabled."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = False
|
|
|
|
with pytest.raises(SemanticSearchDisabledError):
|
|
await search_service.search(
|
|
SearchQuery(
|
|
text="Connected Entity",
|
|
retrieval_mode=SearchRetrievalMode.VECTOR,
|
|
)
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_hybrid_search_fails_when_disabled(search_service, test_graph):
|
|
"""Hybrid mode should fail fast when semantic search is disabled."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = False
|
|
|
|
with pytest.raises(SemanticSearchDisabledError):
|
|
await search_service.search(
|
|
SearchQuery(
|
|
text="Root Entity",
|
|
retrieval_mode=SearchRetrievalMode.HYBRID,
|
|
)
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_vector_search_fails_when_provider_unavailable(search_service, test_graph):
|
|
"""Vector mode should fail fast when semantic provider is unavailable."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = True
|
|
repository._embedding_provider = None
|
|
repository._vector_tables_initialized = False
|
|
|
|
with pytest.raises(SemanticDependenciesMissingError):
|
|
await search_service.search(
|
|
SearchQuery(
|
|
text="Root Entity",
|
|
retrieval_mode=SearchRetrievalMode.VECTOR,
|
|
)
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_vector_mode_rejects_non_text_query(search_service, test_graph):
|
|
"""Vector mode should not silently fall back for title-only queries."""
|
|
with pytest.raises(ValueError):
|
|
await search_service.search(
|
|
SearchQuery(
|
|
title="Root",
|
|
retrieval_mode=SearchRetrievalMode.VECTOR,
|
|
entity_types=[SearchItemType.ENTITY],
|
|
)
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_fts_mode_still_returns_observations(search_service, test_graph):
|
|
"""Explicit FTS mode should preserve existing mixed result behavior."""
|
|
results = await search_service.search(
|
|
SearchQuery(
|
|
text="Root note 1",
|
|
retrieval_mode=SearchRetrievalMode.FTS,
|
|
)
|
|
)
|
|
|
|
assert results
|
|
assert any(result.type == SearchItemType.OBSERVATION.value for result in results)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_vector_sync_skips_embed_opt_out_and_clears_vectors(
|
|
search_service, monkeypatch
|
|
):
|
|
"""Embed opt-out should clear stale vectors instead of regenerating them."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = True
|
|
|
|
monkeypatch.setattr(
|
|
search_service.entity_repository,
|
|
"find_by_id",
|
|
AsyncMock(return_value=SimpleNamespace(id=42, entity_metadata={"embed": False})),
|
|
)
|
|
sync_vectors = AsyncMock()
|
|
delete_entity_vectors = AsyncMock()
|
|
monkeypatch.setattr(repository, "sync_entity_vectors", sync_vectors)
|
|
monkeypatch.setattr(repository, "delete_entity_vector_rows", delete_entity_vectors)
|
|
|
|
await search_service.sync_entity_vectors(42)
|
|
|
|
sync_vectors.assert_not_awaited()
|
|
delete_entity_vectors.assert_awaited_once_with(42)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_vector_sync_resumes_when_embed_opt_out_removed(search_service, monkeypatch):
|
|
"""Removing the opt-out should restore normal embedding sync."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = True
|
|
|
|
monkeypatch.setattr(
|
|
search_service.entity_repository,
|
|
"find_by_id",
|
|
AsyncMock(return_value=SimpleNamespace(id=42, entity_metadata={})),
|
|
)
|
|
sync_vectors = AsyncMock()
|
|
execute_query = AsyncMock()
|
|
monkeypatch.setattr(repository, "sync_entity_vectors", sync_vectors)
|
|
monkeypatch.setattr(repository, "execute_query", execute_query)
|
|
|
|
await search_service.sync_entity_vectors(42)
|
|
|
|
sync_vectors.assert_awaited_once_with(42)
|
|
execute_query.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_vector_sync_batch_skips_embed_opt_out_and_reports_skips(
|
|
search_service, monkeypatch
|
|
):
|
|
"""Batch vector sync should only embed eligible notes and report skipped opt-outs."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = True
|
|
|
|
monkeypatch.setattr(
|
|
search_service.entity_repository,
|
|
"find_by_ids",
|
|
AsyncMock(
|
|
return_value=[
|
|
SimpleNamespace(id=41, entity_metadata={"embed": False}),
|
|
SimpleNamespace(id=42, entity_metadata={}),
|
|
]
|
|
),
|
|
)
|
|
sync_batch = AsyncMock(
|
|
return_value=VectorSyncBatchResult(
|
|
entities_total=1,
|
|
entities_synced=1,
|
|
entities_failed=0,
|
|
)
|
|
)
|
|
delete_entity_vectors = AsyncMock()
|
|
monkeypatch.setattr(repository, "sync_entity_vectors_batch", sync_batch)
|
|
monkeypatch.setattr(repository, "delete_entity_vector_rows", delete_entity_vectors)
|
|
|
|
result = await search_service.sync_entity_vectors_batch([41, 42])
|
|
|
|
sync_batch.assert_awaited_once()
|
|
sync_batch_args = sync_batch.await_args
|
|
assert sync_batch_args is not None
|
|
assert sync_batch_args.args[0] == [42]
|
|
assert result.entities_total == 2
|
|
assert result.entities_synced == 1
|
|
assert result.entities_skipped == 1
|
|
delete_entity_vectors.assert_awaited_once_with(41)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_embed_opt_out_note_still_participates_in_fts(
|
|
search_service, session_maker, test_project
|
|
):
|
|
"""Per-note semantic opt-out should not remove the note from FTS search."""
|
|
entity_repo = EntityRepository(session_maker, project_id=test_project.id)
|
|
entity = await entity_repo.create(
|
|
{
|
|
"title": "FTS Opt Out",
|
|
"note_type": "note",
|
|
"entity_metadata": {"embed": False},
|
|
"content_type": "text/markdown",
|
|
"file_path": "test/fts-opt-out.md",
|
|
"permalink": "test/fts-opt-out",
|
|
"project_id": test_project.id,
|
|
"created_at": datetime.now(),
|
|
"updated_at": datetime.now(),
|
|
}
|
|
)
|
|
|
|
await search_service.index_entity(
|
|
entity,
|
|
content="This note should stay searchable through full text indexing.",
|
|
)
|
|
|
|
results = await search_service.search(
|
|
SearchQuery(
|
|
text="stay searchable",
|
|
retrieval_mode=SearchRetrievalMode.FTS,
|
|
)
|
|
)
|
|
|
|
assert any(result.entity_id == entity.id for result in results)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reindex_vectors_respects_embed_opt_out(search_service, monkeypatch):
|
|
"""Full vector reindex should route through the service-level opt-out filter."""
|
|
monkeypatch.setattr(
|
|
search_service.entity_repository,
|
|
"find_all",
|
|
AsyncMock(
|
|
return_value=[
|
|
SimpleNamespace(id=41, entity_metadata={"embed": False}),
|
|
SimpleNamespace(id=42, entity_metadata={}),
|
|
]
|
|
),
|
|
)
|
|
purge_stale_rows = AsyncMock()
|
|
sync_batch = AsyncMock(
|
|
return_value=VectorSyncBatchResult(
|
|
entities_total=2,
|
|
entities_synced=1,
|
|
entities_failed=0,
|
|
entities_skipped=1,
|
|
)
|
|
)
|
|
monkeypatch.setattr(search_service, "_purge_stale_search_rows", purge_stale_rows)
|
|
monkeypatch.setattr(search_service, "sync_entity_vectors_batch", sync_batch)
|
|
|
|
stats = await search_service.reindex_vectors()
|
|
|
|
purge_stale_rows.assert_awaited_once()
|
|
sync_batch.assert_awaited_once_with([41, 42], progress_callback=None)
|
|
assert stats == {
|
|
"total_entities": 2,
|
|
"embedded": 1,
|
|
"skipped": 1,
|
|
"errors": 0,
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reindex_vectors_purges_sqlite_vectors_before_sync(search_service, monkeypatch):
|
|
"""Regression for #829: stale vec0 rows must be purged before batch sync."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = True
|
|
calls: list[str] = []
|
|
|
|
monkeypatch.setattr(
|
|
search_service.entity_repository,
|
|
"find_all",
|
|
AsyncMock(return_value=[SimpleNamespace(id=42, entity_metadata={})]),
|
|
)
|
|
|
|
async def delete_stale_vector_rows():
|
|
calls.append("purge")
|
|
|
|
async def sync_entity_vectors_batch(entity_ids, progress_callback=None):
|
|
assert entity_ids == [42]
|
|
assert progress_callback is None
|
|
calls.append("sync")
|
|
return VectorSyncBatchResult(entities_total=1, entities_synced=1, entities_failed=0)
|
|
|
|
monkeypatch.setattr(repository, "delete_stale_vector_rows", delete_stale_vector_rows)
|
|
monkeypatch.setattr(search_service, "sync_entity_vectors_batch", sync_entity_vectors_batch)
|
|
|
|
stats = await search_service.reindex_vectors()
|
|
|
|
assert calls == ["purge", "sync"]
|
|
assert stats == {
|
|
"total_entities": 1,
|
|
"embedded": 1,
|
|
"skipped": 0,
|
|
"errors": 0,
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reindex_all_uses_sqlite_vec_aware_drop(search_service, monkeypatch):
|
|
"""Full service reindex should not drop vec0 tables through a raw connection."""
|
|
repository = _sqlite_repo(search_service)
|
|
executed_sql: list[str] = []
|
|
calls: list[str] = []
|
|
|
|
async def execute_query(query, params=None):
|
|
executed_sql.append(str(query))
|
|
|
|
async def drop_vector_tables():
|
|
calls.append("drop_vector_tables")
|
|
|
|
monkeypatch.setattr(repository, "execute_query", execute_query)
|
|
monkeypatch.setattr(repository, "drop_vector_tables", drop_vector_tables)
|
|
monkeypatch.setattr(search_service, "init_search_index", AsyncMock())
|
|
monkeypatch.setattr(search_service.entity_repository, "find_all", AsyncMock(return_value=[]))
|
|
|
|
await search_service.reindex_all()
|
|
|
|
assert calls == ["drop_vector_tables"]
|
|
assert all("search_vector_embeddings" not in sql for sql in executed_sql)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_drop_vector_tables_skips_sqlite_vec_load_for_plain_table(
|
|
search_service,
|
|
monkeypatch,
|
|
):
|
|
"""Plain compatibility tables should not require sqlite-vec to be loaded."""
|
|
repository = _sqlite_repo(search_service)
|
|
await repository.drop_vector_tables()
|
|
|
|
async with db.scoped_session(repository.session_maker) as session:
|
|
await session.execute(
|
|
text("CREATE TABLE search_vector_embeddings (rowid INTEGER PRIMARY KEY)")
|
|
)
|
|
await session.commit()
|
|
|
|
async def fail_if_loaded(_session):
|
|
raise AssertionError("plain tables should not load sqlite-vec")
|
|
|
|
monkeypatch.setattr(repository, "_ensure_sqlite_vec_loaded", fail_if_loaded)
|
|
|
|
await repository.drop_vector_tables()
|
|
|
|
async with db.scoped_session(repository.session_maker) as session:
|
|
result = await session.execute(
|
|
text(
|
|
"SELECT name FROM sqlite_master "
|
|
"WHERE type = 'table' AND name = 'search_vector_embeddings'"
|
|
)
|
|
)
|
|
assert result.scalar() is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reindex_vectors_force_full_clears_project_vectors_before_resync(
|
|
search_service, monkeypatch
|
|
):
|
|
"""Force-full vector reindex should clear derived vectors before batch sync."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = True
|
|
|
|
monkeypatch.setattr(
|
|
search_service.entity_repository,
|
|
"find_all",
|
|
AsyncMock(
|
|
return_value=[
|
|
SimpleNamespace(id=41, entity_metadata={}),
|
|
SimpleNamespace(id=42, entity_metadata={}),
|
|
]
|
|
),
|
|
)
|
|
purge_stale_rows = AsyncMock()
|
|
delete_project_vectors = AsyncMock()
|
|
sync_batch = AsyncMock(
|
|
return_value=VectorSyncBatchResult(
|
|
entities_total=2,
|
|
entities_synced=2,
|
|
entities_failed=0,
|
|
)
|
|
)
|
|
monkeypatch.setattr(search_service, "_purge_stale_search_rows", purge_stale_rows)
|
|
monkeypatch.setattr(repository, "delete_project_vector_rows", delete_project_vectors)
|
|
monkeypatch.setattr(search_service, "sync_entity_vectors_batch", sync_batch)
|
|
|
|
stats = await search_service.reindex_vectors(force_full=True)
|
|
|
|
purge_stale_rows.assert_awaited_once()
|
|
delete_project_vectors.assert_awaited_once()
|
|
sync_batch.assert_awaited_once_with([41, 42], progress_callback=None)
|
|
assert stats == {
|
|
"total_entities": 2,
|
|
"embedded": 2,
|
|
"skipped": 0,
|
|
"errors": 0,
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_semantic_vector_sync_batch_cleans_up_unknown_ids(search_service, monkeypatch):
|
|
"""Deleted entity IDs should still flow through repository cleanup instead of being dropped."""
|
|
repository = _sqlite_repo(search_service)
|
|
repository._semantic_enabled = True
|
|
|
|
monkeypatch.setattr(
|
|
search_service.entity_repository,
|
|
"find_by_ids",
|
|
AsyncMock(return_value=[SimpleNamespace(id=42, entity_metadata={})]),
|
|
)
|
|
sync_batch = AsyncMock(
|
|
side_effect=[
|
|
VectorSyncBatchResult(
|
|
entities_total=1,
|
|
entities_synced=1,
|
|
entities_failed=0,
|
|
entities_skipped=1,
|
|
),
|
|
VectorSyncBatchResult(
|
|
entities_total=1,
|
|
entities_synced=1,
|
|
entities_failed=0,
|
|
),
|
|
]
|
|
)
|
|
monkeypatch.setattr(repository, "sync_entity_vectors_batch", sync_batch)
|
|
progress_callback = AsyncMock()
|
|
|
|
result = await search_service.sync_entity_vectors_batch([41, 42], progress_callback)
|
|
|
|
assert sync_batch.await_count == 2
|
|
called_entity_ids = {tuple(call.args[0]) for call in sync_batch.await_args_list}
|
|
assert called_entity_ids == {(41,), (42,)}
|
|
progress_callback_calls = [
|
|
call
|
|
for call in sync_batch.await_args_list
|
|
if call.kwargs.get("progress_callback") is not None
|
|
]
|
|
assert len(progress_callback_calls) == 1
|
|
assert progress_callback_calls[0].args[0] == [42]
|
|
assert progress_callback_calls[0].kwargs["progress_callback"] is progress_callback
|
|
assert result.entities_total == 2
|
|
assert result.entities_synced == 2
|
|
assert result.entities_failed == 0
|
|
assert result.entities_skipped == 0
|