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
5 changes: 5 additions & 0 deletions engraphis/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from engraphis.inspector.cloud_mount import CLOUD_PREFIXES, mount_cloud_endpoints
from engraphis.config import settings
from engraphis.engines import reweight, thoughts as thoughts_engine
from engraphis.engines.embedder import warmup as _warmup_embedder
from engraphis.logging_setup import configure_logging
from engraphis.netutil import client_ip
from engraphis.routes.memory import router as memory_router
Expand Down Expand Up @@ -59,6 +60,10 @@ async def _lifespan(app: FastAPI):
cancel and await the loop."""
global _background_task
init_db()
# Warm the embedding model eagerly so the first recall call isn't paid
# under request pressure (a cold load + concurrent call used to wedge
# the forked PM2 worker and time out every recall).
await asyncio.get_running_loop().run_in_executor(None, _warmup_embedder)
if settings.loop_interval > 0:
_background_task = asyncio.create_task(_consciousness_loop())
logger.info("Background consciousness loop started (interval=%ds)", settings.loop_interval)
Expand Down
25 changes: 24 additions & 1 deletion engraphis/engines/embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

import logging
import re
import threading
from typing import Optional

import numpy as np
Expand All @@ -17,11 +18,20 @@

_model = None
_dim: Optional[int] = None
# Guards the lazy model load. Without this, concurrent recall calls each
# see `_model is None`, and the forked PM2 worker tries to load the
# 80-400 MB model N times at once — the process wedges and every
# recall atop it times out. One lock + one load for the process lifetime.
_lock = threading.Lock()


def _get_model():
global _model, _dim
if _model is None:
if _model is not None:
return _model
with _lock:
if _model is not None: # double-checked: another thread won the race
return _model
from sentence_transformers import SentenceTransformer

logger.info("Loading embedding model: %s", settings.embed_model)
Expand All @@ -34,6 +44,19 @@ def _get_model():
return _model


def warmup():
"""Load the model eagerly so the first recall call isn't paid under request pressure.

Safe to call from startup; returns True on success. Never raises.
"""
try:
_get_model()
return True
except Exception as exc: # pragma: no cover - defensive; model may be missing
logger.warning("Embedder warmup failed (%s): %s", type(exc).__name__, exc)
return False


def embed_dim() -> int:
"""Return the embedding dimension (loads model if needed)."""
if _dim is None:
Expand Down
98 changes: 98 additions & 0 deletions tests/test_embedder_threading.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,98 @@
"""Tests for embedder thread-safety (double-checked locking) and warmup()."""
import threading
from unittest.mock import MagicMock, patch

import pytest


@pytest.fixture(autouse=True)
def _reset_embedder_state():
"""Reset module-level globals so each test starts clean."""
import engraphis.engines.embedder as emb_mod
old_model, old_dim = emb_mod._model, emb_mod._dim
emb_mod._model = None
emb_mod._dim = None
yield
emb_mod._model = old_model
emb_mod._dim = old_dim


def test_concurrent_get_model_loads_only_once():
"""N threads racing on _get_model() must trigger exactly one model load."""
import engraphis.engines.embedder as emb_mod

load_count = 0
load_lock = threading.Lock()

fake_model = MagicMock()
fake_model.get_embedding_dimension.return_value = 384

def slow_load(*args, **kwargs):
nonlocal load_count
with load_lock:
load_count += 1
# Simulate the 80-400 MB model load taking real time
threading.Event().wait(0.05)
return fake_model

with patch("engraphis.engines.embedder.SentenceTransformer", side_effect=slow_load, create=True):
with patch.dict("sys.modules", {"sentence_transformers": MagicMock(SentenceTransformer=slow_load)}):
barrier = threading.Barrier(8)
results = [None] * 8

def worker(idx):
barrier.wait() # all threads start simultaneously
results[idx] = emb_mod._get_model()

threads = [threading.Thread(target=worker, args=(i,)) for i in range(8)]
for t in threads:
t.start()
for t in threads:
t.join(timeout=10)

assert load_count == 1, f"Model loaded {load_count} times; expected exactly 1"
assert all(r is fake_model for r in results)


def test_warmup_returns_true_on_success():
"""warmup() returns True when the model loads successfully."""
import engraphis.engines.embedder as emb_mod

fake_model = MagicMock()
fake_model.get_embedding_dimension.return_value = 384

with patch.dict("sys.modules", {"sentence_transformers": MagicMock(SentenceTransformer=lambda *a, **kw: fake_model)}):
assert emb_mod.warmup() is True
assert emb_mod._model is fake_model


def test_warmup_returns_false_on_failure():
"""warmup() returns False (never raises) when model loading fails."""
import engraphis.engines.embedder as emb_mod

def exploding_load(*args, **kwargs):
raise RuntimeError("model file missing")

with patch.dict("sys.modules", {"sentence_transformers": MagicMock(SentenceTransformer=exploding_load)}):
assert emb_mod.warmup() is False
assert emb_mod._model is None


def test_warmup_idempotent():
"""Calling warmup() twice loads the model only once."""
import engraphis.engines.embedder as emb_mod

load_count = 0
fake_model = MagicMock()
fake_model.get_embedding_dimension.return_value = 384

def counting_load(*args, **kwargs):
nonlocal load_count
load_count += 1
return fake_model

with patch.dict("sys.modules", {"sentence_transformers": MagicMock(SentenceTransformer=counting_load)}):
assert emb_mod.warmup() is True
assert emb_mod.warmup() is True

assert load_count == 1