|
| 1 | +from concurrent.futures import ThreadPoolExecutor |
1 | 2 | from pathlib import Path, PurePath |
| 3 | +from threading import Event, Lock, Thread |
2 | 4 |
|
3 | 5 | import pytest |
4 | 6 | import json |
5 | 7 | import jwt |
6 | 8 | import requests |
7 | 9 | from unittest.mock import patch, mock_open, Mock |
8 | 10 |
|
| 11 | +from requests import Request |
9 | 12 | from requests.auth import HTTPBasicAuth |
10 | 13 |
|
11 | 14 | from stackit.core.auth_methods.key_auth import KeyAuth, ServiceAccountKey |
@@ -294,3 +297,88 @@ def set_initial_token(auth): |
294 | 297 | auth._KeyAuth__refresh_token() |
295 | 298 |
|
296 | 299 | assert mock_post.call_count == KeyAuth.MAX_REFRESH_RETRIES |
| 300 | + |
| 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 | + |
| 339 | + def test_token_is_expired_before_expiration_with_leeway(self): |
| 340 | + auth = object.__new__(KeyAuth) |
| 341 | + secret = "x" * 32 |
| 342 | + |
| 343 | + with patch("stackit.core.auth_methods.key_auth.time.time", return_value=1_000): |
| 344 | + token_inside_leeway = jwt.encode({"exp": 1_240}, secret, algorithm="HS256") |
| 345 | + token_outside_leeway = jwt.encode({"exp": 1_360}, secret, algorithm="HS256") |
| 346 | + token_without_expiration = jwt.encode({}, secret, algorithm="HS256") |
| 347 | + token_with_zero_expiration = jwt.encode({"exp": 0}, secret, algorithm="HS256") |
| 348 | + |
| 349 | + assert auth._KeyAuth__is_token_expired(token_inside_leeway) |
| 350 | + assert not auth._KeyAuth__is_token_expired(token_outside_leeway) |
| 351 | + assert auth._KeyAuth__is_token_expired(token_without_expiration) |
| 352 | + assert auth._KeyAuth__is_token_expired(token_with_zero_expiration) |
| 353 | + |
| 354 | + def test_call_waits_for_an_in_progress_refresh_before_setting_header(self): |
| 355 | + auth = object.__new__(KeyAuth) |
| 356 | + auth.lock = Lock() |
| 357 | + auth.access_token = jwt.encode({"exp": 0}, "x" * 32, algorithm="HS256") |
| 358 | + fresh_token = jwt.encode({"exp": 4_000_000_000}, "x" * 32, algorithm="HS256") |
| 359 | + refresh_started = Event() |
| 360 | + release_refresh = Event() |
| 361 | + |
| 362 | + class BlockingFuture: |
| 363 | + def done(self): |
| 364 | + return False |
| 365 | + |
| 366 | + def result(self): |
| 367 | + refresh_started.set() |
| 368 | + release_refresh.wait(timeout=1) |
| 369 | + with auth.lock: |
| 370 | + auth.access_token = fresh_token |
| 371 | + |
| 372 | + auth.refresh_future = BlockingFuture() |
| 373 | + request = Request("GET", "https://example.com") |
| 374 | + call_thread = Thread(target=auth, args=(request,)) |
| 375 | + call_thread.start() |
| 376 | + |
| 377 | + assert refresh_started.wait(timeout=1) |
| 378 | + assert call_thread.is_alive() |
| 379 | + |
| 380 | + release_refresh.set() |
| 381 | + call_thread.join(timeout=1) |
| 382 | + |
| 383 | + assert not call_thread.is_alive() |
| 384 | + assert request.headers["Authorization"] == f"Bearer {fresh_token}" |
0 commit comments