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:
@@ -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],
|
||||
)
|
||||
payload = response.json()
|
||||
}
|
||||
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,
|
||||
@@ -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]:
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user