|
| 1 | +from concurrent.futures import ThreadPoolExecutor |
1 | 2 | from pathlib import Path, PurePath |
2 | 3 | from threading import Event, Lock, Thread |
3 | 4 |
|
@@ -299,6 +300,42 @@ def set_initial_token(auth): |
299 | 300 |
|
300 | 301 |
|
301 | 302 | class TestKeyAuth: |
| 303 | + @pytest.mark.parametrize("refresh_already_failed", [False, True]) |
| 304 | + def test_call_propagates_token_refresh_failure(self, refresh_already_failed): |
| 305 | + auth = object.__new__(KeyAuth) |
| 306 | + auth.lock = Lock() |
| 307 | + auth.access_token = jwt.encode({"exp": 0}, "x" * 32, algorithm="HS256") |
| 308 | + auth.refresh_token = jwt.encode({"exp": 4_000_000_000}, "x" * 32, algorithm="HS256") |
| 309 | + auth.token_endpoint = KeyAuth.DEFAULT_TOKEN_ENDPOINT |
| 310 | + auth.refresh_future = None |
| 311 | + original_token = auth.access_token |
| 312 | + request = Request("GET", "https://example.com") |
| 313 | + refresh_error = requests.RequestException("refresh failed") |
| 314 | + |
| 315 | + with ( |
| 316 | + ThreadPoolExecutor(max_workers=1) as executor, |
| 317 | + patch("stackit.core.auth_methods.key_auth.requests.post", side_effect=refresh_error) as mock_post, |
| 318 | + ): |
| 319 | + auth.executor = executor |
| 320 | + if refresh_already_failed: |
| 321 | + auth.refresh_future = executor.submit(auth._KeyAuth__refresh_token) |
| 322 | + with pytest.raises(requests.RequestException, match="Token refresh failed after retries"): |
| 323 | + auth.refresh_future.result(timeout=1) |
| 324 | + |
| 325 | + with pytest.raises(requests.RequestException, match="Token refresh failed after retries") as exc_info: |
| 326 | + auth(request) |
| 327 | + |
| 328 | + assert exc_info.value.__cause__ is refresh_error |
| 329 | + assert auth.refresh_future.done() |
| 330 | + assert mock_post.call_count == KeyAuth.MAX_REFRESH_RETRIES |
| 331 | + mock_post.assert_called_with( |
| 332 | + auth.token_endpoint, |
| 333 | + data={"grant_type": "refresh_token", "refresh_token": auth.refresh_token}, |
| 334 | + timeout=auth.timeout, |
| 335 | + ) |
| 336 | + assert auth.access_token == original_token |
| 337 | + assert "Authorization" not in request.headers |
| 338 | + |
302 | 339 | def test_token_is_expired_before_expiration_with_leeway(self): |
303 | 340 | auth = object.__new__(KeyAuth) |
304 | 341 | secret = "x" * 32 |
|
0 commit comments