|
27 | 27 | import hmac |
28 | 28 | import json |
29 | 29 | import logging |
| 30 | +import math |
| 31 | +import os |
30 | 32 | import re |
31 | | -import time |
32 | 33 | import secrets |
33 | | -import math |
| 34 | +import sys |
| 35 | +import threading |
| 36 | +import time |
34 | 37 |
|
35 | 38 | from collections import OrderedDict |
36 | 39 | from dataclasses import dataclass |
@@ -103,22 +106,27 @@ def set_service(svc: MemoryService) -> None: |
103 | 106 | _service = svc |
104 | 107 |
|
105 | 108 |
|
| 109 | +_service_lock = threading.Lock() |
| 110 | + |
| 111 | + |
106 | 112 | def service() -> MemoryService: |
107 | 113 | """Lazily build the service so server startup is instant (model loads on first use).""" |
108 | 114 | global _service |
109 | 115 | 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 | + ) |
122 | 130 | return _service |
123 | 131 |
|
124 | 132 |
|
@@ -3033,10 +3041,59 @@ def _eager_exact_backend_check() -> None: |
3033 | 3041 | service() |
3034 | 3042 |
|
3035 | 3043 |
|
| 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 | + |
3036 | 3091 | def main() -> None: |
3037 | 3092 | """Console entry point (``engraphis-mcp``). Runs Smart MCP over stdio.""" |
3038 | 3093 | _eager_exact_backend_check() |
3039 | | - mcp.run() |
| 3094 | + _start_background_warmup() |
| 3095 | + import anyio |
| 3096 | + anyio.run(lambda: _safe_run_stdio_async(mcp)) |
3040 | 3097 |
|
3041 | 3098 |
|
3042 | 3099 | if __name__ == "__main__": |
|
0 commit comments