package mls // MLS ratchet tree types (RFC 9420 ยง7). // Data types and serialization only - crypto operations (sign, verify, // parentHash) go in tree_crypto.mx with the cipher suite. import "errors" var ( errInvalidLeafNodeSource error errInvalidNodeType error ) func bytesEqual(a, b []byte) (ok bool) { if len(a) != len(b) { return false } for i := range a { if a[i] != b[i] { return false } } return true } // --- ParentNode --- type parentNode struct { encryptionKey []byte parentHash []byte unmergedLeaves []leafIndex } func (nd *parentNode) unmarshal(r *Reader) (err error) { *nd = parentNode{} var ok bool nd.encryptionKey, ok = r.readOpaqueVec() if !ok { return errUnexpectedEOF } nd.parentHash, ok = r.readOpaqueVec() if !ok { return errUnexpectedEOF } return r.readVector(func(r *Reader) error { v, hasV := r.readUint32() if !hasV { return errUnexpectedEOF } nd.unmergedLeaves = push(nd.unmergedLeaves, leafIndex(v)) return nil }) } func (nd *parentNode) marshal(w *Writer) { w.writeOpaqueVec([]byte(nd.encryptionKey)) w.writeOpaqueVec(nd.parentHash) w.writeVector(len(nd.unmergedLeaves), func(w *Writer, i int32) { w.addUint32(uint32(nd.unmergedLeaves[i])) }) } // --- LeafNodeSource --- type leafNodeSource uint8 const ( leafNodeSourceKeyPackage leafNodeSource = 1 leafNodeSourceUpdate leafNodeSource = 2 leafNodeSourceCommit leafNodeSource = 3 ) func (src *leafNodeSource) unmarshal(r *Reader) (err error) { b, ok := r.readByte() if !ok { return errUnexpectedEOF } *src = leafNodeSource(b) switch *src { case leafNodeSourceKeyPackage, leafNodeSourceUpdate, leafNodeSourceCommit: return nil default: return errInvalidLeafNodeSource } } func (src *leafNodeSource) marshal(w *Writer) { w.addByte(byte(src)) } // --- Capabilities --- type capabilities struct { versions []protocolVersion cipherSuites []CipherSuite extensions []extensionType proposals []proposalType credentials []credentialType } func (caps *capabilities) unmarshal(r *Reader) (err error) { *caps = capabilities{} err := r.readVector(func(r *Reader) error { v, ok := r.readUint16() if !ok { return errUnexpectedEOF } caps.versions = push(caps.versions, protocolVersion(v)) return nil }) if err != nil { return err } err = r.readVector(func(r *Reader) error { v, ok := r.readUint16() if !ok { return errUnexpectedEOF } caps.cipherSuites = push(caps.cipherSuites, CipherSuite(v)) return nil }) if err != nil { return err } err = r.readVector(func(r *Reader) error { v, ok := r.readUint16() if !ok { return errUnexpectedEOF } caps.extensions = push(caps.extensions, extensionType(v)) return nil }) if err != nil { return err } err = r.readVector(func(r *Reader) error { v, ok := r.readUint16() if !ok { return errUnexpectedEOF } caps.proposals = push(caps.proposals, proposalType(v)) return nil }) if err != nil { return err } return r.readVector(func(r *Reader) error { v, ok := r.readUint16() if !ok { return errUnexpectedEOF } caps.credentials = push(caps.credentials, credentialType(v)) return nil }) } func (caps *capabilities) marshal(w *Writer) { w.writeVector(len(caps.versions), func(w *Writer, i int32) { w.addUint16(uint16(caps.versions[i])) }) w.writeVector(len(caps.cipherSuites), func(w *Writer, i int32) { w.addUint16(uint16(caps.cipherSuites[i])) }) w.writeVector(len(caps.extensions), func(w *Writer, i int32) { w.addUint16(uint16(caps.extensions[i])) }) w.writeVector(len(caps.proposals), func(w *Writer, i int32) { w.addUint16(uint16(caps.proposals[i])) }) w.writeVector(len(caps.credentials), func(w *Writer, i int32) { w.addUint16(uint16(caps.credentials[i])) }) } // --- Lifetime --- type lifetime struct { notBefore, notAfter uint64 } func (lt *lifetime) unmarshal(r *Reader) (err error) { *lt = lifetime{} var ok bool lt.notBefore, ok = r.readUint64() if !ok { return errUnexpectedEOF } lt.notAfter, ok = r.readUint64() if !ok { return errUnexpectedEOF } return nil } func (lt *lifetime) marshal(w *Writer) { w.addUint64(lt.notBefore) w.addUint64(lt.notAfter) } // --- Extension --- type extensionType uint16 const ( extensionTypeApplicationID extensionType = 0x0001 extensionTypeRatchetTree extensionType = 0x0002 extensionTypeRequiredCapabilities extensionType = 0x0003 extensionTypeExternalPub extensionType = 0x0004 extensionTypeExternalSenders extensionType = 0x0005 // Marmot: KeyPackage reusable for multiple Welcomes (MIP-00) ExtensionTypeLastResort extensionType = 0x000a // Marmot: Nostr group metadata (MIP-01) ExtensionTypeNostrGroupData extensionType = 0xf2ee ) type Extension = extension type extension struct { extensionType extensionType extensionData []byte } func NewExtension(t extensionType, data []byte) (v extension) { return extension{extensionType: t, extensionData: data} } type ExtensionType = extensionType func unmarshalExtensionVec(r *Reader) (list []extension, err2 error) { var exts []extension err := r.readVector(func(r *Reader) error { var ext extension v, ok := r.readUint16() if !ok { return errUnexpectedEOF } ext.extensionType = extensionType(v) ext.extensionData, ok = r.readOpaqueVec() if !ok { return errUnexpectedEOF } exts = push(exts, ext) return nil }) return exts, err } func marshalExtensionVec(w *Writer, exts []extension) { w.writeVector(len(exts), func(w *Writer, i int32) { ext := exts[i] w.addUint16(uint16(ext.extensionType)) w.writeOpaqueVec(ext.extensionData) }) } func findExtensionData(exts []extension, t extensionType) (buf []byte) { for _, ext := range exts { if ext.extensionType == t { return ext.extensionData } } return nil } // --- LeafNode --- type leafNode struct { encryptionKey []byte signatureKey []byte credential Credential capabilities capabilities leafNodeSource leafNodeSource lifetime *lifetime // for leafNodeSourceKeyPackage parentHash []byte // for leafNodeSourceCommit extensions []extension signature []byte } func (nd *leafNode) unmarshal(r *Reader) (err error) { *nd = leafNode{} var ok bool nd.encryptionKey, ok = r.readOpaqueVec() if !ok { return errUnexpectedEOF } nd.signatureKey, ok = r.readOpaqueVec() if !ok { return errUnexpectedEOF } if e := nd.credential.unmarshal(r); e != nil { return e } if e := nd.capabilities.unmarshal(r); e != nil { return e } if e := nd.leafNodeSource.unmarshal(r); e != nil { return e } var err error switch nd.leafNodeSource { case leafNodeSourceKeyPackage: nd.lifetime = &lifetime{} err = nd.lifetime.unmarshal(r) case leafNodeSourceCommit: nd.parentHash, ok = r.readOpaqueVec() if !ok { err = errUnexpectedEOF } } if err != nil { return err } exts, err := unmarshalExtensionVec(r) if err != nil { return err } nd.extensions = exts nd.signature, ok = r.readOpaqueVec() if !ok { return errUnexpectedEOF } return nil } func (nd *leafNode) marshalBase(w *Writer) { w.writeOpaqueVec([]byte(nd.encryptionKey)) w.writeOpaqueVec([]byte(nd.signatureKey)) nd.credential.marshal(w) nd.capabilities.marshal(w) nd.leafNodeSource.marshal(w) switch nd.leafNodeSource { case leafNodeSourceKeyPackage: nd.lifetime.marshal(w) case leafNodeSourceCommit: w.writeOpaqueVec(nd.parentHash) } marshalExtensionVec(w, nd.extensions) } func (nd *leafNode) marshal(w *Writer) { nd.marshalBase(w) w.writeOpaqueVec(nd.signature) } // --- LeafNodeTBS --- type leafNodeTBS struct { node *leafNode // for leafNodeSourceUpdate and leafNodeSourceCommit groupID []byte leafIndex leafIndex } func (tbs *leafNodeTBS) marshal(w *Writer) { tbs.node.marshalBase(w) switch tbs.node.leafNodeSource { case leafNodeSourceUpdate, leafNodeSourceCommit: w.writeOpaqueVec([]byte(tbs.groupID)) w.addUint32(uint32(tbs.leafIndex)) } } // --- UpdatePathNode --- type updatePathNode struct { encryptionKey []byte encryptedPathSecret []hpkeCiphertext } func (upn *updatePathNode) unmarshal(r *Reader) (err error) { *upn = updatePathNode{} var ok bool upn.encryptionKey, ok = r.readOpaqueVec() if !ok { return errUnexpectedEOF } return r.readVector(func(r *Reader) error { var ct hpkeCiphertext if e := ct.unmarshal(r); e != nil { return e } upn.encryptedPathSecret = push(upn.encryptedPathSecret, ct) return nil }) } func (upn *updatePathNode) marshal(w *Writer) { w.writeOpaqueVec([]byte(upn.encryptionKey)) w.writeVector(len(upn.encryptedPathSecret), func(w *Writer, i int32) { upn.encryptedPathSecret[i].marshal(w) }) } // --- UpdatePath --- type updatePath struct { leafNode leafNode nodes []updatePathNode } func (up *updatePath) unmarshal(r *Reader) (err error) { *up = updatePath{} if e := up.leafNode.unmarshal(r); e != nil { return e } return r.readVector(func(r *Reader) error { var nd updatePathNode if e := nd.unmarshal(r); e != nil { return e } up.nodes = push(up.nodes, nd) return nil }) } func (up *updatePath) marshal(w *Writer) { up.leafNode.marshal(w) w.writeVector(len(up.nodes), func(w *Writer, i int32) { up.nodes[i].marshal(w) }) } // --- NodeType --- type nodeType uint8 const ( nodeTypeLeaf nodeType = 1 nodeTypeParent nodeType = 2 ) func (t *nodeType) unmarshal(r *Reader) (err error) { b, ok := r.readByte() if !ok { return errUnexpectedEOF } *t = nodeType(b) switch *t { case nodeTypeLeaf, nodeTypeParent: return nil default: return errInvalidNodeType } } func (t *nodeType) marshal(w *Writer) { w.addByte(byte(t)) } // --- Node --- type node struct { nodeType nodeType leafNode *leafNode // for nodeTypeLeaf parentNode *parentNode // for nodeTypeParent } func (n *node) unmarshal(r *Reader) (err error) { *n = node{} if e := n.nodeType.unmarshal(r); e != nil { return e } switch n.nodeType { case nodeTypeLeaf: n.leafNode = &leafNode{} return n.leafNode.unmarshal(r) case nodeTypeParent: n.parentNode = &parentNode{} return n.parentNode.unmarshal(r) default: panic("unreachable") } } func (n *node) marshal(w *Writer) { n.nodeType.marshal(w) switch n.nodeType { case nodeTypeLeaf: n.leafNode.marshal(w) case nodeTypeParent: n.parentNode.marshal(w) default: panic("unreachable") } } func (n *node) encryptionKey() (v []byte) { switch n.nodeType { case nodeTypeLeaf: return n.leafNode.encryptionKey case nodeTypeParent: return n.parentNode.encryptionKey default: panic("unreachable") } } // --- RatchetTree --- // // A ratchet tree is a []*node. Moxie forbids named slice types, so the old // methods are free functions whose first parameter is the tree. func ratchetTreeUnmarshal(tree *[]*node, r *Reader) (err error) { *tree = []*node{} err = r.readVector(func(r *Reader) error { present, ok := r.readOptional() if !ok { return errUnexpectedEOF } if present { n := &node{} if e := n.unmarshal(r); e != nil { return e } *tree = push(*tree, n) } else { *tree = push(*tree, nil) } return nil }) if err != nil { return err } // Pad to next power of 2 (width + 1 must be power of 2) for !isPowerOf2(uint32(len(*tree) + 1)) { *tree = push(*tree, nil) } return nil } func ratchetTreeMarshal(tree []*node, w *Writer) { end := len(tree) for end > 0 && tree[end-1] == nil { end-- } w.writeVector(len(tree[:end]), func(w *Writer, i int32) { n := tree[i] w.writeOptional(n != nil) if n != nil { n.marshal(w) } }) } // ratchetTreeMarshalRaw serializes a ratchet tree to bare TLS bytes; a slice // cannot satisfy the marshaler interface, so the raw envelope is explicit. func ratchetTreeMarshalRaw(tree []*node) (out []byte, err error) { var w Writer ratchetTreeMarshal(tree, &w) return w.bytes() } // ratchetTreeUnmarshalRaw reads a ratchet tree from bare TLS bytes. func ratchetTreeUnmarshalRaw(raw []byte, tree *[]*node) (err error) { r := newReader(raw) if e := ratchetTreeUnmarshal(tree, &r); e != nil { return e } if !r.empty() { return errExcessBytes } return nil } func ratchetTreeNumLeaves(tree []*node) (v numLeaves) { return numLeavesFromWidth(uint32(len(tree))) } func ratchetTreeGet(tree []*node, i nodeIndex) (p *node) { return tree[int32(i)] } func ratchetTreeSet(tree []*node, i nodeIndex, nd *node) { tree[int32(i)] = nd } func ratchetTreeGetLeaf(tree []*node, li leafIndex) (p *leafNode) { nd := ratchetTreeGet(tree, li.nodeIndex()) if nd == nil { return nil } return nd.leafNode } func ratchetTreeResolve(tree []*node, x nodeIndex) (ss []nodeIndex) { n := ratchetTreeGet(tree, x) if n == nil { l, r, ok := x.children() if !ok { return nil } return ratchetTreeResolve(tree, l) | ratchetTreeResolve(tree, r) } res := []nodeIndex{x} if n.nodeType == nodeTypeParent { for _, li := range n.parentNode.unmergedLeaves { res = push(res, li.nodeIndex()) } } return res } func ratchetTreeCopy(tree []*node) (v []*node) { newTree := []*node{:len(tree)} for i, nd := range tree { newTree[i] = nd } return newTree } func ratchetTreeAdd(tree *[]*node, ln *leafNode) { li := leafIndex(0) var ni nodeIndex found := false for { ni = li.nodeIndex() if int32(ni) >= len(*tree) { break } if ratchetTreeGet(*tree, ni) == nil { found = true break } li++ } if !found { newLen := ((len(*tree) + 1) * 2) - 1 for len(*tree) < newLen { *tree = push(*tree, nil) } } n := ratchetTreeNumLeaves(*tree) p := ni var ok bool var nd *node for { p, ok = n.parent(p) if !ok { break } nd = ratchetTreeGet(*tree, p) if nd != nil { nd.parentNode.unmergedLeaves = push(nd.parentNode.unmergedLeaves, li) } } ratchetTreeSet(*tree, ni, &node{ nodeType: nodeTypeLeaf, leafNode: ln, }) } func ratchetTreeUpdate(tree []*node, li leafIndex, ln *leafNode) { ni := li.nodeIndex() ratchetTreeSet(tree, ni, &node{ nodeType: nodeTypeLeaf, leafNode: ln, }) n := ratchetTreeNumLeaves(tree) for { var ok bool ni, ok = n.parent(ni) if !ok { break } ratchetTreeSet(tree, ni, nil) } } func ratchetTreeRemove(tree *[]*node, li leafIndex) { ni := li.nodeIndex() n := ratchetTreeNumLeaves(*tree) var ok bool for { ratchetTreeSet(*tree, ni, nil) ni, ok = n.parent(ni) if !ok { break } } li = leafIndex(n - 1) lastPowerOf2 := len(*tree) + 1 for { ni = li.nodeIndex() if ratchetTreeGet(*tree, ni) != nil { break } if isPowerOf2(uint32(ni)) { lastPowerOf2 = int32(ni) } if li == 0 { *tree = nil return } li-- } if lastPowerOf2 < len(*tree)+1 { *tree = (*tree)[:lastPowerOf2-1] } } func ratchetTreeApply(tree *[]*node, proposals []proposal, senders []leafIndex) { for i, prop := range proposals { if prop.proposalType == proposalTypeUpdate { ratchetTreeUpdate(*tree, senders[i], &prop.update.leafNode) } } for _, prop := range proposals { if prop.proposalType == proposalTypeRemove { ratchetTreeRemove(tree, prop.remove.removed) } } for _, prop := range proposals { if prop.proposalType == proposalTypeAdd { ratchetTreeAdd(tree, &prop.add.keyPackage.leafNode) } } } func ratchetTreeFindLeaf(tree []*node, ln *leafNode) (v leafIndex, ok bool) { for li := leafIndex(0); li < leafIndex(ratchetTreeNumLeaves(tree)); li++ { nd := ratchetTreeGetLeaf(tree, li) if nd == nil { continue } if !bytesEqual(nd.encryptionKey, ln.encryptionKey) { continue } raw1, err1 := marshalRaw(ln) raw2, err2 := marshalRaw(nd) return li, err1 == nil && err2 == nil && bytesEqual(raw1, raw2) } return 0, false } func ratchetTreeKeys(tree []*node) (sigKeys, encKeys map[string]bool) { sigKeys = map[string]bool{} encKeys = map[string]bool{} for li := leafIndex(0); li < leafIndex(ratchetTreeNumLeaves(tree)); li++ { nd := ratchetTreeGetLeaf(tree, li) if nd == nil { continue } sigKeys[string(nd.signatureKey)] = true encKeys[string(nd.encryptionKey)] = true } return sigKeys, encKeys } func ratchetTreeSupportedCreds(tree []*node) (m map[credentialType]bool) { numMembers := int32(0) counts := map[credentialType]int32{} for li := leafIndex(0); li < leafIndex(ratchetTreeNumLeaves(tree)); li++ { nd := ratchetTreeGetLeaf(tree, li) if nd == nil { continue } numMembers++ for _, ct := range nd.capabilities.credentials { counts[ct]++ } } result := map[credentialType]bool{} for ct, c := range counts { if c == numMembers { result[ct] = true } } return result } func ratchetTreeFilteredDirectPath(tree []*node, x nodeIndex) (ss []nodeIndex) { n := ratchetTreeNumLeaves(tree) var path []nodeIndex for { p, ok := n.parent(x) if !ok { break } s, ok := n.sibling(x) if !ok { panic("unreachable") } if len(ratchetTreeResolve(tree, s)) > 0 { path = push(path, p) } x = p } return path } func hasUnmergedLeaf(pn *parentNode, target leafIndex) (ok bool) { for _, li := range pn.unmergedLeaves { if li == target { return true } } return false } func ratchetTreeFindParentHash(tree []*node, nodeIndices []nodeIndex, ph []byte) (ok bool) { for _, x := range nodeIndices { nd := ratchetTreeGet(tree, x) if nd == nil { continue } var h []byte switch nd.nodeType { case nodeTypeLeaf: h = nd.leafNode.parentHash case nodeTypeParent: h = nd.parentNode.parentHash } if bytesEqual(h, ph) { return true } } return false } // --- LeafNode verification --- type leafNodeVerifyOptions struct { cipherSuite CipherSuite groupID []byte leafIndex leafIndex supportedCreds map[credentialType]bool signatureKeys map[string]bool encryptionKeys map[string]bool nowUnix int64 // 0 = skip lifetime check } func (ln *leafNode) verify(opts *leafNodeVerifyOptions) (err error) { if !ln.verifySignature(opts.cipherSuite, opts.groupID, opts.leafIndex) { return errors.New("mls: leaf node signature verification failed") } if !opts.supportedCreds[ln.credential.credentialType] { return errors.New("mls: credential type not supported by all members") } if ln.lifetime != nil && opts.nowUnix != 0 { if !ln.lifetime.verifyAt(opts.nowUnix) { return errors.New("mls: lifetime verification failed") } } supportedExts := map[extensionType]bool{} for _, et := range ln.capabilities.extensions { supportedExts[et] = true } for _, ext := range ln.extensions { if !supportedExts[ext.extensionType] { return errors.New("mls: extension type not supported by leaf node") } } if opts.signatureKeys[string(ln.signatureKey)] { return errors.New("mls: duplicate signature key") } if opts.encryptionKeys[string(ln.encryptionKey)] { return errors.New("mls: duplicate encryption key") } return nil } const maxLeafNodeLifetime = 90 * 24 * 3600 // 90 days in seconds func (lt *lifetime) verifyAt(nowUnix int64) (ok bool) { notBefore := int64(lt.notBefore) notAfter := int64(lt.notAfter) duration := notAfter - notBefore if duration <= 0 || duration > maxLeafNodeLifetime { return false } return nowUnix > notBefore && notAfter > nowUnix }