diff --git a/sdk/eventhub/azure-eventhub/CHANGELOG.md b/sdk/eventhub/azure-eventhub/CHANGELOG.md index 307c4bc87e83..d72a775683b3 100644 --- a/sdk/eventhub/azure-eventhub/CHANGELOG.md +++ b/sdk/eventhub/azure-eventhub/CHANGELOG.md @@ -1,6 +1,6 @@ # Release History -## 5.8.0a4 (2022-05-11) +## 5.8.0a4 (Unreleased) ### Features Added diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_client_base.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_client_base.py index 85c57dbe7b81..f2e695ea8560 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_client_base.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_client_base.py @@ -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() diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_consumer.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_consumer.py index d75556e18e91..a8c7efe79638 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_consumer.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_consumer.py @@ -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 @@ -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, diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_producer.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_producer.py index f11aa3fd7c18..d5d8f48d5207 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_producer.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_producer.py @@ -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, diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_connection.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_connection.py index 5f3aeb3e9bf1..c73417d1e56f 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_connection.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_connection.py @@ -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, @@ -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 diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py index 29e506177cd3..594a9da14124 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/_transport.py @@ -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: @@ -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. @@ -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 @@ -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) diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_client_async.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_client_async.py index 957588d2a921..863285f7ca59 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_client_async.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_client_async.py @@ -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 @@ -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 @@ -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: @@ -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): @@ -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() diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_connection_async.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_connection_async.py index 45e1a6e5fa23..790b02bae084 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_connection_async.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_connection_async.py @@ -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 @@ -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 @@ -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 @@ -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: @@ -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 @@ -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): @@ -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, @@ -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( @@ -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: @@ -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 diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_sasl_async.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_sasl_async.py index 014681787c27..88ee25917c7c 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_sasl_async.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/aio/_sasl_async.py @@ -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: @@ -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) @@ -120,3 +127,6 @@ def __init__( **kwargs ) super().__init__(host, port, connect_timeout, ssl, **kwargs) + + async def negotiate(self): + await self._negotiate() 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 39d09213eba3..7f586bec9e5e 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 @@ -79,13 +79,14 @@ def get_running_loop(): _LOGGER.warning('This version of Python is deprecated, please upgrade to >= v3.6') if loop is None: _LOGGER.warning('No running event loop') - loop = self.loop + loop = asyncio.get_event_loop() return loop + class AsyncTransportMixin(): async def receive_frame(self, *args, **kwargs): try: - header, channel, payload = await self.read(**kwargs) + header, channel, payload = await self.read(**kwargs) if not payload: decoded = decode_empty_frame(header) else: @@ -147,6 +148,7 @@ async def send_frame(self, channel, frame, **kwargs): await self.write(data) #_LOGGER.info("OCH%d -> %r", channel, frame) + class AsyncTransport(AsyncTransportMixin): """Common superclass for TCP and SSL transports.""" @@ -160,6 +162,7 @@ def __init__(self, host, port=AMQP_PORT, connect_timeout=None, self.raise_on_initial_eintr = raise_on_initial_eintr self._read_buffer = BytesIO() self.host, self.port = to_host_port(host, port) + self.connect_timeout = connect_timeout self.read_timeout = read_timeout self.write_timeout = write_timeout @@ -452,7 +455,7 @@ async def connect(self): async def _read(self, n, 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) @@ -461,7 +464,7 @@ async def _read(self, n, buffer=None, **kwargs): # pylint: disable=unused-argume while n: data = await self.loop.run_in_executor( None, self.ws.recv - ) + ) if len(data) <= n: view[length: length + len(data)] = data @@ -470,10 +473,15 @@ async def _read(self, n, buffer=None, **kwargs): # pylint: disable=unused-argume view[length: length + n] = data[0:n] self._read_buffer = BytesIO(data[n:]) n = 0 + return view def close(self): """Do any preliminary work in shutting down the connection.""" + # TODO: async close doesn't: + # 1) shutdown socket and close. --> self.sock.shutdown(socket.SHUT_RDWR) and self.sock.close() + # 2) set self.connected = False + # I think we need to do this, like in sync self.ws.close() async def write(self, s): diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/client.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/client.py index 2b6c06070347..25fbb125a4fc 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/client.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/client.py @@ -6,9 +6,7 @@ # pylint: disable=too-many-lines -from collections import namedtuple import logging -import threading import time import uuid import certifi diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/constants.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/constants.py index 66e4ff1ae327..abcc56f0e270 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/constants.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/constants.py @@ -26,6 +26,14 @@ SECURE_PORT = 5671 +# default port for AMQP over Websocket +WEBSOCKET_PORT = 443 + + +# subprotocol for AMQP over Websocket +AMQP_WS_SUBPROTOCOL = 'AMQPWSB10' + + MAJOR = 1 #: Major protocol version. MINOR = 0 #: Minor protocol version. REV = 0 #: Protocol revision. @@ -308,6 +316,7 @@ class MessageDeliveryState(object): MessageDeliveryState.Cancelled ) + class TransportType(Enum): """Transport type The underlying transport protocol type: diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/sasl.py b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/sasl.py index 51848304bfae..6d6d7d98f342 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/sasl.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/sasl.py @@ -94,6 +94,33 @@ def _negotiate(self): else: raise ValueError("SASL negotiation failed.\nOutcome: {}\nDetails: {}".format(*fields)) +class SASLTransportMixin(): + def _negotiate(self): + self.write(SASL_HEADER_FRAME) + _, returned_header = self.receive_frame() + if returned_header[1] != SASL_HEADER_FRAME: + raise ValueError("Mismatching AMQP header protocol. Expected: {}, received: {}".format( + SASL_HEADER_FRAME, returned_header[1])) + + _, supported_mechansisms = self.receive_frame(verify_frame_type=1) + if self.credential.mechanism not in supported_mechansisms[1][0]: # sasl_server_mechanisms + raise ValueError("Unsupported SASL credential type: {}".format(self.credential.mechanism)) + sasl_init = SASLInit( + mechanism=self.credential.mechanism, + initial_response=self.credential.start(), + hostname=self.host) + self.send_frame(0, sasl_init, frame_type=_SASL_FRAME_TYPE) + + _, next_frame = self.receive_frame(verify_frame_type=1) + frame_type, fields = next_frame + if frame_type != 0x00000044: # SASLOutcome + raise NotImplementedError("Unsupported SASL challenge") + if fields[0] == SASLCode.Ok: # code + return + else: + raise ValueError("SASL negotiation failed.\nOutcome: {}\nDetails: {}".format(*fields)) + + class SASLTransport(SSLTransport, SASLTransportMixin): def __init__(self, host, credential, port=AMQPS_PORT, connect_timeout=None, ssl=None, **kwargs): @@ -121,5 +148,5 @@ def __init__(self, host, credential, port=WEBSOCKET_PORT, connect_timeout=None, ) super().__init__(host, port, connect_timeout, ssl, **kwargs) - def negotiate(self): + def negotiate(self): self._negotiate() diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_client_base_async.py b/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_client_base_async.py index 5589af87dc2d..57bb49bf103b 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_client_base_async.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_client_base_async.py @@ -246,8 +246,15 @@ async def _management_request_async(self, mgmt_msg: Message, op_type: bytes) -> last_exception = None while retried_times <= self._config.max_retries: mgmt_auth = await self._create_auth_async() + hostname = self._address.hostname + if self._config.transport_type.name == 'AmqpOverWebsocket': + hostname += '/$servicebus/websocket/' mgmt_client = AMQPClientAsync( - 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: await mgmt_client.open_async() diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_consumer_async.py b/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_consumer_async.py index f8ed74e03d80..da46d8f40166 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_consumer_async.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_consumer_async.py @@ -145,7 +145,7 @@ def _create_handler(self, auth: "JWTTokenAuthAsync") -> None: desired_capabilities = [RECEIVER_RUNTIME_METRIC_SYMBOL] if self._track_last_enqueued_event_properties else None hostname = urlparse(source.address).hostname transport_type = self._client._config.transport_type # pylint:disable=protected-access - if transport_type.name is 'AmqpOverWebsocket': + if transport_type.name == 'AmqpOverWebsocket': hostname += '/$servicebus/websocket/' self._handler = ReceiveClientAsync( hostname, diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_producer_async.py b/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_producer_async.py index cd59c1278621..d23b709d3b42 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_producer_async.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_producer_async.py @@ -100,7 +100,7 @@ def __init__(self, client: "EventHubProducerClient", target: str, **kwargs) -> N def _create_handler(self, auth: "JWTTokenAsync") -> None: hostname = self._client._address.hostname # pylint: disable=protected-access - transport_type = self._client._config.transport_type # pylint:disable=protected-access + transport_type = self._client._config.transport_type # pylint:disable=protected-access if transport_type.name == 'AmqpOverWebsocket': hostname += '/$servicebus/websocket/' self._handler = SendClientAsync( diff --git a/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_producer_client_async.py b/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_producer_client_async.py index f87bb9f8e366..ee421c3ca4f2 100644 --- a/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_producer_client_async.py +++ b/sdk/eventhub/azure-eventhub/azure/eventhub/aio/_producer_client_async.py @@ -189,11 +189,9 @@ def from_connection_string( *, eventhub_name: Optional[str] = None, logging_enable: bool = False, - http_proxy: Optional[Dict[str, Union[str, int]]] = None, auth_timeout: float = 60, user_agent: Optional[str] = None, retry_total: int = 3, - transport_type: Optional["TransportType"] = None, **kwargs: Any ) -> "EventHubProducerClient": """Create an EventHubProducerClient from a connection string. @@ -250,11 +248,9 @@ def from_connection_string( conn_str, eventhub_name=eventhub_name, logging_enable=logging_enable, - http_proxy=http_proxy, auth_timeout=auth_timeout, user_agent=user_agent, retry_total=retry_total, - transport_type=transport_type, **kwargs ) return cls(**constructor_args) diff --git a/sdk/eventhub/azure-eventhub/dev_requirements.txt b/sdk/eventhub/azure-eventhub/dev_requirements.txt index df47262912ac..9c91833e14d8 100644 --- a/sdk/eventhub/azure-eventhub/dev_requirements.txt +++ b/sdk/eventhub/azure-eventhub/dev_requirements.txt @@ -4,5 +4,6 @@ azure-mgmt-eventhub==10.0.0 azure-mgmt-resource==20.0.0 aiohttp>=3.0 +websocket-client -e ../../../tools/azure-devtools -e ../../servicebus/azure-servicebus \ No newline at end of file diff --git a/sdk/eventhub/azure-eventhub/tests/livetest/asynctests/test_send_async.py b/sdk/eventhub/azure-eventhub/tests/livetest/asynctests/test_send_async.py index e3560f6e7e2f..016b1bdc9c86 100644 --- a/sdk/eventhub/azure-eventhub/tests/livetest/asynctests/test_send_async.py +++ b/sdk/eventhub/azure-eventhub/tests/livetest/asynctests/test_send_async.py @@ -273,7 +273,6 @@ async def test_send_multiple_partition_with_app_prop_async(connstr_receivers): @pytest.mark.liveTest @pytest.mark.asyncio async def test_send_over_websocket_async(connstr_receivers): - pytest.skip("websocket unsupported") connection_str, receivers = connstr_receivers client = EventHubProducerClient.from_connection_string(connection_str, transport_type=TransportType.AmqpOverWebsocket) diff --git a/sdk/eventhub/azure-eventhub/tests/livetest/synctests/test_consumer_client.py b/sdk/eventhub/azure-eventhub/tests/livetest/synctests/test_consumer_client.py index 8da5ddeb6ead..9e09cd156ef8 100644 --- a/sdk/eventhub/azure-eventhub/tests/livetest/synctests/test_consumer_client.py +++ b/sdk/eventhub/azure-eventhub/tests/livetest/synctests/test_consumer_client.py @@ -36,7 +36,7 @@ def on_event(partition_context, event): args=(on_event,), kwargs={"starting_position": "-1"}) worker.start() - time.sleep(10) + time.sleep(20) assert on_event.received == 2 checkpoints = list(client._event_processors.values())[0]._checkpoint_store.list_checkpoints( on_event.namespace, on_event.eventhub_name, on_event.consumer_group @@ -133,7 +133,7 @@ def on_event_batch(partition_context, event_batch): worker = threading.Thread(target=client.receive_batch, args=(on_event_batch,), kwargs={"starting_position": "-1"}) worker.start() - time.sleep(10) + time.sleep(20) assert on_event_batch.received == 2 checkpoints = list(client._event_processors.values())[0]._checkpoint_store.list_checkpoints( diff --git a/sdk/eventhub/azure-eventhub/tests/livetest/synctests/test_send.py b/sdk/eventhub/azure-eventhub/tests/livetest/synctests/test_send.py index f9ea0c55c773..9b484340855c 100644 --- a/sdk/eventhub/azure-eventhub/tests/livetest/synctests/test_send.py +++ b/sdk/eventhub/azure-eventhub/tests/livetest/synctests/test_send.py @@ -218,9 +218,22 @@ def test_send_partition(connstr_receivers): client.send_batch(batch) partition_0 = receivers[0].receive_message_batch(timeout=5) - assert len(partition_0) == 0 partition_1 = receivers[1].receive_message_batch(timeout=5) - assert len(partition_1) == 1 + assert len(partition_0) + len(partition_1) == 2 + + with client: + batch = client.create_batch() + batch.add(EventData(b"Data")) + client.send_batch(batch) + + with client: + batch = client.create_batch(partition_id="1") + batch.add(EventData(b"Data")) + client.send_batch(batch) + + partition_0 = receivers[0].receive_message_batch(timeout=5) + partition_1 = receivers[1].receive_message_batch(timeout=5) + assert len(partition_0) + len(partition_1) == 2 @pytest.mark.liveTest @@ -273,7 +286,6 @@ def test_send_multiple_partitions_with_app_prop(connstr_receivers): @pytest.mark.liveTest def test_send_over_websocket_sync(connstr_receivers): - pytest.skip("websocket not supported") connection_str, receivers = connstr_receivers client = EventHubProducerClient.from_connection_string(connection_str, transport_type=TransportType.AmqpOverWebsocket) diff --git a/sdk/eventhub/azure-eventhub/tests/pyamqp_tests/async/test_websocket_async.py b/sdk/eventhub/azure-eventhub/tests/pyamqp_tests/async/test_websocket_async.py deleted file mode 100644 index a76ff7117bab..000000000000 --- a/sdk/eventhub/azure-eventhub/tests/pyamqp_tests/async/test_websocket_async.py +++ /dev/null @@ -1,35 +0,0 @@ -# -------------------------------------------------------------------------------------------- -# Copyright (c) Microsoft Corporation. All rights reserved. -# Licensed under the MIT License. See License.txt in the project root for license information. -# -------------------------------------------------------------------------------------------- - -import pytest -import asyncio -import logging -from azure.eventhub._pyamqp.aio import ReceiveClientAsync, SASTokenAuthAsync -from azure.eventhub._pyamqp.constants import TransportType - -@pytest.mark.asyncio -async def test_event_hubs_client_web_socket(eventhub_config): - uri = "sb://{}/{}".format(eventhub_config['hostname'], eventhub_config['event_hub']) - sas_auth = SASTokenAuthAsync( - uri=uri, - audience=uri, - username=eventhub_config['key_name'], - password=eventhub_config['access_key'] - ) - - source = "amqps://{}/{}/ConsumerGroups/{}/Partitions/{}".format( - eventhub_config['hostname'], - eventhub_config['event_hub'], - eventhub_config['consumer_group'], - eventhub_config['partition']) - - receive_client = ReceiveClientAsync(eventhub_config['hostname'] + '/$servicebus/websocket/', source, auth=sas_auth, debug=False, timeout=5000, prefetch=50, transport_type=TransportType.AmqpOverWebsocket) - await receive_client.open_async() - while not await receive_client.client_ready_async(): - await asyncio.sleep(0.05) - messages = await receive_client.receive_message_batch_async(max_batch_size=1) - logging.info(len(messages)) - logging.info(messages[0]) - await receive_client.close_async() diff --git a/sdk/eventhub/azure-eventhub/tests/pyamqp_tests/synctests/test_websocket.py b/sdk/eventhub/azure-eventhub/tests/pyamqp_tests/synctests/test_websocket.py deleted file mode 100644 index 7dd9e5bfbe9c..000000000000 --- a/sdk/eventhub/azure-eventhub/tests/pyamqp_tests/synctests/test_websocket.py +++ /dev/null @@ -1,27 +0,0 @@ -# -------------------------------------------------------------------------------------------- -# Copyright (c) Microsoft Corporation. All rights reserved. -# Licensed under the MIT License. See License.txt in the project root for license information. -# -------------------------------------------------------------------------------------------- - -import pytest - -from azure.eventhub._pyamqp import authentication, ReceiveClient -from azure.eventhub._pyamqp.constants import TransportType - -def test_event_hubs_client_web_socket(live_eventhub): - uri = "sb://{}/{}".format(live_eventhub['hostname'], live_eventhub['event_hub']) - sas_auth = authentication.SASTokenAuth( - uri=uri, - audience=uri, - username=live_eventhub['key_name'], - password=live_eventhub['access_key'] - ) - - source = "amqps://{}/{}/ConsumerGroups/{}/Partitions/{}".format( - live_eventhub['hostname'], - live_eventhub['event_hub'], - live_eventhub['consumer_group'], - live_eventhub['partition']) - - with ReceiveClient(live_eventhub['hostname'] + '/$servicebus/websocket/', source, auth=sas_auth, debug=False, timeout=5000, prefetch=50, transport_type=TransportType.AmqpOverWebsocket) as receive_client: - receive_client.receive_message_batch(max_batch_size=10)