fix(client): serialise Keycloak token acquisition (#7)

This commit was merged in pull request #7.
This commit is contained in:
2026-09-21 15:15:42 -04:00
parent 1d53029038
commit ebf714cc20
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]: