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
11 changes: 7 additions & 4 deletions server/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,14 +11,16 @@

from server.config import settings
from server.embeddings import close_embedding_provider, get_embedding_provider
from server.embeddings.bm25 import BM25SparseProvider, close_sparse_embedding_provider
from server.embeddings.bm25 import (
close_sparse_embedding_provider,
get_sparse_embedding_provider,
)
from server.state import (
get_commit_store,
get_service_registry,
get_store,
set_commit_store,
set_service_registry,
set_sparse_provider,
set_store,
)
from server.store.commit_store import CommitStore
Expand Down Expand Up @@ -47,8 +49,9 @@ async def lifespan(_: MCPServer) -> AsyncIterator[None]:
await commit_store.ensure_collection()
set_commit_store(commit_store)

sparse_provider = BM25SparseProvider()
set_sparse_provider(sparse_provider)
# Warm the single BM25 model at startup so the first query does not pay the
# load cost; indexing and search share this one instance.
get_sparse_embedding_provider()

set_service_registry(ServiceRegistry())

Expand Down
13 changes: 0 additions & 13 deletions server/state.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,14 +3,12 @@
import asyncio
from collections import defaultdict

from server.embeddings.bm25 import BM25SparseProvider
from server.store.commit_store import CommitStore
from server.store.qdrant import QdrantStore
from server.store.service_registry import ServiceRegistry

_store: QdrantStore | None = None
_commit_store: CommitStore | None = None
_sparse_provider: BM25SparseProvider | None = None
_service_registry: ServiceRegistry | None = None
_reindex_locks: defaultdict[str, asyncio.Lock] = defaultdict(asyncio.Lock)

Expand All @@ -37,17 +35,6 @@ def set_commit_store(store: CommitStore) -> None:
_commit_store = store


def get_sparse_provider() -> BM25SparseProvider:
if _sparse_provider is None:
raise RuntimeError("Sparse embedding provider not initialized")
return _sparse_provider


def set_sparse_provider(provider: BM25SparseProvider) -> None:
global _sparse_provider
_sparse_provider = provider


def get_service_registry() -> ServiceRegistry:
if _service_registry is None:
raise RuntimeError("Service registry not initialized")
Expand Down
7 changes: 4 additions & 3 deletions server/tools/search.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,9 @@

from server.config import settings
from server.embeddings import get_embedding_provider
from server.embeddings.bm25 import get_sparse_embedding_provider
from server.indexer.github_source import fetch_file_content
from server.state import get_service_registry, get_sparse_provider, get_store
from server.state import get_service_registry, get_store
from server.store.service_registry import load_effective_services
from server.tools.file_cache import BlobContentCache

Expand Down Expand Up @@ -44,7 +45,7 @@ async def search_code(
limit: Maximum number of results (default 10)
"""
embedder = get_embedding_provider()
sparse_embedder = get_sparse_provider()
sparse_embedder = get_sparse_embedding_provider()
store = get_store()

dense_vector = await embedder.embed_query(query)
Expand Down Expand Up @@ -147,7 +148,7 @@ async def find_usages(
limit: Maximum number of results (default 10)
"""
embedder = get_embedding_provider()
sparse_embedder = get_sparse_provider()
sparse_embedder = get_sparse_embedding_provider()
store = get_store()

query = f"code that uses or references {symbol_name}"
Expand Down
49 changes: 49 additions & 0 deletions tests/embeddings/test_bm25_singleton.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,49 @@
from __future__ import annotations

import pytest

import server.embeddings.bm25 as bm25_module
import server.indexer.pipeline as pipeline_module
import server.tools.search as search_module
from server.embeddings.bm25 import (
close_sparse_embedding_provider,
get_sparse_embedding_provider,
)


class _StubBm25:
def __init__(self, model_name: str) -> None:
self.model_name = model_name


@pytest.fixture(autouse=True)
def stub_model(monkeypatch):
"""Keep the tests off the real fastembed model download."""
monkeypatch.setattr(bm25_module, "Bm25", _StubBm25)
monkeypatch.setattr(bm25_module, "_provider", None)


def test_get_sparse_embedding_provider_returns_singleton() -> None:
assert get_sparse_embedding_provider() is get_sparse_embedding_provider()


def test_indexing_and_search_share_one_instance() -> None:
"""Guards against reintroducing a second BM25 holder (issue #70): the
pipeline and the search tools must resolve to the same model."""
assert (
pipeline_module.get_sparse_embedding_provider
is search_module.get_sparse_embedding_provider
)
assert (
pipeline_module.get_sparse_embedding_provider()
is search_module.get_sparse_embedding_provider()
)


async def test_close_releases_the_only_instance() -> None:
provider = get_sparse_embedding_provider()

await close_sparse_embedding_provider()

assert bm25_module._provider is None
assert get_sparse_embedding_provider() is not provider
14 changes: 7 additions & 7 deletions tests/tools/test_search.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ async def test_search_code_reports_no_results() -> None:

with (
patch("server.tools.search.get_embedding_provider") as mock_embedder,
patch("server.tools.search.get_sparse_provider") as mock_sparse,
patch("server.tools.search.get_sparse_embedding_provider") as mock_sparse,
patch("server.tools.search.get_store", return_value=store),
):
mock_embedder.return_value.embed_query = AsyncMock(return_value=[0.1])
Expand Down Expand Up @@ -53,7 +53,7 @@ async def test_search_code_formats_hits() -> None:

with (
patch("server.tools.search.get_embedding_provider") as mock_embedder,
patch("server.tools.search.get_sparse_provider") as mock_sparse,
patch("server.tools.search.get_sparse_embedding_provider") as mock_sparse,
patch("server.tools.search.get_store", return_value=store),
):
mock_embedder.return_value.embed_query = AsyncMock(return_value=[0.1])
Expand All @@ -74,7 +74,7 @@ async def test_search_code_passes_chunk_tier_to_store() -> None:

with (
patch("server.tools.search.get_embedding_provider") as mock_embedder,
patch("server.tools.search.get_sparse_provider") as mock_sparse,
patch("server.tools.search.get_sparse_embedding_provider") as mock_sparse,
patch("server.tools.search.get_store", return_value=store),
):
mock_embedder.return_value.embed_query = AsyncMock(return_value=[0.1])
Expand Down Expand Up @@ -169,7 +169,7 @@ async def test_find_usages_excludes_hits_matching_symbol_itself() -> None:

with (
patch("server.tools.search.get_embedding_provider") as mock_embedder,
patch("server.tools.search.get_sparse_provider") as mock_sparse,
patch("server.tools.search.get_sparse_embedding_provider") as mock_sparse,
patch("server.tools.search.get_store", return_value=store),
):
mock_embedder.return_value.embed_query = AsyncMock(return_value=[0.1])
Expand All @@ -187,7 +187,7 @@ async def test_find_usages_over_fetches_to_absorb_self_match_filtering() -> None

with (
patch("server.tools.search.get_embedding_provider") as mock_embedder,
patch("server.tools.search.get_sparse_provider") as mock_sparse,
patch("server.tools.search.get_sparse_embedding_provider") as mock_sparse,
patch("server.tools.search.get_store", return_value=store),
):
mock_embedder.return_value.embed_query = AsyncMock(return_value=[0.1])
Expand Down Expand Up @@ -227,7 +227,7 @@ async def test_find_usages_still_returns_limit_results_when_definition_ranks_fir

with (
patch("server.tools.search.get_embedding_provider") as mock_embedder,
patch("server.tools.search.get_sparse_provider") as mock_sparse,
patch("server.tools.search.get_sparse_embedding_provider") as mock_sparse,
patch("server.tools.search.get_store", return_value=store),
):
mock_embedder.return_value.embed_query = AsyncMock(return_value=[0.1])
Expand Down Expand Up @@ -257,7 +257,7 @@ async def test_find_usages_snippet_windows_around_match() -> None:

with (
patch("server.tools.search.get_embedding_provider") as mock_embedder,
patch("server.tools.search.get_sparse_provider") as mock_sparse,
patch("server.tools.search.get_sparse_embedding_provider") as mock_sparse,
patch("server.tools.search.get_store", return_value=store),
):
mock_embedder.return_value.embed_query = AsyncMock(return_value=[0.1])
Expand Down
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.