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.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user