Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions src/tenantq/benchmark.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,8 +152,8 @@ def run_benchmark(
idx_recall.append(recall_at_k(retrieved, ref, 10))

total_time = sum(latencies) / 1000.0
avg_r5 = sum(r5) / len(r5)
avg_r10 = sum(r10) / len(r10)
avg_r5 = (sum(r5) / len(r5)) if r5 else 0.0
avg_r10 = (sum(r10) / len(r10)) if r10 else 0.0
RECALL_GAUGE.labels(mode=mode, k="5").set(avg_r5)
RECALL_GAUGE.labels(mode=mode, k="10").set(avg_r10)
result.modes.append(
Expand Down
2 changes: 2 additions & 0 deletions src/tenantq/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -123,5 +123,7 @@ def build_embedder(kind: str, dense_model: str, sparse_model: str, dense_dim: in


def batched(seq: Sequence, size: int) -> Iterable[Sequence]:
if size < 1:
raise ValueError("batch_size must be >= 1")
for i in range(0, len(seq), size):
yield seq[i : i + size]
2 changes: 2 additions & 0 deletions src/tenantq/ingest.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,8 @@ def ingest_documents(
parallelism: int = 4,
) -> IngestReport:
"""Embed and upsert ``documents`` in parallel batches."""
if batch_size < 1:
raise ValueError("batch_size must be >= 1")
start = time.perf_counter()
batches = list(batched(list(documents), batch_size))

Expand Down
23 changes: 23 additions & 0 deletions tests/test_benchmark.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,3 +31,26 @@ def test_benchmark_produces_numbers(ingested, settings, embedder, dataset):
assert not math.isnan(modes["dense"].index_recall_at_10)
# hybrid should be at least as good as the weaker single mode on recall@10
assert modes["hybrid"].recall_at_10 >= min(modes["dense"].recall_at_10, modes["sparse"].recall_at_10)


def test_run_benchmark_empty_queries_returns_empty_modes(settings):
from unittest.mock import MagicMock

from tenantq.data import Dataset

embedder = MagicMock()
embedder.embed_dense.return_value = []
result = run_benchmark(
MagicMock(),
settings,
embedder,
Dataset(documents=[], queries=[]),
modes=("dense",),
)
assert result.settings_summary["n_queries"] == 0
assert len(result.modes) == 1
mode = result.modes[0]
assert mode.n_queries == 0
assert mode.recall_at_5 == 0.0
assert mode.recall_at_10 == 0.0
assert mode.qps == 0.0
28 changes: 28 additions & 0 deletions tests/test_ingest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
"""Ingestion input validation."""

from __future__ import annotations

from unittest.mock import MagicMock

import pytest

from tenantq.embeddings import batched
from tenantq.ingest import ingest_documents


@pytest.mark.parametrize("batch_size", [0, -1])
def test_ingest_documents_rejects_batch_size_below_one(settings, batch_size):
with pytest.raises(ValueError, match="batch_size must be >= 1"):
ingest_documents(
MagicMock(),
settings,
MagicMock(),
documents=[],
batch_size=batch_size,
)


@pytest.mark.parametrize("size", [0, -3])
def test_batched_rejects_non_positive_size(size):
with pytest.raises(ValueError, match="batch_size must be >= 1"):
list(batched([1, 2, 3], size))
Loading