main
py 189 lines 5.8 KB
Raw
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"]