Skip to content

Commit eb9972f

Browse files
feat: shared model server for cross-process memory deduplication (#333, #335)
Add a standalone model server process that loads the embedding model and reranker ONCE, serving all TrueMemory processes (MCP server, ingest hooks, CLI) over a Unix domain socket. Reduces memory from ~10GB (5 processes x 2GB each) to ~2.5GB (1 server + 5 lightweight clients). - truememory/model_server.py: UDS listener, lazy model loading, idle timeout auto-shutdown, PID lifecycle management - truememory/model_client.py: EmbeddingProxy/RerankerProxy drop-in replacements, auto-start logic, transparent fallback to local loading - Integration: get_model() and get_reranker() use server when available, fall back to local loading when server isn't running (e.g., in tests) - MCP server startup calls ensure_server_running() to launch the server - Set TRUEMEMORY_NO_MODEL_SERVER=1 to force local loading
1 parent 5f9d5ec commit eb9972f

6 files changed

Lines changed: 512 additions & 6 deletions

File tree

pyproject.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,7 @@ all = ["truememory[agentic]"]
4949
[project.scripts]
5050
truememory-mcp = "truememory.mcp_server:main"
5151
truememory-ingest = "truememory.ingest.cli:main"
52+
truememory-model-server = "truememory.model_server:main"
5253

5354
[project.urls]
5455
Homepage = "https://github.com/buildingjoshbetter/TrueMemory"

truememory/mcp_server.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1482,10 +1482,11 @@ def main():
14821482
except Exception:
14831483
pass
14841484

1485-
# Kick off model preloading before entering the event loop. Models
1486-
# load in background threads (~2.5s) while the MCP handshake
1487-
# completes (~1-3s), so the first search arrives with warm models.
1488-
_preload_models()
1485+
# Start shared model server (loads models once for all processes).
1486+
# Falls back to per-process loading if server can't start.
1487+
from truememory.model_client import ensure_server_running
1488+
if not ensure_server_running():
1489+
_preload_models()
14891490

14901491
# Start background backlog drainer — processes queued session
14911492
# transcripts every 60s while the MCP server is alive, respecting

truememory/model_client.py

Lines changed: 207 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,207 @@
1+
"""Client for the shared model server.
2+
3+
Provides drop-in replacements for get_model() and get_reranker() that
4+
route inference to the shared model_server process over a Unix domain
5+
socket. Auto-starts the server on first request if not running.
6+
7+
Falls back to local model loading if the server cannot be reached.
8+
Set TRUEMEMORY_NO_MODEL_SERVER=1 to force local loading.
9+
"""
10+
11+
import logging
12+
import os
13+
import pickle
14+
import socket
15+
import struct
16+
import subprocess
17+
import sys
18+
import time
19+
from pathlib import Path
20+
21+
import numpy as np
22+
23+
log = logging.getLogger(__name__)
24+
25+
_TRUEMEMORY_DIR = Path.home() / ".truememory"
26+
SOCK_PATH = _TRUEMEMORY_DIR / "model.sock"
27+
PID_PATH = _TRUEMEMORY_DIR / "model_server.pid"
28+
29+
_HEADER_FMT = ">I"
30+
_HEADER_SIZE = struct.calcsize(_HEADER_FMT)
31+
32+
_SERVER_START_TIMEOUT = 30.0
33+
_REQUEST_TIMEOUT = 120.0
34+
35+
36+
def _server_is_alive() -> bool:
37+
if not PID_PATH.exists():
38+
return False
39+
try:
40+
pid = int(PID_PATH.read_text().strip())
41+
os.kill(pid, 0)
42+
return True
43+
except (ProcessLookupError, ValueError, OSError):
44+
return False
45+
46+
47+
def _start_server() -> bool:
48+
"""Start the model server as a detached subprocess."""
49+
_TRUEMEMORY_DIR.mkdir(parents=True, exist_ok=True)
50+
51+
if SOCK_PATH.exists() and not _server_is_alive():
52+
SOCK_PATH.unlink(missing_ok=True)
53+
if PID_PATH.exists() and not _server_is_alive():
54+
PID_PATH.unlink(missing_ok=True)
55+
56+
if _server_is_alive():
57+
return True
58+
59+
log.info("Starting model server...")
60+
try:
61+
subprocess.Popen(
62+
[sys.executable, "-m", "truememory.model_server"],
63+
stdout=subprocess.DEVNULL,
64+
stderr=subprocess.DEVNULL,
65+
start_new_session=True,
66+
)
67+
except Exception as e:
68+
log.warning("Failed to start model server: %s", e)
69+
return False
70+
71+
deadline = time.time() + _SERVER_START_TIMEOUT
72+
while time.time() < deadline:
73+
if SOCK_PATH.exists():
74+
time.sleep(0.2)
75+
return True
76+
time.sleep(0.1)
77+
78+
log.warning("Model server did not start within %.0fs", _SERVER_START_TIMEOUT)
79+
return False
80+
81+
82+
def _send_request(request: dict) -> dict:
83+
"""Send a request to the model server and return the response."""
84+
sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
85+
sock.settimeout(_REQUEST_TIMEOUT)
86+
try:
87+
sock.connect(str(SOCK_PATH))
88+
data = pickle.dumps(request, protocol=pickle.HIGHEST_PROTOCOL)
89+
header = struct.pack(_HEADER_FMT, len(data))
90+
sock.sendall(header + data)
91+
92+
resp_header = _recv_exact(sock, _HEADER_SIZE)
93+
if not resp_header:
94+
raise ConnectionError("Server closed connection")
95+
resp_len = struct.unpack(_HEADER_FMT, resp_header)[0]
96+
resp_data = _recv_exact(sock, resp_len)
97+
if not resp_data:
98+
raise ConnectionError("Incomplete response")
99+
return pickle.loads(resp_data)
100+
finally:
101+
sock.close()
102+
103+
104+
def _recv_exact(sock: socket.socket, n: int) -> bytes | None:
105+
buf = bytearray()
106+
while len(buf) < n:
107+
chunk = sock.recv(n - len(buf))
108+
if not chunk:
109+
return None
110+
buf.extend(chunk)
111+
return bytes(buf)
112+
113+
114+
def _request_with_autostart(request: dict) -> dict:
115+
"""Send request, auto-starting server if needed."""
116+
try:
117+
return _send_request(request)
118+
except (ConnectionRefusedError, FileNotFoundError, OSError):
119+
pass
120+
121+
if not _start_server():
122+
raise ConnectionError("Cannot start model server")
123+
124+
return _send_request(request)
125+
126+
127+
class EmbeddingProxy:
128+
"""Drop-in replacement for the embedding model with .encode() method."""
129+
130+
def __init__(self, tier: str = ""):
131+
self._tier = tier
132+
133+
def encode(self, texts, **kwargs) -> np.ndarray:
134+
if isinstance(texts, str):
135+
texts = [texts]
136+
resp = _request_with_autostart({
137+
"op": "embed",
138+
"texts": list(texts),
139+
"tier": self._tier,
140+
})
141+
if not resp.get("ok"):
142+
raise RuntimeError(f"Model server error: {resp.get('error', 'unknown')}")
143+
return resp["vectors"]
144+
145+
146+
class RerankerProxy:
147+
"""Drop-in replacement for CrossEncoder with .predict() method."""
148+
149+
def __init__(self, model_name: str | None = None):
150+
self._model_name = model_name
151+
152+
def predict(self, pairs, **kwargs) -> np.ndarray:
153+
resp = _request_with_autostart({
154+
"op": "rerank",
155+
"pairs": list(pairs),
156+
"model_name": self._model_name,
157+
})
158+
if not resp.get("ok"):
159+
raise RuntimeError(f"Model server error: {resp.get('error', 'unknown')}")
160+
return resp["scores"]
161+
162+
163+
def use_model_server() -> bool:
164+
"""Check if the model server should be used.
165+
166+
Returns True only if:
167+
1. TRUEMEMORY_NO_MODEL_SERVER is not set
168+
2. The server socket exists (server is running)
169+
170+
Processes that want to ensure the server is running should call
171+
ensure_server_running() first (e.g., during MCP server startup).
172+
"""
173+
if os.environ.get("TRUEMEMORY_NO_MODEL_SERVER", "") == "1":
174+
return False
175+
return SOCK_PATH.exists() and _server_is_alive()
176+
177+
178+
def ensure_server_running() -> bool:
179+
"""Start the model server if it's not already running.
180+
181+
Call from MCP server startup or CLI to enable the shared model server.
182+
Returns True if server is running after this call.
183+
"""
184+
if os.environ.get("TRUEMEMORY_NO_MODEL_SERVER", "") == "1":
185+
return False
186+
if _server_is_alive() and SOCK_PATH.exists():
187+
return True
188+
return _start_server()
189+
190+
191+
def get_embedding_proxy(tier: str = "") -> EmbeddingProxy:
192+
"""Get an embedding proxy connected to the model server."""
193+
return EmbeddingProxy(tier=tier)
194+
195+
196+
def get_reranker_proxy(model_name: str | None = None) -> RerankerProxy:
197+
"""Get a reranker proxy connected to the model server."""
198+
return RerankerProxy(model_name=model_name)
199+
200+
201+
def ping() -> bool:
202+
"""Check if model server is reachable."""
203+
try:
204+
resp = _send_request({"op": "ping"})
205+
return resp.get("ok", False)
206+
except Exception:
207+
return False

0 commit comments

Comments
 (0)