"""Tests for the steering dispatcher (H1/C2/C3/T9/T12) in ``src/claude_session.py`` and its ``__STEERED__`` handling in ``src/router.py``. `ClaudeProcess`/`RunnerRegistry` internals (spawn, stdin/stdout framing, `steer()`'s own race handling) are already covered by ``tests/test_claude_runner.py`` against the real ``tests/fake_claude.py`` subprocess. This file is about the DISPATCHER's decisions — which only need a process double that behaves the right way at the right time, synchronised via ``threading.Event`` rather than real subprocess timing (the plan explicitly warns sleep-based subprocess races are the flaky way to test this). """ import threading import time from unittest.mock import patch import pytest from src import claude_runner, claude_session, router from src.claude_runner import SteerResult, SteerStatus from src.claude_session import send_message from src.sentinels import is_steered def _wait_until(predicate, timeout=5.0, interval=0.01): deadline = time.monotonic() + timeout while time.monotonic() < deadline: if predicate(): return time.sleep(interval) raise AssertionError("condition not met within timeout") # --------------------------------------------------------------------------- # Scriptable ClaudeProcess/RunnerRegistry doubles # --------------------------------------------------------------------------- class FakeProc: """Stand-in for `ClaudeProcess`. `run_turn` blocks on a per-call Event so a test can deterministically control exactly when a turn "finishes".""" def __init__(self, channel_id="ch", session_id="fake-sid"): self.channel_id = channel_id self.session_id = session_id self.inflight = False self._pending_steers: list[str] = [] self.steer_calls: list[str] = [] self.run_turn_calls: list[str] = [] self.stopped = False self.turn_error: Exception | None = None self._release = threading.Event() def run_turn(self, text, on_text=None, timeout=300): self._release.clear() self.inflight = True self.run_turn_calls.append(text) released = self._release.wait(timeout=5.0) self.inflight = False assert released, "test never released the fake turn" if self.turn_error is not None: err, self.turn_error = self.turn_error, None raise err return { "result": f"done:{text}", "session_id": self.session_id, "usage": {}, "total_cost_usd": 0, "cost_usd": 0, "duration_ms": 1, "num_turns": 1, "intermediate_count": 0, "subtype": "success", "is_error": False, } def steer(self, text): self.steer_calls.append(text) assert self.inflight, "dispatcher should only steer an inflight turn" self._pending_steers.append(text) return SteerResult(SteerStatus.STEERED) def pop_pending_steers(self): pending, self._pending_steers = self._pending_steers, [] return pending def release(self): self._release.set() class FakeRegistry: """Stand-in for `RunnerRegistry` — one proc per channel, created on demand, degrading to None past `capacity` (T10).""" def __init__(self, capacity=2): self.capacity = capacity self._procs: dict[str, FakeProc] = {} def get(self, channel_id, model=None, session_id=None, cwd=None): proc = self._procs.get(channel_id) if proc is not None: return proc if len(self._procs) >= self.capacity: return None proc = FakeProc(channel_id, session_id=session_id or "fake-sid") self._procs[channel_id] = proc return proc def stop(self, channel_id): proc = self._procs.pop(channel_id, None) if proc is None: return False proc.stopped = True return True def live_count(self): return len(self._procs) @pytest.fixture(autouse=True) def _clean_state(): """Dispatcher state (`_session_locks`, `_channel_adapter`) and the steering registry singleton are module-level globals — clear them around every test so nothing leaks between tests.""" claude_session._session_locks.clear() claude_session._channel_adapter.clear() claude_runner.reset_registry_for_tests() yield claude_session._session_locks.clear() claude_session._channel_adapter.clear() claude_runner.reset_registry_for_tests() @pytest.fixture def temp_sessions(tmp_path, monkeypatch): """Isolated active.json per test — keeps real session state untouched.""" sessions_dir = tmp_path / "sessions" sessions_dir.mkdir() sf = sessions_dir / "active.json" sf.write_text("{}") monkeypatch.setattr(claude_session, "SESSIONS_DIR", sessions_dir) monkeypatch.setattr(claude_session, "_SESSIONS_FILE", sf) return sf @pytest.fixture def fake_registry(monkeypatch, temp_sessions): """Steering ON, backed by a FakeRegistry instead of real subprocesses.""" registry = FakeRegistry() monkeypatch.setattr(claude_runner, "_registry", registry) monkeypatch.setattr(claude_runner, "get_registry", lambda **kw: registry) monkeypatch.setattr(claude_session, "_steering_config", lambda channel_id=None: (True, 2, 20)) return registry class _FakeConfig: """Stand-in for `src.config.Config` used by `_steering_config`. Takes the full raw dict (`steering`, `channels`, etc.) like the real thing.""" def __init__(self, data: dict): self._data = data def __call__(self, *a, **kw): return self def get(self, key, default=None): return self._data.get(key, default) # --------------------------------------------------------------------------- # `_steering_config` — flag off, kill switch, clamps (X7/X8) # --------------------------------------------------------------------------- class TestSteeringConfig: def test_disabled_by_default_config(self, monkeypatch): monkeypatch.setattr("src.config.Config", _FakeConfig({"steering": {"enabled": False}})) enabled, _, _ = claude_session._steering_config() assert enabled is False def test_kill_switch_wins(self, monkeypatch): monkeypatch.setenv("ECHO_STEERING", "off") monkeypatch.setattr( "src.config.Config", _FakeConfig({"steering": {"enabled": True, "max_live": 2, "idle_minutes": 20}}), ) enabled, _, _ = claude_session._steering_config() assert enabled is False def test_max_live_zero_disables(self, monkeypatch): monkeypatch.setattr( "src.config.Config", _FakeConfig({"steering": {"enabled": True, "max_live": 0, "idle_minutes": 20}}), ) enabled, _, _ = claude_session._steering_config() assert enabled is False def test_idle_minutes_clamped_to_one(self, monkeypatch): monkeypatch.setattr( "src.config.Config", _FakeConfig({"steering": {"enabled": True, "max_live": 2, "idle_minutes": 0}}), ) _, _, idle_minutes = claude_session._steering_config() assert idle_minutes == 1 # --------------------------------------------------------------------------- # Etapa 6 (X16) — per-channel `channels..steering` override # --------------------------------------------------------------------------- class TestPerChannelSteeringOverride: def test_channel_on_overrides_global_off(self, monkeypatch): monkeypatch.setattr("src.config.Config", _FakeConfig({ "steering": {"enabled": False, "max_live": 2, "idle_minutes": 20}, "channels": {"echo-core": {"id": "chan-1", "steering": True}}, })) enabled, _, _ = claude_session._steering_config("chan-1") assert enabled is True def test_channel_off_overrides_global_on(self, monkeypatch): monkeypatch.setattr("src.config.Config", _FakeConfig({ "steering": {"enabled": True, "max_live": 2, "idle_minutes": 20}, "channels": {"echo-core": {"id": "chan-1", "steering": False}}, })) enabled, _, _ = claude_session._steering_config("chan-1") assert enabled is False def test_channel_unset_falls_through_to_global(self, monkeypatch): monkeypatch.setattr("src.config.Config", _FakeConfig({ "steering": {"enabled": True, "max_live": 2, "idle_minutes": 20}, "channels": {"echo-core": {"id": "chan-1"}}, # no "steering" key })) enabled, _, _ = claude_session._steering_config("chan-1") assert enabled is True enabled_other, _, _ = claude_session._steering_config("chan-unrelated") assert enabled_other is True # not present at all -> global default too # --------------------------------------------------------------------------- # Flag off -> unchanged blocking path # --------------------------------------------------------------------------- def test_steering_off_uses_unchanged_blocking_path(monkeypatch, temp_sessions): monkeypatch.setattr(claude_session, "_steering_config", lambda channel_id=None: (False, 0, 20)) calls = [] def fake_run_claude(cmd, timeout, on_text=None, cwd=None, channel_id=None): calls.append(cmd) return { "result": "hi", "session_id": "sid-1", "usage": {}, "total_cost_usd": 0, "cost_usd": 0, "duration_ms": 1, "num_turns": 1, "intermediate_count": 0, "subtype": "success", "is_error": False, } with patch.object(claude_session, "_run_claude", side_effect=fake_run_claude): result = send_message("ch-off", "hello") assert result == "hi" assert len(calls) == 1 # The dispatcher never touched the registry — disabled path is a no-op # with respect to steering, not just "didn't steer". assert claude_runner._registry is None # --------------------------------------------------------------------------- # H1 — contention steers, no separate `inflight` check/TOCTOU # --------------------------------------------------------------------------- def test_contention_steers(fake_registry): outcome = {} def first(): outcome["first"] = send_message("ch-1", "first", adapter_name="discord") t1 = threading.Thread(target=first) t1.start() _wait_until(lambda: fake_registry._procs.get("ch-1") is not None and fake_registry._procs["ch-1"].inflight) second = send_message("ch-1", "second", adapter_name="discord") assert is_steered(second) proc = fake_registry._procs["ch-1"] assert proc.steer_calls == ["second"] proc.release() t1.join(timeout=5) assert outcome["first"] == "done:first" # --------------------------------------------------------------------------- # C2/M1 — adapter mismatch never steers, and never spawns a second writer # --------------------------------------------------------------------------- def test_adapter_mismatch_waits_instead_of_steering(fake_registry): voice_outcome = {} def voice_turn(): voice_outcome["r"] = send_message("ch-2", "voice text", adapter_name="discord-voice") t1 = threading.Thread(target=voice_turn) t1.start() _wait_until(lambda: fake_registry._procs.get("ch-2") is not None and fake_registry._procs["ch-2"].inflight) text_outcome = {} def text_turn(): text_outcome["r"] = send_message("ch-2", "text message", adapter_name="discord") t2 = threading.Thread(target=text_turn) t2.start() proc = fake_registry._procs["ch-2"] time.sleep(0.1) # give t2 a chance to reach the mismatch branch assert proc.steer_calls == [], "a mismatched adapter must never be steered (C2)" proc.release() # finish the voice turn -> releases the lock t1.join(timeout=5) assert voice_outcome["r"] == "done:voice text" _wait_until(lambda: proc.inflight) # t2 now got the lock and started its own turn proc.release() t2.join(timeout=5) assert text_outcome["r"] == "done:text message" assert not is_steered(text_outcome["r"]) assert proc.run_turn_calls == ["voice text", "text message"] assert fake_registry.live_count() == 1, "must never spawn a second writer on the session (M1)" # --------------------------------------------------------------------------- # T10 — registry at capacity degrades to one-shot for a NEW channel # --------------------------------------------------------------------------- def test_capacity_degrades_new_channel_to_one_shot(fake_registry, temp_sessions): fake_registry.capacity = 1 fake_registry._procs["ch-existing"] = FakeProc("ch-existing") calls = [] def fake_run_claude(cmd, timeout, on_text=None, cwd=None, channel_id=None): calls.append(cmd) return { "result": "one-shot reply", "session_id": "sid-2", "usage": {}, "total_cost_usd": 0, "cost_usd": 0, "duration_ms": 1, "num_turns": 1, "intermediate_count": 0, "subtype": "success", "is_error": False, } with patch.object(claude_session, "_run_claude", side_effect=fake_run_claude): result = send_message("ch-new", "hello", adapter_name="discord") assert result == "one-shot reply" assert len(calls) == 1 assert "ch-new" not in fake_registry._procs # --------------------------------------------------------------------------- # T9 — /clear and /model stop the live process # --------------------------------------------------------------------------- def test_clear_session_stops_live_process(fake_registry): proc = fake_registry.get("ch-3") fake_registry._procs["ch-3"] = proc claude_session._save_sessions({"ch-3": {"session_id": "sid", "model": "sonnet"}}) assert claude_session.clear_session("ch-3") is True assert proc.stopped is True assert "ch-3" not in fake_registry._procs def test_set_session_model_stops_live_process(fake_registry): proc = fake_registry.get("ch-4") fake_registry._procs["ch-4"] = proc claude_session._save_sessions({"ch-4": {"session_id": "sid", "model": "sonnet"}}) assert claude_session.set_session_model("ch-4", "opus") is True assert proc.stopped is True assert "ch-4" not in fake_registry._procs # --------------------------------------------------------------------------- # C3 — a steered message must never be lost when its turn then fails # --------------------------------------------------------------------------- def test_pending_steers_redispatched_on_failed_turn(monkeypatch): """router.route_message's exception handler must pop pending steers (claude_session.pop_pending_steers) and redeliver each one via on_text — their own request threads already returned __STEERED__ and are gone, so on_text is the only way left to answer them.""" pending = ["steered one", "steered two"] monkeypatch.setattr(router, "_pop_pending_steers", lambda channel_id: list(pending)) monkeypatch.setattr( router, "send_message", lambda *a, **kw: (_ for _ in ()).throw(RuntimeError( "Claude CLI error (exit 1): You've hit your session limit · resets 10am (UTC)" )), ) fallback_calls = [] def fake_fallback(text, channel_id=None, manual=False): fallback_calls.append(text) return f"fallback:{text}" monkeypatch.setattr(router, "_local_fallback_reply", fake_fallback) streamed = [] result, is_cmd = router.route_message( "ch-c3", "user-1", "original message", on_text=streamed.append, ) assert result == "fallback:original message" assert is_cmd is False assert fallback_calls == ["original message", "steered one", "steered two"] assert streamed == ["fallback:steered one", "fallback:steered two"] def test_pending_steers_logged_without_on_text(monkeypatch, caplog): """No on_text means no way to redeliver — must be logged (T12), not silently dropped.""" monkeypatch.setattr(router, "_pop_pending_steers", lambda channel_id: ["lost message"]) monkeypatch.setattr( router, "send_message", lambda *a, **kw: (_ for _ in ()).throw(RuntimeError("boom")), ) with caplog.at_level("WARNING"): result, is_cmd = router.route_message("ch-c3b", "user-1", "hello") assert result == "Error: boom" assert any("steered message lost" in rec.message for rec in caplog.records)