crypto.py (4609B)
1 """CSRmesh cryptography: key derivation, AES-OFB encryption, HMAC-SHA256 checksums.""" 2 3 from __future__ import annotations 4 5 import hashlib 6 import hmac 7 import secrets 8 9 from Cryptodome.Cipher import AES 10 11 12 class CsrMeshCrypto: 13 """CSRmesh encryption/decryption engine. 14 15 Handles packet encryption, decryption, checksum computation, 16 and sequence number management for the in-lite BLE mesh protocol. 17 """ 18 19 def __init__( 20 self, passphrase: str | None = None, network_key: bytes | None = None 21 ) -> None: 22 """Create crypto from either a mesh passphrase or its direct key.""" 23 if (passphrase is None) == (network_key is None): 24 raise ValueError("provide exactly one of passphrase or network_key") 25 if network_key is not None and len(network_key) != 16: 26 raise ValueError("network_key must be exactly 16 bytes") 27 self._tx_seq_nr = secrets.randbelow(0xFFFFFF) 28 self._controller_address = 0x8000 + secrets.randbelow(0xFFFD - 0x8000) 29 self._enc_key = ( 30 bytes(network_key) 31 if network_key is not None 32 else self._derive_key(passphrase) 33 ) 34 35 @property 36 def controller_address(self) -> int: 37 return self._controller_address 38 39 @staticmethod 40 def derive_key(passphrase: str) -> bytes: 41 """Derive 16-byte AES key: SHA-256(UTF-8(passphrase + '\\x00MCP')), reversed.""" 42 raw = (passphrase + "\x00MCP").encode("utf-8") 43 digest = hashlib.sha256(raw).digest() 44 return bytes(digest[len(digest) - 1 - i] for i in range(16)) 45 46 _derive_key = derive_key 47 48 @staticmethod 49 def _generate_iv(seq_nr: int, src_id: int) -> bytes: 50 iv = bytearray(16) 51 iv[0] = seq_nr & 0xFF 52 iv[1] = (seq_nr >> 8) & 0xFF 53 iv[2] = (seq_nr >> 16) & 0xFF 54 iv[4] = src_id & 0xFF 55 iv[5] = (src_id >> 8) & 0xFF 56 return bytes(iv) 57 58 def _get_checksum(self, seq_nr: int, src_id: int, encrypted: bytes) -> bytes: 59 buf = bytearray(8) # 8 zero bytes 60 buf.append(seq_nr & 0xFF) 61 buf.append((seq_nr >> 8) & 0xFF) 62 buf.append((seq_nr >> 16) & 0xFF) 63 buf.append(src_id & 0xFF) 64 buf.append((src_id >> 8) & 0xFF) 65 buf.extend(encrypted) 66 h = hmac.new(self._enc_key, bytes(buf), hashlib.sha256).digest() 67 return bytes(h[len(h) - 1 - i] for i in range(8)) 68 69 def encrypt_packet( # noqa: D417 70 self, dest_id: int, pkt_type: int, data: bytes, ttl: int = 5 71 ) -> bytes: 72 """Encrypt and build a CSRmesh packet ready for BLE transmission.""" 73 self._tx_seq_nr = (self._tx_seq_nr + 1) % 0x1000000 74 seq = self._tx_seq_nr 75 src = self._controller_address 76 77 # Build header: seq(3 LE) + src(2 LE) 78 header = bytearray() 79 header.append(seq & 0xFF) 80 header.append((seq >> 8) & 0xFF) 81 header.append((seq >> 16) & 0xFF) 82 header.append(src & 0xFF) 83 header.append((src >> 8) & 0xFF) 84 85 # Build plaintext: dest(2 LE) + pkt_type(1) + data 86 plaintext = bytearray() 87 plaintext.append(dest_id & 0xFF) 88 plaintext.append((dest_id >> 8) & 0xFF) 89 plaintext.append(pkt_type) 90 plaintext.extend(data) 91 92 # Encrypt with AES-OFB 93 iv = self._generate_iv(seq, src) 94 cipher = AES.new(self._enc_key, AES.MODE_OFB, iv=iv) 95 encrypted = cipher.encrypt(bytes(plaintext)) 96 97 # Build checksum 98 checksum = self._get_checksum(seq, src, encrypted) 99 100 # Assemble: header + encrypted + checksum + ttl 101 packet = bytearray(header) 102 packet.extend(encrypted) 103 packet.extend(checksum) 104 packet.append(ttl) 105 return bytes(packet) 106 107 def decrypt_packet(self, raw: bytes) -> dict | None: 108 """Decrypt a CSRmesh packet. Returns dict or None if checksum fails.""" 109 if len(raw) < 14: 110 return None 111 112 seq = raw[0] | (raw[1] << 8) | (raw[2] << 16) 113 src = raw[3] | (raw[4] << 8) 114 ttl = raw[-1] 115 checksum = raw[-9:-1] 116 encrypted = raw[5:-9] 117 118 if len(encrypted) < 3: 119 return None 120 121 expected = self._get_checksum(seq, src, encrypted) 122 if checksum != expected: 123 return None 124 125 iv = self._generate_iv(seq, src) 126 cipher = AES.new(self._enc_key, AES.MODE_OFB, iv=iv) 127 dec = cipher.decrypt(encrypted) 128 129 return { 130 "seq_nr": seq, 131 "src_id": src, 132 "dest_id": dec[0] | (dec[1] << 8), 133 "pkt_type": dec[2], 134 "ttl": ttl, 135 "data": dec[3:], 136 }