Reuse preloaded local embedding models
Cache the underlying SentenceTransformer by its effective constructor options so startup preload and first-chat initialization share one instance while wrappers keep their own runtime configuration. Pass the runtime embedding configuration through preload and cover reuse, concurrency, option changes, and failure handling.
Alessandro committed
Aug 26, 2026 at 23:04 UTC
5a96eda0787342ac90265f4f25e7df9236fd07a0
4 files changed
+214
-2
AGENTS.md
+1
@@ -29,6 +29,7 @@
29
- Use Linux paths and commands in examples.
30
- When a live Dockerized Agent Zero target is explicitly named, verify that exact runtime instead of assuming a fixed localhost port.
31
- Message-loop completion flows through a response tool with `break_loop`; plain or malformed Chat Completions text enters repair, and native Responses output text is normalized through the same response-tool path.
32
+- Reuse the startup-preloaded local embedding model for matching runtime configurations; wrappers retain their own rate-limit configuration while sharing the underlying inference model.
33
- Prompt Markdown may retain fenced JSON examples for readability; final system-prompt rendering removes only their JSON fence markers before model calls and preserves non-JSON fences.
34
- Copy live core-plugin changes back into tracked source under `plugins/`.
35
- Develop new custom plugins under ignored `usr/plugins/`; tracked bundled plugins live under `plugins/`.
models.py
+20
-1
@@ -1,7 +1,9 @@
1
from dataclasses import dataclass, field
2
from enum import Enum
3
+import json
4
import logging
5
import os
6
+import threading
7
from typing import (
8
Any,
9
Awaitable,
@@ -828,6 +830,23 @@ class LiteLLMEmbeddingWrapper(Embeddings):
830
return item.get("embedding") if isinstance(item, dict) else item.embedding # type: ignore
831
832
833
+_LOCAL_EMBEDDING_MODELS: dict[tuple[str, str], SentenceTransformer] = {}
834
+_LOCAL_EMBEDDING_MODELS_LOCK = threading.Lock()
835
+
836
+
837
+def _get_local_embedding_model(
838
+ model: str, kwargs: dict[str, Any]
839
+) -> SentenceTransformer:
840
+ key = (model, json.dumps(kwargs, sort_keys=True, default=repr))
841
+ with _LOCAL_EMBEDDING_MODELS_LOCK:
842
+ cached = _LOCAL_EMBEDDING_MODELS.get(key)
843
+ if cached is None:
844
+ cached = SentenceTransformer(model, **kwargs)
845
+ _LOCAL_EMBEDDING_MODELS.clear()
846
+ _LOCAL_EMBEDDING_MODELS[key] = cached
847
+ return cached
848
+
849
+
850
class LocalSentenceTransformerWrapper(Embeddings):
851
"""Local wrapper for sentence-transformers models to avoid HuggingFace API calls"""
852
@@ -856,7 +875,7 @@ class LocalSentenceTransformerWrapper(Embeddings):
875
}
876
st_kwargs = {k: v for k, v in (kwargs or {}).items() if k in st_allowed_keys}
877
859
- self.model = SentenceTransformer(model, **st_kwargs)
878
+ self.model = _get_local_embedding_model(model, st_kwargs)
879
self.model_name = model
880
self.a0_model_conf = model_config
881
preload.py
+4
-1
@@ -25,7 +25,10 @@ async def preload():
25
emb_cfg = get_embedding_model_config_object()
26
if emb_cfg.provider.lower() == "huggingface":
27
emb_mod = models.get_embedding_model(
28
- "huggingface", emb_cfg.name
28
+ emb_cfg.provider,
29
+ emb_cfg.name,
30
+ model_config=emb_cfg,
31
+ **emb_cfg.build_kwargs(),
32
)
33
emb_txt = await emb_mod.aembed_query("test")
34
return emb_txt
tests/test_embedding_model_preload.py
new
+189
@@ -0,0 +1,189 @@
1
+from concurrent.futures import ThreadPoolExecutor
2
+import threading
3
+from types import SimpleNamespace
4
+
5
+import pytest
6
+
7
+import models
8
+
9
+
10
+def _clear_local_embedding_models():
11
+ with models._LOCAL_EMBEDDING_MODELS_LOCK:
12
+ models._LOCAL_EMBEDDING_MODELS.clear()
13
+
14
+
15
+def test_local_embedding_preload_is_reused_with_runtime_model_config(monkeypatch):
16
+ created = []
17
+
18
+ class FakeSentenceTransformer:
19
+ def __init__(self, model, **kwargs):
20
+ created.append((model, kwargs))
21
+
22
+ monkeypatch.setattr(models, "SentenceTransformer", FakeSentenceTransformer)
23
+ _clear_local_embedding_models()
24
+
25
+ try:
26
+ preload = models.LocalSentenceTransformerWrapper(
27
+ "huggingface",
28
+ "sentence-transformers/example",
29
+ device="cpu",
30
+ model_kwargs={"revision": "stable", "trust_remote_code": False},
31
+ )
32
+ runtime_config = SimpleNamespace(name="runtime")
33
+ runtime = models.LocalSentenceTransformerWrapper(
34
+ "huggingface",
35
+ "sentence-transformers/example",
36
+ model_config=runtime_config,
37
+ model_kwargs={"trust_remote_code": False, "revision": "stable"},
38
+ device="cpu",
39
+ )
40
+
41
+ assert runtime.model is preload.model
42
+ assert runtime.a0_model_conf is runtime_config
43
+ assert created == [
44
+ (
45
+ "example",
46
+ {
47
+ "device": "cpu",
48
+ "model_kwargs": {
49
+ "revision": "stable",
50
+ "trust_remote_code": False,
51
+ },
52
+ },
53
+ )
54
+ ]
55
+ finally:
56
+ _clear_local_embedding_models()
57
+
58
+
59
+def test_local_embedding_cache_tracks_effective_constructor_options(monkeypatch):
60
+ created = []
61
+
62
+ class FakeSentenceTransformer:
63
+ def __init__(self, model, **kwargs):
64
+ created.append((model, kwargs))
65
+
66
+ monkeypatch.setattr(models, "SentenceTransformer", FakeSentenceTransformer)
67
+ _clear_local_embedding_models()
68
+
69
+ try:
70
+ first = models.LocalSentenceTransformerWrapper(
71
+ "huggingface", "sentence-transformers/example", device="cpu"
72
+ )
73
+ second = models.LocalSentenceTransformerWrapper(
74
+ "huggingface", "sentence-transformers/example", device="cuda"
75
+ )
76
+
77
+ assert second.model is not first.model
78
+ assert created == [
79
+ ("example", {"device": "cpu"}),
80
+ ("example", {"device": "cuda"}),
81
+ ]
82
+ assert len(models._LOCAL_EMBEDDING_MODELS) == 1
83
+ finally:
84
+ _clear_local_embedding_models()
85
+
86
+
87
+def test_concurrent_preload_and_runtime_share_one_model(monkeypatch):
88
+ created = []
89
+ construction_started = threading.Event()
90
+ release_construction = threading.Event()
91
+
92
+ class FakeSentenceTransformer:
93
+ def __init__(self, model, **kwargs):
94
+ created.append((model, kwargs))
95
+ construction_started.set()
96
+ assert release_construction.wait(timeout=2)
97
+
98
+ monkeypatch.setattr(models, "SentenceTransformer", FakeSentenceTransformer)
99
+ _clear_local_embedding_models()
100
+
101
+ try:
102
+ with ThreadPoolExecutor(max_workers=2) as executor:
103
+ first = executor.submit(
104
+ models.LocalSentenceTransformerWrapper,
105
+ "huggingface",
106
+ "sentence-transformers/example",
107
+ )
108
+ assert construction_started.wait(timeout=2)
109
+ second = executor.submit(
110
+ models.LocalSentenceTransformerWrapper,
111
+ "huggingface",
112
+ "sentence-transformers/example",
113
+ )
114
+ release_construction.set()
115
+
116
+ assert second.result().model is first.result().model
117
+
118
+ assert created == [("example", {})]
119
+ finally:
120
+ release_construction.set()
121
+ _clear_local_embedding_models()
122
+
123
+
124
+def test_failed_model_change_keeps_the_working_cached_model(monkeypatch):
125
+ created = []
126
+
127
+ class FakeSentenceTransformer:
128
+ def __init__(self, model, **kwargs):
129
+ created.append((model, kwargs))
130
+ if model == "broken":
131
+ raise RuntimeError("model unavailable")
132
+
133
+ monkeypatch.setattr(models, "SentenceTransformer", FakeSentenceTransformer)
134
+ _clear_local_embedding_models()
135
+
136
+ try:
137
+ working = models.LocalSentenceTransformerWrapper(
138
+ "huggingface", "sentence-transformers/working"
139
+ )
140
+ with pytest.raises(RuntimeError, match="model unavailable"):
141
+ models.LocalSentenceTransformerWrapper(
142
+ "huggingface", "sentence-transformers/broken"
143
+ )
144
+ reused = models.LocalSentenceTransformerWrapper(
145
+ "huggingface", "sentence-transformers/working"
146
+ )
147
+
148
+ assert reused.model is working.model
149
+ assert created == [("working", {}), ("broken", {})]
150
+ finally:
151
+ _clear_local_embedding_models()
152
+
153
+
154
+@pytest.mark.asyncio
155
+async def test_preload_uses_the_runtime_embedding_configuration(monkeypatch):
156
+ import preload
157
+ from plugins._model_config.helpers import model_config
158
+
159
+ config = SimpleNamespace(
160
+ provider="huggingface",
161
+ name="sentence-transformers/example",
162
+ build_kwargs=lambda: {"device": "cpu"},
163
+ )
164
+ calls = []
165
+ embedded = []
166
+
167
+ class FakeEmbeddings:
168
+ async def aembed_query(self, text):
169
+ embedded.append(text)
170
+
171
+ def get_embedding_model(provider, name, **kwargs):
172
+ calls.append((provider, name, kwargs))
173
+ return FakeEmbeddings()
174
+
175
+ monkeypatch.setattr(
176
+ model_config, "get_embedding_model_config_object", lambda: config
177
+ )
178
+ monkeypatch.setattr(preload.models, "get_embedding_model", get_embedding_model)
179
+
180
+ await preload.preload()
181
+
182
+ assert calls == [
183
+ (
184
+ "huggingface",
185
+ "sentence-transformers/example",
186
+ {"model_config": config, "device": "cpu"},
187
+ )
188
+ ]
189
+ assert embedded == ["test"]