From d07ba472be4ff0af072c26f8ff2f6c5e0e5faa66 Mon Sep 17 00:00:00 2001 From: Robert Resch Date: Sat, 20 Aug 2022 14:53:52 +0200 Subject: [PATCH] allow self signed certs --- bumper/mqtt/proxy.py | 138 ++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 136 insertions(+), 2 deletions(-) diff --git a/bumper/mqtt/proxy.py b/bumper/mqtt/proxy.py index 8b15711..ee786b7 100644 --- a/bumper/mqtt/proxy.py +++ b/bumper/mqtt/proxy.py @@ -1,11 +1,27 @@ """Mqtt proxy module.""" import asyncio import re +import ssl +import typing from typing import Any, MutableMapping +from urllib.parse import urlparse, urlunparse -from amqtt.client import MQTTClient +import websockets +from amqtt.adapters import ( + StreamReaderAdapter, + StreamWriterAdapter, + WebSocketsReader, + WebSocketsWriter, +) +from amqtt.client import ConnectException, MQTTClient +from amqtt.mqtt.connack import CONNECTION_ACCEPTED from amqtt.mqtt.constants import QOS_0, QOS_1, QOS_2 +from amqtt.mqtt.protocol.client_handler import ClientProtocolHandler +from amqtt.mqtt.protocol.handler import ProtocolHandlerException from cachetools import TTLCache +from websockets.exceptions import InvalidHandshake, InvalidURI +from websockets.legacy.client import connect +from websockets.typing import Subprotocol from ..util import get_logger @@ -36,7 +52,7 @@ class ProxyClient: self.request_mapper: MutableMapping[str, str] = TTLCache( maxsize=timeout * 60, ttl=timeout * 1.1 ) - self._client = MQTTClient(client_id=client_id, config=config) + self._client = _NoCertVerifyClient(client_id=client_id, config=config) self._host = host self._port = port @@ -87,3 +103,121 @@ class ProxyClient: async def publish(self, topic: str, message: bytes, qos: int | None = None) -> None: await self._client.publish(topic, message, qos) + + +class _NoCertVerifyClient(MQTTClient): + """ + Mqtt client, which is not verify the certificate. + + Purpose is only to add "sc.verify_mode = ssl.CERT_NONE # Ignore verify of cert" + """ + + @typing.no_type_check + async def _connect_coro(self): + kwargs = dict() + + # Decode URI attributes + uri_attributes = urlparse(self.session.broker_uri) + scheme = uri_attributes.scheme + secure = True if scheme in ("mqtts", "wss") else False + self.session.username = ( + self.session.username if self.session.username else uri_attributes.username + ) + self.session.password = ( + self.session.password if self.session.password else uri_attributes.password + ) + self.session.remote_address = uri_attributes.hostname + self.session.remote_port = uri_attributes.port + if scheme in ("mqtt", "mqtts") and not self.session.remote_port: + self.session.remote_port = 8883 if scheme == "mqtts" else 1883 + if scheme in ("ws", "wss") and not self.session.remote_port: + self.session.remote_port = 443 if scheme == "wss" else 80 + if scheme in ("ws", "wss"): + # Rewrite URI to conform to https://tools.ietf.org/html/rfc6455#section-3 + uri = ( + scheme, + self.session.remote_address + ":" + str(self.session.remote_port), + uri_attributes[2], + uri_attributes[3], + uri_attributes[4], + uri_attributes[5], + ) + self.session.broker_uri = urlunparse(uri) + # Init protocol handler + # if not self._handler: + self._handler = ClientProtocolHandler(self.plugins_manager) + + if secure: + sc = ssl.create_default_context( + ssl.Purpose.SERVER_AUTH, + cafile=self.session.cafile, + capath=self.session.capath, + cadata=self.session.cadata, + ) + if "certfile" in self.config and "keyfile" in self.config: + sc.load_cert_chain(self.config["certfile"], self.config["keyfile"]) + if "check_hostname" in self.config and isinstance( + self.config["check_hostname"], bool + ): + sc.check_hostname = self.config["check_hostname"] + + sc.verify_mode = ssl.CERT_NONE # Ignore verify of cert + kwargs["ssl"] = sc + + try: + reader = None + writer = None + self._connected_state.clear() + # Open connection + if scheme in ("mqtt", "mqtts"): + conn_reader, conn_writer = await asyncio.open_connection( + self.session.remote_address, self.session.remote_port, **kwargs + ) + reader = StreamReaderAdapter(conn_reader) + writer = StreamWriterAdapter(conn_writer) + elif scheme in ("ws", "wss"): + websocket = await websockets.connect( + self.session.broker_uri, + subprotocols=["mqtt"], + extra_headers=self.extra_headers, + **kwargs, + ) + reader = WebSocketsReader(websocket) + writer = WebSocketsWriter(websocket) + # Start MQTT protocol + self._handler.attach(self.session, reader, writer) + return_code = await self._handler.mqtt_connect() + if return_code is not CONNECTION_ACCEPTED: + self.session.transitions.disconnect() + self.logger.warning("Connection rejected with code '%s'" % return_code) + exc = ConnectException("Connection rejected by broker") + exc.return_code = return_code + raise exc + else: + # Handle MQTT protocol + await self._handler.start() + self.session.transitions.connect() + self._connected_state.set() + self.logger.debug( + "connected to %s:%s" + % (self.session.remote_address, self.session.remote_port) + ) + return return_code + except InvalidURI as iuri: + self.logger.warning( + "connection failed: invalid URI '%s'" % self.session.broker_uri + ) + self.session.transitions.disconnect() + raise ConnectException( + "connection failed: invalid URI '%s'" % self.session.broker_uri, iuri + ) + except InvalidHandshake as ihs: + self.logger.warning("connection failed: invalid websocket handshake") + self.session.transitions.disconnect() + raise ConnectException( + "connection failed: invalid websocket handshake", ihs + ) + except (ProtocolHandlerException, ConnectionError, OSError) as e: + self.logger.warning("MQTT connection failed: %r" % e) + self.session.transitions.disconnect() + raise ConnectException(e)