The refresh burst test started its provider deadline before four application lifespans were ready. The console timeout test started its deadline before the curator fork reached the blocking request. Both races failed under CI load without violating the production single-flight or cancellation contracts. Synchronize the refresh requests after client startup. Give the console worker wait a narrow seam so the timeout case fires only after the provider request starts. Pin curator candidate discovery because it is not part of the cancellation contract. Validation: both focused files passed twice with retries disabled. The Windows footgun and plugin compatibility checks also passed.
220 lines
9.3 KiB
Python
220 lines
9.3 KiB
Python
"""Native HTTP replay boundaries and deterministic provider-level concurrency."""
|
|
import logging
|
|
import threading
|
|
import time
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
|
|
import pytest
|
|
from fastapi import FastAPI
|
|
from fastapi.testclient import TestClient
|
|
|
|
from hermes_cli.dashboard_auth import clear_providers, register_provider
|
|
from hermes_cli.dashboard_auth import refresh_singleflight as replay
|
|
from hermes_cli.dashboard_auth.base import ProviderError, RefreshExpiredError, Session
|
|
from hermes_cli.dashboard_auth.routes import router
|
|
from tests.hermes_cli.conftest_dashboard_auth import StubAuthProvider
|
|
|
|
|
|
class Provider(StubAuthProvider):
|
|
def __init__(self, name, outcome="success"):
|
|
super().__init__()
|
|
self.name = name
|
|
self.outcome = outcome
|
|
self.calls = 0
|
|
self.entered = threading.Event()
|
|
self.release = threading.Event()
|
|
self.release.set()
|
|
|
|
def refresh_session(self, *, refresh_token):
|
|
self.calls += 1
|
|
self.entered.set()
|
|
assert self.release.wait(5), "test provider timed out"
|
|
if self.outcome == "expired":
|
|
raise RefreshExpiredError("expired")
|
|
if self.outcome == "outage":
|
|
raise ProviderError("temporarily unavailable")
|
|
return Session(user_id=self.name, email="test@example.test", display_name=self.name,
|
|
org_id="test", provider=self.name, expires_at=int(time.time()) + 3600,
|
|
access_token=f"access-{self.name}", refresh_token=f"rotated-{self.name}")
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def isolated_registry():
|
|
clear_providers()
|
|
with replay._guard:
|
|
replay._cache.clear()
|
|
replay._flights.clear()
|
|
yield
|
|
clear_providers()
|
|
with replay._guard:
|
|
assert not replay._flights
|
|
replay._cache.clear()
|
|
|
|
|
|
@pytest.mark.parametrize("case", ["hint-fallback", "negative", "outage", "replacement",
|
|
"ttl", "capacity", "independent", "network-hop"])
|
|
def test_native_http_refresh_boundaries(case, monkeypatch):
|
|
owner = Provider("owner", "expired" if case == "negative" else "outage" if case == "outage" else "success")
|
|
other = Provider("other", "success" if case == "independent" else "expired")
|
|
register_provider(owner)
|
|
register_provider(other)
|
|
app = FastAPI()
|
|
app.include_router(router)
|
|
now = [100.0]
|
|
monkeypatch.setattr(replay.time, "monotonic", lambda: now[0])
|
|
with TestClient(app) as client:
|
|
def request(hint="owner", token="opaque-old-token", **kwargs):
|
|
return client.post("/auth/native/refresh", json={"refresh_token": token, "provider": hint}, **kwargs)
|
|
|
|
first = request()
|
|
assert first.status_code == (401 if case == "negative" else 503 if case == "outage" else 200)
|
|
if case in {"hint-fallback", "negative", "outage"}:
|
|
for hint in ("other", "", "unknown-a", "unknown-b", "owner"):
|
|
response = request(hint)
|
|
assert response.status_code == first.status_code
|
|
if first.status_code == 200:
|
|
assert response.json() == first.json()
|
|
assert owner.calls == (6 if case == "outage" else 1)
|
|
assert other.calls == 1
|
|
elif case == "replacement":
|
|
clear_providers()
|
|
replacement = Provider("owner")
|
|
register_provider(replacement)
|
|
assert request().status_code == 200
|
|
assert replacement.calls == 1
|
|
elif case == "ttl":
|
|
assert request().json() == first.json()
|
|
now[0] += replay._SUCCESS_TTL
|
|
assert request().status_code == 200
|
|
assert owner.calls == 2
|
|
elif case == "capacity":
|
|
monkeypatch.setattr(replay, "_MAX_ENTRIES", 2)
|
|
for token in ("new-token-1", "new-token-2", "new-token-3"):
|
|
assert request(token=token).status_code == 200
|
|
assert len(replay._cache) == 2
|
|
assert all(isinstance(key[1], bytes) and len(key[1]) == 32 for key in replay._cache)
|
|
elif case == "independent":
|
|
response = request("other")
|
|
assert response.status_code == 200
|
|
assert response.json()["provider"] == "other"
|
|
assert first.json()["provider"] == "owner"
|
|
assert owner.calls == other.calls == 1
|
|
else:
|
|
# A burst that straddles a network change (laptop wakes on another Wi-Fi) still
|
|
# coalesces: the RT identifies the session, the peer address does not.
|
|
with TestClient(app, client=("192.0.2.12", 2345)) as another_client:
|
|
assert another_client.post("/auth/native/refresh", json={"refresh_token": "opaque-old-token"}).json() == first.json()
|
|
assert owner.calls == 1
|
|
|
|
|
|
@pytest.mark.parametrize("outcome, independent", [("success", False), ("expired", False),
|
|
("outage", False), ("success", True)])
|
|
def test_concurrent_refresh_uses_concrete_provider_identity(outcome, independent):
|
|
def coalesced(token, hint):
|
|
return replay.refresh_session_coalesced(
|
|
token, hint, phase="test", log=logging.getLogger(__name__))
|
|
|
|
owner = Provider("owner", outcome)
|
|
other = Provider("other", "success" if independent else "expired")
|
|
owner.release.clear()
|
|
if independent:
|
|
other.release.clear()
|
|
register_provider(owner)
|
|
register_provider(other)
|
|
with ThreadPoolExecutor(max_workers=3) as pool:
|
|
first = pool.submit(coalesced, "same-token", "owner")
|
|
assert owner.entered.wait(3)
|
|
second = pool.submit(coalesced, "same-token", "other")
|
|
try:
|
|
if independent:
|
|
assert other.entered.wait(3), "unrelated providers must not share a lock"
|
|
else:
|
|
deadline = time.monotonic() + 3
|
|
while time.monotonic() < deadline:
|
|
with replay._guard:
|
|
if any(key[0] == id(owner) and flight.users == 2 for key, flight in replay._flights.items()):
|
|
break
|
|
time.sleep(0.005)
|
|
else:
|
|
pytest.fail("different hints did not converge on the concrete issuer lock")
|
|
assert owner.calls == 1
|
|
finally:
|
|
owner.release.set()
|
|
other.release.set()
|
|
if outcome == "outage":
|
|
for future in (first, second):
|
|
with pytest.raises(ProviderError):
|
|
future.result(timeout=3)
|
|
assert owner.calls == 2
|
|
else:
|
|
results = [first.result(timeout=3), second.result(timeout=3)]
|
|
if outcome == "expired":
|
|
assert results == [None, None]
|
|
else:
|
|
assert [result[1] for result in results] == ["owner", "other" if independent else "owner"]
|
|
assert owner.calls == 1
|
|
assert not replay._flights
|
|
|
|
|
|
class _RotatingReuseDetectingProvider(Provider):
|
|
"""A rotating-RT IdP with reuse detection: replaying a rotated RT kills the session."""
|
|
|
|
def __init__(self):
|
|
super().__init__("stub")
|
|
self.rotated: set[str] = set()
|
|
|
|
def verify_session(self, *, access_token):
|
|
return None # every AT presented is expired -> the gate must refresh
|
|
|
|
def refresh_session(self, *, refresh_token):
|
|
self.calls += 1
|
|
if refresh_token in self.rotated:
|
|
raise RefreshExpiredError("refresh token reuse detected")
|
|
self.rotated.add(refresh_token)
|
|
self.entered.set()
|
|
assert self.release.wait(5), "test provider timed out"
|
|
return Session(user_id="u", email="u@example.test", display_name="u", org_id="o",
|
|
provider=self.name, expires_at=int(time.time()) + 900,
|
|
access_token="fresh-at", refresh_token=f"rt-{self.calls}")
|
|
|
|
|
|
@pytest.fixture
|
|
def gated_web_app():
|
|
from hermes_cli import web_server
|
|
|
|
prev = {k: getattr(web_server.app.state, k, None) for k in ("bound_host", "bound_port", "auth_required")}
|
|
web_server.app.state.bound_host = "gw.example.test"
|
|
web_server.app.state.bound_port = 443
|
|
web_server.app.state.auth_required = True
|
|
yield web_server.app
|
|
for k, v in prev.items():
|
|
setattr(web_server.app.state, k, v)
|
|
|
|
|
|
def test_cookie_gate_burst_with_stale_rt_rotates_once(gated_web_app):
|
|
"""#55712: a browser burst after AT expiry carries one stale RT in N requests; exactly one
|
|
reaches the provider and every sibling is served under the rotated session."""
|
|
provider = _RotatingReuseDetectingProvider()
|
|
provider.release.clear()
|
|
register_provider(provider)
|
|
cookies = {"hermes_session_at": "expired-at", "hermes_session_rt": "stale-rt",
|
|
"hermes_session_provider": "stub"}
|
|
clients_ready = threading.Barrier(5)
|
|
|
|
def call():
|
|
# One TestClient per request: a shared jar would hand later requests the rotated RT.
|
|
with TestClient(gated_web_app, base_url="http://gw.example.test") as client:
|
|
clients_ready.wait(timeout=60)
|
|
return client.get("/api/auth/me", cookies=cookies)
|
|
|
|
with ThreadPoolExecutor(max_workers=4) as pool:
|
|
futures = [pool.submit(call) for _ in range(4)]
|
|
clients_ready.wait(timeout=60)
|
|
try:
|
|
assert provider.entered.wait(10)
|
|
finally:
|
|
provider.release.set()
|
|
statuses = sorted(f.result(timeout=10).status_code for f in futures)
|
|
assert statuses == [200, 200, 200, 200]
|
|
assert provider.calls == 1
|