Files
matrix-screen-controller/核桃派软件源代码/app/mobile/protocol.py
T

166 lines
7.0 KiB
Python

"""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):
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, 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:], LABEL + self.transcript + bytes([d]) + seq)
self.rx_sequence += 1
return result
except Exception:
self.failed = True
raise ProtocolError("BAD_REQUEST") from None