Skip to content
30 changes: 17 additions & 13 deletions sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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?
Expand All @@ -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)
Expand Down Expand Up @@ -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):
Expand All @@ -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()
Expand Down Expand Up @@ -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."""
Expand Down