From ebf714cc208e9c58f60467317d566587f1dd75aa Mon Sep 17 00:00:00 2001 From: "Claude (agent)" Date: Mon, 21 Sep 2026 15:15:42 -0400 Subject: [PATCH] fix(client): serialise Keycloak token acquisition (#7) --- src/zerto_rewind_mcp/client.py | 124 +++++++++++++++++++++++++-------- tests/test_client_auth.py | 120 +++++++++++++++++++++++++++++++ 2 files changed, 214 insertions(+), 30 deletions(-) create mode 100644 tests/test_client_auth.py diff --git a/src/zerto_rewind_mcp/client.py b/src/zerto_rewind_mcp/client.py index 9d68d5e..6d3ad20 100644 --- a/src/zerto_rewind_mcp/client.py +++ b/src/zerto_rewind_mcp/client.py @@ -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]: diff --git a/tests/test_client_auth.py b/tests/test_client_auth.py new file mode 100644 index 0000000..022df45 --- /dev/null +++ b/tests/test_client_auth.py @@ -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