fix(client): serialise Keycloak token acquisition

Two concurrent tool calls each found the token stale, each POSTed the
password grant, and the appliance rejected one of the simultaneous
grants. That call died with "Keycloak token failed HTTP 401" -- nothing
to do with the operation it was performing. Hit while testing the guard:
two concurrent zerto_guard_before_mutate calls, one came back 401.

ensure_token now double-checks the cache under an asyncio.Lock, so N
concurrent callers produce exactly one token request and the rest reuse
the result. Measured against the lab ZVM with a cold client: 8 concurrent
reads made 8 token POSTs before, 1 after.

The 401-retry path had the same shape and was worse: every in-flight
request that got a 401 set _token = None and re-authed independently, so
one expiry became a thundering herd. Requests now capture a token
generation, and _reauth re-fetches only if nothing else has already
moved past it.

Token fetch also gets a bounded retry for transient failures (5xx,
network) with linear backoff. 401 and 403 are NOT retried: those are the
credentials themselves, and hammering Keycloak can trip its brute-force
lockout on a real service account.

Per the Zerto API lessons, auth failures now name the cause Keycloak
reported instead of a generic hint -- invalid_client means the client_id
is wrong for this appliance (10.x zerto-client, 9.x may be zerto-api),
invalid_grant means the username or password is. That is the first thing
to check and it was previously guesswork.

Also clears the two long-standing lint findings in this file (PIE810,
E501) while it is open. ruff check now passes across src/ and tests/.

pytest 48 passed (7 new, covering single-request concurrency, cache
reuse, both generation branches, no-retry-on-bad-credentials, the
invalid_client message, and transient retry).

Co-Authored-By: Claude Opus 5 (1M context) <[email protected]>
Claude-Session: https://claude.ai/code/session_016yVfC5nvZowoLFnEGWhLGn
This commit is contained in:
2026-09-21 15:12:28 -04:00
co-authored by Claude Opus 5
parent 1d53029038
commit 8ac89c7f82
2 changed files with 214 additions and 30 deletions
+94 -30
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import asyncio
import time
from typing import Any
from urllib.parse import urljoin
@@ -10,6 +11,9 @@ import httpx
from zerto_rewind_mcp.util import pick
TOKEN_FETCH_ATTEMPTS = 3
TOKEN_RETRY_BACKOFF_S = 0.5
class ZertoError(Exception):
def __init__(self, message: str, status_code: int | None = None, body: str | None = None):
@@ -34,6 +38,14 @@ class ZertoClient:
self.client_id = client_id
self._token: str | None = None
self._token_exp = 0.0
# Serialises token acquisition. Without it, concurrent tool calls each
# see a stale token and each POST to Keycloak; the ZVM rejects one of
# the simultaneous password grants and that call dies with HTTP 401.
self._token_lock = asyncio.Lock()
# Bumped on every successful fetch. A caller that got a 401 uses it to
# tell "my token is stale" from "another task already replaced it",
# so one expiry causes one re-auth, not one per in-flight request.
self._token_gen = 0
self._http = httpx.AsyncClient(verify=verify_tls, timeout=timeout)
async def aclose(self) -> None:
@@ -44,33 +56,83 @@ class ZertoClient:
return path
return urljoin(self.base_url + "/", path.lstrip("/"))
async def ensure_token(self) -> str:
if self._token and time.time() < self._token_exp - 30:
return self._token
def _token_is_fresh(self) -> bool:
return bool(self._token) and time.time() < self._token_exp - 30
async def _fetch_token(self) -> str:
"""POST the password grant. Caller must hold _token_lock."""
url = self._url("/auth/realms/zerto/protocol/openid-connect/token")
response = await self._http.post(
url,
data={
"grant_type": "password",
"username": self.username,
"password": self.password,
"client_id": self.client_id,
"scope": "openid",
},
headers={"Content-Type": "application/x-www-form-urlencoded"},
)
if response.status_code >= 400:
raise ZertoError(
f"Keycloak token failed HTTP {response.status_code}. "
"Check username/password/client_id "
"(10.x uses zerto-client; 9.x may use zerto-api).",
status_code=response.status_code,
body=response.text[:500],
data = {
"grant_type": "password",
"username": self.username,
"password": self.password,
"client_id": self.client_id,
"scope": "openid",
}
headers = {"Content-Type": "application/x-www-form-urlencoded"}
last: httpx.Response | None = None
for attempt in range(TOKEN_FETCH_ATTEMPTS):
try:
last = await self._http.post(url, data=data, headers=headers)
except httpx.HTTPError as exc:
if attempt == TOKEN_FETCH_ATTEMPTS - 1:
raise ZertoError(f"Keycloak token request failed: {exc}") from exc
await asyncio.sleep(TOKEN_RETRY_BACKOFF_S * (attempt + 1))
continue
if last.status_code < 400:
payload = last.json()
self._token = payload["access_token"]
self._token_exp = time.time() + float(payload.get("expires_in") or 60)
self._token_gen += 1
return self._token
# 401/403 is the credentials themselves. Retrying hammers Keycloak
# and can trip its brute-force lockout on a real account, so do not.
if last.status_code in (401, 403) or attempt == TOKEN_FETCH_ATTEMPTS - 1:
break
await asyncio.sleep(TOKEN_RETRY_BACKOFF_S * (attempt + 1))
assert last is not None
# Keycloak names the cause: invalid_client is the wrong client_id,
# invalid_grant is the wrong username/password. Saying which saves
# chasing the wrong one.
detail = ""
try:
error = (last.json() or {}).get("error") or ""
except ValueError:
error = ""
if error == "invalid_client":
detail = (
f" Keycloak said invalid_client: client_id {self.client_id!r} is wrong "
"for this appliance (10.x uses zerto-client; 9.x may use zerto-api)."
)
payload = response.json()
self._token = payload["access_token"]
self._token_exp = time.time() + float(payload.get("expires_in") or 60)
return self._token
elif error == "invalid_grant":
detail = " Keycloak said invalid_grant: the username or password is wrong."
elif error:
detail = f" Keycloak said {error}."
raise ZertoError(
f"Keycloak token failed HTTP {last.status_code}."
f"{detail or ' Check username/password/client_id.'}",
status_code=last.status_code,
body=last.text[:500],
)
async def ensure_token(self) -> str:
if self._token_is_fresh():
assert self._token is not None
return self._token
async with self._token_lock:
# Another task may have fetched one while we waited for the lock.
if self._token_is_fresh():
assert self._token is not None
return self._token
return await self._fetch_token()
async def _reauth(self, seen_gen: int) -> str:
"""Re-auth after a 401, unless another task already did it."""
async with self._token_lock:
if self._token_gen != seen_gen and self._token:
return self._token
self._token = None
return await self._fetch_token()
async def request(
self,
@@ -81,6 +143,7 @@ class ZertoClient:
json_body: Any = None,
) -> httpx.Response:
token = await self.ensure_token()
gen = self._token_gen
headers = {"Authorization": f"Bearer {token}"}
if json_body is not None:
headers["Content-Type"] = "application/json"
@@ -92,8 +155,7 @@ class ZertoClient:
headers=headers,
)
if response.status_code == 401:
self._token = None
token = await self.ensure_token()
token = await self._reauth(gen)
headers["Authorization"] = f"Bearer {token}"
response = await self._http.request(
method,
@@ -217,7 +279,7 @@ class ZertoClient:
async def fetch_download(self, token: str) -> bytes:
token = str(token).strip().strip('"')
if token.startswith("v1/") or token.startswith("/v1/"):
if token.startswith(("v1/", "/v1/")):
paths = [token if token.startswith("/") else f"/{token}"]
else:
paths = [f"/v1/downloads/{token}", f"/v1/flrs/{token}"]
@@ -226,9 +288,11 @@ class ZertoClient:
last = await self.request("GET", path)
if last.status_code < 400:
return last.content
status = last.status_code if last else None
body = last.text[:300] if last else ""
raise ZertoError(
f"FLR download HTTP {last.status_code if last else '?'}: {(last.text[:300] if last else '')}",
status_code=last.status_code if last else None,
f"FLR download HTTP {status if status is not None else '?'}: {body}",
status_code=status,
)
async def get_localsite(self) -> dict[str, Any]:
+120
View File
@@ -0,0 +1,120 @@
import asyncio
import pytest
from zerto_rewind_mcp.client import ZertoClient, ZertoError
class _FakeResponse:
def __init__(self, status_code, payload=None, text=""):
self.status_code = status_code
self._payload = payload
self.text = text
def json(self):
if self._payload is None:
raise ValueError("no json")
return self._payload
class _FakeHttp:
"""Counts token POSTs so we can prove they are serialised."""
def __init__(self, token_responses=None, delay=0.01):
self.token_posts = 0
self.delay = delay
self.token_responses = list(token_responses or [])
async def post(self, url, data=None, headers=None):
self.token_posts += 1
await asyncio.sleep(self.delay) # widen the race window
if self.token_responses:
return self.token_responses.pop(0)
return _FakeResponse(200, {"access_token": f"tok{self.token_posts}", "expires_in": 300})
def _client(http):
c = ZertoClient(base_url="https://zvm", username="u", password="p")
c._http = http
return c
def test_concurrent_ensure_token_makes_one_request():
"""The bug: N concurrent calls each POST to Keycloak and one gets 401."""
http = _FakeHttp()
client = _client(http)
async def run():
return await asyncio.gather(*(client.ensure_token() for _ in range(10)))
tokens = asyncio.run(run())
assert http.token_posts == 1, f"expected 1 token request, got {http.token_posts}"
assert set(tokens) == {"tok1"}
def test_cached_token_is_reused_without_a_request():
http = _FakeHttp()
client = _client(http)
asyncio.run(client.ensure_token())
asyncio.run(client.ensure_token())
assert http.token_posts == 1
def test_reauth_is_skipped_when_another_task_already_refreshed():
http = _FakeHttp()
client = _client(http)
first = asyncio.run(client.ensure_token())
stale_gen = client._token_gen - 1 # pretend we held the previous token
async def run():
return await client._reauth(stale_gen)
again = asyncio.run(run())
# someone else already moved past our generation: reuse, do not re-auth
assert again == first
assert http.token_posts == 1
def test_reauth_refetches_when_our_token_is_the_current_one():
http = _FakeHttp()
client = _client(http)
asyncio.run(client.ensure_token())
gen = client._token_gen
token = asyncio.run(client._reauth(gen))
assert token == "tok2"
assert http.token_posts == 2
def test_bad_credentials_are_not_retried():
"""Retrying a 401 risks tripping Keycloak brute-force lockout."""
http = _FakeHttp(token_responses=[_FakeResponse(401, {"error": "invalid_grant"}, "denied")])
client = _client(http)
with pytest.raises(ZertoError) as err:
asyncio.run(client.ensure_token())
assert http.token_posts == 1
assert "invalid_grant" in str(err.value)
assert "username or password" in str(err.value)
def test_wrong_client_id_is_named():
http = _FakeHttp(token_responses=[_FakeResponse(401, {"error": "invalid_client"}, "nope")])
client = _client(http)
with pytest.raises(ZertoError) as err:
asyncio.run(client.ensure_token())
assert "invalid_client" in str(err.value)
assert "zerto-client" in str(err.value)
def test_transient_server_error_is_retried():
http = _FakeHttp(
token_responses=[
_FakeResponse(503, None, "busy"),
_FakeResponse(200, {"access_token": "ok", "expires_in": 300}),
]
)
client = _client(http)
import zerto_rewind_mcp.client as mod
mod.TOKEN_RETRY_BACKOFF_S = 0
assert asyncio.run(client.ensure_token()) == "ok"
assert http.token_posts == 2