diff --git a/src/neo4j_agent_memory/cli/main.py b/src/neo4j_agent_memory/cli/main.py index 5ba2b858..e2a0a5c4 100644 --- a/src/neo4j_agent_memory/cli/main.py +++ b/src/neo4j_agent_memory/cli/main.py @@ -919,6 +919,26 @@ def mcp_serve( ) sys.exit(1) + # Windows deadlock guard: when a local sentence-transformers embedder is + # configured, import its native stack (scipy/torch) up-front on the main + # thread, before the event loop starts. Importing it lazily inside the + # running loop on the first embed() call deadlocks in the Windows DLL loader + # lock (see SentenceTransformerEmbedder / preload()). No-op for cloud + # embedders and when sentence-transformers isn't installed. + if embedding: + from neo4j_agent_memory.embeddings.sentence_transformers import preload + from neo4j_agent_memory.llm.adapters.sentence_transformers import ( + SentenceTransformersProvider, + ) + from neo4j_agent_memory.llm.factory import from_provider + + try: + resolved_embedder = from_provider(embedding, kind="embedding") + except Exception: + resolved_embedder = None + if isinstance(resolved_embedder, SentenceTransformersProvider): + preload() + asyncio.run( run_server( neo4j_uri=uri, diff --git a/src/neo4j_agent_memory/embeddings/sentence_transformers.py b/src/neo4j_agent_memory/embeddings/sentence_transformers.py index f2efdd7d..11553819 100644 --- a/src/neo4j_agent_memory/embeddings/sentence_transformers.py +++ b/src/neo4j_agent_memory/embeddings/sentence_transformers.py @@ -1,4 +1,12 @@ -"""Sentence Transformers embedding provider for local embeddings.""" +"""Sentence Transformers embedding provider for local embeddings. + +Note on Windows: importing ``sentence_transformers`` (which pulls in scipy's and +torch's native extensions) for the first time *inside a running asyncio event +loop* deadlocks in the Windows DLL loader lock — regardless of which thread runs +the import — hanging the very first ``embed`` call forever. Startup code should +call :func:`preload` from synchronous context before the event loop starts to +do that native import safely on the main thread. +""" from __future__ import annotations @@ -12,6 +20,19 @@ from sentence_transformers import SentenceTransformer +def preload() -> None: + """Eagerly import the sentence-transformers stack on the current thread. + + Call this from synchronous startup code *before* any asyncio event loop is + running (see the module docstring for why). No-op if sentence-transformers + is not installed. + """ + import importlib.util + + if importlib.util.find_spec("sentence_transformers") is not None: + importlib.import_module("sentence_transformers") + + # Model dimensions mapping (common models) MODEL_DIMENSIONS = { "all-MiniLM-L6-v2": 384, @@ -69,11 +90,16 @@ def dimensions(self) -> int: async def embed(self, text: str) -> list[float]: """Generate embedding for a single text.""" - model = self._ensure_model() - try: - # Run in thread pool since sentence-transformers is sync + # Run in thread pool since sentence-transformers is sync. + # _ensure_model() must ALSO run in the executor, not on the event + # loop thread: the first call lazily imports sentence_transformers + # (-> scipy/torch native extensions) and loads the model. Doing that + # on the loop thread blocks the whole server for the load; on Windows + # the first-time native import even deadlocks in the DLL loader lock + # (see the module docstring / preload()). loop = asyncio.get_event_loop() + model = await loop.run_in_executor(None, self._ensure_model) embedding = await loop.run_in_executor( None, lambda: model.encode(text, convert_to_numpy=True) ) @@ -86,11 +112,11 @@ async def embed_batch(self, texts: list[str]) -> list[list[float]]: if not texts: return [] - model = self._ensure_model() - try: - # Run in thread pool since sentence-transformers is sync + # Run in thread pool since sentence-transformers is sync. + # _ensure_model() offloaded too — see the note in embed() above. loop = asyncio.get_event_loop() + model = await loop.run_in_executor(None, self._ensure_model) embeddings = await loop.run_in_executor( None, lambda: model.encode(texts, convert_to_numpy=True) ) diff --git a/tests/unit/embeddings/test_sentence_transformers.py b/tests/unit/embeddings/test_sentence_transformers.py new file mode 100644 index 00000000..81fd2fa3 --- /dev/null +++ b/tests/unit/embeddings/test_sentence_transformers.py @@ -0,0 +1,94 @@ +"""Unit tests for the local sentence-transformers embedding provider.""" + +import threading +from typing import Any +from unittest.mock import patch + +from neo4j_agent_memory.embeddings.sentence_transformers import ( + SentenceTransformerEmbedder, + preload, +) + + +class _FakeEncoding: + def tolist(self) -> list[float]: + return [0.1, 0.2, 0.3] + + +class _FakeModel: + def encode( + self, text_or_texts: Any, *args: Any, **kwargs: Any + ) -> _FakeEncoding | list[_FakeEncoding]: + # Mirror sentence-transformers: a batch (list input) yields one encoding + # per text; a single string yields a single encoding. + if isinstance(text_or_texts, (list, tuple)): + return [_FakeEncoding() for _ in text_or_texts] + return _FakeEncoding() + + def get_sentence_embedding_dimension(self) -> int: + return 3 + + +class TestSentenceTransformerEmbedderThreading: + """The heavy model load must never run on the asyncio event-loop thread.""" + + async def test_embed_loads_model_off_the_event_loop_thread(self) -> None: + """Regression guard for the Windows first-call hang. + + The first ``embed`` call lazily imports ``sentence_transformers`` + (-> scipy native libs) via ``_ensure_model``. Doing that import on the + event-loop thread deadlocks in the Windows DLL loader lock, so + ``_ensure_model`` must be offloaded to a worker thread. + """ + embedder = SentenceTransformerEmbedder("all-MiniLM-L6-v2") + loop_thread_id = threading.get_ident() + seen: dict[str, int] = {} + + def fake_ensure() -> _FakeModel: + seen["thread_id"] = threading.get_ident() + return _FakeModel() + + with patch.object(embedder, "_ensure_model", side_effect=fake_ensure): + result = await embedder.embed("hello") + + assert result == [0.1, 0.2, 0.3] + assert seen["thread_id"] != loop_thread_id + + async def test_embed_batch_loads_model_off_the_event_loop_thread(self) -> None: + """``embed_batch`` offloads ``_ensure_model`` for the same reason.""" + embedder = SentenceTransformerEmbedder("all-MiniLM-L6-v2") + loop_thread_id = threading.get_ident() + seen: dict[str, int] = {} + + def fake_ensure() -> _FakeModel: + seen["thread_id"] = threading.get_ident() + return _FakeModel() + + with patch.object(embedder, "_ensure_model", side_effect=fake_ensure): + result = await embedder.embed_batch(["a", "b"]) + + assert result == [[0.1, 0.2, 0.3], [0.1, 0.2, 0.3]] + assert seen["thread_id"] != loop_thread_id + + +class TestPreload: + """``preload`` imports the native stack up-front, or no-ops if absent.""" + + def test_preload_is_noop_when_not_installed(self) -> None: + with ( + patch("importlib.util.find_spec", return_value=None) as find_spec, + patch("importlib.import_module") as import_module, + ): + preload() + + find_spec.assert_called_once_with("sentence_transformers") + import_module.assert_not_called() + + def test_preload_imports_when_installed(self) -> None: + with ( + patch("importlib.util.find_spec", return_value=object()), + patch("importlib.import_module") as import_module, + ): + preload() + + import_module.assert_called_once_with("sentence_transformers")