Skip to content
Merged
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 @@ -222,7 +222,10 @@ def _batch_send(
process_storage_error(error)

class TransportWrapper(HttpTransport):

"""Wrapper class that ensures that an inner client created
by a `get_client` method does not close the outer transport for the parent
when used in a context manager.
"""
def __init__(self, transport):
self._transport = transport

Expand All @@ -235,10 +238,10 @@ def open(self):
def close(self):
pass

def __enter__(self, *args): # pylint: disable=arguments-differ
def __enter__(self):
pass

def __exit__(self, *args): # pylint: disable=arguments-differ
def __exit__(self, *args): # pylint: disable=arguments-differ
pass


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,10 @@ async def _batch_send(

class AsyncTransportWrapper(AsyncHttpTransport):

"""Wrapper class that ensures that an inner client created
by a `get_client` method does not close the outer transport for the parent
when used in a context manager.
"""
def __init__(self, async_transport):
self._transport = async_transport

Expand All @@ -143,8 +147,8 @@ async def open(self):
async def close(self):
pass

async def __aenter__(self, *args): # pylint: disable=arguments-differ
async def __aenter__(self):
pass

async def __aexit__(self, *args): # pylint: disable=arguments-differ
async def __aexit__(self, *args): # pylint: disable=arguments-differ
pass
7 changes: 6 additions & 1 deletion sdk/storage/azure-storage-blob/tests/test_common_blob.py
Original file line number Diff line number Diff line change
Expand Up @@ -1836,15 +1836,20 @@ def test_set_blob_permission(self):
self.assertEqual(permission._str, 'wrdx')

def test_transport_closed_only_once(self):
if TestMode.need_recording_file(self.test_mode):
return
transport = RequestsTransport()
url = self._get_account_url()
credential = self._get_shared_key_credential()
blob_name = self._get_blob_reference()
with BlobServiceClient(url, credential=credential, transport=transport) as bsc:
bsc.get_service_properties()
assert transport.session is not None
with bsc.get_blob_client(self.container_name, blob_name) as bc:
assert transport.session is not None
assert transport.session is not None # Right now it's None
bsc.get_service_properties()
assert transport.session is not None


#------------------------------------------------------------------------------
if __name__ == '__main__':
Expand Down
16 changes: 13 additions & 3 deletions sdk/storage/azure-storage-blob/tests/test_common_blob_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -2281,16 +2281,26 @@ def test_upload_to_url_file_with_credential(self):
loop = asyncio.get_event_loop()
loop.run_until_complete(self._test_upload_to_url_file_with_credential())

async def test_transport_closed_only_once(self):
transport = AsyncioRequestsTransport()
async def _test_transport_closed_only_once(self):
if TestMode.need_recording_file(self.test_mode):
return
transport = AioHttpTransport()
url = self._get_account_url()
credential = self._get_shared_key_credential()
blob_name = self._get_blob_reference()
async with BlobServiceClient(url, credential=credential, transport=transport) as bsc:
await bsc.get_service_properties()
assert transport.session is not None
async with bsc.get_blob_client(self.container_name, blob_name) as bc:
assert transport.session is not None
assert transport.session is not None # Right now it's None
await bsc.get_service_properties()
assert transport.session is not None

@record
def test_transport_closed_only_once(self):
loop = asyncio.get_event_loop()
loop.run_until_complete(self._test_transport_closed_only_once())

# ------------------------------------------------------------------------------
if __name__ == '__main__':
unittest.main()
Original file line number Diff line number Diff line change
Expand Up @@ -223,7 +223,10 @@ def _batch_send(


class TransportWrapper(HttpTransport):

"""Wrapper class that ensures that an inner client created
by a `get_client` method does not close the outer transport for the parent
when used in a context manager.
"""
def __init__(self, transport):
self._transport = transport

Expand All @@ -236,10 +239,10 @@ def open(self):
def close(self):
pass

def __enter__(self, *args): # pylint: disable=arguments-differ
def __enter__(self):
pass

def __exit__(self, *args): # pylint: disable=arguments-differ
def __exit__(self, *args): # pylint: disable=arguments-differ
pass


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -127,7 +127,10 @@ async def _batch_send(


class AsyncTransportWrapper(AsyncHttpTransport):

"""Wrapper class that ensures that an inner client created
by a `get_client` method does not close the outer transport for the parent
when used in a context manager.
"""
def __init__(self, async_transport):
self._transport = async_transport

Expand All @@ -140,8 +143,8 @@ async def open(self):
async def close(self):
pass

async def __aenter__(self, *args): # pylint: disable=arguments-differ
async def __aenter__(self):
pass

async def __aexit__(self, *args): # pylint: disable=arguments-differ
async def __aexit__(self, *args): # pylint: disable=arguments-differ
pass
11 changes: 8 additions & 3 deletions sdk/storage/azure-storage-file/tests/test_share.py
Original file line number Diff line number Diff line change
Expand Up @@ -765,15 +765,20 @@ def test_create_permission_for_share(self):

@record
def test_transport_closed_only_once(self):
if TestMode.need_recording_file(self.test_mode):
return
transport = RequestsTransport()
url = self.get_file_url()
credential = self.get_shared_key_credential()
share = self._get_share_reference()
prefix = TEST_SHARE_PREFIX
share_name = self.get_resource_name(prefix)
with FileServiceClient(url, credential=credential, transport=transport) as fsc:
fsc.get_service_properties()
assert transport.session is not None
with fsc.get_share_client(share.share_name) as fc:
with fsc.get_share_client(share_name) as fc:
assert transport.session is not None
assert transport.session is not None # Right now it's None
fsc.get_service_properties()
assert transport.session is not None

# ------------------------------------------------------------------------------
if __name__ == '__main__':
Expand Down
20 changes: 15 additions & 5 deletions sdk/storage/azure-storage-file/tests/test_share_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -919,16 +919,26 @@ def test_create_permission_for_share_async(self):
loop = asyncio.get_event_loop()
loop.run_until_complete(self._test_create_permission_for_share())

async def test_transport_closed_only_once_async(self):
transport = AsyncioRequestsTransport()
async def _test_transport_closed_only_once_async(self):
if TestMode.need_recording_file(self.test_mode):
return
transport = AioHttpTransport()
url = self.get_file_url()
credential = self.get_shared_key_credential()
share = self._get_share_reference()
prefix = TEST_SHARE_PREFIX
share_name = self.get_resource_name(prefix)
async with FileServiceClient(url, credential=credential, transport=transport) as fsc:
await fsc.get_service_properties()
assert transport.session is not None
async with fsc.get_share_client(share.share_name) as fc:
async with fsc.get_share_client(share_name) as fc:
assert transport.session is not None
assert transport.session is not None # Right now it's None
await fsc.get_service_properties()
assert transport.session is not None

@record
def test_transport_closed_only_once_async(self):
loop = asyncio.get_event_loop()
loop.run_until_complete(self._test_transport_closed_only_once_async())

# ------------------------------------------------------------------------------
if __name__ == '__main__':
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -222,7 +222,10 @@ def _batch_send(


class TransportWrapper(HttpTransport):

"""Wrapper class that ensures that an inner client created
by a `get_client` method does not close the outer transport for the parent
when used in a context manager.
"""
def __init__(self, transport):
self._transport = transport

Expand All @@ -235,10 +238,10 @@ def open(self):
def close(self):
pass

def __enter__(self, *args): # pylint: disable=arguments-differ
def __enter__(self):
pass

def __exit__(self, *args): # pylint: disable=arguments-differ
def __exit__(self, *args): # pylint: disable=arguments-differ
pass


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -127,7 +127,10 @@ async def _batch_send(


class AsyncTransportWrapper(AsyncHttpTransport):

"""Wrapper class that ensures that an inner client created
by a `get_client` method does not close the outer transport for the parent
when used in a context manager.
"""
def __init__(self, async_transport):
self._transport = async_transport

Expand All @@ -140,8 +143,8 @@ async def open(self):
async def close(self):
pass

async def __aenter__(self, *args): # pylint: disable=arguments-differ
async def __aenter__(self):
pass

async def __aexit__(self, *args): # pylint: disable=arguments-differ
async def __aexit__(self, *args): # pylint: disable=arguments-differ
pass
6 changes: 5 additions & 1 deletion sdk/storage/azure-storage-queue/tests/test_queue.py
Original file line number Diff line number Diff line change
Expand Up @@ -967,14 +967,18 @@ def test_unicode_update_message_unicode_data(self, resource_group, location, sto
@ResourceGroupPreparer()
@StorageAccountPreparer(name_prefix='pyacrstorage')
def test_transport_closed_only_once(self, resource_group, location, storage_account, storage_account_key):
if not self.is_live:
return
transport = RequestsTransport()
prefix = TEST_QUEUE_PREFIX
queue_name = self.get_resource_name(prefix)
with QueueServiceClient(self._account_url(storage_account.name), credential=storage_account_key, transport=transport) as qsc:
qsc.get_service_properties()
assert transport.session is not None
with qsc.get_queue_client(queue_name) as qc:
assert transport.session is not None
assert transport.session is not None # Right now it's None
qsc.get_service_properties()
assert transport.session is not None

# ------------------------------------------------------------------------------
if __name__ == '__main__':
Expand Down
8 changes: 6 additions & 2 deletions sdk/storage/azure-storage-queue/tests/test_queue_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -1030,14 +1030,18 @@ async def test_unicode_update_message_unicode_data(self, resource_group, locatio
@StorageAccountPreparer(name_prefix='pyacrstorage')
@AsyncQueueTestCase.await_prepared_test
async def test_transport_closed_only_once_async(self, resource_group, location, storage_account, storage_account_key):
transport = AsyncioRequestsTransport()
if not self.is_live:
return
transport = AioHttpTransport()
prefix = TEST_QUEUE_PREFIX
queue_name = self.get_resource_name(prefix)
async with QueueServiceClient(self._account_url(storage_account.name), credential=storage_account_key, transport=transport) as qsc:
await qsc.get_service_properties()
assert transport.session is not None
async with qsc.get_queue_client(queue_name) as qc:
assert transport.session is not None
assert transport.session is not None # Right now it's None
await qsc.get_service_properties()
assert transport.session is not None

# ------------------------------------------------------------------------------
if __name__ == '__main__':
Expand Down