personal memory agent
1# SPDX-License-Identifier: AGPL-3.0-only
2# Copyright (c) 2026 sol pbc
3
4import asyncio
5import logging
6from collections import OrderedDict
7from unittest.mock import Mock
8
9import httpx
10import pytest
11
12import solstone.think.supervisor as mod
13from solstone.think.providers import local_server
14from solstone.think.providers.shared import classify_provider_error
15
16
17@pytest.fixture(autouse=True)
18def isolate_supervisor_wedge_state(monkeypatch):
19 monkeypatch.setattr(
20 mod,
21 "_wedge_state",
22 {
23 "providers": OrderedDict(),
24 "failures": set(),
25 "cooldown_until": 0.0,
26 "awaiting_recovery": False,
27 },
28 )
29 monkeypatch.setattr(
30 mod,
31 "_recovery_state",
32 {
33 "local": mod.ProviderRecoveryState(),
34 "parakeet": mod.ProviderRecoveryState(),
35 },
36 )
37 monkeypatch.setattr(mod, "_managed_procs", [])
38 monkeypatch.setattr(mod, "_SERVICE_STATE", {})
39 monkeypatch.setattr(mod, "_RESTART_POLICIES", {})
40 monkeypatch.setattr(mod, "_is_remote_mode", False)
41 monkeypatch.setattr(mod, "shutdown_requested", False)
42 monkeypatch.setattr(mod, "_supervisor_callosum", None)
43
44
45class _ProcessStub:
46 def __init__(self, returncode: int | None = None):
47 self.returncode = returncode
48 self.pid = 12345
49
50 def poll(self):
51 return self.returncode
52
53
54class _ManagedStub:
55 def __init__(self, name: str, cmd: list[str], returncode: int | None = None):
56 self.name = name
57 self.cmd = cmd
58 self.process = _ProcessStub(returncode)
59 self.ref = f"{name}-ref"
60 self.cleanup = Mock()
61
62
63def _start(use_id: str, provider: str = "local") -> dict:
64 return {
65 "tract": "cortex",
66 "event": "start",
67 "use_id": use_id,
68 "provider": provider,
69 }
70
71
72def _error(use_id: str, reason_code: str | None = "provider_unavailable") -> dict:
73 message = {
74 "tract": "cortex",
75 "event": "error",
76 "use_id": use_id,
77 "error": "generation failed",
78 }
79 if reason_code is not None:
80 message["reason_code"] = reason_code
81 return message
82
83
84def _finish(use_id: str) -> dict:
85 return {
86 "tract": "cortex",
87 "event": "finish",
88 "use_id": use_id,
89 "result": {"ok": True},
90 }
91
92
93def _drive_wedge(handler=mod._handle_cortex_outcome, prefix: str = "fail") -> None:
94 for idx in range(mod.LOCAL_WEDGE_THRESHOLD):
95 use_id = f"{prefix}-{idx}"
96 handler(_start(use_id))
97 handler(_error(use_id))
98
99
100def _ready_local_server(monkeypatch, port: int = 9999) -> None:
101 monkeypatch.setattr(mod, "read_service_port", Mock(return_value=port))
102 monkeypatch.setattr(
103 local_server,
104 "_probe_health",
105 Mock(return_value=(local_server.STATE_READY, None)),
106 )
107
108
109def test_remote_mode_ignores_cortex_events_without_state_or_io(monkeypatch):
110 class FailingState(dict):
111 def __getitem__(self, key):
112 raise AssertionError("remote mode should not read wedge state")
113
114 monkeypatch.setattr(mod, "_is_remote_mode", True)
115 monkeypatch.setattr(mod, "_wedge_state", FailingState())
116 monkeypatch.setattr(
117 mod,
118 "read_service_port",
119 Mock(side_effect=AssertionError("should not read service port")),
120 )
121 monkeypatch.setattr(
122 mod,
123 "_request_provider_runtime_recycle",
124 Mock(side_effect=AssertionError("should not request recycle")),
125 )
126
127 mod._handle_cortex_outcome(_start("u1"))
128 mod._handle_cortex_outcome(_error("u1"))
129 mod._handle_cortex_outcome(_finish("u1"))
130
131
132def test_unknown_use_id_terminal_does_not_count_or_reset():
133 mod._handle_cortex_outcome(_start("known"))
134 mod._handle_cortex_outcome(_error("known"))
135
136 mod._handle_cortex_outcome(_finish("unknown"))
137 mod._handle_cortex_outcome(_error("missing"))
138
139 assert mod._wedge_state["failures"] == {"known"}
140
141
142def test_local_finish_resets_failure_counter_by_use_id_attribution():
143 mod._handle_cortex_outcome(_start("u1"))
144 mod._handle_cortex_outcome(_error("u1"))
145 mod._handle_cortex_outcome(_start("u2"))
146 mod._handle_cortex_outcome(_error("u2"))
147 mod._handle_cortex_outcome(_start("remote", provider="google"))
148 mod._handle_cortex_outcome(_finish("remote"))
149
150 assert mod._wedge_state["failures"] == {"u1", "u2"}
151
152 mod._handle_cortex_outcome(_start("ok"))
153 mod._handle_cortex_outcome(_finish("ok"))
154
155 assert mod._wedge_state["failures"] == set()
156
157
158def test_only_real_500_provider_unavailable_counts(monkeypatch):
159 monkeypatch.setattr(
160 mod,
161 "read_service_port",
162 Mock(side_effect=AssertionError("should not probe below threshold")),
163 )
164 req = httpx.Request("POST", "http://localhost:8080/v1/chat/completions")
165 err500 = httpx.HTTPStatusError(
166 "server error",
167 request=req,
168 response=httpx.Response(500, request=req),
169 )
170 err400 = httpx.HTTPStatusError(
171 "bad request",
172 request=req,
173 response=httpx.Response(400, request=req),
174 )
175 errto = httpx.TimeoutException("timed out")
176
177 reason500 = classify_provider_error(err500, "local")
178 reason400 = classify_provider_error(err400, "local")
179 reasonto = classify_provider_error(errto, "local")
180 assert reason500 == "provider_unavailable"
181 assert reason400 == "unknown"
182 assert reasonto == "chat_timeout"
183
184 for use_id, reason in (
185 ("u500", reason500),
186 ("u400", reason400),
187 ("uto", reasonto),
188 ("umissing", None),
189 ):
190 mod._handle_cortex_outcome(_start(use_id))
191 mod._handle_cortex_outcome(_error(use_id, reason))
192
193 assert mod._wedge_state["failures"] == {"u500"}
194
195
196@pytest.mark.parametrize(
197 ("platform", "proctitle"),
198 [
199 ("darwin", mod.MLX_SERVER_PROCESS_NAME),
200 ("linux", mod.LOCAL_SERVER_PROCESS_NAME),
201 ],
202)
203def test_recycles_through_provider_reconciler_by_platform(
204 monkeypatch, platform, proctitle
205):
206 monkeypatch.setattr(mod.sys, "platform", platform)
207 _ready_local_server(monkeypatch)
208 managed = _ManagedStub(proctitle, ["/tmp/server"])
209 mod._managed_procs.append(managed)
210 recycle = Mock(return_value=True)
211 monkeypatch.setattr(mod, "_request_provider_runtime_recycle", recycle)
212 monkeypatch.setattr(
213 mod,
214 "_restart_service",
215 Mock(side_effect=AssertionError("wedge must not use generic restart")),
216 )
217
218 _drive_wedge()
219
220 recycle.assert_called_once()
221 assert recycle.call_args.args == ("local",)
222 assert recycle.call_args.kwargs["reason_code"] == (
223 "local-wedge-provider-unavailable"
224 )
225 assert (
226 recycle.call_args.kwargs["detail"]["health_state"] == local_server.STATE_READY
227 )
228 assert recycle.call_args.kwargs["detail"]["port"] == 9999
229 assert mod._SERVICE_STATE == {}
230
231 managed.process.returncode = 1
232 states = {
233 "local": mod.ProviderRuntimeState("local"),
234 "parakeet": mod.ProviderRuntimeState("parakeet"),
235 }
236 monkeypatch.setattr(mod, "_provider_runtime_states", states)
237 monkeypatch.setattr(mod, "_write_provider_runtime", lambda _state, **_kwargs: None)
238 monkeypatch.setattr(
239 mod,
240 "_launch_process",
241 Mock(side_effect=AssertionError("provider exit must not relaunch")),
242 )
243
244 procs = [managed]
245 asyncio.run(mod.handle_runner_exits(procs))
246
247 assert procs == []
248 assert states["local"].latest_phase == "stopped"
249 assert mod._recovery_state["local"].down_generation == states["local"].generation
250
251
252def test_cooldown_ignores_terminal_events_and_prevents_rerecycle(monkeypatch):
253 now = [1000.0]
254 monkeypatch.setattr(mod.time, "monotonic", lambda: now[0])
255 _ready_local_server(monkeypatch)
256 recycle = Mock(return_value=True)
257 monkeypatch.setattr(mod, "_request_provider_runtime_recycle", recycle)
258
259 _drive_wedge()
260
261 assert recycle.call_count == 1
262 assert mod._wedge_state["cooldown_until"] == 1120.0
263 assert mod._wedge_state["awaiting_recovery"] is True
264
265 now[0] = 1050.0
266 for idx in range(mod.LOCAL_WEDGE_THRESHOLD):
267 use_id = f"cooldown-{idx}"
268 mod._handle_cortex_outcome(_start(use_id))
269 mod._handle_cortex_outcome(_error(use_id))
270 mod._handle_cortex_outcome(_start("cooldown-finish"))
271 mod._handle_cortex_outcome(_finish("cooldown-finish"))
272
273 assert recycle.call_count == 1
274 assert mod._wedge_state["failures"] == set()
275 assert mod._wedge_state["awaiting_recovery"] is True
276
277
278def test_wedge_logs_declared_recycling_and_recovered(monkeypatch, caplog):
279 now = [2000.0]
280 monkeypatch.setattr(mod.time, "monotonic", lambda: now[0])
281 _ready_local_server(monkeypatch)
282 monkeypatch.setattr(
283 mod, "_request_provider_runtime_recycle", Mock(return_value=True)
284 )
285 caplog.set_level(logging.INFO)
286
287 _drive_wedge()
288 now[0] = mod._wedge_state["cooldown_until"] + 1.0
289 mod._handle_cortex_outcome(_start("recovered"))
290 mod._handle_cortex_outcome(_finish("recovered"))
291
292 assert "local server wedge: declared" in caplog.text
293 assert "local server wedge: requested local provider recycle" in caplog.text
294 assert "local server wedge: recovered after recycle" in caplog.text
295
296
297def test_dispatch_routes_cortex_events_and_duplicate_errors_are_idempotent(
298 monkeypatch, mock_callosum
299):
300 _ready_local_server(monkeypatch)
301 recycle = Mock(return_value=True)
302 monkeypatch.setattr(mod, "_request_provider_runtime_recycle", recycle)
303
304 mod._handle_callosum_message(_start("dupe"))
305 mod._handle_callosum_message(_error("dupe"))
306 mod._handle_callosum_message(_error("dupe"))
307 mod._handle_callosum_message(_start("u2"))
308 mod._handle_callosum_message(_error("u2"))
309 mod._handle_callosum_message(_start("u3"))
310 mod._handle_callosum_message(_error("u3"))
311
312 recycle.assert_called_once()
313 assert recycle.call_args.args == ("local",)
314 assert recycle.call_args.kwargs["reason_code"] == (
315 "local-wedge-provider-unavailable"
316 )
317
318
319def test_provider_map_cap_evicts_oldest_and_evicted_terminal_is_ignored(monkeypatch):
320 monkeypatch.setattr(mod, "LOCAL_WEDGE_PROVIDER_MAP_CAP", 2)
321
322 mod._handle_cortex_outcome(_start("u1"))
323 mod._handle_cortex_outcome(_start("u2"))
324 mod._handle_cortex_outcome(_start("u3"))
325 mod._handle_cortex_outcome(_error("u1"))
326
327 assert list(mod._wedge_state["providers"].keys()) == ["u2", "u3"]
328 assert mod._wedge_state["failures"] == set()
329
330
331@pytest.mark.parametrize(
332 ("port", "health_state", "expected_log"),
333 [
334 (None, local_server.STATE_READY, "local service port unavailable"),
335 (9999, local_server.STATE_LOADING, "health state=loading"),
336 ],
337)
338def test_probe_deferral_clears_failures_without_cooldown(
339 monkeypatch, caplog, port, health_state, expected_log
340):
341 monkeypatch.setattr(mod, "read_service_port", Mock(return_value=port))
342 probe_health = Mock(return_value=(health_state, None))
343 monkeypatch.setattr(local_server, "_probe_health", probe_health)
344 recycle = Mock()
345 monkeypatch.setattr(mod, "_request_provider_runtime_recycle", recycle)
346 caplog.set_level(logging.WARNING)
347
348 _drive_wedge()
349
350 assert expected_log in caplog.text
351 assert mod._wedge_state["failures"] == set()
352 assert mod._wedge_state["cooldown_until"] == 0.0
353 assert mod._wedge_state["awaiting_recovery"] is False
354 recycle.assert_not_called()
355
356
357def test_recycle_request_false_defers_without_cooldown(monkeypatch, caplog):
358 _ready_local_server(monkeypatch)
359 recycle = Mock(return_value=False)
360 monkeypatch.setattr(mod, "_request_provider_runtime_recycle", recycle)
361 caplog.set_level(logging.WARNING)
362
363 _drive_wedge()
364
365 assert "local server wedge: recycle request failed" in caplog.text
366 assert mod._wedge_state["failures"] == set()
367 assert mod._wedge_state["cooldown_until"] == 0.0
368 assert mod._wedge_state["awaiting_recovery"] is False