Skip to content
Open
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
20 changes: 20 additions & 0 deletions src/neo4j_agent_memory/cli/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
40 changes: 33 additions & 7 deletions src/neo4j_agent_memory/embeddings/sentence_transformers.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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,
Expand Down Expand Up @@ -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)
)
Expand All @@ -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)
)
Expand Down
94 changes: 94 additions & 0 deletions tests/unit/embeddings/test_sentence_transformers.py
Original file line number Diff line number Diff line change
@@ -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")
Loading