gitrag-free-mcp / packages /core /tests /test_persistence_libsql.py
jacklalane's picture
Tune hybrid retrieval weights and rerank instruction
4971dc1 verified
Raw
History Blame Contribute Delete
5.63 kB
import pytest
from unittest.mock import MagicMock, patch
from sqlalchemy import text
from Persistence.Persist import DBConfig, LibsqlConfig, create_persistence_adapter, PersistInLibsql
from Chunker.Chunk import Chunk
@pytest.fixture
def mock_engine():
engine = MagicMock()
conn = MagicMock()
engine.begin.return_value.__enter__.return_value = conn
engine.connect.return_value.__enter__.return_value = conn
return engine, conn
def test_libsql_config_from_parts():
cfg = LibsqlConfig.from_parts(database_url="libsql://test", auth_token="token", table="my_chunks")
assert cfg.provider == "libsql"
assert cfg.url == "libsql://test"
assert cfg.auth_token == "token"
assert cfg.table == "my_chunks"
assert cfg.fts_table == "my_chunks_fts"
assert cfg.database_url == "libsql://test"
assert cfg.resolved_fts_table == "my_chunks_fts"
def test_libsql_persist_batch(mock_engine):
engine, conn = mock_engine
cfg = LibsqlConfig.from_parts(database_url="libsql://test")
# Mocking _bootstrap to avoid DDL execution
with patch.object(PersistInLibsql, "_bootstrap"):
persist = PersistInLibsql(cfg=cfg, engine=engine)
chunk = Chunk(
chunk="test content",
repo="repo",
path="file.py",
language="python",
start_rc=(1, 0),
end_rc=(2, 0),
start_bytes=0,
end_bytes=12,
embeddings=b"0" * 4096
)
persist.persist_batch([chunk])
# Verify execute called for insert and FTS refresh
# 1. conn.execute(insert_stmt, params)
# 2. conn.execute(delete_fts, params)
# 3. conn.execute(insert_fts, params)
assert conn.execute.call_count >= 3
def test_libsql_delete_batch(mock_engine):
engine, conn = mock_engine
cfg = LibsqlConfig.from_parts(database_url="libsql://test")
with patch.object(PersistInLibsql, "_bootstrap"):
persist = PersistInLibsql(cfg=cfg, engine=engine)
persist.delete_batch(["path1", "path2"], repo="repo")
assert conn.execute.call_count == 2
def test_libsql_close(mock_engine):
engine, conn = mock_engine
cfg = LibsqlConfig.from_parts(database_url="libsql://test")
with patch.object(PersistInLibsql, "_bootstrap"):
persist = PersistInLibsql(cfg=cfg, engine=engine)
persist.close()
# engine.dispose should NOT be called because engine was passed in (owns_engine=False)
assert engine.dispose.call_count == 0
def test_create_persistence_adapter_unsupported():
cfg = DBConfig(provider="invalid", url="...")
with pytest.raises(ValueError, match="Unsupported persistence adapter"):
create_persistence_adapter("invalid", cfg=cfg)
def test_libsql_invalid_provider():
cfg = DBConfig(provider="postgres", url="...")
with pytest.raises(TypeError, match="cfg.provider must be 'libsql'"):
PersistInLibsql(cfg=cfg)
def test_libsql_persist_batch_missing_embeddings(mock_engine):
engine, conn = mock_engine
cfg = LibsqlConfig.from_parts(database_url="libsql://test")
with patch.object(PersistInLibsql, "_bootstrap"):
persist = PersistInLibsql(cfg=cfg, engine=engine)
chunk = Chunk(chunk="c", repo="r", path="p", language="l", start_rc=(0, 0), end_rc=(0, 0), start_bytes=0,
end_bytes=0)
with pytest.raises(ValueError, match="missing embeddings"):
persist.persist_batch([chunk])
def test_libsql_ignores_read_only_bootstrap_error(mock_engine):
engine, _ = mock_engine
cfg = LibsqlConfig.from_parts(database_url="libsql://test")
with patch.object(PersistInLibsql, "_bootstrap", side_effect=ValueError("write operations are forbidden")):
PersistInLibsql(cfg=cfg, engine=engine)
def test_libsql_search_uses_weighted_hybrid_scores(mock_engine):
engine, _ = mock_engine
cfg = LibsqlConfig.from_parts(database_url="libsql://test")
with patch.object(PersistInLibsql, "_bootstrap"):
persist = PersistInLibsql(cfg=cfg, engine=engine)
with patch.object(persist, "_keyword_search", return_value=[
{"id": "broad", "path": "USongs.pas", "repo": "r", "branch": "b", "language": "pascal", "start_row": 1,
"start_col": 0, "end_row": 2, "end_col": 0, "start_bytes": 0, "end_bytes": 10, "chunk": "song list", "embedding": None,
"keyword_score": 10.0, "vector_score": 0.0},
{"id": "impl", "path": "USong.pas", "repo": "r", "branch": "b", "language": "pascal", "start_row": 3,
"start_col": 0, "end_row": 4, "end_col": 0, "start_bytes": 11, "end_bytes": 40, "chunk": "LoadSong ReadTXTHeader", "embedding": None,
"keyword_score": 6.0, "vector_score": 0.0},
]), patch.object(persist, "_vector_search", return_value=[
{"id": "impl", "path": "USong.pas", "repo": "r", "branch": "b", "language": "pascal", "start_row": 3,
"start_col": 0, "end_row": 4, "end_col": 0, "start_bytes": 11, "end_bytes": 40, "chunk": "LoadSong ReadTXTHeader", "embedding": None,
"keyword_score": 0.0, "vector_score": 0.95},
{"id": "broad", "path": "USongs.pas", "repo": "r", "branch": "b", "language": "pascal", "start_row": 1,
"start_col": 0, "end_row": 2, "end_col": 0, "start_bytes": 0, "end_bytes": 10, "chunk": "song list", "embedding": None,
"keyword_score": 0.0, "vector_score": 0.20},
]):
results = persist.search(query_embedding=[1.0, 0.0], query_text="where is the song file parsed", limit=2, repo="r", branch="b")
assert [chunk.path for chunk in results] == ["USong.pas", "USongs.pas"]