personal memory agent
0

Configure Feed

Select the types of activity you want to include in your feed.

solstone / tests / test_rerank_scorer.py
7.4 kB 246 lines
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]