Skip to content

Commit dc4538c

Browse files
fix(auth): refresh token before expiration, use locks while refreshing (#9045)
STACKITSDK-412
1 parent c704ae3 commit dc4538c

2 files changed

Lines changed: 113 additions & 12 deletions

File tree

‎core/src/stackit/core/auth_methods/key_auth.py‎

Lines changed: 25 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -79,6 +79,16 @@ def __call__(self, r: Request) -> Request:
7979
if self.__is_token_expired(self.access_token):
8080
if self.refresh_future is None or self.refresh_future.done():
8181
self.refresh_future = self.executor.submit(self.__refresh_token)
82+
refresh_future = self.refresh_future
83+
else:
84+
refresh_future = None
85+
86+
# Do not hold the lock while waiting. The refresh worker acquires it when
87+
# updating the access token.
88+
if refresh_future is not None:
89+
refresh_future.result()
90+
91+
with self.lock:
8292
r.headers["Authorization"] = f"Bearer {self.access_token}"
8393
return r
8494

@@ -108,8 +118,9 @@ def __fetch_token_from_endpoint(self) -> None:
108118
response = requests.post(self.token_endpoint, data=body, timeout=self.timeout)
109119
response.raise_for_status()
110120
response_json = response.json()
111-
self.access_token = response_json["access_token"]
112-
self.refresh_token = response_json["refresh_token"]
121+
with self.lock:
122+
self.access_token = response_json["access_token"]
123+
self.refresh_token = response_json["refresh_token"]
113124
except requests.RequestException as e:
114125
raise requests.RequestException("Initial token fetch failed") from e
115126

@@ -130,14 +141,17 @@ def token_refresh_task():
130141
thread.start()
131142

132143
def __refresh_token(self):
133-
if self.__is_token_expired(self.refresh_token):
144+
with self.lock:
145+
refresh_token = self.refresh_token
146+
147+
if self.__is_token_expired(refresh_token):
134148
self.__create_initial_token()
135149
self.__fetch_token_from_endpoint()
136150
return
137151

138152
body = {
139153
"grant_type": "refresh_token",
140-
"refresh_token": self.refresh_token,
154+
"refresh_token": refresh_token,
141155
}
142156

143157
last_exception = None
@@ -147,24 +161,23 @@ def __refresh_token(self):
147161
response.raise_for_status()
148162
response_data = response.json()
149163
new_token = response_data.get("access_token")
150-
self.access_token = new_token
164+
with self.lock:
165+
self.access_token = new_token
151166
return
152167
except requests.RequestException as e:
153168
last_exception = e
154169

155170
raise requests.RequestException("Token refresh failed after retries") from last_exception
156171

157-
def __is_token_expired(self, token: str) -> bool:
172+
def __is_token_expired(self, token: Optional[str]) -> bool:
158173
try:
159174
decoded_token = jwt.decode(token, options={"verify_signature": False})
160175
exp = decoded_token.get("exp")
161-
if exp:
162-
return time.time() > (exp + self.EXPIRATION_LEEWAY.total_seconds())
163-
except jwt.ExpiredSignatureError:
164-
return True
165-
except jwt.DecodeError:
176+
if exp is None:
177+
return True
178+
return time.time() > (float(exp) - self.EXPIRATION_LEEWAY.total_seconds())
179+
except (jwt.InvalidTokenError, TypeError, ValueError):
166180
return True
167-
return False
168181

169182
def __shutdown(self):
170183
self.executor.shutdown(wait=False)

‎core/tests/core/test_auth.py‎

Lines changed: 88 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,14 @@
1+
from concurrent.futures import ThreadPoolExecutor
12
from pathlib import Path, PurePath
3+
from threading import Event, Lock, Thread
24

35
import pytest
46
import json
57
import jwt
68
import requests
79
from unittest.mock import patch, mock_open, Mock
810

11+
from requests import Request
912
from requests.auth import HTTPBasicAuth
1013

1114
from stackit.core.auth_methods.key_auth import KeyAuth, ServiceAccountKey
@@ -294,3 +297,88 @@ def set_initial_token(auth):
294297
auth._KeyAuth__refresh_token()
295298

296299
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

Comments
 (0)