2323"""
2424from typing import List , Optional , Tuple
2525from google .cloud import _storage_v2
26+ from google .cloud .storage ._experimental .asyncio import _utils
2627from google .cloud .storage ._experimental .asyncio .async_grpc_client import AsyncGrpcClient
2728from google .cloud .storage ._experimental .asyncio .async_abstract_object_stream import (
2829 _AsyncAbstractObjectStream ,
@@ -59,7 +60,7 @@ class _AsyncWriteObjectStream(_AsyncAbstractObjectStream):
5960 same name already exists, it will be overwritten the moment
6061 `writer.open()` is called.
6162
62- :type write_handle: bytes
63+ :type write_handle: _storage_v2.BidiWriteHandle
6364 :param write_handle: (Optional) An existing handle for writing the object.
6465 If provided, opening the bidi-gRPC connection will be faster.
6566 """
@@ -70,7 +71,7 @@ def __init__(
7071 bucket_name : str ,
7172 object_name : str ,
7273 generation_number : Optional [int ] = None , # None means new object
73- write_handle : Optional [bytes ] = None ,
74+ write_handle : Optional [_storage_v2 . BidiWriteHandle ] = None ,
7475 routing_token : Optional [str ] = None ,
7576 ) -> None :
7677 if client is None :
@@ -86,7 +87,7 @@ def __init__(
8687 generation_number = generation_number ,
8788 )
8889 self .client : AsyncGrpcClient .grpc_client = client
89- self .write_handle : Optional [bytes ] = write_handle
90+ self .write_handle : Optional [_storage_v2 . BidiWriteHandle ] = write_handle
9091 self .routing_token : Optional [str ] = routing_token
9192
9293 self ._full_bucket_name = f"projects/_/buckets/{ self .bucket_name } "
@@ -117,8 +118,6 @@ async def open(self, metadata: Optional[List[Tuple[str, str]]] = None) -> None:
117118 if self ._is_stream_open :
118119 raise ValueError ("Stream is already open" )
119120
120- write_handle = self .write_handle if self .write_handle else None
121-
122121 # Create a new object or overwrite existing one if generation_number
123122 # is None. This makes it consistent with GCS JSON API behavior.
124123 # Created object type would be Appendable Object.
@@ -140,27 +139,26 @@ async def open(self, metadata: Optional[List[Tuple[str, str]]] = None) -> None:
140139 bucket = self ._full_bucket_name ,
141140 object = self .object_name ,
142141 generation = self .generation_number ,
143- write_handle = write_handle ,
142+ write_handle = self . write_handle if self . write_handle else None ,
144143 routing_token = self .routing_token if self .routing_token else None ,
145144 ),
146145 )
147146
148- request_params = [f"bucket={ self ._full_bucket_name } " ]
149- other_metadata = []
147+ request_param_values = [f"bucket={ self ._full_bucket_name } " ]
148+ final_metadata = []
150149 if metadata :
151150 for key , value in metadata :
152151 if key == "x-goog-request-params" :
153- request_params .append (value )
152+ request_param_values .append (value )
154153 else :
155- other_metadata .append ((key , value ))
154+ final_metadata .append ((key , value ))
156155
157- current_metadata = other_metadata
158- current_metadata .append (("x-goog-request-params" , "," .join (request_params )))
156+ final_metadata .append (("x-goog-request-params" , "," .join (request_param_values )))
159157
160158 self .socket_like_rpc = AsyncBidiRpc (
161159 self .rpc ,
162160 initial_request = self .first_bidi_write_req ,
163- metadata = current_metadata ,
161+ metadata = final_metadata ,
164162 )
165163
166164 await self .socket_like_rpc .open () # this is actually 1 send
@@ -194,7 +192,7 @@ async def requests_done(self):
194192 """Signals that all requests have been sent."""
195193
196194 await self .socket_like_rpc .send (None )
197- await self .socket_like_rpc .recv ()
195+ _utils . update_write_handle_if_exists ( self , await self .socket_like_rpc .recv () )
198196
199197 async def send (
200198 self , bidi_write_object_request : _storage_v2 .BidiWriteObjectRequest
@@ -236,4 +234,3 @@ async def recv(self) -> _storage_v2.BidiWriteObjectResponse:
236234 @property
237235 def is_stream_open (self ) -> bool :
238236 return self ._is_stream_open
239-
0 commit comments