fix(gateway): coalesce native refresh per concrete provider across hints

(cherry picked from commit 18ce843daba52c427ab5d4b5f21990efc1202486)
This commit is contained in:
Doud-FR
2026-09-06 23:06:15 +02:00
committed by Teknium
parent 6c27a289c3
commit f561155a70
3 changed files with 241 additions and 5 deletions

View File

@@ -0,0 +1,78 @@
"""Short-lived native refresh replay, scoped to the concrete provider and caller.
A provider hint only orders discovery: it must neither split one rotating credential's
lock nor let an unrelated provider reuse its result. Raw refresh tokens are never keys.
"""
from __future__ import annotations
import hashlib
import logging
import threading
import time
from dataclasses import dataclass, field
from hermes_cli.dashboard_auth.base import DashboardAuthProvider, RefreshExpiredError, Session
from hermes_cli.dashboard_auth.request_utils import scan_session_providers
_SUCCESS_TTL = 30.0
_FAILURE_TTL = 5.0
_MAX_ENTRIES = 256
_guard = threading.Lock()
@dataclass
class _Flight:
lock: threading.Lock = field(default_factory=threading.Lock)
users: int = 0
# Hold the provider itself while caching: replacement (including same-name scoped
# registrations) invalidates identity, and Python cannot recycle its id under a live entry.
_cache: dict[tuple[int, bytes], tuple[float, DashboardAuthProvider, Session | None]] = {}
_flights: dict[tuple[int, bytes], _Flight] = {}
def _prune(now: float) -> None:
for key, (expires, _, _) in list(_cache.items()):
if expires <= now:
del _cache[key]
while len(_cache) > _MAX_ENTRIES:
del _cache[min(_cache, key=lambda key: _cache[key][0])]
def _refresh_provider(provider: DashboardAuthProvider, token: str, client_ip: str) -> Session | None:
digest = hashlib.sha256(client_ip.encode() + b"\0" + token.encode()).digest()
key = (id(provider), digest)
with _guard:
_prune(time.monotonic())
flight = _flights.setdefault(key, _Flight())
flight.users += 1
try:
with flight.lock:
with _guard:
cached = _cache.get(key)
if cached is not None and cached[0] > time.monotonic():
return cached[2]
try:
session = provider.refresh_session(refresh_token=token)
except RefreshExpiredError:
session = None
# ProviderError and unexpected execution failures are deliberately not cached.
with _guard:
now = time.monotonic()
_cache[key] = (now + (_SUCCESS_TTL if session is not None else _FAILURE_TTL), provider, session)
_prune(now)
return session
finally:
with _guard:
flight.users -= 1
if flight.users == 0:
_flights.pop(key, None)
def refresh_native_session(token: str, provider_hint: str, client_ip: str) -> Session | None:
"""Preserve upstream provider fallback/503 behavior while coalescing each actual issuer."""
return scan_session_providers(
provider_hint, lambda provider: _refresh_provider(provider, token, client_ip),
phase="native refresh", log=logging.getLogger(__name__),
)

View File

@@ -27,21 +27,23 @@ from urllib.parse import quote, unquote, urlencode, urlparse, urlunparse
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
from pydantic import BaseModel
from starlette.concurrency import run_in_threadpool
from hermes_cli.dashboard_auth import (
get_provider, list_providers, list_session_providers, native_flow)
from hermes_cli.dashboard_auth import prefix as _prefix_mod
from hermes_cli.dashboard_auth.audit import AuditEvent, audit_log
from hermes_cli.dashboard_auth.base import (
InvalidCodeError, InvalidCredentialsError, ProviderError, RefreshExpiredError, Session)
InvalidCodeError, InvalidCredentialsError, ProviderError, Session)
from hermes_cli.dashboard_auth.cookies import (
clear_pkce_cookie, clear_session_cookies, clear_sso_attempt_cookie, detect_https,
parse_pkce_payload, read_pkce_cookie, read_session_cookies, set_pkce_cookie,
set_session_cookies)
from hermes_cli.dashboard_auth.login_page import (
render_login_html, render_native_provider_choice_html)
from hermes_cli.dashboard_auth.native_refresh import refresh_native_session
from hermes_cli.dashboard_auth.request_utils import (
access_token_max_age, client_ip as _client_ip, is_safe_next_path, scan_session_providers)
access_token_max_age, client_ip as _client_ip, is_safe_next_path)
_log = logging.getLogger(__name__)
@@ -499,9 +501,11 @@ async def auth_native_refresh(request: Request, body: _NativeRefreshBody):
if not body.refresh_token:
raise _http(400, "refresh_token required")
try:
session = scan_session_providers(
body.provider, lambda p: p.refresh_session(refresh_token=body.refresh_token),
phase="native refresh", log=_log, swallow=(RefreshExpiredError,))
# Uvicorn validates trusted proxy peers before updating the ASGI client.
# Never split replay keys on caller-controlled X-Forwarded-For prefixes.
session = await run_in_threadpool(
refresh_native_session, body.refresh_token, body.provider,
request.client.host if request.client else "")
except ProviderError as e:
raise _http(503, f"Auth provider {str(e)!r} unreachable")
if session is not None:

View File

@@ -0,0 +1,154 @@
"""Native HTTP replay boundaries and deterministic provider-level concurrency."""
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 native_refresh 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", "client",
"ttl", "capacity", "independent", "xff"])
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 == "client":
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"}).status_code == 200
assert owner.calls == 2
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:
for prefix in ("192.0.2.1", "192.0.2.2"):
assert request(headers={"x-forwarded-for": f"{prefix}, 192.0.2.100"}).status_code == 200
# Only the ASGI peer (validated by Uvicorn), never an arbitrary header, scopes replay.
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):
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(replay.refresh_native_session, "same-token", "owner", "client")
assert owner.entered.wait(3)
second = pool.submit(replay.refresh_native_session, "same-token", "other", "client")
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.provider for result in results] == ["owner", "other" if independent else "owner"]
assert owner.calls == 1
assert not replay._flights