#!/usr/bin/env python3

"""Minimal MQTT 3.1.1 over WebSocket pub/sub smoke test using stdlib only."""

import argparse
import base64
import hashlib
import os
import socket
import struct
import sys


def mqtt_length(value):
    encoded = bytearray()
    while True:
        byte = value % 128
        value //= 128
        if value:
            byte |= 0x80
        encoded.append(byte)
        if not value:
            return bytes(encoded)


def mqtt_string(value):
    data = value.encode()
    return struct.pack("!H", len(data)) + data


def mqtt_packet(kind, body):
    return bytes([kind]) + mqtt_length(len(body)) + body


class WebSocket:
    def __init__(self, host, port):
        self.sock = socket.create_connection((host, port), timeout=8)
        key = base64.b64encode(os.urandom(16)).decode()
        request = (
            f"GET /mqtt HTTP/1.1\r\n"
            f"Host: {host}:{port}\r\n"
            "Upgrade: websocket\r\n"
            "Connection: Upgrade\r\n"
            f"Sec-WebSocket-Key: {key}\r\n"
            "Sec-WebSocket-Version: 13\r\n"
            "Sec-WebSocket-Protocol: mqtt\r\n\r\n"
        )
        self.sock.sendall(request.encode())
        response = self._headers()
        if b" 101 " not in response.split(b"\r\n", 1)[0]:
            raise RuntimeError("WebSocket upgrade failed")
        accept = base64.b64encode(
            hashlib.sha1(
                (key + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11").encode()
            ).digest()
        )
        if b"sec-websocket-accept: " + accept.lower() not in response.lower():
            raise RuntimeError("WebSocket accept mismatch")

    def _headers(self):
        data = bytearray()
        while b"\r\n\r\n" not in data:
            chunk = self.sock.recv(4096)
            if not chunk:
                raise RuntimeError("connection closed during WebSocket upgrade")
            data.extend(chunk)
            if len(data) > 65536:
                raise RuntimeError("oversized WebSocket headers")
        return bytes(data)

    def _exact(self, length):
        data = bytearray()
        while len(data) < length:
            chunk = self.sock.recv(length - len(data))
            if not chunk:
                raise RuntimeError("WebSocket connection closed")
            data.extend(chunk)
        return bytes(data)

    def send(self, payload, opcode=2):
        mask = os.urandom(4)
        length = len(payload)
        header = bytearray([0x80 | opcode])
        if length < 126:
            header.append(0x80 | length)
        elif length < 65536:
            header.append(0x80 | 126)
            header.extend(struct.pack("!H", length))
        else:
            header.append(0x80 | 127)
            header.extend(struct.pack("!Q", length))
        masked = bytes(byte ^ mask[index % 4] for index, byte in enumerate(payload))
        self.sock.sendall(bytes(header) + mask + masked)

    def receive(self):
        while True:
            first, second = self._exact(2)
            opcode = first & 0x0F
            length = second & 0x7F
            if length == 126:
                length = struct.unpack("!H", self._exact(2))[0]
            elif length == 127:
                length = struct.unpack("!Q", self._exact(8))[0]
            mask = self._exact(4) if second & 0x80 else b""
            payload = self._exact(length)
            if mask:
                payload = bytes(
                    byte ^ mask[index % 4] for index, byte in enumerate(payload)
                )
            if opcode == 9:
                self.send(payload, opcode=10)
                continue
            if opcode == 8:
                raise RuntimeError("WebSocket closed by server")
            if opcode in (0, 2):
                return payload

    def close(self):
        try:
            self.send(b"", opcode=8)
        finally:
            self.sock.close()


def mqtt_connect(ws, client_id):
    variable = mqtt_string("MQTT") + bytes([4, 2]) + struct.pack("!H", 30)
    ws.send(mqtt_packet(0x10, variable + mqtt_string(client_id)))
    if not ws.receive().startswith(b"\x20\x02\x00\x00"):
        raise RuntimeError("MQTT CONNACK failed")


def parse_publish(packet):
    if not packet.startswith(b"\x30"):
        raise RuntimeError("MQTT PUBLISH was not received")
    index = 1
    multiplier = 1
    remaining = 0
    while True:
        byte = packet[index]
        index += 1
        remaining += (byte & 127) * multiplier
        if not byte & 128:
            break
        multiplier *= 128
    topic_length = struct.unpack("!H", packet[index:index + 2])[0]
    index += 2
    topic = packet[index:index + topic_length].decode()
    index += topic_length
    payload_length = remaining - topic_length - 2
    return topic, packet[index:index + payload_length].decode()


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--subscriber-host", required=True)
    parser.add_argument("--publisher-host", required=True)
    parser.add_argument("--topic", required=True)
    parser.add_argument("--payload", required=True)
    args = parser.parse_args()
    subscriber = WebSocket(args.subscriber_host, 8083)
    publisher = WebSocket(args.publisher_host, 8083)
    try:
        mqtt_connect(subscriber, "q240-sub-" + os.urandom(6).hex())
        mqtt_connect(publisher, "q240-pub-" + os.urandom(6).hex())
        subscribe = struct.pack("!H", 1) + mqtt_string(args.topic) + b"\x00"
        subscriber.send(mqtt_packet(0x82, subscribe))
        if not subscriber.receive().startswith(b"\x90"):
            raise RuntimeError("MQTT SUBACK failed")
        publisher.send(mqtt_packet(0x30, mqtt_string(args.topic) + args.payload.encode()))
        topic, payload = parse_publish(subscriber.receive())
        if topic != args.topic or payload != args.payload:
            raise RuntimeError("MQTT WebSocket topic or payload mismatch")
        print("WS_MQTT_PUBSUB_MATCH=PASS")
    finally:
        subscriber.close()
        publisher.close()


if __name__ == "__main__":
    try:
        main()
    except Exception as exc:
        print(f"ERROR: {exc}", file=sys.stderr)
        sys.exit(1)
