personal memory agent
0

Configure Feed

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

solstone / tests / test_supervisor_wedge.py
12 kB 368 lines
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