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
357 lines
13 KiB
Python
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,
|
|
)
|