diff --git a/packages/google-auth/docs/reference/google.oauth2.sts.rst b/packages/google-auth/docs/reference/google.oauth2.sts.rst index 49d99dfe66f0..a9bd3e3f2d0b 100644 --- a/packages/google-auth/docs/reference/google.oauth2.sts.rst +++ b/packages/google-auth/docs/reference/google.oauth2.sts.rst @@ -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: diff --git a/packages/google-auth/google/oauth2/sts.py b/packages/google-auth/google/oauth2/sts.py index a48db0a4580b..c7047cf45d1d 100644 --- a/packages/google-auth/google/oauth2/sts.py +++ b/packages/google-auth/google/oauth2/sts.py @@ -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"} @@ -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(): + 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 "") + ) + + 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, diff --git a/packages/google-auth/google/oauth2/utils.py b/packages/google-auth/google/oauth2/utils.py index d72ff1916631..f57afd979c94 100644 --- a/packages/google-auth/google/oauth2/utils.py +++ b/packages/google-auth/google/oauth2/utils.py @@ -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 @@ -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) diff --git a/packages/google-auth/tests/oauth2/test_sts.py b/packages/google-auth/tests/oauth2/test_sts.py index acbe485f3d2f..fa3809e894f5 100644 --- a/packages/google-auth/tests/oauth2/test_sts.py +++ b/packages/google-auth/tests/oauth2/test_sts.py @@ -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() diff --git a/packages/google-auth/tests/oauth2/test_utils.py b/packages/google-auth/tests/oauth2/test_utils.py index ea845aa6046a..36df73ec4238 100644 --- a/packages/google-auth/tests/oauth2/test_utils.py +++ b/packages/google-auth/tests/oauth2/test_utils.py @@ -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