|
| 1 | +import unittest |
| 2 | +from unittest import mock |
| 3 | + |
| 4 | +import pytest |
| 5 | +from google.api_core import exceptions |
| 6 | +from google.api_core.retry.retry_streaming_async import AsyncStreamingRetry |
| 7 | + |
| 8 | +from google.cloud.storage._experimental.asyncio.retry import manager |
| 9 | +from google.cloud.storage._experimental.asyncio.retry import strategy |
| 10 | + |
| 11 | + |
| 12 | +def _is_retriable(exc): |
| 13 | + return isinstance(exc, exceptions.ServiceUnavailable) |
| 14 | + |
| 15 | + |
| 16 | +DEFAULT_TEST_RETRY = AsyncStreamingRetry(predicate=_is_retriable, deadline=1) |
| 17 | + |
| 18 | + |
| 19 | +class TestBidiStreamRetryManager(unittest.IsolatedAsyncioTestCase): |
| 20 | + async def test_execute_success_on_first_try(self): |
| 21 | + """Verify the manager correctly handles a stream that succeeds immediately.""" |
| 22 | + mock_strategy = mock.AsyncMock(spec=strategy._BaseResumptionStrategy) |
| 23 | + |
| 24 | + async def mock_stream_opener(*args, **kwargs): |
| 25 | + yield "response_1" |
| 26 | + |
| 27 | + retry_manager = manager._BidiStreamRetryManager( |
| 28 | + strategy=mock_strategy, |
| 29 | + stream_opener=mock_stream_opener, |
| 30 | + retry_policy=DEFAULT_TEST_RETRY, |
| 31 | + ) |
| 32 | + |
| 33 | + await retry_manager.execute(initial_state={}) |
| 34 | + |
| 35 | + mock_strategy.generate_requests.assert_called_once() |
| 36 | + mock_strategy.update_state_from_response.assert_called_once_with( |
| 37 | + "response_1", {} |
| 38 | + ) |
| 39 | + mock_strategy.recover_state_on_failure.assert_not_called() |
| 40 | + |
| 41 | + async def test_execute_retries_and_succeeds(self): |
| 42 | + """Verify the manager retries on a transient error and then succeeds.""" |
| 43 | + mock_strategy = mock.AsyncMock(spec=strategy._BaseResumptionStrategy) |
| 44 | + |
| 45 | + attempt_count = 0 |
| 46 | + |
| 47 | + async def mock_stream_opener(*args, **kwargs): |
| 48 | + nonlocal attempt_count |
| 49 | + attempt_count += 1 |
| 50 | + if attempt_count == 1: |
| 51 | + raise exceptions.ServiceUnavailable("Service is down") |
| 52 | + else: |
| 53 | + yield "response_2" |
| 54 | + |
| 55 | + retry_manager = manager._BidiStreamRetryManager( |
| 56 | + strategy=mock_strategy, |
| 57 | + stream_opener=mock_stream_opener, |
| 58 | + retry_policy=AsyncStreamingRetry(predicate=_is_retriable, initial=0.01), |
| 59 | + ) |
| 60 | + |
| 61 | + await retry_manager.execute(initial_state={}) |
| 62 | + |
| 63 | + self.assertEqual(attempt_count, 2) |
| 64 | + self.assertEqual(mock_strategy.generate_requests.call_count, 2) |
| 65 | + mock_strategy.recover_state_on_failure.assert_called_once() |
| 66 | + mock_strategy.update_state_from_response.assert_called_once_with( |
| 67 | + "response_2", {} |
| 68 | + ) |
| 69 | + |
| 70 | + async def test_execute_fails_after_deadline_exceeded(self): |
| 71 | + """Verify the manager raises RetryError if the deadline is exceeded.""" |
| 72 | + mock_strategy = mock.AsyncMock(spec=strategy._BaseResumptionStrategy) |
| 73 | + |
| 74 | + async def mock_stream_opener(*args, **kwargs): |
| 75 | + raise exceptions.ServiceUnavailable("Service is always down") |
| 76 | + |
| 77 | + # Use a very short deadline to make the test fast. |
| 78 | + fast_retry = AsyncStreamingRetry( |
| 79 | + predicate=_is_retriable, deadline=0.1, initial=0.05 |
| 80 | + ) |
| 81 | + retry_manager = manager._BidiStreamRetryManager( |
| 82 | + strategy=mock_strategy, |
| 83 | + stream_opener=mock_stream_opener, |
| 84 | + retry_policy=fast_retry, |
| 85 | + ) |
| 86 | + |
| 87 | + with pytest.raises(exceptions.RetryError, match="Deadline of 0.1s exceeded"): |
| 88 | + await retry_manager.execute(initial_state={}) |
| 89 | + |
| 90 | + # Verify it attempted to recover state after each failure. |
| 91 | + self.assertGreater(mock_strategy.recover_state_on_failure.call_count, 1) |
| 92 | + |
| 93 | + async def test_execute_fails_immediately_on_non_retriable_error(self): |
| 94 | + """Verify the manager aborts immediately on a non-retriable error.""" |
| 95 | + mock_strategy = mock.AsyncMock(spec=strategy._BaseResumptionStrategy) |
| 96 | + |
| 97 | + async def mock_stream_opener(*args, **kwargs): |
| 98 | + raise exceptions.PermissionDenied("Auth error") |
| 99 | + |
| 100 | + retry_manager = manager._BidiStreamRetryManager( |
| 101 | + strategy=mock_strategy, |
| 102 | + stream_opener=mock_stream_opener, |
| 103 | + retry_policy=DEFAULT_TEST_RETRY, |
| 104 | + ) |
| 105 | + |
| 106 | + with pytest.raises(exceptions.PermissionDenied): |
| 107 | + await retry_manager.execute(initial_state={}) |
| 108 | + |
| 109 | + # Verify that it did not try to recover or update state. |
| 110 | + mock_strategy.recover_state_on_failure.assert_not_called() |
| 111 | + mock_strategy.update_state_from_response.assert_not_called() |
0 commit comments