aboutsummaryrefslogtreecommitdiffstats
path: root/packages/meshbay-common/src/meshbay_common/senderkeys.py
diff options
context:
space:
mode:
Diffstat (limited to 'packages/meshbay-common/src/meshbay_common/senderkeys.py')
-rw-r--r--packages/meshbay-common/src/meshbay_common/senderkeys.py287
1 files changed, 287 insertions, 0 deletions
diff --git a/packages/meshbay-common/src/meshbay_common/senderkeys.py b/packages/meshbay-common/src/meshbay_common/senderkeys.py
new file mode 100644
index 0000000..932e2e6
--- /dev/null
+++ b/packages/meshbay-common/src/meshbay_common/senderkeys.py
@@ -0,0 +1,287 @@
+"""
+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)