diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py index 594a9da14124..c89beb8286de 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py @@ -411,7 +411,7 @@ def read(self, verify_frame_type=0, **kwargs): # TODO: verify frame type? read_frame_buffer.write(read(size - SIGNED_INT_MAX, buffer=payload[SIGNED_INT_MAX:])) else: read_frame_buffer.write(read(payload_size, buffer=payload)) - except socket.timeout: + except (socket.timeout, TimeoutError): read_frame_buffer.write(self._read_buffer.getvalue()) self._read_buffer = read_frame_buffer self._read_buffer.seek(0) @@ -446,7 +446,7 @@ def receive_frame(self, *args, **kwargs): decoded = decode_frame(payload) # TODO: Catch decode error and return amqp:decode-error return channel, decoded - except socket.timeout: + except (socket.timeout, TimeoutError): return None, None def send_frame(self, channel, frame, **kwargs): @@ -658,7 +658,7 @@ def Transport(host, transport_type, connect_timeout=None, ssl=False, **kwargs): return transport(host, connect_timeout=connect_timeout, ssl=ssl, **kwargs) class WebSocketTransport(_AbstractTransport): - def __init__(self, host, port=WEBSOCKET_PORT, connect_timeout=None, ssl=None, **kwargs): + def __init__(self, host, port=WEBSOCKET_PORT, connect_timeout=1, ssl=None, **kwargs): self.sslopts = ssl if isinstance(ssl, dict) else {} self._connect_timeout = connect_timeout self._host = host @@ -694,23 +694,27 @@ def connect(self): def _read(self, n, initial=False, buffer=None, **kwargs): # pylint: disable=unused-arguments """Read exactly n bytes from the peer.""" + from websocket import WebSocketTimeoutException length = 0 view = buffer or memoryview(bytearray(n)) nbytes = self._read_buffer.readinto(view) length += nbytes n -= nbytes - while n: - data = self.ws.recv() + try: + while n: + data = self.ws.recv() - if len(data) <= n: - view[length: length + len(data)] = data - n -= len(data) - else: - view[length: length + n] = data[0:n] - self._read_buffer = BytesIO(data[n:]) - n = 0 - return view + if len(data) <= n: + view[length: length + len(data)] = data + n -= len(data) + else: + view[length: length + n] = data[0:n] + self._read_buffer = BytesIO(data[n:]) + n = 0 + return view + except WebSocketTimeoutException: + raise TimeoutError() def _shutdown_transport(self): """Do any preliminary work in shutting down the connection.""" diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_transport_async.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_transport_async.py index 7f586bec9e5e..ea679ebd392c 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_transport_async.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_transport_async.py @@ -94,7 +94,7 @@ async def receive_frame(self, *args, **kwargs): # TODO: Catch decode error and return amqp:decode-error #_LOGGER.info("ICH%d <- %r", channel, decoded) return channel, decoded - except (socket.timeout, asyncio.IncompleteReadError, asyncio.TimeoutError): + except (TimeoutError, socket.timeout, asyncio.IncompleteReadError, asyncio.TimeoutError): return None, None async def read(self, verify_frame_type=0, **kwargs): # TODO: verify frame type? @@ -121,7 +121,7 @@ async def read(self, verify_frame_type=0, **kwargs): # TODO: verify frame type? read_frame_buffer.write(await self._read(size - SIGNED_INT_MAX, buffer=payload[SIGNED_INT_MAX:])) else: read_frame_buffer.write(await self._read(payload_size, buffer=payload)) - except (socket.timeout, asyncio.IncompleteReadError): + except (TimeoutError, socket.timeout, asyncio.IncompleteReadError): read_frame_buffer.write(self._read_buffer.getvalue()) self._read_buffer = read_frame_buffer self._read_buffer.seek(0) @@ -404,7 +404,7 @@ async def receive_frame_with_lock(self, *args, **kwargs): else: decoded = decode_frame(payload) return channel, decoded - except socket.timeout: + except (socket.timeout, TimeoutError): return None, None async def negotiate(self): @@ -418,7 +418,7 @@ async def negotiate(self): class WebSocketTransportAsync(AsyncTransportMixin): - def __init__(self, host, port=WEBSOCKET_PORT, connect_timeout=None, ssl=None, **kwargs + def __init__(self, host, port=WEBSOCKET_PORT, connect_timeout=1, ssl=None, **kwargs ): self._read_buffer = BytesIO() self.loop = get_running_loop() @@ -455,26 +455,30 @@ async def connect(self): async def _read(self, n, buffer=None, **kwargs): # pylint: disable=unused-arguments """Read exactly n bytes from the peer.""" + from websocket import WebSocketTimeoutException length = 0 view = buffer or memoryview(bytearray(n)) nbytes = self._read_buffer.readinto(view) length += nbytes n -= nbytes - while n: - data = await self.loop.run_in_executor( - None, self.ws.recv - ) - - if len(data) <= n: - view[length: length + len(data)] = data - n -= len(data) - else: - view[length: length + n] = data[0:n] - self._read_buffer = BytesIO(data[n:]) - n = 0 + try: + while n: + data = await self.loop.run_in_executor( + None, self.ws.recv + ) - return view + if len(data) <= n: + view[length: length + len(data)] = data + n -= len(data) + else: + view[length: length + n] = data[0:n] + self._read_buffer = BytesIO(data[n:]) + n = 0 + + return view + except WebSocketTimeoutException as wex: + raise TimeoutError() def close(self): """Do any preliminary work in shutting down the connection."""