Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions packages/google-auth/docs/reference/google.oauth2.sts.rst
Original file line number Diff line number Diff line change
@@ -1,6 +1,12 @@
google.oauth2.sts module
========================

The client retries transient HTTP and OAuth error responses with exponential
backoff, for up to three attempts. If these attempts are exhausted, the raised
``google.auth.exceptions.OAuthError`` has ``retryable=True``. Permanent error
responses, such as invalid credentials, are not retried. Exceptions raised by
the transport propagate unchanged.

.. automodule:: google.oauth2.sts
:members:
:inherited-members:
Expand Down
54 changes: 29 additions & 25 deletions packages/google-auth/google/oauth2/sts.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,8 @@
import json
import urllib

from google.oauth2 import utils
from google.auth import _exponential_backoff
from google.oauth2 import _client, utils

_URLENCODED_HEADERS = {"Content-Type": "application/x-www-form-urlencoded"}

Expand Down Expand Up @@ -71,30 +72,33 @@ def _make_request(self, request, headers, request_body, url=None):
# Use default token exchange endpoint if no url is provided.
url = url or self._token_exchange_endpoint

# Execute request.
response = request(
url=url,
method="POST",
headers=request_headers,
body=urllib.parse.urlencode(request_body).encode("utf-8"),
)

response_body = (
response.data.decode("utf-8")
if hasattr(response.data, "decode")
else response.data
)

# If non-200 response received, translate to OAuthError exception.
if response.status != http_client.OK:
utils.handle_error_response(response_body)

# A successful token revocation returns an empty response body.
if not response_body:
return {}

# Other successful responses should be valid JSON.
return json.loads(response_body)
encoded_body = urllib.parse.urlencode(request_body).encode("utf-8")
for _ in _exponential_backoff.ExponentialBackoff():
Comment on lines +75 to +76

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

To prevent potential UnboundLocalError exceptions if the _exponential_backoff.ExponentialBackoff() generator is empty or mocked to return no elements, initialize response_body and retryable before entering the loop.

        encoded_body = urllib.parse.urlencode(request_body).encode("utf-8")
        response_body = ""
        retryable = False
        for _ in _exponential_backoff.ExponentialBackoff():

response = request(
url=url,
method="POST",
headers=request_headers,
body=encoded_body,
)
response_body = (
response.data.decode("utf-8")
if hasattr(response.data, "decode")
else (response.data or "")
)
Comment on lines +83 to +87

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

If response.data is None (which can happen with empty responses or in certain mock/error scenarios), response_body will be assigned None. This will cause a TypeError when json.loads(response_body) is called later in this function, or in utils.handle_error_response(), because json.loads does not accept None. Falling back to an empty string "" avoids this issue and allows the JSON parser to raise a ValueError (or JSONDecodeError), which is already correctly handled by the try/except blocks.

Suggested change
response_body = (
response.data.decode("utf-8")
if hasattr(response.data, "decode")
else response.data
)
response_body = (
response.data.decode("utf-8")
if hasattr(response.data, "decode")
else (response.data or "")
)


if response.status == http_client.OK:
# A successful token revocation returns an empty body.
return json.loads(response_body) if response_body else {}

try:
response_data = json.loads(response_body)
except ValueError:
response_data = response_body
retryable = _client._can_retry(response.status, response_data)
if not retryable:
break

utils.handle_error_response(response_body, retryable=retryable)

def exchange_token(
self,
Expand Down
5 changes: 3 additions & 2 deletions packages/google-auth/google/oauth2/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -141,12 +141,13 @@ def _inject_authenticated_request_body(self, request_body):
)


def handle_error_response(response_body):
def handle_error_response(response_body, retryable=False):
"""Translates an error response from an OAuth operation into an
OAuthError exception.

Args:
response_body (str): The decoded response data.
retryable (bool): Whether the operation may be retried. Defaults to False.

Raises:
google.auth.exceptions.OAuthError
Expand All @@ -165,4 +166,4 @@ def handle_error_response(response_body):
except (KeyError, ValueError):
error_details = response_body

raise exceptions.OAuthError(error_details, response_body)
raise exceptions.OAuthError(error_details, response_body, retryable=retryable)
77 changes: 77 additions & 0 deletions packages/google-auth/tests/oauth2/test_sts.py
Original file line number Diff line number Diff line change
Expand Up @@ -547,3 +547,80 @@ def test__make_request_empty_response(self):
response = client._make_request(request, {}, {})

assert response == {}

@pytest.mark.parametrize("status,retryable", [(400, False), (503, True)])
@mock.patch("time.sleep", return_value=None)
def test_missing_response_body(self, sleep, status, retryable):
request = self.make_mock_request("", status, use_json=False)
request.return_value.data = None

with pytest.raises(exceptions.OAuthError) as caught:
self.make_client()._make_request(request, {}, {})

assert caught.value.args[1] == ""
assert caught.value.retryable is retryable
assert request.call_count == (3 if retryable else 1)
assert sleep.call_count == (2 if retryable else 0)

@pytest.mark.parametrize(
"status,data,use_json",
[
(500, {"error": "server_error"}, True),
(503, "Service unavailable", False),
(400, {"error": "temporarily_unavailable"}, True),
],
)
@mock.patch("time.sleep", return_value=None)
def test_transient_response_retried(self, sleep, status, data, use_json):
client = self.make_client(self.CLIENT_AUTH_REQUEST_BODY)
request = self.make_mock_request(data, status, use_json)
success = self.make_mock_request(self.SUCCESS_RESPONSE).return_value
request.side_effect = [request.return_value, success]

assert (
client._make_request(request, {"a": "b"}, {"c": "d"})
== self.SUCCESS_RESPONSE
)
assert request.call_count == 2
assert request.call_args_list[0] == request.call_args_list[1]
assert sleep.call_count == 1

@pytest.mark.parametrize(
"status,data,use_json,retryable",
[
(500, {"error": "server_error"}, True, True),
(503, "Service unavailable", False, True),
(400, {"error": "temporarily_unavailable"}, True, True),
(400, {"error": "invalid_grant"}, True, False),
(401, "Unauthorized", False, False),
],
)
@mock.patch("time.sleep", return_value=None)
def test_response_retryability(self, sleep, status, data, use_json, retryable):
request = self.make_mock_request(data, status, use_json)
with pytest.raises(exceptions.OAuthError) as caught:
self.make_client()._make_request(request, {}, {})

assert caught.value.retryable is retryable
assert caught.value.args[1] == request.return_value.data.decode("utf-8")
assert request.call_count == (3 if retryable else 1)
assert sleep.call_count == (2 if retryable else 0)

@mock.patch("time.sleep", return_value=None)
def test_retry_stops_on_permanent_error(self, sleep):
request = self.make_mock_request({"error": "server_error"}, 500)
permanent = self.make_mock_request(self.ERROR_RESPONSE, 400).return_value
request.side_effect = [request.return_value, permanent]
with pytest.raises(exceptions.OAuthError) as caught:
self.make_client()._make_request(request, {}, {})
assert caught.value.retryable is False
assert request.call_count == 2
assert sleep.call_count == 1

def test_transport_error_propagates_unchanged(self):
error = exceptions.TransportError("connection failed")
request = mock.Mock(side_effect=error)
with pytest.raises(exceptions.TransportError) as caught:
self.make_client()._make_request(request, {}, {})
assert caught.value is error
request.assert_called_once()
8 changes: 8 additions & 0 deletions packages/google-auth/tests/oauth2/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -261,3 +261,11 @@ def test__handle_error_response_non_json():
utils.handle_error_response(response_data)

assert excinfo.match(r"Oops, something wrong happened")


@pytest.mark.parametrize("retryable", [True, False])
def test__handle_error_response_retryable(retryable):
response_data = json.dumps({"error": "server_error"})
with pytest.raises(exceptions.OAuthError) as caught:
utils.handle_error_response(response_data, retryable=retryable)
assert caught.value.retryable is retryable
Loading