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