diff --git a/tests/tui_gateway/test_readiness_singleflight.py b/tests/tui_gateway/test_readiness_singleflight.py new file mode 100644 index 0000000000..76d38d563e --- /dev/null +++ b/tests/tui_gateway/test_readiness_singleflight.py @@ -0,0 +1,194 @@ +"""Overlapping readiness polls must not saturate the shared RPC pool (#65151). + +``setup.runtime_check`` / ``setup.status`` are Desktop-polled and execute on the +shared RPC executor (``_LONG_HANDLERS``). Before the single-flight, every +overlapping poll ran its own provider resolution while a slow one (blocked +keyring, OAuth refresh, GIL pressure) was still in flight — one shared worker +per poll, all resolving the same state, until unrelated RPCs starved. + +These tests drive the REAL dispatch path (``server.dispatch`` → the shared +pool) with a recording transport and a blocking resolver. On the old behavior +the pool saturates and an unrelated long-handler RPC is never answered; with +the single-flight there is exactly one probe and the pool stays responsive. +""" + +from __future__ import annotations + +import threading +import time + +from tui_gateway import server + + +class _RecordingTransport: + """Collect worker-written responses; the tests wait on them by request id.""" + + def __init__(self): + self._lock = threading.Lock() + self._changed = threading.Event() + self.responses = {} + + def write(self, response): + with self._lock: + self.responses[response.get("id")] = response + self._changed.set() + return True + + def wait_for(self, request_id, timeout=2.0): + deadline = time.monotonic() + timeout + while True: + with self._lock: + response = self.responses.get(request_id) + if response is not None: + return response + self._changed.clear() + remaining = deadline - time.monotonic() + assert remaining > 0, f"timed out waiting for response {request_id!r}" + self._changed.wait(timeout=remaining) + + +def _dispatch(transport, request_id, method, params=None): + assert server.dispatch( + {"id": request_id, "method": method, "params": params or {}}, transport) is None + + +def _patch_fast_probe_env(monkeypatch, resolve): + monkeypatch.setattr("hermes_cli.runtime_provider.resolve_runtime_provider", resolve) + monkeypatch.setattr("hermes_cli.main._has_any_provider_configured", lambda **_kw: True) + monkeypatch.setattr(server, "_resolve_startup_runtime", lambda: ("custom/m", None)) + + +def test_overlapping_runtime_checks_share_one_probe_and_keep_pool_responsive(monkeypatch): + """Same-key polls join the in-flight probe; the shared pool answers an unrelated RPC.""" + started = threading.Event() + release = threading.Event() + calls = [] + + def slow_resolve(requested=None, **kwargs): + calls.append(requested) + started.set() + release.wait(timeout=10) + return {"provider": "custom", "api_key": "no-key-required", "source": "config"} + + _patch_fast_probe_env(monkeypatch, slow_resolve) + transport = _RecordingTransport() + + _dispatch(transport, "owner", "setup.runtime_check") + assert started.wait(timeout=2) + + try: + # Saturate every remaining shared RPC worker with overlapping polls. + for index in range(1, server._rpc_pool_workers): + _dispatch(transport, f"poll-{index}", "setup.runtime_check") + + # An unrelated long-handler RPC must still be answered promptly: with the + # single-flight the overlapping polls joined instead of each running a + # resolver, so they did not occupy the pool. + monkeypatch.setitem( + server._methods, "process.list", + lambda rid, _params: server._ok(rid, {"processes": []})) + before = time.monotonic() + _dispatch(transport, "unrelated", "process.list", {"session_id": "x"}) + unrelated = transport.wait_for("unrelated", timeout=2) + elapsed = time.monotonic() - before + + assert unrelated["result"] == {"processes": []} + assert elapsed < 1.0, f"unrelated RPC waited {elapsed:.2f}s — the shared pool saturated" + assert len(calls) == 1, "overlapping polls must not each run provider resolution" + + # An overlapping poll answered with the retryable unknown — a JSON-RPC + # error, never a fabricated ok=False result. + joiner = transport.wait_for("poll-1", timeout=2) + assert "result" not in joiner + assert joiner["error"]["code"] == server._READINESS_IN_PROGRESS_ERR + finally: + # Never leave pool workers blocked in the fake resolver, even on failure. + release.set() + + owner = transport.wait_for("owner", timeout=3) + assert owner["result"]["ok"] is True + assert owner["result"]["provider"] == "custom" + + # The in-flight entry is cleared once the probe settles: a later poll + # starts a fresh probe instead of reading a stale answer. + deadline = time.monotonic() + 2 + while server._readiness_inflight and time.monotonic() < deadline: + time.sleep(0.01) + assert not server._readiness_inflight + + +def test_blocked_provider_does_not_serialize_a_healthy_provider_check(monkeypatch): + """The single-flight key includes the requested provider; distinct providers probe in parallel.""" + provider_a_started = threading.Event() + release_a = threading.Event() + calls = {"provider-a": 0, "provider-b": 0} + + def resolve_by_provider(requested=None, **kwargs): + calls[requested] = calls.get(requested, 0) + 1 + if requested == "provider-a": + provider_a_started.set() + release_a.wait(timeout=10) + return {"provider": requested, "api_key": "no-key-required", "source": "config"} + + _patch_fast_probe_env(monkeypatch, resolve_by_provider) + transport = _RecordingTransport() + + _dispatch(transport, "a", "setup.runtime_check", {"provider": "provider-a"}) + assert provider_a_started.wait(timeout=2) + + try: + before = time.monotonic() + _dispatch(transport, "b", "setup.runtime_check", {"provider": "provider-b"}) + provider_b = transport.wait_for("b", timeout=2) + elapsed = time.monotonic() - before + + assert provider_b["result"]["ok"] is True + assert provider_b["result"]["provider"] == "provider-b" + assert elapsed < 1.0 + assert calls == {"provider-a": 1, "provider-b": 1} + finally: + release_a.set() + provider_a = transport.wait_for("a", timeout=3) + assert provider_a["result"]["ok"] is True + assert provider_a["result"]["provider"] == "provider-a" + + +def test_probe_outliving_its_budget_answers_retryable_unknown_then_reprobes(monkeypatch): + """A probe slower than the join budget answers the retryable error (unknown, not ok=False) + while it keeps running; the next poll starts a fresh probe with the real answer.""" + monkeypatch.setattr(server, "_READINESS_SHARE_WAIT_SECONDS", 0.2) + started = threading.Event() + release = threading.Event() + calls = [] + + def slow_resolve(requested=None, **kwargs): + calls.append(requested) + started.set() + release.wait(timeout=10) + return {"provider": "custom", "api_key": "no-key-required", "source": "config"} + + _patch_fast_probe_env(monkeypatch, slow_resolve) + transport = _RecordingTransport() + + _dispatch(transport, "first", "setup.runtime_check") + assert started.wait(timeout=2) + try: + first = transport.wait_for("first", timeout=3) + + # Unknown readiness, expressed as an unanswered RPC — the desktop keeps the + # last authoritative result and setup.status stays the credential source. + assert "result" not in first + assert first["error"]["code"] == server._READINESS_IN_PROGRESS_ERR + finally: + release.set() + + deadline = time.monotonic() + 2 + while server._readiness_inflight and time.monotonic() < deadline: + time.sleep(0.01) + assert not server._readiness_inflight + + _dispatch(transport, "second", "setup.runtime_check") + second = transport.wait_for("second", timeout=3) + assert second["result"]["ok"] is True + assert second["result"]["provider"] == "custom" + assert len(calls) == 2 diff --git a/tui_gateway/methods_config.py b/tui_gateway/methods_config.py index 63489d8cad..9605fa27a3 100644 --- a/tui_gateway/methods_config.py +++ b/tui_gateway/methods_config.py @@ -2,7 +2,12 @@ (method_ctx.bind_module) and reference them bare. ``config.set`` lives in methods_config_set. """ +import atexit +import concurrent.futures +import threading + from .method_ctx import HandlerRegistry, bind_module +from ._env import env_int from hermes_constants import DEFAULT_INDICATOR_STYLE, INDICATOR_STYLES from hermes_constants import display_hermes_home as _display_hermes_home @@ -12,6 +17,39 @@ method = _registry.method _profile_scoped = _registry.profile_scoped +# ── setup readiness single-flight (#65151) ───────────────────────────────── +# +# Readiness probes are Desktop-polled and execute on the shared RPC pool via +# ``_LONG_HANDLERS``. A slow probe (blocked keyring, OAuth refresh, GIL pressure) +# used to run one provider-resolution call per overlapping poll, each occupying +# another shared worker while they all resolved the same state. The probe now +# runs on this small dedicated executor, single-flighted per +# ``(kind, profile, requested provider)``: +# +# * the first caller submits the probe and waits a bounded budget; +# * an overlapping poll for a still-running probe answers a retryable error +# immediately — a JSON-RPC error, never a fabricated ``ok`` (the result +# contract requires the real shape, and the desktop already treats an errored +# runtime_check as unknown, keeping setup.status authoritative); +# * a probe that outlives the budget answers the same retryable error while it +# keeps running in the background; the in-flight entry is cleared when it +# settles, so the next poll starts a fresh probe and never reads a stale one. +_readiness_pool = concurrent.futures.ThreadPoolExecutor( + max_workers=max(2, min(4, env_int("HERMES_TUI_RPC_POOL_WORKERS", 8))), + thread_name_prefix="tui-readiness") +atexit.register(lambda: _readiness_pool.shutdown(wait=False, cancel_futures=True)) +_readiness_lock = threading.Lock() +_readiness_inflight: dict[tuple, concurrent.futures.Future] = {} +_READINESS_SHARE_WAIT_SECONDS = 4.0 +# setup.status's probe legitimately blocks on the boot bootstrap's record +# (free_tier_bootstrap.SETUP_READY_WAIT_SECONDS = 8s): its join budget must +# cover that wait or every boot poll would answer the retryable error. +_READINESS_STATUS_SHARE_WAIT_SECONDS = 12.0 +# Retryable-transient code for "the probe is still running / outlived its +# budget"; the client treats an errored runtime_check as unknown readiness. +_READINESS_IN_PROGRESS_ERR = 5097 + + def _projects_handler(name: str): """``@method(name)`` (profile-scoped) whose body's uncaught exception becomes ``_err(rid, 5061)``.""" def deco(fn): @@ -246,12 +284,54 @@ def _(rid, params: dict) -> dict: # ── setup readiness -def _readiness_check(rid, params, probe): +def _readiness_cleared(key): + """Done-callback for a single-flighted probe: forget the entry (the next poll + re-probes — completed answers are never cached), and surface a failure nobody + waited for (every caller timed out) in the log instead of dropping it.""" + def _clear(future): + with _readiness_lock: + if _readiness_inflight.get(key) is future: + _readiness_inflight.pop(key, None) + if not future.cancelled() and future.exception() is not None: + logger.debug("readiness probe %s failed after its callers returned: %s", + key, future.exception()) + return _clear + + +def _readiness_share(rid, key, run_probe, wait_seconds): + """Run ``run_probe`` single-flighted under ``key`` on the dedicated readiness pool. + + The first caller submits the probe and waits up to ``wait_seconds``; a caller that + finds a still-running probe answers the retryable error immediately (its shared RPC + worker is freed at once — the probe keeps running for the first caller), and one that + finds it settled reads the shared result. A probe that outlives the budget answers + the same retryable error while it continues in the background.""" + with _readiness_lock: + future = _readiness_inflight.get(key) + owner = future is None + if owner: + future = _readiness_pool.submit(run_probe) + _readiness_inflight[key] = future + future.add_done_callback(_readiness_cleared(key)) + if not owner and not future.done(): + return _err(rid, _READINESS_IN_PROGRESS_ERR, + "readiness check still in progress; retrying next tick") + try: + return _ok(rid, future.result(timeout=wait_seconds)) + except concurrent.futures.TimeoutError: + logger.warning("readiness probe %s exceeded %.1fs; it continues in the background", + key, wait_seconds) + return _err(rid, _READINESS_IN_PROGRESS_ERR, + f"readiness check timed out after {wait_seconds:.0f}s; retrying next tick") + + +def _readiness_check(rid, params, probe, *, probe_key, wait_seconds): """Shared shell of setup.status / setup.runtime_check. ``probe(profile, scoped)`` runs inside the optional ``profile`` param's HERMES_HOME + ``.env`` secret scope (ContextVars: concurrent checks stay isolated); ``scoped`` is the ``{"profile": ...}`` payload stamp (``{}`` for the launch profile). An unknown profile answers ``ok=False`` (never a JSON-RPC error, never a quiet answer - for the launch profile instead).""" + for the launch profile instead). ``probe_key`` + the profile single-flight the probe, and + ``wait_seconds`` bounds how long this shared-RPC worker waits for it (#65151).""" profile = str(params.get("profile") or "").strip() if isinstance(params, dict) else "" home = None if profile: @@ -264,9 +344,14 @@ def _readiness_check(rid, params, probe): # run under its own frozen secret scope too (``_profile_runtime_scope_tokens`` binds nothing in # a single-profile process), or the first profile-scoped read inside the resolver # (``HERMES_CODEX_BASE_URL`` for openai-codex) fails closed and the UI shows onboarding. - with _session_profile_runtime_scope({"profile_home": str(home) if home is not None else None}): - payload = probe(profile, {"profile": profile} if profile else {}) - return _ok(rid, payload) + def run_probe(): + # Applied on the readiness pool thread: ContextVars do not cross threads, and + # concurrent probes (different profiles) stay isolated exactly as they did when + # each ran on its caller's handler thread. + with _session_profile_runtime_scope({"profile_home": str(home) if home is not None else None}): + return probe(profile, {"profile": profile} if profile else {}) + + return _readiness_share(rid, (probe_key, profile), run_probe, wait_seconds) @method("setup.status") @@ -301,7 +386,8 @@ def _(rid, params: dict) -> dict: return {"provider_configured": record.provider_configured, "ready": True, "free_tier": record.free_tier, "other_providers": record.other_providers, "inference_provider": record.inference_provider, **record.failure_fields(), **scoped} - return _readiness_check(rid, params, probe) + return _readiness_check(rid, params, probe, probe_key="status", + wait_seconds=_READINESS_STATUS_SHARE_WAIT_SECONDS) except Exception as e: return _err(rid, 5016, str(e)) @@ -351,7 +437,8 @@ def _(rid, params: dict) -> dict: "source": runtime.get("source"), "free_tier": provider == "nous" and route_is_welcome_host(runtime.get("base_url")), **scoped} - return _readiness_check(rid, params, probe) + return _readiness_check(rid, params, probe, probe_key=f"runtime:{requested or ''}", + wait_seconds=_READINESS_SHARE_WAIT_SECONDS) except Exception as e: return _ok(rid, {"ok": False, "error": str(e)})