121 lines
3.7 KiB
Python
121 lines
3.7 KiB
Python
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
|