""" MeshBay — Sender Keys protocol for group messaging. Signal Groups approach: each member maintains their own sending chain. Advantages over shared Double Ratchet: - O(N) state per group (one chain per member) vs O(N^2) pairwise - Single encrypt per message (not N encryptions) - No key/nonce reuse — each sender has an independent chain Key components: - Chain key ratchet: HKDF per message, provides forward secrecy - Message key derivation: separate HKDF from chain key - Ed25519 signing: each sender signs their ciphertext - AES-256-GCM encryption: browser-compatible symmetric cipher Key distribution: - On join: admin wraps each sender's SenderKeyDistribution with GEK - On leave: all remaining members rotate their chain keys """ import os import struct from dataclasses import dataclass, field from cryptography.hazmat.primitives.asymmetric.ed25519 import ( Ed25519PrivateKey, Ed25519PublicKey, ) from cryptography.hazmat.primitives.ciphers.aead import AESGCM from cryptography.hazmat.primitives.kdf.hkdf import HKDF from cryptography.hazmat.primitives import hashes, serialization CHAIN_INFO = b"meshbay:sk:chain:v1" MSG_KEY_INFO = b"meshbay:sk:msg:v1" CHAIN_KEY_LEN = 32 MSG_KEY_LEN = 32 MAX_SKIP = 256 def _hkdf(ikm: bytes, info: bytes, length: int = 32) -> bytes: return HKDF( algorithm=hashes.SHA256(), length=length, salt=None, info=info, ).derive(ikm) def _ratchet_chain(chain_key: bytes) -> tuple[bytes, bytes]: """Advance chain key → (new_chain_key, message_key).""" new_ck = _hkdf(chain_key, CHAIN_INFO, CHAIN_KEY_LEN) mk = _hkdf(chain_key, MSG_KEY_INFO, MSG_KEY_LEN) return new_ck, mk # ── Data structures ────────────────────────────────────────────────────────── @dataclass class SenderKeyDistribution: """Sent to group members when a sender joins or rotates.""" sender_id: str chain_key: bytes # 32-byte initial chain key iteration: int # current message counter signing_pk: bytes # 32-byte raw Ed25519 public key def serialize(self) -> bytes: sender_bytes = self.sender_id.encode() return ( struct.pack(">H", len(sender_bytes)) + sender_bytes + self.chain_key + struct.pack(">I", self.iteration) + self.signing_pk ) @classmethod def deserialize(cls, data: bytes) -> "SenderKeyDistribution": sender_len = struct.unpack(">H", data[:2])[0] offset = 2 sender_id = data[offset:offset + sender_len].decode() offset += sender_len chain_key = data[offset:offset + 32] offset += 32 iteration = struct.unpack(">I", data[offset:offset + 4])[0] offset += 4 signing_pk = data[offset:offset + 32] return cls(sender_id=sender_id, chain_key=chain_key, iteration=iteration, signing_pk=signing_pk) @dataclass class SenderKeyState: """One sender's chain state as seen by any group member.""" sender_id: str chain_key: bytes iteration: int signing_key: Ed25519PublicKey _skipped_keys: dict[int, bytes] = field(default_factory=dict) @classmethod def from_distribution(cls, dist: SenderKeyDistribution) -> "SenderKeyState": pk = Ed25519PublicKey.from_public_bytes(dist.signing_pk) return cls( sender_id=dist.sender_id, chain_key=dist.chain_key, iteration=dist.iteration, signing_key=pk, ) def advance_to(self, target: int) -> bytes: """Advance chain to target iteration, caching skipped keys. Returns message key.""" if target < self.iteration: mk = self._skipped_keys.pop(target, None) if mk is None: raise ValueError(f"Message key {target} already consumed or too old") return mk skip_count = target - self.iteration if skip_count > MAX_SKIP: raise ValueError(f"Too many skipped messages: {skip_count}") for i in range(skip_count): new_ck, mk = _ratchet_chain(self.chain_key) self._skipped_keys[self.iteration] = mk self.chain_key = new_ck self.iteration += 1 new_ck, mk = _ratchet_chain(self.chain_key) self.chain_key = new_ck self.iteration += 1 return mk @dataclass class SenderKeyRecord: """Sender's own key state (includes signing private key).""" sender_id: str chain_key: bytes iteration: int signing_sk: Ed25519PrivateKey @classmethod def create(cls, sender_id: str) -> "SenderKeyRecord": return cls( sender_id=sender_id, chain_key=os.urandom(CHAIN_KEY_LEN), iteration=0, signing_sk=Ed25519PrivateKey.generate(), ) def distribution(self) -> SenderKeyDistribution: pk_raw = self.signing_sk.public_key().public_bytes( serialization.Encoding.Raw, serialization.PublicFormat.Raw) return SenderKeyDistribution( sender_id=self.sender_id, chain_key=self.chain_key, iteration=self.iteration, signing_pk=pk_raw, ) def rotate(self) -> "SenderKeyRecord": """Create a new record with fresh chain key (call on member removal).""" return SenderKeyRecord( sender_id=self.sender_id, chain_key=os.urandom(CHAIN_KEY_LEN), iteration=0, signing_sk=Ed25519PrivateKey.generate(), ) # ── Group store ────────────────────────────────────────────────────────────── class GroupSenderKeyStore: """All sender key states for one group, held by one member.""" def __init__(self, group_id: str): self.group_id = group_id self._states: dict[str, SenderKeyState] = {} def add_sender(self, dist: SenderKeyDistribution) -> None: self._states[dist.sender_id] = SenderKeyState.from_distribution(dist) def remove_sender(self, sender_id: str) -> None: self._states.pop(sender_id, None) def get_state(self, sender_id: str) -> SenderKeyState | None: return self._states.get(sender_id) @property def sender_count(self) -> int: return len(self._states) # ── Encrypt / Decrypt ──────────────────────────────────────────────────────── @dataclass class SenderKeyMessage: """Wire format for a Sender Keys encrypted message.""" sender_id: str iteration: int ciphertext: bytes nonce: bytes signature: bytes def serialize(self) -> bytes: sender_bytes = self.sender_id.encode() return ( struct.pack(">H", len(sender_bytes)) + sender_bytes + struct.pack(">I", self.iteration) + struct.pack(">I", len(self.ciphertext)) + self.ciphertext + self.nonce + self.signature ) @classmethod def deserialize(cls, data: bytes) -> "SenderKeyMessage": offset = 0 sender_len = struct.unpack(">H", data[offset:offset + 2])[0] offset += 2 sender_id = data[offset:offset + sender_len].decode() offset += sender_len iteration = struct.unpack(">I", data[offset:offset + 4])[0] offset += 4 ct_len = struct.unpack(">I", data[offset:offset + 4])[0] offset += 4 ciphertext = data[offset:offset + ct_len] offset += ct_len nonce = data[offset:offset + 12] offset += 12 signature = data[offset:offset + 64] return cls(sender_id=sender_id, iteration=iteration, ciphertext=ciphertext, nonce=nonce, signature=signature) def encrypt_message( record: SenderKeyRecord, plaintext: bytes, aad: bytes = b"", ) -> tuple[SenderKeyMessage, SenderKeyRecord]: """ Encrypt a message with the sender's chain key. Returns (message, updated_record). """ new_ck, mk = _ratchet_chain(record.chain_key) iteration = record.iteration nonce = os.urandom(12) ct = AESGCM(mk).encrypt(nonce, plaintext, aad or None) sig_payload = struct.pack(">I", iteration) + nonce + ct signature = record.signing_sk.sign(sig_payload) msg = SenderKeyMessage( sender_id=record.sender_id, iteration=iteration, ciphertext=ct, nonce=nonce, signature=signature, ) updated = SenderKeyRecord( sender_id=record.sender_id, chain_key=new_ck, iteration=iteration + 1, signing_sk=record.signing_sk, ) return msg, updated def decrypt_message( store: GroupSenderKeyStore, msg: SenderKeyMessage, aad: bytes = b"", ) -> bytes: """ Decrypt and verify a Sender Keys message. Advances the sender's chain state in the store. """ state = store.get_state(msg.sender_id) if state is None: raise ValueError(f"Unknown sender: {msg.sender_id}") sig_payload = struct.pack(">I", msg.iteration) + msg.nonce + msg.ciphertext state.signing_key.verify(msg.signature, sig_payload) mk = state.advance_to(msg.iteration) return AESGCM(mk).decrypt(msg.nonce, msg.ciphertext, aad or None)