Skip to content

Commit c9c9ca2

Browse files
fix(mcp): isolate stdio wire protocol, add cache-first embedder fast path and background warmup (#194)
* fix(mcp): isolate stdio wire protocol, add cache-first embedder fast path and background warmup - Redirect sys.stdout to sys.stderr in stdio MCP mode so library prints/warnings never corrupt the JSON-RPC wire - Load cached SentenceTransformer models with local_files_only=True first to eliminate network roundtrips and timeouts - Make service() singleton initialization thread-safe and warm up in a background daemon thread - Add embedder health check to 'engraphis-init --check' and provide 'engraphis-init --prefetch' command * test(backends): mock sentence_transformers via sys.modules for offline test suites
1 parent a15ed41 commit c9c9ca2

6 files changed

Lines changed: 221 additions & 15 deletions

File tree

‎engraphis/backends/embedder_st.py‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -319,6 +319,18 @@ def get_embedder(
319319
}
320320
if local_files_only:
321321
factory_kwargs["local_files_only"] = True
322+
else:
323+
# Optimize cold start: if the model is already in local Hugging Face cache,
324+
# loading with local_files_only=True avoids network roundtrips to huggingface.co.
325+
# If cached, it returns immediately; if not, it seamlessly falls through to download.
326+
try:
327+
cached_kwargs = dict(factory_kwargs)
328+
cached_kwargs["local_files_only"] = True
329+
emb = SentenceTransformerEmbedder(resolved_model_name, **cached_kwargs)
330+
LAST_EMBEDDER_ERROR = ""
331+
return emb
332+
except Exception:
333+
pass
322334
emb = SentenceTransformerEmbedder(resolved_model_name, **factory_kwargs)
323335
LAST_EMBEDDER_ERROR = ""
324336
return emb

‎engraphis/mcp_server.py‎

Lines changed: 72 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -27,10 +27,13 @@
2727
import hmac
2828
import json
2929
import logging
30+
import math
31+
import os
3032
import re
31-
import time
3233
import secrets
33-
import math
34+
import sys
35+
import threading
36+
import time
3437

3538
from collections import OrderedDict
3639
from dataclasses import dataclass
@@ -103,22 +106,27 @@ def set_service(svc: MemoryService) -> None:
103106
_service = svc
104107

105108

109+
_service_lock = threading.Lock()
110+
111+
106112
def service() -> MemoryService:
107113
"""Lazily build the service so server startup is instant (model loads on first use)."""
108114
global _service
109115
if _service is None:
110-
_service = MemoryService.create(
111-
settings.db_path,
112-
embed_model=settings.embed_model or None,
113-
embed_revision=getattr(settings, "embed_revision", "") or None,
114-
require_immutable_models=bool(getattr(settings, "require_immutable_models", False)),
115-
require_exact_backends=bool(getattr(settings, "require_exact_backends", False)),
116-
embed_dim=settings.embed_dim if settings.embed_dim is not None else 384,
117-
vector_backend=settings.vector_backend,
118-
rerank_model=getattr(settings, "rerank_model", "") or None,
119-
rerank_revision=getattr(settings, "rerank_revision", "") or None,
120-
extractor=settings.extractor,
121-
)
116+
with _service_lock:
117+
if _service is None:
118+
_service = MemoryService.create(
119+
settings.db_path,
120+
embed_model=settings.embed_model or None,
121+
embed_revision=getattr(settings, "embed_revision", "") or None,
122+
require_immutable_models=bool(getattr(settings, "require_immutable_models", False)),
123+
require_exact_backends=bool(getattr(settings, "require_exact_backends", False)),
124+
embed_dim=settings.embed_dim if settings.embed_dim is not None else 384,
125+
vector_backend=settings.vector_backend,
126+
rerank_model=getattr(settings, "rerank_model", "") or None,
127+
rerank_revision=getattr(settings, "rerank_revision", "") or None,
128+
extractor=settings.extractor,
129+
)
122130
return _service
123131

124132

@@ -3033,10 +3041,59 @@ def _eager_exact_backend_check() -> None:
30333041
service()
30343042

30353043

3044+
def _start_background_warmup() -> None:
3045+
"""Warm up the memory service in a background daemon thread.
3046+
3047+
Allows the initial MCP handshake (initialize, tools/list) to respond in
3048+
milliseconds while warming SQLite and the embedding model before the agent's
3049+
first tool invocation. Can be disabled via ENGRAPHIS_MCP_WARMUP=0.
3050+
"""
3051+
warmup_env = os.environ.get("ENGRAPHIS_MCP_WARMUP", "1").strip().lower()
3052+
if warmup_env in {"0", "false", "no", "off"}:
3053+
return
3054+
thread = threading.Thread(target=service, name="engraphis-warmup", daemon=True)
3055+
thread.start()
3056+
3057+
3058+
async def _safe_run_stdio_async(server: FastMCP) -> None:
3059+
"""Run stdio transport with pure wire protocol isolation.
3060+
3061+
In stdio MCP, standard output is exclusively the JSON-RPC wire. Redirect
3062+
Python's global `sys.stdout` to `sys.stderr` so that any prints, warnings,
3063+
or dependency output (PyTorch, transformers, tqdm, pydantic) flow safely to
3064+
stderr without corrupting JSON-RPC messages on the client pipe.
3065+
"""
3066+
import anyio
3067+
from io import TextIOWrapper
3068+
from mcp.server.stdio import stdio_server
3069+
3070+
real_stdout_buffer = getattr(sys.stdout, "buffer", None)
3071+
real_stdin_buffer = getattr(sys.stdin, "buffer", None)
3072+
if real_stdout_buffer is not None and real_stdin_buffer is not None:
3073+
sys.stdout = sys.stderr
3074+
wrapped_stdout = anyio.wrap_file(TextIOWrapper(real_stdout_buffer, encoding="utf-8"))
3075+
wrapped_stdin = anyio.wrap_file(TextIOWrapper(real_stdin_buffer, encoding="utf-8", errors="replace"))
3076+
async with stdio_server(stdin=wrapped_stdin, stdout=wrapped_stdout) as (read_stream, write_stream):
3077+
await server._mcp_server.run(
3078+
read_stream,
3079+
write_stream,
3080+
server._mcp_server.create_initialization_options(),
3081+
)
3082+
else:
3083+
async with stdio_server() as (read_stream, write_stream):
3084+
await server._mcp_server.run(
3085+
read_stream,
3086+
write_stream,
3087+
server._mcp_server.create_initialization_options(),
3088+
)
3089+
3090+
30363091
def main() -> None:
30373092
"""Console entry point (``engraphis-mcp``). Runs Smart MCP over stdio."""
30383093
_eager_exact_backend_check()
3039-
mcp.run()
3094+
_start_background_warmup()
3095+
import anyio
3096+
anyio.run(lambda: _safe_run_stdio_async(mcp))
30403097

30413098

30423099
if __name__ == "__main__":

‎scripts/init.py‎

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
engraphis-init --encrypted # require SQLCipher and provision a private DB key file
1313
engraphis-init --force # overwrite the trusted config file
1414
engraphis-init --check # doctor: verify install, extras, DB writability
15+
engraphis-init --prefetch # pre-cache embedding model weights for instant MCP startup
1516
1617
Non-interactive by design (no prompts): safe in scripts, CI, and agent shells.
1718
"""
@@ -104,10 +105,38 @@ def cmd_check() -> int:
104105
except Exception:
105106
_miss("Engraphis Cloud", "saved session unavailable; reconnect if needed")
106107

108+
try:
109+
from engraphis.backends.embedder_st import get_embedder
110+
emb = get_embedder(settings.embed_model or None, dim=settings.embed_dim or 384)
111+
emb.embed(["engraphis doctor check"])
112+
_ok("embedder functional", f"{type(emb).__name__} ({getattr(emb, 'dim', 384)}d)")
113+
except Exception as exc:
114+
_fail("embedder functional", f"{type(exc).__name__}: {exc}")
115+
failures += 1
116+
107117
print("all good" if failures == 0 else f"{failures} problem(s) found")
108118
return 0 if failures == 0 else 1
109119

110120

121+
def cmd_prefetch() -> int:
122+
"""Download and warm up the configured embedding model ahead of time."""
123+
from engraphis.config import settings
124+
model_name = (settings.embed_model or "").strip()
125+
if not model_name:
126+
print(" [--] No remote embedding model configured; deterministic offline embedder is active.")
127+
return 0
128+
print(f"engraphis prefetch - model '{model_name}'")
129+
try:
130+
from engraphis.backends.embedder_st import get_embedder
131+
emb = get_embedder(model_name, dim=settings.embed_dim or 384, require_exact=True)
132+
emb.embed(["engraphis prefetch warmup"])
133+
_ok("model prefetch", f"{type(emb).__name__} ({getattr(emb, 'dim', 384)}d) ready")
134+
return 0
135+
except Exception as exc:
136+
_fail("model prefetch", f"{type(exc).__name__}: {exc}")
137+
return 1
138+
139+
111140
def _env_content(db_path: Path, token: str, key_path: Optional[Path] = None) -> str:
112141
lines = [
113142
"# Engraphis - generated by engraphis-init. Full reference: .env.example",
@@ -243,10 +272,14 @@ def main(argv=None) -> int:
243272
)
244273
ap.add_argument("--check", action="store_true",
245274
help="doctor mode: verify the installation without writing config")
275+
ap.add_argument("--prefetch", action="store_true",
276+
help="pre-cache the configured embedding model for instant MCP startup")
246277
args = ap.parse_args(argv)
247278

248279
if args.check:
249280
return cmd_check()
281+
if args.prefetch:
282+
return cmd_prefetch()
250283

251284
db_path = Path(args.db).expanduser().resolve()
252285
try:

‎tests/test_backends_factories.py‎

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -115,6 +115,52 @@ def __init__(
115115
}
116116

117117

118+
def test_get_embedder_prefers_local_cache(monkeypatch):
119+
calls = []
120+
121+
class FakeST:
122+
def __init__(self, model_name, **kwargs):
123+
calls.append((model_name, kwargs))
124+
125+
def get_embedding_dimension(self):
126+
return 384
127+
128+
monkeypatch.setitem(
129+
sys.modules,
130+
"sentence_transformers",
131+
SimpleNamespace(SentenceTransformer=FakeST),
132+
)
133+
emb = get_embedder("sentence-transformers/all-MiniLM-L6-v2", 384)
134+
assert emb.dim == 384
135+
assert len(calls) == 1
136+
assert calls[0][1].get("local_files_only") is True
137+
138+
139+
def test_get_embedder_falls_back_when_not_cached(monkeypatch):
140+
calls = []
141+
142+
class FakeST:
143+
def __init__(self, model_name, **kwargs):
144+
calls.append((model_name, kwargs))
145+
if kwargs.get("local_files_only") is True:
146+
raise OSError("not in cache")
147+
148+
def get_embedding_dimension(self):
149+
return 384
150+
151+
monkeypatch.setitem(
152+
sys.modules,
153+
"sentence_transformers",
154+
SimpleNamespace(SentenceTransformer=FakeST),
155+
)
156+
emb = get_embedder("sentence-transformers/all-MiniLM-L6-v2", 384)
157+
assert emb.dim == 384
158+
assert len(calls) == 2
159+
assert calls[0][1].get("local_files_only") is True
160+
assert not calls[1][1].get("local_files_only")
161+
162+
163+
118164
@pytest.mark.parametrize("revision", [None, "main", "A" * 40, "a" * 39])
119165
def test_embedder_strict_mode_rejects_mutable_remote_revision_before_load(monkeypatch, revision):
120166
import engraphis.backends.embedder_st as embedder_st

‎tests/test_init.py‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -229,3 +229,22 @@ def test_doctor_reports_connected_cloud_install(tmp_path, monkeypatch, capsys):
229229
assert main(["--check"]) == 0
230230
out = capsys.readouterr().out
231231
assert "Engraphis Cloud - installation connected" in out
232+
233+
234+
def test_doctor_reports_functional_embedder(tmp_path, monkeypatch, capsys):
235+
_fresh_settings(monkeypatch, tmp_path)
236+
assert main(["--check"]) == 0
237+
out = capsys.readouterr().out
238+
assert "embedder functional" in out
239+
240+
241+
def test_prefetch_command_reports_ready_or_offline(tmp_path, monkeypatch, capsys):
242+
_fresh_settings(monkeypatch, tmp_path)
243+
# Test prefetch with offline deterministic model
244+
monkeypatch.setenv("ENGRAPHIS_EMBED_MODEL", "")
245+
import engraphis.config as cfg
246+
monkeypatch.setattr(cfg, "settings", cfg.Settings())
247+
assert main(["--prefetch"]) == 0
248+
out = capsys.readouterr().out
249+
assert "deterministic offline embedder is active" in out
250+

‎tests/test_mcp_server.py‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1267,3 +1267,42 @@ def test_classic_remember_persists_subject_key_and_claim_kind_to_chain(monkeypat
12671267
assert head["id"] == payload["id"]
12681268
assert head["subject_key"] == "deploy.timeout"
12691269
assert head["claim_kind"] == "configured_value"
1270+
1271+
1272+
def test_service_singleton_is_thread_safe():
1273+
import engraphis.mcp_server as srv
1274+
from concurrent.futures import ThreadPoolExecutor
1275+
# Concurrently call srv.service() across 5 threads
1276+
with ThreadPoolExecutor(max_workers=5) as executor:
1277+
instances = list(executor.map(lambda _: srv.service(), range(5)))
1278+
assert len(instances) == 5
1279+
for inst in instances[1:]:
1280+
assert inst is instances[0]
1281+
1282+
1283+
def test_background_warmup_honors_env(monkeypatch):
1284+
import engraphis.mcp_server as server
1285+
called = []
1286+
monkeypatch.setattr(server, "service", lambda: called.append(True))
1287+
1288+
# Disabled by env
1289+
monkeypatch.setenv("ENGRAPHIS_MCP_WARMUP", "0")
1290+
server._start_background_warmup()
1291+
assert not called
1292+
1293+
# Enabled by default or env
1294+
started_threads = []
1295+
real_thread = server.threading.Thread
1296+
1297+
def fake_thread(*args, **kwargs):
1298+
t = real_thread(*args, **kwargs)
1299+
started_threads.append(t)
1300+
return t
1301+
1302+
monkeypatch.setattr(server.threading, "Thread", fake_thread)
1303+
monkeypatch.setenv("ENGRAPHIS_MCP_WARMUP", "1")
1304+
server._start_background_warmup()
1305+
assert len(started_threads) == 1
1306+
assert started_threads[0].daemon is True
1307+
assert started_threads[0].name == "engraphis-warmup"
1308+

0 commit comments

Comments
 (0)