fix(client): serialise Keycloak token acquisition #7
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from urllib.parse import urljoin
|
from urllib.parse import urljoin
|
||||||
@@ -10,6 +11,9 @@ import httpx
|
|||||||
|
|
||||||
from zerto_rewind_mcp.util import pick
|
from zerto_rewind_mcp.util import pick
|
||||||
|
|
||||||
|
TOKEN_FETCH_ATTEMPTS = 3
|
||||||
|
TOKEN_RETRY_BACKOFF_S = 0.5
|
||||||
|
|
||||||
|
|
||||||
class ZertoError(Exception):
|
class ZertoError(Exception):
|
||||||
def __init__(self, message: str, status_code: int | None = None, body: str | None = None):
|
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.client_id = client_id
|
||||||
self._token: str | None = None
|
self._token: str | None = None
|
||||||
self._token_exp = 0.0
|
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)
|
self._http = httpx.AsyncClient(verify=verify_tls, timeout=timeout)
|
||||||
|
|
||||||
async def aclose(self) -> None:
|
async def aclose(self) -> None:
|
||||||
@@ -44,33 +56,83 @@ class ZertoClient:
|
|||||||
return path
|
return path
|
||||||
return urljoin(self.base_url + "/", path.lstrip("/"))
|
return urljoin(self.base_url + "/", path.lstrip("/"))
|
||||||
|
|
||||||
async def ensure_token(self) -> str:
|
def _token_is_fresh(self) -> bool:
|
||||||
if self._token and time.time() < self._token_exp - 30:
|
return bool(self._token) and time.time() < self._token_exp - 30
|
||||||
return self._token
|
|
||||||
|
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")
|
url = self._url("/auth/realms/zerto/protocol/openid-connect/token")
|
||||||
response = await self._http.post(
|
data = {
|
||||||
url,
|
"grant_type": "password",
|
||||||
data={
|
"username": self.username,
|
||||||
"grant_type": "password",
|
"password": self.password,
|
||||||
"username": self.username,
|
"client_id": self.client_id,
|
||||||
"password": self.password,
|
"scope": "openid",
|
||||||
"client_id": self.client_id,
|
}
|
||||||
"scope": "openid",
|
headers = {"Content-Type": "application/x-www-form-urlencoded"}
|
||||||
},
|
last: httpx.Response | None = None
|
||||||
headers={"Content-Type": "application/x-www-form-urlencoded"},
|
for attempt in range(TOKEN_FETCH_ATTEMPTS):
|
||||||
)
|
try:
|
||||||
if response.status_code >= 400:
|
last = await self._http.post(url, data=data, headers=headers)
|
||||||
raise ZertoError(
|
except httpx.HTTPError as exc:
|
||||||
f"Keycloak token failed HTTP {response.status_code}. "
|
if attempt == TOKEN_FETCH_ATTEMPTS - 1:
|
||||||
"Check username/password/client_id "
|
raise ZertoError(f"Keycloak token request failed: {exc}") from exc
|
||||||
"(10.x uses zerto-client; 9.x may use zerto-api).",
|
await asyncio.sleep(TOKEN_RETRY_BACKOFF_S * (attempt + 1))
|
||||||
status_code=response.status_code,
|
continue
|
||||||
body=response.text[:500],
|
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()
|
elif error == "invalid_grant":
|
||||||
self._token = payload["access_token"]
|
detail = " Keycloak said invalid_grant: the username or password is wrong."
|
||||||
self._token_exp = time.time() + float(payload.get("expires_in") or 60)
|
elif error:
|
||||||
return self._token
|
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(
|
async def request(
|
||||||
self,
|
self,
|
||||||
@@ -81,6 +143,7 @@ class ZertoClient:
|
|||||||
json_body: Any = None,
|
json_body: Any = None,
|
||||||
) -> httpx.Response:
|
) -> httpx.Response:
|
||||||
token = await self.ensure_token()
|
token = await self.ensure_token()
|
||||||
|
gen = self._token_gen
|
||||||
headers = {"Authorization": f"Bearer {token}"}
|
headers = {"Authorization": f"Bearer {token}"}
|
||||||
if json_body is not None:
|
if json_body is not None:
|
||||||
headers["Content-Type"] = "application/json"
|
headers["Content-Type"] = "application/json"
|
||||||
@@ -92,8 +155,7 @@ class ZertoClient:
|
|||||||
headers=headers,
|
headers=headers,
|
||||||
)
|
)
|
||||||
if response.status_code == 401:
|
if response.status_code == 401:
|
||||||
self._token = None
|
token = await self._reauth(gen)
|
||||||
token = await self.ensure_token()
|
|
||||||
headers["Authorization"] = f"Bearer {token}"
|
headers["Authorization"] = f"Bearer {token}"
|
||||||
response = await self._http.request(
|
response = await self._http.request(
|
||||||
method,
|
method,
|
||||||
@@ -217,7 +279,7 @@ class ZertoClient:
|
|||||||
|
|
||||||
async def fetch_download(self, token: str) -> bytes:
|
async def fetch_download(self, token: str) -> bytes:
|
||||||
token = str(token).strip().strip('"')
|
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}"]
|
paths = [token if token.startswith("/") else f"/{token}"]
|
||||||
else:
|
else:
|
||||||
paths = [f"/v1/downloads/{token}", f"/v1/flrs/{token}"]
|
paths = [f"/v1/downloads/{token}", f"/v1/flrs/{token}"]
|
||||||
@@ -226,9 +288,11 @@ class ZertoClient:
|
|||||||
last = await self.request("GET", path)
|
last = await self.request("GET", path)
|
||||||
if last.status_code < 400:
|
if last.status_code < 400:
|
||||||
return last.content
|
return last.content
|
||||||
|
status = last.status_code if last else None
|
||||||
|
body = last.text[:300] if last else ""
|
||||||
raise ZertoError(
|
raise ZertoError(
|
||||||
f"FLR download HTTP {last.status_code if last else '?'}: {(last.text[:300] if last else '')}",
|
f"FLR download HTTP {status if status is not None else '?'}: {body}",
|
||||||
status_code=last.status_code if last else None,
|
status_code=status,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def get_localsite(self) -> dict[str, Any]:
|
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