Skip to content
Closed
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
8 changes: 7 additions & 1 deletion sdk/eventhub/azure-eventhub/azure/eventhub/_consumer.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,10 @@ 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 @@ -151,11 +155,13 @@ def _create_handler(self, auth):
desired_capabilities = [RECEIVER_RUNTIME_METRIC_SYMBOL] if self._track_last_enqueued_event_properties else None

self._handler = ReceiveClient(
urlparse(source.address).hostname,
hostname,
source,
auth=auth,
idle_timeout=self._idle_timeout,
network_trace=self._client._config.network_tracing, # pylint:disable=protected-access
transport_type=transport_type,
http_proxy=self._client._config.http_proxy, # pylint:disable=protected-access
link_credit=self._prefetch,
link_properties=self._link_properties,
retry_policy=self._retry_policy,
Expand Down
8 changes: 7 additions & 1 deletion sdk/eventhub/azure-eventhub/azure/eventhub/_producer.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,12 +125,18 @@ 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
hostname = self._client._address.hostname # pylint: disable=protected-access
if transport_type.name is 'AmqpOverWebsocket':
hostname += '/$servicebus/websocket/'
self._handler = SendClient(
self._client._address.hostname, # pylint: disable=protected-access
hostname, # pylint: disable=protected-access
self._target,
auth=auth,
idle_timeout=self._idle_timeout,
network_trace=self._client._config.network_tracing, # pylint: disable=protected-access
transport_type=transport_type,
http_proxy=self._client._config.http_proxy, # pylint:disable=protected-access
retry_policy=self._retry_policy,
keep_alive_interval=self._keep_alive,
client_name=self._name,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
from ssl import SSLError

from ._transport import Transport
from .sasl import SASLTransport
from .sasl import SASLTransport, SASLWithWebSocket
from .session import Session
from .performatives import OpenFrame, CloseFrame
from .constants import (
Expand All @@ -22,7 +22,8 @@
MAX_FRAME_SIZE_BYTES,
HEADER_FRAME,
ConnectionState,
EMPTY_FRAME
EMPTY_FRAME,
TransportType
)

from .error import (
Expand Down Expand Up @@ -83,6 +84,7 @@ def __init__(self, endpoint, **kwargs):
# type(str, Any) -> None
parsed_url = urlparse(endpoint)
self._hostname = parsed_url.hostname
endpoint = self._hostname
if parsed_url.port:
self._port = parsed_url.port
elif parsed_url.scheme == 'amqps':
Expand All @@ -92,16 +94,22 @@ def __init__(self, endpoint, **kwargs):
self.state = None # type: Optional[ConnectionState]

transport = kwargs.get('transport')
self._transport_type = kwargs.pop('transport_type', TransportType.Amqp)
if transport:
self._transport = transport
elif 'sasl_credential' in kwargs:
self._transport = SASLTransport(
host=parsed_url.netloc,
sasl_transport = SASLTransport
if self._transport_type.name is 'AmqpOverWebsocket' or kwargs.get("http_proxy"):
sasl_transport = SASLWithWebSocket
endpoint = parsed_url.hostname + parsed_url.path
self._transport = sasl_transport(
host=endpoint,
port=self._port,
credential=kwargs['sasl_credential'],
**kwargs
)
else:
self._transport = Transport(parsed_url.netloc, **kwargs)
self._transport = Transport(parsed_url.netloc, 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 @@ -51,7 +51,7 @@
from ._platform import KNOWN_TCP_OPTS, SOL_TCP, pack, unpack
from ._encode import encode_frame
from ._decode import decode_frame, decode_empty_frame
from .constants import TLS_HEADER_FRAME
from .constants import TLS_HEADER_FRAME, WEBSOCKET_PORT, TransportType, AMQP_WS_SUBPROTOCOL


try:
Expand Down Expand Up @@ -647,11 +647,82 @@ def _read(self, n, initial=False, _errnos=(errno.EAGAIN, errno.EINTR)):
return result


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

Given a few parameters from the Connection constructor,
select and create a subclass of _AbstractTransport.
"""
transport = SSLTransport if ssl else TCPTransport
if transport_type == TransportType.AmqpOverWebsocket:
transport = WebSocketTransport
else:
transport = SSLTransport if ssl else TCPTransport
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
):
self.sslopts = ssl if isinstance(ssl, dict) else {}
self._connect_timeout = connect_timeout
self._host = host
super().__init__(
host, port, connect_timeout, **kwargs
)
self.ws = None
self._http_proxy = kwargs.get('http_proxy', None)

def connect(self):
http_proxy_host, http_proxy_port, http_proxy_auth = None, None, None
if self._http_proxy:
http_proxy_host = self._http_proxy['proxy_hostname']
http_proxy_port = self._http_proxy['proxy_port']
username = self._http_proxy.get('username', None)
password = self._http_proxy.get('password', None)
if username or password:
http_proxy_auth = (username, password)
try:
from websocket import create_connection
self.ws = create_connection(
url="wss://{}".format(self._host),
subprotocols=[AMQP_WS_SUBPROTOCOL],
timeout=self._connect_timeout,
skip_utf8_validation=True,
sslopt=self.sslopts,
http_proxy_host=http_proxy_host,
http_proxy_port=http_proxy_port,
http_proxy_auth=http_proxy_auth
)
except ImportError:
raise ValueError("Please install websocket-client library to use websocket transport.")

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)
length += nbytes
n -= nbytes
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

def _shutdown_transport(self):
"""Do any preliminary work in shutting down the connection."""
self.ws.close()

def _write(self, s):
"""Completely write a string to the peer.
ABNF, OPCODE_BINARY = 0x2
See http://tools.ietf.org/html/rfc5234
http://tools.ietf.org/html/rfc6455#section-5.2
"""
self.ws.send_binary(s)
11 changes: 10 additions & 1 deletion sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@
SenderSettleMode,
ReceiverSettleMode,
LinkDeliverySettleReason,
TransportType,
SEND_DISPOSITION_ACCEPT,
SEND_DISPOSITION_REJECT,
AUTH_TYPE_CBS,
Expand Down Expand Up @@ -155,6 +156,12 @@ def __init__(self, hostname, auth=None, **kwargs):
self._receive_settle_mode = kwargs.pop('receive_settle_mode', ReceiverSettleMode.Second)
self._desired_capabilities = kwargs.pop('desired_capabilities', None)

# transport
if kwargs.get('transport_type') is TransportType.Amqp and kwargs.get('http_proxy') is not None:
raise ValueError("Http proxy settings can't be passed if transport_type is explicitly set to Amqp")
self._transport_type = kwargs.pop('transport_type', TransportType.Amqp)
self._http_proxy = kwargs.pop('http_proxy', None)

def __enter__(self):
"""Run Client in a context manager."""
self.open()
Expand Down Expand Up @@ -240,7 +247,9 @@ def open(self):
channel_max=self._channel_max,
idle_timeout=self._idle_timeout,
properties=self._properties,
network_trace=self._network_trace
network_trace=self._network_trace,
transport_type=self._transport_type,
http_proxy=self._http_proxy
)
self._connection.open()
if not self._session:
Expand Down
16 changes: 16 additions & 0 deletions sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,12 @@
#: The port number is reserved for future transport mappings to these protocols.
PORT = 5672

# default port for AMQP over Websocket
WEBSOCKET_PORT = 443

# subprotocol for AMQP over Websocket
AMQP_WS_SUBPROTOCOL = 'AMQPWSB10'


#: The IANA assigned port number for secure AMQP (amqps).The standard AMQP port number that has been assigned
#: by IANA for secure TCP using TLS. Implementations listening on this port should NOT expect a protocol
Expand Down Expand Up @@ -302,3 +308,13 @@ class MessageDeliveryState(object):
MessageDeliveryState.Timeout,
MessageDeliveryState.Cancelled
)

class TransportType(Enum):
"""Transport type
The underlying transport protocol type:
Amqp: AMQP over the default TCP transport protocol, it uses port 5671.
AmqpOverWebsocket: Amqp over the Web Sockets transport protocol, it uses
port 443.
"""
Amqp = 1
AmqpOverWebsocket = 2
74 changes: 48 additions & 26 deletions sdk/eventhub/azure-eventhub/azure/eventhub/_pyamqp/sasl.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,9 @@
import struct
from enum import Enum

from ._transport import SSLTransport, AMQPS_PORT
from ._transport import SSLTransport, WebSocketTransport, AMQPS_PORT
from .types import AMQPTypes, TYPE, VALUE
from .constants import FIELD, SASLCode, SASL_HEADER_FRAME
from .constants import FIELD, SASLCode, SASL_HEADER_FRAME, TransportType, WEBSOCKET_PORT
from .performatives import (
SASLOutcome,
SASLResponse,
Expand Down Expand Up @@ -68,8 +68,33 @@ class SASLExternalCredential(object):
def start(self):
return b''

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):
class SASLTransport(SSLTransport, SASLTransportMixin):

def __init__(self, host, credential, port=AMQPS_PORT, connect_timeout=None, ssl=None, **kwargs):
self.credential = credential
Expand All @@ -78,26 +103,23 @@ def __init__(self, host, credential, port=AMQPS_PORT, connect_timeout=None, ssl=

def negotiate(self):
with self.block():
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))
self._negotiate()

class SASLWithWebSocket(WebSocketTransport, SASLTransportMixin):

def __init__(self, host, credential, port=WEBSOCKET_PORT, connect_timeout=None, ssl=None, **kwargs):
self.credential = credential
ssl = ssl or True
http_proxy = kwargs.pop('http_proxy', None)
self._transport = WebSocketTransport(
host,
port=port,
connect_timeout=connect_timeout,
ssl=ssl,
http_proxy=http_proxy,
**kwargs
)
super().__init__(host, port, connect_timeout, ssl, **kwargs)

def negotiate(self):
self._negotiate()