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
2 changes: 1 addition & 1 deletion sdk/eventhub/azure-eventhub/CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
# Release History

## 5.8.0a4 (2022-05-11)
## 5.8.0a4 (Unreleased)

### Features Added

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -367,8 +367,15 @@ def _management_request(self, mgmt_msg, op_type):
last_exception = None
while retried_times <= self._config.max_retries:
mgmt_auth = self._create_auth()
hostname = self._address.hostname
if self._config.transport_type.name == 'AmqpOverWebsocket':
hostname += '/$servicebus/websocket/'
mgmt_client = AMQPClient(
self._address.hostname, auth=mgmt_auth, debug=self._config.network_tracing
hostname,
auth=mgmt_auth,
debug=self._config.network_tracing,
transport_type=self._config.transport_type,
http_proxy=self._config.http_proxy
)
try:
mgmt_client.open()
Expand Down
9 changes: 5 additions & 4 deletions sdk/eventhub/azure-eventhub/azure/eventhub/_consumer.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,10 +136,6 @@ def __init__(self, client, source, **kwargs):

def _create_handler(self, auth):
# type: (JWTTokenAuth) -> None
transport_type = self._client._config.transport_type # pylint:disable=protected-access
hostname = urlparse(source.address).hostname
if transport_type.name is 'AmqpOverWebsocket':
hostname += '/$servicebus/websocket/'
source = Source(address=self._source, filters={})
if self._offset is not None:
filter_key = ApacheFilters.selector_filter
Expand All @@ -154,6 +150,11 @@ def _create_handler(self, auth):
)
desired_capabilities = [RECEIVER_RUNTIME_METRIC_SYMBOL] if self._track_last_enqueued_event_properties else None

transport_type = self._client._config.transport_type # pylint:disable=protected-access
hostname = urlparse(source.address).hostname
if transport_type.name == 'AmqpOverWebsocket':
hostname += '/$servicebus/websocket/'

self._handler = ReceiveClient(
hostname,
source,
Expand Down
4 changes: 2 additions & 2 deletions sdk/eventhub/azure-eventhub/azure/eventhub/_producer.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,9 +125,9 @@ def __init__(self, client, target, **kwargs):

def _create_handler(self, auth):
# type: (JWTTokenAuth) -> None
transport_type=self._client._config.transport_type # pylint:disable=protected-access
transport_type = self._client._config.transport_type # pylint:disable=protected-access
hostname = self._client._address.hostname # pylint: disable=protected-access
if transport_type.name is 'AmqpOverWebsocket':
if transport_type.name == 'AmqpOverWebsocket':
hostname += '/$servicebus/websocket/'
self._handler = SendClient(
hostname,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -78,7 +78,7 @@ class Connection(object):
Default value is `0.1`.
:keyword bool network_trace: Whether to log the network traffic. Default value is `False`. If enabled, frames
will be logged at the logging.INFO level.
:keyword str transport_type: Determines if the transport type is Amqp or AmqpOverWebSocket.
:keyword str transport_type: Determines if the transport type is Amqp or AmqpOverWebSocket.
Defaults to TransportType.Amqp. It will be AmqpOverWebSocket if using http_proxy.
:keyword Dict http_proxy: HTTP proxy settings. This must be a dictionary with the following
keys: `'proxy_hostname'` (str value) and `'proxy_port'` (int value). When using these settings,
Expand Down Expand Up @@ -114,7 +114,7 @@ def __init__(self, endpoint, **kwargs):
**kwargs
)
else:
self._transport = Transport(parsed_url.netloc, self._transport_type, **kwargs)
self._transport = Transport(parsed_url.netloc, transport_type=self._transport_type, **kwargs)

self._container_id = kwargs.pop('container_id', None) or str(uuid.uuid4()) # type: str
self._max_frame_size = kwargs.pop('max_frame_size', MAX_FRAME_SIZE_BYTES) # type: int
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -439,7 +439,7 @@ def write(self, s):

def receive_frame(self, *args, **kwargs):
try:
header, channel, payload = self.read(**kwargs)
header, channel, payload = self.read(**kwargs)
if not payload:
decoded = decode_empty_frame(header)
else:
Expand Down Expand Up @@ -645,7 +645,6 @@ def _read(self, n, initial=False, _errnos=(errno.EAGAIN, errno.EINTR)):
result, self._read_buffer = rbuf[:n], rbuf[n:]
return result


def Transport(host, transport_type, connect_timeout=None, ssl=False, **kwargs):
"""Create transport.

Expand All @@ -659,8 +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=None, ssl=None, **kwargs):
self.sslopts = ssl if isinstance(ssl, dict) else {}
self._connect_timeout = connect_timeout
self._host = host
Expand Down Expand Up @@ -694,9 +692,9 @@ def connect(self):
except ImportError:
raise ValueError("Please install websocket-client library to use websocket transport.")

def _read(self, n, buffer=None, **kwargs): # pylint: disable=unused-arguments
def _read(self, n, initial=False, buffer=None, **kwargs): # pylint: disable=unused-arguments
"""Read exactly n bytes from the peer."""

length = 0
view = buffer or memoryview(bytearray(n))
nbytes = self._read_buffer.readinto(view)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@
from ._receiver_async import ReceiverLink
from ._sender_async import SenderLink
from ._session_async import Session
from ._sasl_async import SASLTransport
from ._cbs_async import CBSAuthenticator
from ..client import AMQPClient as AMQPClientSync
from ..client import ReceiveClient as ReceiveClientSync
Expand Down Expand Up @@ -176,23 +175,6 @@ async def _do_retryable_operation_async(self, operation, *args, **kwargs):
absolute_timeout -= (end_time - start_time)
raise retry_settings['history'][-1]

async def _keep_alive_worker_async(self):
interval = 10 if self._keep_alive is True else self._keep_alive
start_time = time.time()
try:
while self._connection and not self._shutdown:
current_time = time.time()
elapsed_time = (current_time - start_time)
if elapsed_time >= interval:
_logger.info("Keeping %r connection alive. %r",
self.__class__.__name__,
self._connection._container_id)
await self._connection._get_remote_timeout(current_time)
start_time = current_time
await asyncio.sleep(1)
except Exception as e: # pylint: disable=broad-except
_logger.info("Connection keep-alive for %r failed: %r.", self.__class__.__name__, e)

async def open_async(self):
"""Asynchronously open the client. The client can create a new Connection
or an existing Connection can be passed in. This existing Connection
Expand All @@ -217,10 +199,10 @@ async def open_async(self):
max_frame_size=self._max_frame_size,
channel_max=self._channel_max,
idle_timeout=self._idle_timeout,
transport_type=self._transport_type,
http_proxy=self._http_proxy,
properties=self._properties,
network_trace=self._network_trace
network_trace=self._network_trace,
transport_type=self._transport_type,
http_proxy=self._http_proxy
)
await self._connection.open()
if not self._session:
Expand All @@ -236,8 +218,6 @@ async def open_async(self):
auth_timeout=self._auth_timeout
)
await self._cbs_authenticator.open()
if self._keep_alive:
self._keep_alive_thread = asyncio.ensure_future(self._keep_alive_worker_async())
self._shutdown = False

async def close_async(self):
Expand All @@ -249,9 +229,6 @@ async def close_async(self):
self._shutdown = True
if not self._session:
return # already closed.
if self._keep_alive_thread:
await self._keep_alive_thread
self._keep_alive_thread = None
await self._close_link_async(close=True)
if self._cbs_authenticator:
await self._cbs_authenticator.close()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,15 @@
# license information.
# --------------------------------------------------------------------------

import threading
import struct
import uuid
import logging
import time
from urllib.parse import urlparse
import socket
from ssl import SSLError
from enum import Enum
import asyncio

from ._transport_async import AsyncTransport
Expand Down Expand Up @@ -66,8 +69,8 @@ class Connection(object):

def __init__(self, endpoint, **kwargs):
parsed_url = urlparse(endpoint)
self._hostname = parsed_url.hostname
endpoint = self._hostname
self.hostname = parsed_url.hostname
endpoint = self.hostname
self._transport_type = kwargs.pop('transport_type', TransportType.Amqp)
if parsed_url.port:
self.port = parsed_url.port
Expand All @@ -76,21 +79,22 @@ def __init__(self, endpoint, **kwargs):
else:
self.port = PORT
self.state = None

transport = kwargs.get('transport')
if transport:
self._transport = transport
self.transport = transport
elif 'sasl_credential' in kwargs:
sasl_transport = SASLTransport
if self._transport_type.name == 'AmqpOverWebsocket' or kwargs.get("http_proxy"):
if self._transport_type.name == "AmqpOverWebsocket" or kwargs.get("http_proxy"):
sasl_transport = SASLWithWebSocket
endpoint = parsed_url.hostname + parsed_url.path
self._transport = sasl_transport(
self.transport = sasl_transport(
host=endpoint,
credential=kwargs['sasl_credential'],
**kwargs
)
else:
self._transport = AsyncTransport(parsed_url.netloc, **kwargs)
self.transport = AsyncTransport(parsed_url.netloc, **kwargs)
self._container_id = kwargs.get('container_id') or str(uuid.uuid4())
self.max_frame_size = kwargs.get('max_frame_size', MAX_FRAME_SIZE_BYTES)
self._remote_max_frame_size = None
Expand Down Expand Up @@ -141,9 +145,9 @@ async def _set_state(self, new_state):
async def _connect(self):
try:
if not self.state:
await self._transport.connect()
await self.transport.connect()
await self._set_state(ConnectionState.START)
await self._transport.negotiate()
await self.transport.negotiate()
await self._outgoing_header()
await self._set_state(ConnectionState.HDR_SENT)
if not self.allow_pipelined_open:
Expand All @@ -164,7 +168,7 @@ async def _disconnect(self, *args):
if self.state == ConnectionState.END:
return
await self._set_state(ConnectionState.END)
self._transport.close()
self.transport.close()

def _can_read(self):
# type: () -> bool
Expand All @@ -173,7 +177,7 @@ def _can_read(self):

async def _read_frame(self, **kwargs):
if self._can_read():
return await self._transport.receive_frame(**kwargs)
return await self.transport.receive_frame(**kwargs)
_LOGGER.warning("Cannot read frame in current state: %r", self.state)

def _can_write(self):
Expand All @@ -190,7 +194,7 @@ async def _send_frame(self, channel, frame, timeout=None, **kwargs):
if self._can_write():
try:
self.last_frame_sent_time = time.time()
await self._transport.send_frame(channel, frame, **kwargs)
await self.transport.send_frame(channel, frame, **kwargs)
except (OSError, IOError, SSLError, socket.error) as exc:
self._error = AMQPConnectionError(
ErrorCondition.SocketError,
Expand Down Expand Up @@ -218,7 +222,7 @@ async def _outgoing_empty(self):
_LOGGER.info("-> empty()", extra=self.network_trace_params)
try:
if self._can_write():
await self._transport.write(EMPTY_FRAME)
await self.transport.write(EMPTY_FRAME)
self.last_frame_sent_time = time.time()
except (OSError, IOError, SSLError, socket.error) as exc:
self._error = AMQPConnectionError(
Expand All @@ -231,7 +235,7 @@ async def _outgoing_header(self):
self.last_frame_sent_time = time.time()
if self.network_trace:
_LOGGER.info("-> header(%r)", HEADER_FRAME, extra=self.network_trace_params)
await self._transport.write(HEADER_FRAME)
await self.transport.write(HEADER_FRAME)

async def _incoming_header(self, channel, frame):
if self.network_trace:
Expand All @@ -246,7 +250,7 @@ async def _incoming_header(self, channel, frame):
async def _outgoing_open(self):
open_frame = OpenFrame(
container_id=self._container_id,
hostname=self._hostname,
hostname=self.hostname,
max_frame_size=self.max_frame_size,
channel_max=self.channel_max,
idle_timeout=self.idle_timeout * 1000 if self.idle_timeout else None, # Convert to milliseconds
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -72,8 +72,9 @@ class SASLExternalCredential(object):
def start(self):
return b''


class SASLTransportMixinAsync():
async def negotiate(self):
async def _negotiate(self):
await self.write(SASL_HEADER_FRAME)
_, returned_header = await self.receive_frame()
if returned_header[1] != SASL_HEADER_FRAME:
Expand All @@ -98,16 +99,22 @@ async def negotiate(self):
else:
raise ValueError("SASL negotiation failed.\nOutcome: {}\nDetails: {}".format(*fields))


class SASLTransport(AsyncTransport, SASLTransportMixinAsync):
def __init__(self, host, credential, connect_timeout=None, ssl=None, **kwargs):

def __init__(self, host, credential, port=AMQPS_PORT, connect_timeout=None, ssl=None, **kwargs):
self.credential = credential
ssl = ssl or True
super(SASLTransport, self).__init__(host, connect_timeout=connect_timeout, ssl=ssl, **kwargs)
super(SASLTransport, self).__init__(host, port=port, connect_timeout=connect_timeout, ssl=ssl, **kwargs)

async def negotiate(self):
await self._negotiate()


class SASLWithWebSocket(WebSocketTransportAsync, SASLTransportMixinAsync):
def __init__(
self, host, credential, port=WEBSOCKET_PORT, connect_timeout=None, ssl=None, **kwargs
): # pylint: disable=super-init-not-called
):
self.credential = credential
ssl = ssl or True
http_proxy = kwargs.pop('http_proxy', None)
Expand All @@ -120,3 +127,6 @@ def __init__(
**kwargs
)
super().__init__(host, port, connect_timeout, ssl, **kwargs)

async def negotiate(self):
await self._negotiate()
Loading