Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,8 @@
# See the License for the specific language governing permissions and
# limitations under the License.

import asyncio
import logging
from typing import List, Optional, Tuple

from google.api_core.bidi_async import AsyncBidiRpc
Expand All @@ -22,6 +24,8 @@
)
from google.cloud.storage.asyncio.async_grpc_client import AsyncGrpcClient

logger = logging.getLogger(__name__)


class _AsyncReadObjectStream(_AsyncAbstractObjectStream):
"""Class representing a gRPC bidi-stream for reading data from a GCS ``Object``.
Expand Down Expand Up @@ -126,40 +130,53 @@ async def open(self, metadata: Optional[List[Tuple[str, str]]] = None) -> None:
initial_request=self.first_bidi_read_req,
metadata=current_metadata,
)
await self.socket_like_rpc.open() # this is actually 1 send
response = await self.socket_like_rpc.recv()
# populated only in the first response of bidi-stream and when opened
# without using `read_handle`
if hasattr(response, "metadata") and response.metadata:
if self.generation_number is None:
self.generation_number = response.metadata.generation
# update persisted size
self.persisted_size = response.metadata.size
self.object_metadata = response.metadata
if (
hasattr(response.metadata, "finalize_time")
and response.metadata.finalize_time
and response.metadata.finalize_time.second > 0
):
self.is_finalized = True
try:
await self.socket_like_rpc.open() # this is actually 1 send
response = await self.socket_like_rpc.recv()
# populated only in the first response of bidi-stream and when opened
# without using `read_handle`
if hasattr(response, "metadata") and response.metadata:
if self.generation_number is None:
self.generation_number = response.metadata.generation
# update persisted size
self.persisted_size = response.metadata.size
self.object_metadata = response.metadata
if (
hasattr(response.metadata, "checksums")
and response.metadata.checksums
hasattr(response.metadata, "finalize_time")
and response.metadata.finalize_time
and response.metadata.finalize_time.second > 0
):
self.full_obj_server_crc32c = response.metadata.checksums.crc32c

if response and response.read_handle:
self.read_handle = response.read_handle

self._is_stream_open = True
self.is_finalized = True
if (
hasattr(response.metadata, "checksums")
and response.metadata.checksums
):
self.full_obj_server_crc32c = response.metadata.checksums.crc32c

if response and response.read_handle:
self.read_handle = response.read_handle

self._is_stream_open = True
except (asyncio.CancelledError, Exception):
await self._close_socket_like_rpc()
raise

async def _close_socket_like_rpc(self) -> None:
try:
await self.socket_like_rpc.close()
except Exception as exc:
logger.debug("Error while closing the read bidi-gRPC stream: %s", exc)

async def close(self) -> None:
"""Closes the bidi-gRPC connection."""
if not self._is_stream_open:
raise ValueError("Stream is not open")
await self.requests_done()
await self.socket_like_rpc.close()
self._is_stream_open = False
try:
if self.socket_like_rpc.is_active:
await self.requests_done()
finally:
self._is_stream_open = False
await self._close_socket_like_rpc()

async def requests_done(self):
"""Signals that all requests have been sent."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,10 +12,12 @@
# See the License for the specific language governing permissions and
# limitations under the License.

import asyncio
from unittest import mock
from unittest.mock import AsyncMock

import pytest
from google.api_core.exceptions import Aborted

from google.cloud import _storage_v2
from google.cloud.storage.asyncio import async_read_object_stream
Expand Down Expand Up @@ -203,6 +205,123 @@ async def test_close(mock_client, mock_cls_async_bidi_rpc):
assert not read_obj_stream.is_stream_open


@mock.patch("google.cloud.storage.asyncio.async_read_object_stream.AsyncBidiRpc")
@mock.patch(
"google.cloud.storage.asyncio.async_grpc_client.AsyncGrpcClient.grpc_client"
)
@pytest.mark.asyncio
async def test_open_closes_rpc_when_first_recv_fails(
mock_client, mock_cls_async_bidi_rpc
):
read_obj_stream = await instantiate_read_obj_stream(
mock_client, mock_cls_async_bidi_rpc, open=False
)
socket_like_rpc = mock_cls_async_bidi_rpc.return_value
socket_like_rpc.recv = AsyncMock(
side_effect=Aborted("Idle stream has been closed.")
)

with pytest.raises(Aborted):
await read_obj_stream.open()

socket_like_rpc.open.assert_awaited_once()
socket_like_rpc.close.assert_awaited_once()
assert not read_obj_stream.is_stream_open


@mock.patch("google.cloud.storage.asyncio.async_read_object_stream.AsyncBidiRpc")
@mock.patch(
"google.cloud.storage.asyncio.async_grpc_client.AsyncGrpcClient.grpc_client"
)
@pytest.mark.parametrize("cancelled_await", ["open", "recv"])
@pytest.mark.asyncio
async def test_open_closes_rpc_when_cancelled(
mock_client, mock_cls_async_bidi_rpc, cancelled_await
):
read_obj_stream = await instantiate_read_obj_stream(
mock_client, mock_cls_async_bidi_rpc, open=False
)
socket_like_rpc = mock_cls_async_bidi_rpc.return_value
setattr(
socket_like_rpc,
cancelled_await,
AsyncMock(side_effect=asyncio.CancelledError),
)

with pytest.raises(asyncio.CancelledError):
await read_obj_stream.open()

socket_like_rpc.close.assert_awaited_once()
assert not read_obj_stream.is_stream_open


@mock.patch("google.cloud.storage.asyncio.async_read_object_stream.AsyncBidiRpc")
@mock.patch(
"google.cloud.storage.asyncio.async_grpc_client.AsyncGrpcClient.grpc_client"
)
@pytest.mark.asyncio
async def test_open_propagates_close_failure_from_failed_open(
mock_client, mock_cls_async_bidi_rpc
):
read_obj_stream = await instantiate_read_obj_stream(
mock_client, mock_cls_async_bidi_rpc, open=False
)
socket_like_rpc = mock_cls_async_bidi_rpc.return_value
socket_like_rpc.recv = AsyncMock(
side_effect=Aborted("Idle stream has been closed.")
)
socket_like_rpc.close = AsyncMock(side_effect=RuntimeError("close blew up"))

with pytest.raises(Aborted):
await read_obj_stream.open()

assert not read_obj_stream.is_stream_open


@mock.patch("google.cloud.storage.asyncio.async_read_object_stream.AsyncBidiRpc")
@mock.patch(
"google.cloud.storage.asyncio.async_grpc_client.AsyncGrpcClient.grpc_client"
)
@pytest.mark.asyncio
async def test_close_closes_rpc_when_requests_done_fails(
mock_client, mock_cls_async_bidi_rpc
):
read_obj_stream = await instantiate_read_obj_stream(
mock_client, mock_cls_async_bidi_rpc, open=True
)
socket_like_rpc = read_obj_stream.socket_like_rpc
read_obj_stream.requests_done = AsyncMock(
side_effect=Aborted("Idle stream has been closed.")
)

with pytest.raises(Aborted):
await read_obj_stream.close()

socket_like_rpc.close.assert_awaited_once()
assert not read_obj_stream.is_stream_open


@mock.patch("google.cloud.storage.asyncio.async_read_object_stream.AsyncBidiRpc")
@mock.patch(
"google.cloud.storage.asyncio.async_grpc_client.AsyncGrpcClient.grpc_client"
)
@pytest.mark.asyncio
async def test_close_skips_requests_done_on_inactive_rpc(
mock_client, mock_cls_async_bidi_rpc
):
read_obj_stream = await instantiate_read_obj_stream(
mock_client, mock_cls_async_bidi_rpc, open=True
)
read_obj_stream.socket_like_rpc.is_active = False
read_obj_stream.requests_done = AsyncMock()

await read_obj_stream.close()

read_obj_stream.requests_done.assert_not_called()
read_obj_stream.socket_like_rpc.close.assert_awaited_once()
assert not read_obj_stream.is_stream_open


@mock.patch("google.cloud.storage.asyncio.async_read_object_stream.AsyncBidiRpc")
@mock.patch(
"google.cloud.storage.asyncio.async_grpc_client.AsyncGrpcClient.grpc_client"
Expand Down
Loading