diff --git a/packages/google-api-core/google/api_core/resumable_transfer/upload_async.py b/packages/google-api-core/google/api_core/resumable_transfer/upload_async.py index 13cfcb5c6fbd..d4c3e5da5ecb 100644 --- a/packages/google-api-core/google/api_core/resumable_transfer/upload_async.py +++ b/packages/google-api-core/google/api_core/resumable_transfer/upload_async.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Asynchronous Resumable Upload session and helpers using aiohttp.""" +"""Asynchronous Resumable Upload session and helpers.""" import asyncio import datetime @@ -28,6 +28,7 @@ Awaitable, BinaryIO, Callable, + Dict, Generator, Generic, Iterable, @@ -37,6 +38,7 @@ Tuple, TypeVar, Union, + cast, ) try: @@ -46,6 +48,14 @@ except ImportError: # pragma: NO COVER _HAS_AIOHTTP = False +try: + from google.auth import exceptions as auth_exceptions + from google.auth.aio.transport.sessions import AsyncAuthorizedSession + + _HAS_GOOGLE_AUTH_AIO = True +except ImportError: # pragma: NO COVER + _HAS_GOOGLE_AUTH_AIO = False + import google.api_core.retry from google.api_core import exceptions from google.api_core.resumable_transfer import common, upload_state @@ -60,6 +70,20 @@ _monotonic_clock = time.monotonic +AsyncTransport = Union["aiohttp.ClientSession", "AsyncAuthorizedSession"] +"""Asynchronous HTTP transports accepted by :class:`AsyncResumableUploadSession`. + +``google.auth.aio.transport.sessions.AsyncAuthorizedSession`` is the transport +used by generated asynchronous REST clients. +""" + +# google-auth >= 2.49.0 accepts ``total_attempts`` to cap AsyncAuthorizedSession's +# own status-code retries; older releases forward unknown kwargs to aiohttp. +_HAS_TOTAL_ATTEMPTS = _HAS_GOOGLE_AUTH_AIO and ( + "total_attempts" in inspect.signature(AsyncAuthorizedSession.request).parameters +) + + def _get_buffer_size(stream: object) -> Optional[int]: """Returns buffer size in bytes if stream exposes getbuffer(), else None.""" getbuffer_fn = getattr(stream, "getbuffer", None) @@ -166,7 +190,7 @@ def __init__( self, upload_url: Optional[str] = None, config: Optional[ResumableUploadConfig] = None, - transport: Optional[Any] = None, + transport: Optional[AsyncTransport] = None, content_type: Optional[str] = None, response_type: Optional[Any] = None, start_retry: Optional[google.api_core.retry.AsyncRetry] = None, @@ -179,7 +203,7 @@ def __init__( new upload, or the pre-existing upload session URL when resuming. config: Optional upload configuration parameters. Defaults to ``ResumableUploadConfig()`` when ``None``. - transport: Optional aiohttp.ClientSession. When ``None``, a + transport: Optional :data:`AsyncTransport`. When ``None``, a transport must be provided to ``upload()`` or ``resume()``. content_type: Optional MIME type of the stream payload. When ``None``, no content-type header is sent unless overridden. @@ -276,6 +300,93 @@ def _ensure_aiohttp(self) -> None: "Please install google-api-core[async_rest]." ) + def _get_transport(self, transport: Optional[AsyncTransport]) -> AsyncTransport: + """Returns the transport to use, raising ValueError when none is set.""" + sess = transport or self._transport + if sess is None: + raise ValueError( + "An aiohttp.ClientSession or AsyncAuthorizedSession transport " + "must be provided." + ) + return sess + + async def _send_request( + self, + transport: AsyncTransport, + method: str, + url: str, + payload: bytes, + headers: Mapping[str, str], + timeout: float, + ) -> Tuple[int, Mapping[str, str], bytes]: + """Sends one request over ``transport`` and reads the whole response. + + Args: + transport: The transport to dispatch the request over. + method: HTTP method verb. + url: Request URL. + payload: Request body bytes. + headers: Request headers. + timeout: Total timeout in seconds for this attempt. + + Returns: + Tuple of (status code, response headers, response body bytes). + + Raises: + exceptions.GoogleAPICallError: If the status is not 200 or 201. + asyncio.TimeoutError: If the attempt times out, whichever transport + is used. + """ + if _HAS_GOOGLE_AUTH_AIO and isinstance(transport, AsyncAuthorizedSession): + # google-auth applies its own overall deadline (``max_allowed_time``, + # 180s by default) on top of the per-request ``timeout``; use the + # caller's timeout for both. Where supported, a single attempt also + # disables google-auth's own status-code retries so that this module + # is the only retry layer. + retry_kwargs: Dict[str, Any] = ( + {"total_attempts": 1} if _HAS_TOTAL_ATTEMPTS else {} + ) + try: + response = await transport.request( + method, + url, + data=payload, + headers=headers, + timeout=timeout, + max_allowed_time=timeout, + **retry_kwargs, + ) + try: + status_code = response.status_code + resp_headers: Mapping[str, str] = dict(response.headers) + body = await response.read() + finally: + await response.close() + except auth_exceptions.TimeoutError as exc: + # google-auth raises its own TimeoutError (not an asyncio one, and + # not reliably chained to one); re-raise it as asyncio.TimeoutError + # so stall control sees the same type as with aiohttp.ClientSession. + raise asyncio.TimeoutError(str(exc)) from exc + else: + # mypy cannot narrow the negative of the guarded check above. + session = cast("aiohttp.ClientSession", transport) + async with session.request( + method, + url, + data=payload, + headers=headers, + timeout=aiohttp.ClientTimeout(total=timeout), + ) as resp: + status_code = resp.status + resp_headers = dict(resp.headers) + body = await resp.read() + + if status_code not in (200, 201): + raise exceptions.from_http_status( + status_code, body.decode("utf-8", errors="replace") + ) + return status_code, resp_headers, body + def _enrich_exception(self, exc: BaseException) -> None: """Attaches session diagnostic metadata to an active exception. @@ -366,7 +477,9 @@ def _get_retry_predicate( ``UnseekableStreamError``) always return ``False``. 2. Protocol-recoverable errors during chunk transfer (``RECOVERABLE_STATUS_CODES`` and ``MissingStatusHeaderError``) - and transport errors (``aiohttp.ClientError`` and + and transport errors (``aiohttp.ClientError``, + ``google.auth.exceptions.TransportError``, + ``google.auth.exceptions.ResponseError``, and ``asyncio.TimeoutError``) always return ``True`` so the session can query server state and recover. @@ -389,6 +502,10 @@ def should_retry(exc: Exception) -> bool: _HAS_AIOHTTP and isinstance(exc, aiohttp.ClientError) ): return True + if _HAS_GOOGLE_AUTH_AIO and isinstance( + exc, (auth_exceptions.TransportError, auth_exceptions.ResponseError) + ): + return True if ( custom_predicate is not None and custom_predicate is not google.api_core.retry.if_transient_error @@ -453,7 +570,7 @@ def _get_async_streaming_retry( async def _initiate( self, - transport: Any, + transport: AsyncTransport, request_body: Union[str, bytes] = "", size: Optional[int] = None, progress_queue: Optional[List[UploadProgress]] = None, @@ -462,7 +579,7 @@ async def _initiate( """Initiates the upload session asynchronously. Args: - transport: The aiohttp client session. + transport: The asynchronous transport to dispatch requests over. request_body: Initial metadata payload sent with start request. size: Total stream size in bytes, if known. progress_queue: Optional queue to receive progress event. @@ -483,21 +600,15 @@ async def _initiate( ) async def do_initiate() -> str: - timeout_sec = self._get_start_timeout() - client_timeout = aiohttp.ClientTimeout(total=timeout_sec) - async with transport.request( - method, url, data=payload, headers=headers, timeout=client_timeout - ) as resp: - resp_headers = dict(resp.headers) - body = await resp.read() - if resp.status not in (200, 201): - raise exceptions.from_http_status( - resp.status, body.decode("utf-8", errors="replace") - ) - session_url = self._state.process_start_response( - resp.status, resp_headers - ) - return session_url + status_code, resp_headers, _ = await self._send_request( + transport, + method, + url, + payload, + headers, + timeout=self._get_start_timeout(), + ) + return self._state.process_start_response(status_code, resp_headers) retry_policy = self._get_async_retry() retryable_initiate = retry_policy(do_initiate) @@ -507,7 +618,7 @@ async def do_initiate() -> str: async def _transmit_chunk( self, - transport: Any, + transport: AsyncTransport, reader_fn: Callable[[int], Awaitable[bytes]], size: Optional[int], progress_queue: Optional[List[UploadProgress]] = None, @@ -516,7 +627,7 @@ async def _transmit_chunk( """Transmits the next data chunk asynchronously with stall control. Args: - transport: The aiohttp client session. + transport: The asynchronous transport to dispatch requests over. reader_fn: Async callable returning chunk bytes. size: Total stream size in bytes, if known. progress_queue: Optional queue to receive progress updates. @@ -563,23 +674,16 @@ async def _transmit_chunk( per_attempt_timeout = self._compute_chunk_timeout( data_len, timeout_override=timeout ) - client_timeout = aiohttp.ClientTimeout(total=per_attempt_timeout) t_start = _monotonic_clock() try: - async with transport.request( + status_code, resp_headers, resp_body = await self._send_request( + transport, method, url, - data=payload, - headers=headers, - timeout=client_timeout, - ) as resp: - resp_headers = dict(resp.headers) - resp_body = await resp.read() - if resp.status not in (200, 201): - raise exceptions.from_http_status( - resp.status, resp_body.decode("utf-8", errors="replace") - ) - status_code = resp.status + payload, + headers, + timeout=per_attempt_timeout, + ) t_elapsed = _monotonic_clock() - t_start except Exception as exc: self._enrich_exception(exc) @@ -684,14 +788,14 @@ def _update_stall_control(self, data_len: int, t_elapsed: float) -> None: async def _recover( self, - transport: Any, + transport: AsyncTransport, stream_obj: Optional[object] = None, progress_queue: Optional[List[UploadProgress]] = None, ) -> Tuple[int, Mapping[str, str], bytes]: """Queries server for committed byte offset and adjusts buffer. Args: - transport: The aiohttp client session. + transport: The asynchronous transport to dispatch requests over. stream_obj: Underlying stream object to rewind if seekable. progress_queue: Optional queue to receive progress updates. @@ -703,18 +807,14 @@ async def _recover( exceptions.GoogleAPICallError: If query request fails on the server. """ method, url, headers, payload = self._state.build_query_request() - timeout_sec = self._get_start_timeout() - client_timeout = aiohttp.ClientTimeout(total=timeout_sec) - async with transport.request( - method, url, data=payload, headers=headers, timeout=client_timeout - ) as resp: - resp_headers = dict(resp.headers) - body = await resp.read() - if resp.status not in (200, 201): - raise exceptions.from_http_status( - resp.status, body.decode("utf-8", errors="replace") - ) - status_code = resp.status + status_code, resp_headers, body = await self._send_request( + transport, + method, + url, + payload, + headers, + timeout=self._get_start_timeout(), + ) received = self._state.process_query_response(status_code, resp_headers) self._notify_progress(common.ProgressState.OFFSET_RECEIVED, progress_queue) @@ -760,37 +860,32 @@ async def _recover( chunk_size=self.chunk_size, ) - async def cancel(self, transport: Optional[Any] = None) -> None: + async def cancel(self, transport: Optional[AsyncTransport] = None) -> None: """Cancels the resumable upload session asynchronously. Args: - transport: Optional aiohttp client session. + transport: Optional aiohttp.ClientSession or AsyncAuthorizedSession. Raises: ValueError: If transport is missing. exceptions.GoogleAPICallError: If cancellation request fails on the server. """ self._ensure_aiohttp() - sess = transport or self._transport - if sess is None: - raise ValueError("An aiohttp.ClientSession transport must be provided.") + sess = self._get_transport(transport) method, url, headers, payload = self._state.build_cancel_request() - timeout_sec = self._get_start_timeout() - client_timeout = aiohttp.ClientTimeout(total=timeout_sec) - async with sess.request( - method, url, data=payload, headers=headers, timeout=client_timeout - ) as resp: - resp_headers = dict(resp.headers) - body = await resp.read() - if resp.status not in (200, 201): - raise exceptions.from_http_status( - resp.status, body.decode("utf-8", errors="replace") - ) - self._state.process_cancel_response(resp.status, resp_headers) + status_code, resp_headers, _ = await self._send_request( + sess, + method, + url, + payload, + headers, + timeout=self._get_start_timeout(), + ) + self._state.process_cancel_response(status_code, resp_headers) async def _transmit_all_chunks( self, - transport: Any, + transport: AsyncTransport, reader_fn: Callable[[int], Awaitable[bytes]], computed_size: Optional[int], progress_queue: Optional[List[UploadProgress]] = None, @@ -801,7 +896,7 @@ async def _transmit_all_chunks( """Transmits chunks until completion using a single outer AsyncStreamingRetry coordinator. Args: - transport: The aiohttp client session. + transport: The asynchronous transport to dispatch requests over. reader_fn: Async callable returning chunk bytes. computed_size: Total stream size in bytes, if known. progress_queue: Optional list receiving UploadProgress snapshots. @@ -887,7 +982,7 @@ def upload( stream: Union[AsyncIterable[bytes], BinaryIO, bytes, Iterable[bytes]], request_body: Union[str, bytes] = "", size: Optional[int] = None, - transport: Optional[Any] = None, + transport: Optional[AsyncTransport] = None, content_type: Optional[str] = None, retry: Optional[google.api_core.retry.AsyncStreamingRetry] = None, timeout: Optional[float] = None, @@ -898,7 +993,7 @@ def upload( stream: Data payload to upload (async iterable, binary stream, bytes, or iterable). request_body: Initial metadata payload sent with the start request. size: Total stream size in bytes, if known. - transport: Optional aiohttp client session. + transport: Optional aiohttp.ClientSession or AsyncAuthorizedSession. content_type: Optional MIME type of the stream payload. retry: Optional retry configuration (``AsyncStreamingRetry``) for chunk upload requests. Use this to customize exponential backoff timing between chunk retries or to @@ -919,9 +1014,7 @@ def upload( ValueError: If transport is missing. """ self._ensure_aiohttp() - sess = transport or self._transport - if sess is None: - raise ValueError("An aiohttp.ClientSession transport must be provided.") + sess = self._get_transport(transport) if content_type is not None: self._content_type = content_type @@ -963,7 +1056,7 @@ def resume( stream: Union[AsyncIterable[bytes], BinaryIO, bytes, Iterable[bytes]], size: Optional[int] = None, chunk_size: Optional[int] = None, - transport: Optional[Any] = None, + transport: Optional[AsyncTransport] = None, retry: Optional[google.api_core.retry.AsyncStreamingRetry] = None, timeout: Optional[float] = None, ) -> AsyncUploadOperation: @@ -975,7 +1068,7 @@ def resume( Data payload to resume uploading. size (Optional[int]): Total stream size in bytes, if known. chunk_size (Optional[int]): Optional chunk size override in bytes. - transport (Optional[Any]): Optional aiohttp client session. + transport (Optional[AsyncTransport]): Optional transport override. retry (Optional[google.api_core.retry.AsyncStreamingRetry]): Optional retry configuration (``AsyncStreamingRetry``) for chunk upload requests. Use this to customize exponential backoff timing @@ -998,9 +1091,7 @@ def resume( ValueError: If transport, upload_url, or stream is missing. """ self._ensure_aiohttp() - sess = transport or self._transport - if sess is None: - raise ValueError("An aiohttp.ClientSession transport must be provided.") + sess = self._get_transport(transport) actual_url = upload_url or self.upload_url if not actual_url: raise ValueError("An upload URL must be provided to resume.") diff --git a/packages/google-api-core/tests/asyncio/test_resumable_transfer_async.py b/packages/google-api-core/tests/asyncio/test_resumable_transfer_async.py index b16211795ee6..a55941495d01 100644 --- a/packages/google-api-core/tests/asyncio/test_resumable_transfer_async.py +++ b/packages/google-api-core/tests/asyncio/test_resumable_transfer_async.py @@ -14,8 +14,11 @@ """Asynchronous tests for Resumable Upload protocol implementation.""" +# mypy: disable-error-code="arg-type" + import asyncio import datetime +import inspect import io import json from typing import ( @@ -32,6 +35,7 @@ from unittest import mock import pytest +from google.auth import exceptions as auth_exceptions from google.protobuf import empty_pb2 from google.api_core import exceptions @@ -52,6 +56,8 @@ try: import aiohttp # noqa: F401 import google.auth.aio.transport # noqa: F401 + from google.auth.aio import credentials as aio_credentials + from google.auth.aio.transport import sessions as aio_sessions GOOGLE_AUTH_AIO_INSTALLED = True except ImportError: @@ -1325,12 +1331,12 @@ def test_async_transport_missing_errors() -> None: upload_url="https://api.example.com/start", ) with pytest.raises( - ValueError, match="An aiohttp.ClientSession transport must be provided" + ValueError, match="or AsyncAuthorizedSession transport must be provided" ): session.upload(stream=b"data") with pytest.raises( - ValueError, match="An aiohttp.ClientSession transport must be provided" + ValueError, match="or AsyncAuthorizedSession transport must be provided" ): session.resume(upload_url="https://upload.example.com/123", stream=b"data") @@ -1342,7 +1348,7 @@ async def test_async_cancel_missing_transport_and_error() -> None: upload_url="https://upload.example.com/123", ) with pytest.raises( - ValueError, match="An aiohttp.ClientSession transport must be provided" + ValueError, match="or AsyncAuthorizedSession transport must be provided" ): await session.cancel() @@ -2292,3 +2298,204 @@ def test_async_retry_predicate_includes_asyncio_timeout_error() -> None: session = AsyncResumableUploadSession(upload_url="https://api.example.com/start") predicate = session._get_retry_predicate() assert predicate(asyncio.TimeoutError()) is True + + +# ===================================================================== +# 11. google.auth AsyncAuthorizedSession Transport Tests +# ===================================================================== + + +START_HEADERS = { + "X-Goog-Upload-Status": "active", + "X-Goog-Upload-URL": "https://upload.example.com/123", +} + + +def _auth_response( + status_code: int = 200, + headers: Optional[Mapping[str, str]] = None, + body: bytes = b"", +) -> mock.Mock: + """Builds a ``google.auth.aio.transport.Response`` double.""" + return mock.Mock( + status_code=status_code, + headers=dict(headers or {}), + read=mock.AsyncMock(return_value=body), + close=mock.AsyncMock(), + ) + + +def _authorized_session( + *outcomes: Union[mock.Mock, BaseException], +) -> Tuple["aio_sessions.AsyncAuthorizedSession", mock.AsyncMock]: + """Builds a real AsyncAuthorizedSession whose HTTP adapter replays ``outcomes``. + + google-auth invokes the adapter positionally as + ``(url, method, body, headers, timeout, **kwargs)``. + """ + auth_request = mock.AsyncMock(side_effect=list(outcomes)) + transport = aio_sessions.AsyncAuthorizedSession( + aio_credentials.AnonymousCredentials(), auth_request=auth_request + ) + return transport, auth_request + + +def _adapter_commands(auth_request: mock.AsyncMock) -> List[str]: + return [c.args[3]["X-Goog-Upload-Command"] for c in auth_request.call_args_list] + + +def _supports_total_attempts() -> bool: + """Returns True if google-auth (>= 2.49.0) accepts ``total_attempts``.""" + params = inspect.signature(aio_sessions.AsyncAuthorizedSession.request).parameters + return "total_attempts" in params + + +@pytest.mark.asyncio +async def test_async_authorized_session_upload() -> None: + """Uploads through a real AsyncAuthorizedSession (the generated REST client transport). + + Unlike aiohttp.ClientSession, its ``request()`` is a coroutine returning a + ``google.auth.aio.transport.Response`` (``status_code``, explicit ``close()``). + """ + responses = [ + _auth_response(200, START_HEADERS), + _auth_response(200, {"X-Goog-Upload-Status": "active"}), + _auth_response( + 200, {"X-Goog-Upload-Status": "final"}, b'{"name": "auth.txt", "size": 6}' + ), + ] + transport, auth_request = _authorized_session(*responses) + session = AsyncResumableUploadSession( + upload_url="https://api.example.com/start", + config=ResumableUploadConfig(chunk_size=4), + transport=transport, + response_type=DummyResponse, + start_timeout=7.5, + ) + + with mock.patch.object(transport, "request", wraps=transport.request) as spy: + result = await session.upload(stream=b"012345") + + assert isinstance(result, DummyResponse) + assert (result.name, result.size) == ("auth.txt", 6) + assert session.bytes_uploaded == 6 + assert _adapter_commands(auth_request) == ["start", "upload", "upload, finalize"] + assert [c.args[2] for c in auth_request.call_args_list] == [b"", b"0123", b"45"] + assert [c.args[0] for c in auth_request.call_args_list] == [ + "https://api.example.com/start", + "https://upload.example.com/123", + "https://upload.example.com/123", + ] + for response in responses: + response.close.assert_awaited_once() + # api_core owns the timeout: google-auth's wall-clock guard is bounded by the + # same per-attempt timeout, and nothing extra leaks through to the adapter. + assert spy.call_args_list[0].kwargs["timeout"] == 7.5 + expected_total_attempts = 1 if _supports_total_attempts() else None + for call in spy.call_args_list: + assert call.kwargs["max_allowed_time"] == call.kwargs["timeout"] + assert call.kwargs.get("total_attempts") == expected_total_attempts + assert all(c.kwargs == {} for c in auth_request.call_args_list) + + +@pytest.mark.asyncio +async def test_async_authorized_session_server_error_not_retried_internally() -> None: + """Verifies a 5xx reaches api_core's retry layer instead of google-auth's. + + AsyncAuthorizedSession retries retryable status codes itself by default; + api_core passes ``total_attempts=1`` (google-auth >= 2.49.0) so that its + protocol-aware retry and offset recovery is the single retry layer. + """ + if not _supports_total_attempts(): + pytest.skip("google-auth < 2.49.0 retries 5xx responses internally") + response = _auth_response(503, body=b"Service Unavailable") + transport, auth_request = _authorized_session(response) + session = AsyncResumableUploadSession( + upload_url="https://api.example.com/start", + transport=transport, + start_retry=google.api_core.retry.AsyncRetry(predicate=lambda exc: False), + ) + + with pytest.raises(exceptions.ServiceUnavailable, match="Service Unavailable"): + await session.upload(stream=b"data") + + assert auth_request.await_count == 1 + response.close.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("on_read", [False, True], ids=["request", "read"]) +async def test_async_authorized_session_timeout_raises_transfer_stalled( + on_read: bool, +) -> None: + """Verifies google.auth.exceptions.TimeoutError feeds stall control. + + It is not a subclass of asyncio.TimeoutError, so it is normalized to one + (keeping the original as the cause) to surface as TransferStalledError. + """ + read_resp = _auth_response(200, {"X-Goog-Upload-Status": "final"}) + read_resp.read.side_effect = auth_exceptions.TimeoutError("read timed out") + chunk_outcome: Any = ( + read_resp if on_read else auth_exceptions.TimeoutError("timed out") + ) + transport, auth_request = _authorized_session( + _auth_response(200, START_HEADERS), chunk_outcome + ) + session = AsyncResumableUploadSession( + upload_url="https://api.example.com/start", transport=transport + ) + no_retry = google.api_core.retry.AsyncStreamingRetry( + predicate=lambda e: False, timeout=0.001 + ) + + with pytest.raises(TransferStalledError) as exc_info: + await session.upload(stream=b"data", retry=no_retry) + + cause = exc_info.value.__cause__ + assert isinstance(cause, asyncio.TimeoutError) + assert isinstance(cause.__cause__, auth_exceptions.TimeoutError) + assert _adapter_commands(auth_request) == ["start", "upload, finalize"] + if on_read: + read_resp.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_async_authorized_session_transport_error_recovers() -> None: + """Verifies google.auth.exceptions.TransportError triggers offset recovery.""" + transport, auth_request = _authorized_session( + _auth_response(200, START_HEADERS), + auth_exceptions.TransportError("connection reset"), + _auth_response( + 200, {"X-Goog-Upload-Status": "active", "X-Goog-Upload-Size-Received": "0"} + ), + _auth_response(200, {"X-Goog-Upload-Status": "final"}, b'{"name": "f"}'), + ) + session = AsyncResumableUploadSession( + upload_url="https://api.example.com/start", + transport=transport, + response_type=DummyResponse, + ) + + result = await session.upload( + stream=b"data", + retry=google.api_core.retry.AsyncStreamingRetry(initial=0.01, maximum=0.01), + ) + + assert isinstance(result, DummyResponse) + assert result.name == "f" + assert _adapter_commands(auth_request) == [ + "start", + "upload, finalize", + "query", + "upload, finalize", + ] + + +def test_async_retry_predicate_includes_google_auth_errors() -> None: + """Verifies google-auth transport errors are retryable, unlike its other errors.""" + session = AsyncResumableUploadSession(upload_url="https://api.example.com/start") + for is_start in (True, False): + predicate = session._get_retry_predicate(is_start=is_start) + assert predicate(auth_exceptions.TransportError()) is True + assert predicate(auth_exceptions.ResponseError()) is True + assert predicate(auth_exceptions.RefreshError()) is False