194 lines
5.6 KiB
Python
194 lines
5.6 KiB
Python
"""Tests for LanceDBVectorStore query projection behavior."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Any
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
import pytest
|
|
|
|
from src.vector_store.lancedb import LanceDBVectorStore
|
|
|
|
|
|
def _build_query_chain(rows: list[dict[str, Any]]) -> MagicMock:
|
|
"""Build a chainable mock that mirrors LanceDB's async query builder."""
|
|
chain = MagicMock()
|
|
chain.distance_type.return_value = chain
|
|
chain.limit.return_value = chain
|
|
chain.select.return_value = chain
|
|
chain.where.return_value = chain
|
|
chain.to_list = AsyncMock(return_value=rows)
|
|
return chain
|
|
|
|
|
|
def _patch_table(
|
|
store: LanceDBVectorStore, rows: list[dict[str, Any]]
|
|
) -> tuple[MagicMock, MagicMock]:
|
|
"""Patch _get_table to return a mock whose vector_search yields the chain."""
|
|
chain = _build_query_chain(rows)
|
|
table = MagicMock()
|
|
table.vector_search = MagicMock(return_value=chain)
|
|
store._get_table = AsyncMock(return_value=table) # pyright: ignore[reportPrivateUsage]
|
|
return table, chain
|
|
|
|
|
|
@pytest.fixture
|
|
def store() -> LanceDBVectorStore:
|
|
return LanceDBVectorStore()
|
|
|
|
|
|
def test_build_where_clause_membership(store: LanceDBVectorStore) -> None:
|
|
"""Both the dict `in` form and the bare-list sugar produce an IN clause."""
|
|
assert (
|
|
store._build_where_clause({"session_name": {"in": ["s1", "s2"]}}) # pyright: ignore[reportPrivateUsage]
|
|
== "session_name IN ('s1', 's2')"
|
|
)
|
|
assert (
|
|
store._build_where_clause({"session_name": ["s1", "s2"]}) # pyright: ignore[reportPrivateUsage]
|
|
== "session_name IN ('s1', 's2')"
|
|
)
|
|
|
|
|
|
def test_build_where_clause_empty_membership_fails_closed(
|
|
store: LanceDBVectorStore,
|
|
) -> None:
|
|
"""An empty membership list must emit an always-false predicate, never an
|
|
omitted condition that would widen scope (fail-open)."""
|
|
assert store._build_where_clause({"session_name": {"in": []}}) == "1 = 0" # pyright: ignore[reportPrivateUsage]
|
|
assert store._build_where_clause({"session_name": []}) == "1 = 0" # pyright: ignore[reportPrivateUsage]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_query_returns_empty_when_table_missing(
|
|
store: LanceDBVectorStore,
|
|
) -> None:
|
|
store._get_table = AsyncMock(return_value=None) # pyright: ignore[reportPrivateUsage]
|
|
|
|
results = await store.query("honcho.msg.missing", [0.1, 0.2, 0.3, 0.4])
|
|
|
|
assert results == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_query_default_does_not_project(store: LanceDBVectorStore) -> None:
|
|
_table, chain = _patch_table(store, rows=[])
|
|
|
|
await store.query("honcho.msg.test", [0.1, 0.2, 0.3, 0.4])
|
|
|
|
chain.select.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_query_with_include_attributes_false_selects_only_id(
|
|
store: LanceDBVectorStore,
|
|
) -> None:
|
|
_table, chain = _patch_table(store, rows=[])
|
|
|
|
await store.query(
|
|
"honcho.doc.test",
|
|
[0.1, 0.2, 0.3, 0.4],
|
|
include_attributes=False,
|
|
)
|
|
|
|
chain.select.assert_called_once_with(["id"])
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_query_with_attribute_list_projects_id_plus_listed(
|
|
store: LanceDBVectorStore,
|
|
) -> None:
|
|
_table, chain = _patch_table(store, rows=[])
|
|
|
|
await store.query(
|
|
"honcho.msg.test",
|
|
[0.1, 0.2, 0.3, 0.4],
|
|
include_attributes=["message_id"],
|
|
)
|
|
|
|
chain.select.assert_called_once_with(["id", "message_id"])
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_query_attribute_list_dedupes_explicit_id(
|
|
store: LanceDBVectorStore,
|
|
) -> None:
|
|
_table, chain = _patch_table(store, rows=[])
|
|
|
|
await store.query(
|
|
"honcho.msg.test",
|
|
[0.1, 0.2, 0.3, 0.4],
|
|
include_attributes=["id", "message_id"],
|
|
)
|
|
|
|
chain.select.assert_called_once_with(["id", "message_id"])
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_query_converts_rows_to_results_with_score_and_metadata(
|
|
store: LanceDBVectorStore,
|
|
) -> None:
|
|
rows: list[dict[str, Any]] = [
|
|
{
|
|
"id": "vec_1",
|
|
"_distance": 0.12,
|
|
"vector": [0.0, 0.0, 0.0, 0.0],
|
|
"message_id": "msg_1",
|
|
"session_name": "sess_a",
|
|
},
|
|
{
|
|
"id": "vec_2",
|
|
"_distance": 0.34,
|
|
"message_id": "msg_2",
|
|
},
|
|
]
|
|
_patch_table(store, rows=rows)
|
|
|
|
results = await store.query("honcho.msg.test", [0.1, 0.2, 0.3, 0.4])
|
|
|
|
assert [r.id for r in results] == ["vec_1", "vec_2"]
|
|
assert [r.score for r in results] == [0.12, 0.34]
|
|
# id, vector, _distance must not leak into metadata
|
|
assert results[0].metadata == {
|
|
"message_id": "msg_1",
|
|
"session_name": "sess_a",
|
|
}
|
|
assert results[1].metadata == {"message_id": "msg_2"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_query_filters_by_max_distance(store: LanceDBVectorStore) -> None:
|
|
rows: list[dict[str, Any]] = [
|
|
{"id": "vec_close", "_distance": 0.05, "message_id": "msg_1"},
|
|
{"id": "vec_far", "_distance": 0.9, "message_id": "msg_2"},
|
|
]
|
|
_patch_table(store, rows=rows)
|
|
|
|
results = await store.query(
|
|
"honcho.msg.test",
|
|
[0.1, 0.2, 0.3, 0.4],
|
|
max_distance=0.5,
|
|
)
|
|
|
|
assert [r.id for r in results] == ["vec_close"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_query_returns_empty_without_opening_table_on_nonpositive_top_k(
|
|
store: LanceDBVectorStore,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
get_table = AsyncMock(return_value=MagicMock())
|
|
monkeypatch.setattr(store, "_get_table", get_table)
|
|
|
|
for top_k in (0, -1):
|
|
assert (
|
|
await store.query(
|
|
"honcho.msg.test",
|
|
[0.1, 0.2, 0.3, 0.4],
|
|
top_k=top_k,
|
|
)
|
|
== []
|
|
)
|
|
|
|
get_table.assert_not_awaited()
|