Files
zerto-ai-rewind/src/zerto_rewind_mcp/client.py
T
justinandClaude Opus 5 8ac89c7f82 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
2026-09-21 15:12:28 -04:00

357 lines
13 KiB
Python

"""ZVM / ZCA REST client. Same paths on both (zerto_api_lessons)."""
from __future__ import annotations
import asyncio
import time
from typing import Any
from urllib.parse import urljoin
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):
super().__init__(message)
self.status_code = status_code
self.body = body
class ZertoClient:
def __init__(
self,
base_url: str,
username: str,
password: str,
client_id: str = "zerto-client",
verify_tls: bool = False,
timeout: float = 60.0,
):
self.base_url = base_url.rstrip("/")
self.username = username
self.password = password
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:
await self._http.aclose()
def _url(self, path: str) -> str:
if path.startswith("http"):
return path
return urljoin(self.base_url + "/", path.lstrip("/"))
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")
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)."
)
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,
method: str,
path: str,
*,
params: dict[str, Any] | None = None,
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"
response = await self._http.request(
method,
self._url(path),
params=params,
json=json_body,
headers=headers,
)
if response.status_code == 401:
token = await self._reauth(gen)
headers["Authorization"] = f"Bearer {token}"
response = await self._http.request(
method,
self._url(path),
params=params,
json=json_body,
headers=headers,
)
return response
async def json(
self,
method: str,
path: str,
*,
params: dict[str, Any] | None = None,
json_body: Any = None,
) -> Any:
response = await self.request(method, path, params=params, json_body=json_body)
if response.status_code >= 400:
raise ZertoError(
f"{method} {path} HTTP {response.status_code}: {response.text[:400]}",
status_code=response.status_code,
body=response.text[:1000],
)
if response.status_code == 204 or not response.content:
return None
try:
return response.json()
except ValueError:
return response.text
async def get_vms(
self,
*,
vm_name: str | None = None,
vm_identifier: str | None = None,
) -> list[dict[str, Any]]:
params: dict[str, Any] = {}
if vm_identifier:
params["vmIdentifier"] = vm_identifier
if vm_name:
params["vmName"] = vm_name
data = await self.json("GET", "/v1/vms", params=params or None)
if data is None:
return []
if isinstance(data, list):
return data
if isinstance(data, dict):
inner = pick(data, "value", "items", "vms")
if isinstance(inner, list):
return inner
raise ZertoError(f"GET /v1/vms returned unexpected shape: {type(data).__name__}")
async def get_vpg(self, vpg_identifier: str) -> dict[str, Any]:
data = await self.json("GET", f"/v1/vpgs/{vpg_identifier}")
if not isinstance(data, dict):
raise ZertoError(f"GET /v1/vpgs/{vpg_identifier} did not return an object")
return data
async def list_checkpoints(self, vpg_identifier: str) -> list[dict[str, Any]]:
data = await self.json("GET", f"/v1/vpgs/{vpg_identifier}/checkpoints")
if data is None:
return []
if isinstance(data, list):
return data
raise ZertoError("GET checkpoints returned unexpected shape")
async def insert_checkpoint(self, vpg_identifier: str, tag: str) -> Any:
"""POST tagged checkpoint. Body key is CheckpointName on documented 9.x API."""
path = f"/v1/vpgs/{vpg_identifier}/checkpoints"
try:
return await self.json("POST", path, json_body={"CheckpointName": tag})
except ZertoError as exc:
if exc.status_code != 400:
raise
return await self.json("POST", path, json_body={"checkpointName": tag})
async def get_task(self, task_id: str) -> dict[str, Any]:
data = await self.json("GET", f"/v1/tasks/{task_id}")
if not isinstance(data, dict):
raise ZertoError("GET /v1/tasks/{id} did not return an object")
return data
async def start_flr(
self,
vpg_identifier: str,
vm_identifier: str,
checkpoint_identifier: str,
initial_download_path: str,
) -> Any:
return await self.json(
"POST",
"/v1/flrs",
json_body={
"jflr": {
"vpgIdentifier": vpg_identifier,
"vmIdentifier": vm_identifier,
"CheckpointIdentifier": checkpoint_identifier,
"initialDownloadPath": initial_download_path,
}
},
)
async def get_flr(self, session_id: str) -> Any:
return await self.json("GET", f"/v1/flrs/{session_id}")
async def browse_flr(self, session_id: str, path: str = "", recursive: bool = False) -> Any:
return await self.json(
"POST",
f"/v1/flrs/{session_id}/browse",
json_body={"path": path, "recursive": recursive},
)
async def download_flr(self, session_id: str, path_list: list[str]) -> Any:
return await self.json(
"POST",
f"/v1/flrs/{session_id}/download",
json_body={"pathList": path_list},
)
async def fetch_download(self, token: str) -> bytes:
token = str(token).strip().strip('"')
if token.startswith(("v1/", "/v1/")):
paths = [token if token.startswith("/") else f"/{token}"]
else:
paths = [f"/v1/downloads/{token}", f"/v1/flrs/{token}"]
last = None
for path in paths:
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 {status if status is not None else '?'}: {body}",
status_code=status,
)
async def get_localsite(self) -> dict[str, Any]:
data = await self.json("GET", "/v1/localsite")
return data if isinstance(data, dict) else {}
async def get_peersites(self) -> list[dict[str, Any]]:
data = await self.json("GET", "/v1/peersites")
return [r for r in data if isinstance(r, dict)] if isinstance(data, list) else []
async def list_flrs(self) -> Any:
"""GET /v1/flrs. Every FLR session the ZVM currently knows about."""
return await self.json("GET", "/v1/flrs")
async def end_flr(self, session_id: str) -> None:
try:
await self.json("DELETE", f"/v1/flrs/{session_id}")
except ZertoError:
await self.json("POST", f"/v1/flrs/{session_id}")
async def start_failover_test(
self,
vpg_identifier: str,
checkpoint_identifier: str | None = None,
vm_identifiers: list[str] | None = None,
) -> Any:
body: dict[str, Any] = {}
if checkpoint_identifier:
body["checkpointIdentifier"] = checkpoint_identifier
if vm_identifiers:
body["vmIdentifiers"] = vm_identifiers
return await self.json(
"POST",
f"/v1/vpgs/{vpg_identifier}/FailoverTest",
json_body=body or None,
)
async def stop_failover_test(self, vpg_identifier: str, success: bool, summary: str) -> Any:
return await self.json(
"POST",
f"/v1/vpgs/{vpg_identifier}/FailoverTestStop",
json_body={"failoverTestSuccess": success, "failoverTestSummary": summary},
)
async def start_clone(
self,
vpg_identifier: str,
checkpoint_identifier: str,
datastore_identifier: str | None = None,
vm_identifiers: list[str] | None = None,
) -> Any:
body: dict[str, Any] = {"checkpointIdentifier": checkpoint_identifier}
if datastore_identifier:
body["datastoreIdentifier"] = datastore_identifier
if vm_identifiers:
body["vmIdentifiers"] = vm_identifiers
return await self.json(
"POST",
f"/v1/vpgs/{vpg_identifier}/CloneStart",
json_body=body,
)