personal memory agent
1# SPDX-License-Identifier: AGPL-3.0-only
2# Copyright (c) 2026 sol pbc
3
4from __future__ import annotations
5
6import sys
7import types
8from pathlib import Path
9from typing import Any
10
11import numpy as np
12import pytest
13
14from solstone.think.indexer import rerank_scorer
15from solstone.think.providers import rerank_install
16
17MODEL_PATH = "onnx/model.onnx"
18TOKENIZER_PATH = "tokenizer.json"
19FIXTURE_REPO = "Xenova/ms-marco-MiniLM-L-6-v2"
20FIXTURE_REVISION = "a09144355adeed5f58c8ed011d209bf8ee5a1fec"
21
22
23@pytest.fixture(autouse=True)
24def _reset_scorer_state():
25 _reset_scorer()
26 yield
27 _reset_scorer()
28
29
30def _reset_scorer() -> None:
31 rerank_scorer._session = None
32 rerank_scorer._tokenizer = None
33 rerank_scorer._np = None
34 rerank_scorer._disabled = False
35
36
37def _stage_fixture_assets(
38 tmp_path: Path, monkeypatch: pytest.MonkeyPatch
39) -> rerank_install.RerankModelSpec:
40 monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path))
41 model_path = (
42 tmp_path / "cache" / "providers" / "rerank" / FIXTURE_REVISION / MODEL_PATH
43 )
44 tokenizer_path = (
45 tmp_path / "cache" / "providers" / "rerank" / FIXTURE_REVISION / TOKENIZER_PATH
46 )
47 model_path.parent.mkdir(parents=True, exist_ok=True)
48 tokenizer_path.parent.mkdir(parents=True, exist_ok=True)
49 model_path.write_bytes(b"stub-model")
50 tokenizer_path.write_bytes(b"stub-tokenizer")
51 spec = rerank_install.RerankModelSpec(
52 repo=FIXTURE_REPO,
53 revision=FIXTURE_REVISION,
54 files=(
55 rerank_install.RerankFileSpec(
56 path=MODEL_PATH,
57 sha256="0" * 64,
58 size_bytes=model_path.stat().st_size,
59 ),
60 rerank_install.RerankFileSpec(
61 path=TOKENIZER_PATH,
62 sha256="1" * 64,
63 size_bytes=tokenizer_path.stat().st_size,
64 ),
65 ),
66 )
67 monkeypatch.setattr(rerank_install, "RERANK_MODEL_SPEC", spec)
68 return spec
69
70
71class _Input:
72 def __init__(self, name: str) -> None:
73 self.name = name
74
75
76class _Encoding:
77 ids = [2, 4, 3, 5, 3]
78 attention_mask = [1, 1, 1, 1, 1]
79 type_ids = [0, 0, 0, 1, 1]
80
81
82class _FakeTokenizer:
83 def enable_truncation(self, *, max_length):
84 self.max_length = max_length
85
86 def enable_padding(self):
87 self.padding_enabled = True
88
89 def encode_batch(self, pairs):
90 return [_Encoding() for _pair in pairs]
91
92
93def test_missing_assets_fail_closed_and_latch(tmp_path, monkeypatch) -> None:
94 monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path))
95 spec = rerank_install.RerankModelSpec(
96 repo=FIXTURE_REPO,
97 revision=FIXTURE_REVISION,
98 files=(
99 rerank_install.RerankFileSpec(
100 path=MODEL_PATH,
101 sha256="0" * 64,
102 size_bytes=1,
103 ),
104 rerank_install.RerankFileSpec(
105 path=TOKENIZER_PATH,
106 sha256="1" * 64,
107 size_bytes=1,
108 ),
109 ),
110 )
111 monkeypatch.setattr(rerank_install, "RERANK_MODEL_SPEC", spec)
112 calls = 0
113 original_load = rerank_scorer._load
114
115 def spy_load():
116 nonlocal calls
117 calls += 1
118 return original_load()
119
120 monkeypatch.setattr(rerank_scorer, "_load", spy_load)
121
122 assert rerank_scorer.score("query", ["doc"]) is None
123 assert rerank_scorer.score("query", ["doc"]) is None
124 assert calls == 1
125 assert rerank_scorer._disabled is True
126
127
128def test_inference_exception_fails_closed_and_latches(monkeypatch) -> None:
129 class FailingSession:
130 def get_inputs(self):
131 return [_Input("input_ids")]
132
133 def run(self, *_args: Any, **_kwargs: Any):
134 raise RuntimeError("inference broke")
135
136 calls = 0
137
138 def fake_load():
139 nonlocal calls
140 calls += 1
141 return _FakeTokenizer(), FailingSession(), np
142
143 monkeypatch.setattr(rerank_scorer, "_load", fake_load)
144
145 assert rerank_scorer.score("query", ["doc"]) is None
146 assert rerank_scorer.score("query", ["doc"]) is None
147 assert calls == 1
148 assert rerank_scorer._disabled is True
149
150
151def test_unexpected_output_shape_fails_closed_and_latches(monkeypatch) -> None:
152 class BadShapeSession:
153 def get_inputs(self):
154 return [_Input("input_ids")]
155
156 def run(self, _outputs: Any, feed: dict[str, Any]):
157 return [np.zeros((feed["input_ids"].shape[0], 2), dtype=np.float32)]
158
159 monkeypatch.setattr(
160 rerank_scorer,
161 "_load",
162 lambda: (_FakeTokenizer(), BadShapeSession(), np),
163 )
164
165 assert rerank_scorer.score("query", ["doc"]) is None
166 assert rerank_scorer.score("query", ["doc"]) is None
167 assert rerank_scorer._disabled is True
168
169
170def test_scoring_path_never_calls_installer_or_downloader(
171 tmp_path, monkeypatch
172) -> None:
173 _stage_fixture_assets(tmp_path, monkeypatch)
174
175 class SessionOptions:
176 pass
177
178 class StubSession:
179 def get_inputs(self):
180 return [_Input("input_ids")]
181
182 def run(self, _outputs: Any, feed: dict[str, Any]):
183 batch = feed["input_ids"].shape[0]
184 return [np.arange(batch, dtype=np.float32).reshape(-1, 1)]
185
186 class RuntimeTokenizer:
187 @staticmethod
188 def from_file(_path: str):
189 return _FakeTokenizer()
190
191 fake_ort = types.ModuleType("onnxruntime")
192 fake_ort.SessionOptions = SessionOptions
193 fake_ort.InferenceSession = lambda *_args, **_kwargs: StubSession()
194 fake_tok = types.ModuleType("tokenizers")
195 fake_tok.Tokenizer = RuntimeTokenizer
196 monkeypatch.setitem(sys.modules, "onnxruntime", fake_ort)
197 monkeypatch.setitem(sys.modules, "tokenizers", fake_tok)
198 monkeypatch.setattr(
199 rerank_install,
200 "_download_file",
201 lambda *_args, **_kwargs: pytest.fail("scoring should not download"),
202 )
203 monkeypatch.setattr(
204 rerank_install,
205 "install_rerank_model",
206 lambda *_args, **_kwargs: pytest.fail("scoring should not install"),
207 )
208
209 assert rerank_scorer.score("query one", ["doc two"]) is not None
210 _reset_scorer()
211 missing_spec = rerank_install.RerankModelSpec(
212 repo=FIXTURE_REPO,
213 revision="missing-revision",
214 files=rerank_install.RERANK_MODEL_SPEC.files,
215 )
216 monkeypatch.setattr(rerank_install, "RERANK_MODEL_SPEC", missing_spec)
217 assert rerank_scorer.score("query", ["doc"]) is None
218
219
220def test_mocked_session_batches_and_feeds_declared_subset(monkeypatch) -> None:
221 class RecordingSession:
222 def __init__(self) -> None:
223 self.feeds: list[dict[str, Any]] = []
224
225 def get_inputs(self):
226 return [_Input("input_ids")]
227
228 def run(self, _outputs: Any, feed: dict[str, Any]):
229 self.feeds.append(feed)
230 assert set(feed) == {"input_ids"}
231 assert feed["input_ids"].dtype == np.int64
232 batch = feed["input_ids"].shape[0]
233 return [np.arange(batch, dtype=np.float32).reshape(batch, 1)]
234
235 session = RecordingSession()
236 monkeypatch.setattr(
237 rerank_scorer,
238 "_load",
239 lambda: (_FakeTokenizer(), session, np),
240 )
241
242 result = rerank_scorer.score("query", [f"doc {i}" for i in range(17)])
243
244 assert result is not None
245 assert len(result) == 17
246 assert [feed["input_ids"].shape[0] for feed in session.feeds] == [16, 1]