"""QMS BLE v1 framing and ephemeral authenticated encryption. This encrypts traffic, but does not authenticate the identity of an App. Each instance belongs to exactly one connection; discard it after any error. """ from __future__ import annotations import base64 import hashlib import json import os import struct import time from cryptography.hazmat.primitives import hashes, serialization from cryptography.hazmat.primitives.asymmetric import ec from cryptography.hazmat.primitives.ciphers.aead import AESGCM from cryptography.hazmat.primitives.kdf.hkdf import HKDF MAX_MESSAGE = 65536 HEADER = struct.Struct("!HBBIHHI") LABEL = b"QMS-BLE-1" class ProtocolError(ValueError): """Fixed public error: never include untrusted payload or crypto details.""" def fragments(kind: int, message_id: int, payload: bytes, mtu: int): if kind not in (1, 2, 3, 4) or not 0 <= message_id <= 0xFFFFFFFF: raise ProtocolError("BAD_REQUEST") if not 1 <= len(payload) <= MAX_MESSAGE or not 23 <= mtu <= 517: raise ProtocolError("BAD_REQUEST") capacity = mtu - 19 count = (len(payload) + capacity - 1) // capacity for index in range(count): yield HEADER.pack(0x514D, 1, kind, message_id, index, count, len(payload)) + payload[index * capacity:(index + 1) * capacity] class Reassembler: def __init__(self, clock=time.monotonic): self.clock = clock self.next_message = 0 self.pending = None self.buffer = bytearray() self.index = 0 self.started = 0.0 def accept(self, part: bytes): if not 16 < len(part) <= 514: raise ProtocolError("BAD_REQUEST") magic, wire, kind, mid, index, count, total = HEADER.unpack(part[:16]) if magic != 0x514D or wire != 1 or kind not in (1, 2, 3, 4) or not 1 <= count <= total <= MAX_MESSAGE: raise ProtocolError("BAD_REQUEST") signature = (kind, mid, count, total) if self.pending is None: if index != 0 or mid != self.next_message: raise ProtocolError("BAD_REQUEST") self.pending = signature self.started = self.clock() if self.clock() - self.started >= 15: raise ProtocolError("TIMEOUT") if self.pending != signature or index != self.index or index >= count: raise ProtocolError("BAD_REQUEST") self.buffer.extend(part[16:]) self.index += 1 if len(self.buffer) > total or (self.index < count and len(self.buffer) >= total): raise ProtocolError("BAD_REQUEST") if self.index != count: return None if len(self.buffer) != total: raise ProtocolError("BAD_REQUEST") result = kind, bytes(self.buffer) self.pending = None self.buffer.clear() self.index = 0 self.next_message += 1 return result def json_object(raw: bytes) -> dict: def unique(pairs): result = {} for key, value in pairs: if key in result: raise ValueError("duplicate") result[key] = value return result try: value = json.loads(raw.decode("utf-8"), object_pairs_hook=unique, parse_constant=lambda _: (_ for _ in ()).throw(ValueError())) if not isinstance(value, dict): raise ValueError() return value except (ValueError, UnicodeError, RecursionError): raise ProtocolError("BAD_REQUEST") from None class Handshake: def __init__(self, *, private_key=None, random_bytes=None): # Explicit values are exclusively for public interoperability fixtures. self.key = private_key or ec.generate_private_key(ec.SECP256R1()) self.random = random_bytes if random_bytes is not None else os.urandom(32) if len(self.random) != 32: raise ValueError("random length") public = self.key.public_key().public_bytes(serialization.Encoding.X962, serialization.PublicFormat.UncompressedPoint) self.hello = json.dumps(dict(protocol_major=1, protocol_minor=0, public_key=base64.b64encode(public).decode("ascii"), random=base64.b64encode(self.random).decode("ascii")), separators=(",", ":")).encode("utf-8") def finish(self, peer_hello: bytes, *, server: bool): if len(peer_hello) > 1024: raise ProtocolError("BAD_REQUEST") peer = json_object(peer_hello) if type(peer.get("protocol_major")) is not int or peer["protocol_major"] != 1: raise ProtocolError("UNSUPPORTED_VERSION") if type(peer.get("protocol_minor")) is not int or peer["protocol_minor"] < 0: raise ProtocolError("BAD_REQUEST") try: public = base64.b64decode(peer["public_key"], validate=True) random = base64.b64decode(peer["random"], validate=True) if len(public) != 65 or public[0] != 4 or len(random) != 32: raise ValueError() peer_key = ec.EllipticCurvePublicKey.from_encoded_point(ec.SECP256R1(), public) secret = self.key.exchange(ec.ECDH(), peer_key) except (KeyError, TypeError, ValueError): raise ProtocolError("BAD_REQUEST") from None c, s = (peer_hello, self.hello) if server else (self.hello, peer_hello) cr, sr = (random, self.random) if server else (self.random, random) transcript = hashlib.sha256(struct.pack("!I", len(c)) + c + struct.pack("!I", len(s)) + s).digest() material = HKDF(algorithm=hashes.SHA256(), length=72, salt=hashlib.sha256(cr + sr).digest(), info=LABEL + transcript).derive(secret) return Cipher(material, transcript, server=server) class Cipher: def __init__(self, material: bytes, transcript: bytes, *, server: bool, label=LABEL): self.label = label self.tx_direction = 1 if server else 0 self.rx_direction = 1 - self.tx_direction self.keys = (AESGCM(material[:32]), AESGCM(material[32:64])) self.prefixes = (material[64:68], material[68:72]) self.transcript = transcript self.tx_sequence = self.rx_sequence = 0 self.failed = False def encrypt(self, payload: bytes) -> bytes: if self.failed or not 1 <= len(payload) <= MAX_MESSAGE - 24 or self.tx_sequence >= 2**64: raise ProtocolError("BAD_REQUEST") seq = struct.pack("!Q", self.tx_sequence) d = self.tx_direction result = seq + self.keys[d].encrypt(self.prefixes[d] + seq, payload, self.label + self.transcript + bytes([d]) + seq) self.tx_sequence += 1 return result def decrypt(self, record: bytes) -> bytes: try: if self.failed or not 25 <= len(record) <= MAX_MESSAGE or int.from_bytes(record[:8], "big") != self.rx_sequence: raise ValueError() seq, d = record[:8], self.rx_direction result = self.keys[d].decrypt(self.prefixes[d] + seq, record[8:], self.label + self.transcript + bytes([d]) + seq) self.rx_sequence += 1 return result except Exception: self.failed = True raise ProtocolError("BAD_REQUEST") from None