From a3d33fe22f3883ae440b437c910fca79fd34f008 Mon Sep 17 00:00:00 2001 From: kshitijk4poor <82637225+kshitijk4poor@users.noreply.github.com> Date: Thu, 3 Sep 2026 01:40:49 +0530 Subject: [PATCH] fix(copilot): single-flight the token exchange and close abandoned responses Follow-up to the off-loop move: once the credential-pool handlers run on worker threads, the dashboard's periodic /api/credentials/pool polls can overlap, and during a DNS outage each poll would have started its own exchange and abandoned its own hung resolver thread. - Per-fingerprint threading.Lock around the exchange: concurrent callers wait on the one in-flight attempt, then hit the positive or negative cache (bounded worker count, no duplicate network calls). - _urlopen_bounded: when the hard cap fires and the abandoned worker later succeeds, close the HTTPResponse instead of leaking the socket. - Tests (none shipped with the original PR): hard cap + late-close, single-flight success and failure paths, and the pool endpoint running off-loop / keeping the loop responsive under a 200 ms blocking read. --- hermes_cli/copilot_auth.py | 47 +++- .../test_credential_pool_off_loop.py | 209 ++++++++++++++++++ 2 files changed, 252 insertions(+), 4 deletions(-) create mode 100644 tests/hermes_cli/test_credential_pool_off_loop.py diff --git a/hermes_cli/copilot_auth.py b/hermes_cli/copilot_auth.py index 6780fe4186..95955e6547 100644 --- a/hermes_cli/copilot_auth.py +++ b/hermes_cli/copilot_auth.py @@ -381,6 +381,20 @@ _JWT_DISK_MAX_BYTES = 1_048_576 # 1 MiB cap on the persisted JWT store read # Maps raw-token fingerprint -> epoch until which exchange attempts are # skipped (raise immediately). Success clears the entry. _exchange_failure_cache: dict[str, float] = {} +# Single-flight guard per token fingerprint: concurrent callers (the dashboard +# polls /api/credentials/pool every few seconds, each poll off-loop) wait on +# the ONE in-flight exchange and then hit the positive/negative cache, instead +# of each spawning their own hung resolver thread during a DNS outage. +_exchange_locks: dict[str, threading.Lock] = {} +_exchange_locks_guard = threading.Lock() + + +def _exchange_lock_for(fp: str) -> threading.Lock: + with _exchange_locks_guard: + lock = _exchange_locks.get(fp) + if lock is None: + lock = _exchange_locks[fp] = threading.Lock() + return lock _EXCHANGE_FAILURE_TTL_TRANSIENT_SECONDS = 60.0 # network blips: retry soon _EXCHANGE_FAILURE_TTL_PERMANENT_SECONDS = 1800.0 # 401/403/404: won't heal # HTTP statuses that indicate the token itself is rejected — retrying with @@ -539,12 +553,23 @@ def _urlopen_bounded(req, timeout: float): import urllib.request box: dict = {} + abandoned = threading.Event() def _worker() -> None: try: - box["resp"] = urllib.request.urlopen(req, timeout=timeout) + resp = urllib.request.urlopen(req, timeout=timeout) except BaseException as exc: # re-raised on the caller's thread box["exc"] = exc + return + if abandoned.is_set(): + # The caller already timed out; nobody will read this response, + # so release its socket instead of leaking it with the thread. + try: + resp.close() + except Exception: + pass + return + box["resp"] = resp t = threading.Thread( target=_worker, name="copilot-token-exchange", daemon=True @@ -552,6 +577,7 @@ def _urlopen_bounded(req, timeout: float): t.start() t.join(timeout + _DNS_GRACE_SECONDS) if t.is_alive(): + abandoned.set() raise TimeoutError( "copilot token exchange exceeded hard cap of " f"{timeout + _DNS_GRACE_SECONDS:.0f}s (DNS/getaddrinfo hang?)" @@ -579,11 +605,24 @@ def exchange_copilot_token(raw_token: str, *, timeout: float = 10.0) -> tuple[st Results are cached in-process and reused until close to expiry. Raises ``ValueError`` on failure. """ - import urllib.request - fp = _token_fingerprint(raw_token) - # Check in-process cache first + # Fast path outside the lock: a valid in-process JWT needs no exchange. + cached = _jwt_cache.get(fp) + if cached and time.time() < cached[1] - _JWT_REFRESH_MARGIN_SECONDS: + return cached + + with _exchange_lock_for(fp): + return _exchange_copilot_token_locked(raw_token, fp, timeout=timeout) + + +def _exchange_copilot_token_locked( + raw_token: str, fp: str, *, timeout: float +) -> tuple[str, float, Optional[str]]: + import urllib.request + + # Re-check the caches under the lock: a concurrent caller may have just + # completed (or just failed) the exchange we were queued behind. cached = _jwt_cache.get(fp) if cached: api_token, expires_at, base_url = cached diff --git a/tests/hermes_cli/test_credential_pool_off_loop.py b/tests/hermes_cli/test_credential_pool_off_loop.py new file mode 100644 index 0000000000..881d861b54 --- /dev/null +++ b/tests/hermes_cli/test_credential_pool_off_loop.py @@ -0,0 +1,209 @@ +"""Regression tests for the #91912 salvage — credential-pool handlers off-loop +and the bounded Copilot token exchange. + +The 2026-08-22 incident: ``GET /api/credentials/pool`` ran ``load_pool()`` on +the uvicorn event loop; for the copilot provider that reaches +``urllib.request.urlopen`` whose ``timeout`` does not bound ``getaddrinfo``, +so a networkless host froze the whole dashboard backend for 17 minutes. +""" + +from __future__ import annotations + +import asyncio +import threading +import time +from unittest.mock import patch + +import pytest + +from hermes_cli import copilot_auth + + +# --------------------------------------------------------------------------- +# _urlopen_bounded +# --------------------------------------------------------------------------- + + +class TestUrlopenBounded: + def test_returns_response_when_worker_completes(self): + sentinel = object() + with patch("urllib.request.urlopen", return_value=sentinel): + assert copilot_auth._urlopen_bounded("req", 1.0) is sentinel + + def test_reraises_worker_exception(self): + with patch("urllib.request.urlopen", side_effect=OSError("boom")): + with pytest.raises(OSError, match="boom"): + copilot_auth._urlopen_bounded("req", 1.0) + + def test_hard_cap_fires_on_hung_resolver_and_closes_late_response(self, monkeypatch): + """A urlopen that hangs past timeout + grace must raise TimeoutError + promptly, and when the abandoned worker later *succeeds* it must close + the response instead of leaking the socket.""" + monkeypatch.setattr(copilot_auth, "_DNS_GRACE_SECONDS", 0.05) + release = threading.Event() + closed = threading.Event() + + class _LateResponse: + def close(self): + closed.set() + + def hung_urlopen(req, timeout): + release.wait(timeout=5) + return _LateResponse() + + with patch("urllib.request.urlopen", side_effect=hung_urlopen): + started = time.monotonic() + with pytest.raises(TimeoutError, match="hard cap"): + copilot_auth._urlopen_bounded("req", 0.05) + elapsed = time.monotonic() - started + assert elapsed < 2.0 + release.set() + assert closed.wait(timeout=2), "late response was not closed" + + +# --------------------------------------------------------------------------- +# single-flight exchange +# --------------------------------------------------------------------------- + + +class TestExchangeSingleFlight: + @pytest.fixture(autouse=True) + def _clean_caches(self, monkeypatch, tmp_path): + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + copilot_auth._jwt_cache.clear() + copilot_auth._exchange_failure_cache.clear() + copilot_auth._exchange_locks.clear() + yield + copilot_auth._jwt_cache.clear() + copilot_auth._exchange_failure_cache.clear() + copilot_auth._exchange_locks.clear() + + def test_concurrent_callers_share_one_exchange(self, monkeypatch): + """N concurrent callers for the same token must perform ONE network + exchange; the rest wait on the lock and hit the populated cache.""" + monkeypatch.setattr(copilot_auth, "_load_jwt_from_disk", lambda fp: None) + monkeypatch.setattr(copilot_auth, "_save_jwt_to_disk", lambda *a, **k: None) + calls = [] + gate = threading.Event() + + class _Resp: + def __enter__(self): + return self + + def __exit__(self, *exc): + return False + + def read(self): + return b'{"token": "tid=1;exp=9", "expires_at": 4102444800}' + + def fake_bounded(req, timeout): + calls.append(threading.get_ident()) + gate.wait(timeout=5) # hold the first exchange open while others queue + return _Resp() + + monkeypatch.setattr(copilot_auth, "_urlopen_bounded", fake_bounded) + + results = [] + threads = [ + threading.Thread(target=lambda: results.append(copilot_auth.exchange_copilot_token("ghu_" + "x" * 30))) + for _ in range(8) + ] + for t in threads: + t.start() + time.sleep(0.2) # let every caller reach the lock + assert len(calls) == 1 + gate.set() + for t in threads: + t.join(timeout=5) + + assert len(calls) == 1 + assert len(results) == 8 + assert {r[0] for r in results} == {"tid=1;exp=9"} + + def test_waiters_observe_negative_cache_after_failed_exchange(self, monkeypatch): + """When the single in-flight exchange fails, queued callers must not + each start their own exchange — they see the negative cache.""" + monkeypatch.setattr(copilot_auth, "_load_jwt_from_disk", lambda fp: None) + monkeypatch.setattr(copilot_auth, "_EXCHANGE_MAX_ATTEMPTS", 1) + calls = [] + gate = threading.Event() + + def fake_bounded(req, timeout): + calls.append(1) + gate.wait(timeout=5) + raise TimeoutError("hard cap") + + monkeypatch.setattr(copilot_auth, "_urlopen_bounded", fake_bounded) + errors = [] + + def run(): + try: + copilot_auth.exchange_copilot_token("ghu_" + "y" * 30) + except ValueError as exc: + errors.append(str(exc)) + + threads = [threading.Thread(target=run) for _ in range(5)] + for t in threads: + t.start() + time.sleep(0.2) + gate.set() + for t in threads: + t.join(timeout=5) + + assert len(calls) == 1 + assert len(errors) == 5 + assert any("recently failed" in e for e in errors) + + +# --------------------------------------------------------------------------- +# web_server credential-pool handlers off the loop +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_list_credential_pool_runs_off_event_loop(monkeypatch): + import hermes_cli.auth as auth_mod + from hermes_cli import web_server + + loop_thread = threading.get_ident() + seen = {} + + def fake_read_pool(*args, **kwargs): + seen["thread"] = threading.get_ident() + return {} + + monkeypatch.setattr(auth_mod, "read_credential_pool", fake_read_pool) + result = await web_server.list_credential_pool() + + assert result == {"providers": []} + assert seen["thread"] != loop_thread + + +@pytest.mark.asyncio +async def test_list_credential_pool_keeps_loop_responsive(monkeypatch): + """A 200 ms blocking pool read must not freeze a concurrent ticker.""" + import hermes_cli.auth as auth_mod + from hermes_cli import web_server + + def slow_read(*args, **kwargs): + time.sleep(0.2) + return {} + + monkeypatch.setattr(auth_mod, "read_credential_pool", slow_read) + + gaps = [] + stop = asyncio.Event() + + async def ticker(): + last = time.perf_counter() + while not stop.is_set(): + await asyncio.sleep(0.005) + now = time.perf_counter() + gaps.append(now - last) + last = now + + t = asyncio.create_task(ticker()) + await web_server.list_credential_pool() + stop.set() + await t + assert max(gaps) < 0.1, f"event loop stalled for {max(gaps) * 1000:.0f} ms"