Skip to content
This repository was archived by the owner on Mar 31, 2026. It is now read-only.

Commit 92c852c

Browse files
committed
feat(experimental): Add bidi stream retry manager
1 parent dad4405 commit 92c852c

2 files changed

Lines changed: 178 additions & 0 deletions

File tree

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,67 @@
1+
import asyncio
2+
from typing import Any, AsyncIterator, Callable
3+
4+
from google.api_core import exceptions
5+
from google.cloud.storage._experimental.asyncio.retry.base_strategy import (
6+
_BaseResumptionStrategy,
7+
)
8+
9+
10+
class _BidiStreamRetryManager:
11+
"""Manages the generic retry loop for a bidi streaming operation."""
12+
13+
def __init__(
14+
self,
15+
strategy: _BaseResumptionStrategy,
16+
stream_opener: Callable[..., AsyncIterator[Any]],
17+
retry_policy,
18+
):
19+
"""Initializes the retry manager."""
20+
self._strategy = strategy
21+
self._stream_opener = stream_opener
22+
self._retry_policy = retry_policy
23+
24+
async def execute(self, initial_state: Any):
25+
"""
26+
Executes the bidi operation with the configured retry policy.
27+
28+
This method implements a manual retry loop that provides the necessary
29+
control points to manage state between attempts, which is not possible
30+
with a simple retry decorator.
31+
"""
32+
state = initial_state
33+
retry_policy = self._retry_policy
34+
35+
while True:
36+
try:
37+
# 1. Generate requests based on the current state.
38+
requests = self._strategy.generate_requests(state)
39+
40+
# 2. Open and consume the stream.
41+
stream = self._stream_opener(requests, state)
42+
async for response in stream:
43+
self._strategy.update_state_from_response(response, state)
44+
45+
# 3. If the stream completes without error, exit the loop.
46+
return
47+
48+
except Exception as e:
49+
# 4. If an error occurs, check if it's retriable.
50+
if not retry_policy.predicate(e):
51+
# If not retriable, fail fast.
52+
raise
53+
54+
# 5. If retriable, allow the strategy to recover state.
55+
# This is where routing tokens are extracted or QueryWriteStatus is called.
56+
try:
57+
await self._strategy.recover_state_on_failure(e, state)
58+
except Exception as recovery_exc:
59+
# If state recovery itself fails, we must abort.
60+
raise exceptions.RetryError(
61+
"Failed to recover state after a transient error.",
62+
cause=recovery_exc,
63+
) from recovery_exc
64+
65+
# 6. Use the policy to sleep and check for deadline expiration.
66+
# This will raise a RetryError if the deadline is exceeded.
67+
await asyncio.sleep(await retry_policy.sleep(e))
Lines changed: 111 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,111 @@
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

Comments
 (0)