diff --git a/src/basic_memory/repository/sqlite_search_repository.py b/src/basic_memory/repository/sqlite_search_repository.py index d00a6e19..2a869783 100644 --- a/src/basic_memory/repository/sqlite_search_repository.py +++ b/src/basic_memory/repository/sqlite_search_repository.py @@ -438,12 +438,17 @@ class SQLiteSearchRepository(SearchRepositoryBase): """Load sqlite-vec extension for the session.""" await self._ensure_sqlite_vec_loaded(session) + # sqlite-vec hard limit for knn k parameter + SQLITE_VEC_MAX_K = 4096 + async def _run_vector_query( self, session: AsyncSession, query_embedding: list[float], candidate_limit: int, ) -> list[dict]: + # Constraint: sqlite-vec enforces k <= 4096 for knn queries + vector_k = min(candidate_limit, self.SQLITE_VEC_MAX_K) query_embedding_json = json.dumps(query_embedding) vector_result = await session.execute( text( @@ -458,12 +463,13 @@ class SQLiteSearchRepository(SearchRepositoryBase): "JOIN search_vector_chunks c ON c.id = vector_matches.rowid " "WHERE c.project_id = :project_id " "ORDER BY best_distance ASC " - "LIMIT :vector_k" + "LIMIT :candidate_limit" ), { "query_embedding": query_embedding_json, "project_id": self.project_id, - "vector_k": candidate_limit, + "vector_k": vector_k, + "candidate_limit": candidate_limit, }, ) return [dict(row) for row in vector_result.mappings().all()] diff --git a/tests/repository/test_sqlite_vector_search_repository.py b/tests/repository/test_sqlite_vector_search_repository.py index e032e2a1..d0177e17 100644 --- a/tests/repository/test_sqlite_vector_search_repository.py +++ b/tests/repository/test_sqlite_vector_search_repository.py @@ -1,6 +1,8 @@ """SQLite sqlite-vec search repository tests.""" +import json from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock import pytest from sqlalchemy import text @@ -264,3 +266,52 @@ async def test_sqlite_hybrid_search_combines_fts_and_vector(search_repository): assert results assert any(result.permalink == "specs/search-index" for result in results) + + +@pytest.mark.asyncio +async def test_run_vector_query_caps_k_at_sqlite_vec_limit(search_repository): + """_run_vector_query must cap the knn k param at SQLITE_VEC_MAX_K (4096). + + sqlite-vec raises OperationalError when k > 4096. The candidate_limit + passed from the base class can exceed this for large projects, so + _run_vector_query clamps k while keeping the outer LIMIT unclamped. + """ + if not isinstance(search_repository, SQLiteSearchRepository): + pytest.skip("sqlite-vec k limit is SQLite-specific.") + + _enable_semantic(search_repository) + await search_repository.init_search_index() + + # Track the parameters passed to session.execute + captured_params: list[dict] = [] + original_execute = None + + async def capturing_execute(stmt, params=None): + if params and "vector_k" in params: + captured_params.append(dict(params)) + # Return empty result set + mock_result = MagicMock() + mock_result.mappings.return_value.all.return_value = [] + return mock_result + + async with db.scoped_session(search_repository.session_maker) as session: + await search_repository._prepare_vector_session(session) + original_execute = session.execute + session.execute = capturing_execute + + query_embedding = [0.1] * search_repository._vector_dimensions + + # candidate_limit exceeds sqlite-vec limit + await search_repository._run_vector_query(session, query_embedding, 10000) + + assert len(captured_params) == 1 + assert captured_params[0]["vector_k"] == SQLiteSearchRepository.SQLITE_VEC_MAX_K + assert captured_params[0]["candidate_limit"] == 10000 + + # candidate_limit within limit should pass through unchanged + captured_params.clear() + await search_repository._run_vector_query(session, query_embedding, 500) + + assert len(captured_params) == 1 + assert captured_params[0]["vector_k"] == 500 + assert captured_params[0]["candidate_limit"] == 500