diff --git a/sdk/core/azure-core/azure/core/pipeline/policies/_retry.py b/sdk/core/azure-core/azure/core/pipeline/policies/_retry.py index 7e073f61f1ae..77c7f3f17f64 100644 --- a/sdk/core/azure-core/azure/core/pipeline/policies/_retry.py +++ b/sdk/core/azure-core/azure/core/pipeline/policies/_retry.py @@ -28,6 +28,7 @@ This module is the requests implementation of Pipeline ABC """ from __future__ import absolute_import # we have a "requests" module that conflicts with "requests" on Py2.7 +from io import SEEK_SET, UnsupportedOperation import logging import time import email @@ -317,7 +318,21 @@ def increment(self, settings, response=None, error=None): settings['status'] -= 1 settings['history'].append(RequestHistory(response.http_request, http_response=response.http_response)) - return not self.is_exhausted(settings) + if not self.is_exhausted(settings): + if response.http_request.body and hasattr(response.http_request.body, 'read'): + try: + body_position = response.request.http_request.body.tell() + except (AttributeError, UnsupportedOperation): + # if body position cannot be obtained, then retries will not work + return False + try: + # attempt to rewind the body to the initial position + response.http_request.body.seek(settings['body_position'], SEEK_SET) + except (UnsupportedOperation, ValueError, AttributeError) as err: + # if body is not seekable, then retry would not work + return False + return True + return False def update_context(self, context, retry_settings): """Updates retry history in pipeline context. diff --git a/sdk/core/azure-core/tests/test_universal_pipeline.py b/sdk/core/azure-core/tests/test_universal_pipeline.py index 32f73c2d5daf..d20ad1f5e075 100644 --- a/sdk/core/azure-core/tests/test_universal_pipeline.py +++ b/sdk/core/azure-core/tests/test_universal_pipeline.py @@ -31,7 +31,10 @@ import mock import requests - +try: + from io import BytesIO +except ImportError: + from cStringIO import StringIO as BytesIO import pytest from azure.core.exceptions import DecodeError @@ -52,6 +55,7 @@ ContentDecodePolicy, UserAgentPolicy, HttpLoggingPolicy, + RetryPolicy, ) def test_user_agent(): @@ -130,6 +134,34 @@ def test_no_log(mock_http_logger): second_count = mock_http_logger.debug.call_count assert second_count == first_count * 2 +def test_retry_seekable_body(): + def build_response(body, content_type=None): + class MockResponse(HttpResponse): + def __init__(self): + super(MockResponse, self).__init__(None, None) + self._body = 'test' + + def body(self): + return self._body + + data = BytesIO(b"Lots of dataaaa") + universal_request = HttpRequest('GET', 'http://127.0.0.1/', data=data) + universal_request.set_streamed_data_body(data) + return PipelineResponse(universal_request, MockResponse(), PipelineContext(None, stream=True)) + + response = build_response(b"", content_type="application/xml") + http_retry = RetryPolicy() + setting = { + 'total': 3, + 'status': 3, + 'history': [], + 'connect': 3, + 'read': 3, + 'body_position': 10, + } + increment = http_retry.increment(setting, response) + assert increment + def test_raw_deserializer(): raw_deserializer = ContentDecodePolicy()