package mls // MLS tree crypto operations (RFC 9420 §7). import "errors" var ( errNodeKeyMismatch error ) func (nd *parentNode) computeParentHash(cs CipherSuite, originalSiblingTreeHash []byte) (out []byte, err2 error) { raw, err := marshalParentHashInput(nd.encryptionKey, nd.parentHash, originalSiblingTreeHash) if err != nil { return nil, err } return cs.hash(raw), nil } func marshalParentHashInput(encryptionKey []byte, parentHash, originalSiblingTreeHash []byte) (out []byte, err error) { var w Writer w.writeOpaqueVec([]byte(encryptionKey)) w.writeOpaqueVec(parentHash) w.writeOpaqueVec(originalSiblingTreeHash) return w.bytes() } func (nd *leafNode) sign(cs CipherSuite, groupID []byte, li leafIndex, signerPriv []byte) (err error) { tbs, err := marshalRaw(&leafNodeTBS{ node: nd, groupID: groupID, leafIndex: li, }) if err != nil { return err } sig, err := cs.signWithLabel(signerPriv, []byte("LeafNodeTBS"), tbs) if err != nil { return err } nd.signature = sig return nil } func (nd *leafNode) verifySignature(cs CipherSuite, groupID []byte, li leafIndex) (ok bool) { tbs, err := marshalRaw(&leafNodeTBS{ node: nd, groupID: groupID, leafIndex: li, }) if err != nil { return false } return cs.verifyWithLabel(nd.signatureKey, []byte("LeafNodeTBS"), tbs, nd.signature) } func decryptPathSecret(cs CipherSuite, nodePriv []byte, ctx *groupContext, ct hpkeCiphertext) (out []byte, err2 error) { rawCtx, err := marshalRaw(ctx) if err != nil { return nil, err } return cs.decryptWithLabel(nodePriv, []byte("UpdatePathNode"), rawCtx, ct.kemOutput, ct.ciphertext) } func nodePrivFromPathSecret(cs CipherSuite, pathSecret []byte, nodePub []byte) (v []byte, err2 error) { nodeSecret, err := cs.deriveSecret(pathSecret, []byte("node")) if err != nil { return nil, err } pub, priv, err := cs.deriveEncryptionKeyPair(nodeSecret) if err != nil { return nil, err } if !bytesEqual(pub, nodePub) { return nil, errNodeKeyMismatch } return priv, nil } // computeRootTreeHash computes the tree hash for integrity verification. func ratchetTreeComputeRootTreeHash(tree []*node, cs CipherSuite) (out []byte, err error) { n := ratchetTreeNumLeaves(tree) return ratchetTreeComputeTreeHash(tree, cs, n.root(), n, nil) } func ratchetTreeComputeTreeHash(tree []*node, cs CipherSuite, x nodeIndex, n numLeaves, exclude map[leafIndex]bool) (out []byte, err error) { if x.isLeaf() { return ratchetTreeComputeLeafTreeHash(tree, cs, x, exclude) } return ratchetTreeComputeParentTreeHash(tree, cs, x, n, exclude) } func ratchetTreeComputeLeafTreeHash(tree []*node, cs CipherSuite, x nodeIndex, exclude map[leafIndex]bool) (out []byte, err2 error) { var w Writer li, _ := x.leafIndex() w.addUint32(uint32(li)) nd := ratchetTreeGet(tree, x) if exclude[li] { nd = nil } w.writeOptional(nd != nil) if nd != nil { nd.leafNode.marshal(&w) } input, err := w.bytes() if err != nil { return nil, err } return cs.hash(input), nil } func ratchetTreeComputeParentTreeHash(tree []*node, cs CipherSuite, x nodeIndex, n numLeaves, exclude map[leafIndex]bool) (out []byte, err2 error) { l, r, ok := x.children() if !ok { return nil, errUnexpectedEOF } leftHash, err := ratchetTreeComputeTreeHash(tree, cs, l, n, exclude) if err != nil { return nil, err } rightHash, err := ratchetTreeComputeTreeHash(tree, cs, r, n, exclude) if err != nil { return nil, err } var w Writer nd := ratchetTreeGet(tree, x) if nd != nil && len(exclude) > 0 { // Filter unmerged leaves for parent hash verification pn := nd.parentNode if pn != nil && len(pn.unmergedLeaves) > 0 { filtered := []leafIndex{:0:len(pn.unmergedLeaves)} for _, li := range pn.unmergedLeaves { if !exclude[li] { filtered = push(filtered, li) } } filteredNode := *pn filteredNode.unmergedLeaves = filtered w.writeOptional(true) filteredNode.marshal(&w) w.writeOpaqueVec(leftHash) w.writeOpaqueVec(rightHash) in, e := w.bytes() if e != nil { return nil, e } return cs.hash(in), nil } } w.writeOptional(nd != nil) if nd != nil { nd.parentNode.marshal(&w) } w.writeOpaqueVec(leftHash) w.writeOpaqueVec(rightHash) input, err := w.bytes() if err != nil { return nil, err } return cs.hash(input), nil } // verifyParentHashes verifies the parent hash chain (RFC 9420 §7.9.2). func ratchetTreeVerifyParentHashes(tree []*node, cs CipherSuite) (ok bool) { n := ratchetTreeNumLeaves(tree) for i, nd := range tree { if nd == nil { continue } x := nodeIndex(i) l, r, hasChildren := x.children() if !hasChildren { continue } pn := nd.parentNode exclude := map[leafIndex]bool{} for _, li := range pn.unmergedLeaves { exclude[li] = true } leftTreeHash, err := ratchetTreeComputeTreeHash(tree, cs, l, n, exclude) if err != nil { return false } rightTreeHash, err := ratchetTreeComputeTreeHash(tree, cs, r, n, exclude) if err != nil { return false } leftParentHash, err := pn.computeParentHash(cs, rightTreeHash) if err != nil { return false } rightParentHash, err := pn.computeParentHash(cs, leftTreeHash) if err != nil { return false } isLeft := ratchetTreeFindParentHash(tree, ratchetTreeResolve(tree, l), leftParentHash) isRight := ratchetTreeFindParentHash(tree, ratchetTreeResolve(tree, r), rightParentHash) if isLeft == isRight { return false } } return true } // verifyIntegrity verifies the ratchet tree (RFC 9420 §12.4.3.1). func ratchetTreeVerifyIntegrity(tree []*node, ctx *groupContext, nowUnix int64) (err error) { cs := ctx.cipherSuite n := ratchetTreeNumLeaves(tree) h, err := ratchetTreeComputeRootTreeHash(tree, cs) if err != nil { return err } if !bytesEqual(h, ctx.treeHash) { return errors.New("mls: tree hash verification failed") } if !ratchetTreeVerifyParentHashes(tree, cs) { return errors.New("mls: parent hash verification failed") } supportedCreds := ratchetTreeSupportedCreds(tree) sigKeys := map[string]bool{} encKeys := map[string]bool{} for li := leafIndex(0); li < leafIndex(n); li++ { nd := ratchetTreeGetLeaf(tree, li) if nd == nil { continue } e := nd.verify(&leafNodeVerifyOptions{ cipherSuite: cs, groupID: ctx.groupID, leafIndex: li, supportedCreds: supportedCreds, signatureKeys: sigKeys, encryptionKeys: encKeys, nowUnix: nowUnix, }) if e != nil { return e } sigKeys[string(nd.signatureKey)] = true encKeys[string(nd.encryptionKey)] = true } // Check unmerged leaf ancestry for i, nd := range tree { if nd == nil || nd.nodeType != nodeTypeParent { continue } p := nodeIndex(i) for _, ul := range nd.parentNode.unmergedLeaves { x := ul.nodeIndex() for { var ok bool x, ok = n.parent(x) if !ok { return errors.New("mls: unmerged leaf not descendant of parent") } if x == p { break } intermediate := ratchetTreeGet(tree, x) if intermediate != nil && !hasUnmergedLeaf(intermediate.parentNode, ul) { return errors.New("mls: intermediate node missing unmerged leaf") } } } if encKeys[string(nd.parentNode.encryptionKey)] { return errors.New("mls: duplicate encryption key in tree") } encKeys[string(nd.parentNode.encryptionKey)] = true } return nil } // mergeUpdatePath applies an update path to the tree (RFC 9420 §7.5). func ratchetTreeMergeUpdatePath(tree []*node, cs CipherSuite, senderLI leafIndex, path *updatePath) (err error) { senderNI := senderLI.nodeIndex() n := ratchetTreeNumLeaves(tree) directPath := n.directPath(senderNI) for _, ni := range directPath { ratchetTreeSet(tree, ni, nil) } filteredDP := ratchetTreeFilteredDirectPath(tree, senderNI) if len(filteredDP) != len(path.nodes) { return errors.New("mls: update path length mismatch") } for i, ni := range filteredDP { ratchetTreeSet(tree, ni, &node{ nodeType: nodeTypeParent, parentNode: &parentNode{ encryptionKey: path.nodes[i].encryptionKey, }, }) } // Compute parent hashes root-to-leaf var prevParentHash []byte for i := len(filteredDP) - 1; i >= 0; i-- { ni := filteredDP[i] pn := ratchetTreeGet(tree, ni).parentNode l, r, ok := ni.children() if !ok { panic("unreachable") } s := l found := false for _, dp := range directPath { if dp == s { found = true break } } if s == senderNI || found { s = r } treeHash, e := ratchetTreeComputeTreeHash(tree, cs, s, n, nil) if e != nil { return e } pn.parentHash = prevParentHash h, e := pn.computeParentHash(cs, treeHash) if e != nil { return e } prevParentHash = h } if !bytesEqual(path.leafNode.parentHash, prevParentHash) { return errors.New("mls: parent hash mismatch for update path leaf node") } ratchetTreeSet(tree, senderNI, &node{ nodeType: nodeTypeLeaf, leafNode: &path.leafNode, }) return nil } // decryptPathSecrets decrypts path secrets from an update path (RFC 9420 §7.6). func ratchetTreeDecryptPathSecrets(tree []*node, cs CipherSuite, ctx *groupContext, senderLI, recipientLI leafIndex, path *updatePath, privTree [][]byte) (out []byte, err2 error) { senderNI := senderLI.nodeIndex() recipientNI := recipientLI.nodeIndex() senderFDP := ratchetTreeFilteredDirectPath(tree, senderNI) if len(path.nodes) != len(senderFDP) { return nil, errors.New("mls: invalid update path length") } // Find the common ancestor in the filtered direct path recipientAncestor := commonAncestor(senderNI, recipientNI) recipientAncestorIdx := -1 for i, ni := range senderFDP { if ni == recipientAncestor { recipientAncestorIdx = i break } } if recipientAncestorIdx < 0 { return nil, errors.New("mls: cannot find recipient ancestor") } upNode := path.nodes[recipientAncestorIdx] // Find the copath node ancestor := commonAncestor(senderNI, recipientNI) var copathNode nodeIndex var ok bool if recipientNI < senderNI { copathNode, ok = ancestor.left() } else { copathNode, ok = ancestor.right() } if !ok { panic("unreachable") } copathRes := ratchetTreeResolve(tree, copathNode) if len(upNode.encryptedPathSecret) != len(copathRes) { return nil, errors.New("mls: invalid encrypted path secret length") } // Find a node in the resolution for which we have a private key var nodePriv []byte resIdx := -1 for i, ni := range copathRes { if p := privTree[int32(ni)]; p != nil { nodePriv = p resIdx = i break } } if nodePriv == nil { return nil, errors.New("mls: no private key found") } pathSecret, err := decryptPathSecret(cs, nodePriv, ctx, upNode.encryptedPathSecret[resIdx]) if err != nil { return nil, err } nodePub := ratchetTreeGet(tree, recipientAncestor).encryptionKey() nodePriv, err = nodePrivFromPathSecret(cs, pathSecret, nodePub) if err != nil { return nil, err } privTree[int32(recipientAncestor)] = nodePriv // Derive path secrets for remaining ancestors for _, ni := range senderFDP[recipientAncestorIdx+1:] { pathSecret, err = cs.deriveSecret(pathSecret, []byte("path")) if err != nil { return nil, err } nodePriv, err = nodePrivFromPathSecret(cs, pathSecret, ratchetTreeGet(tree, ni).encryptionKey()) if err != nil { return nil, err } privTree[int32(ni)] = nodePriv } commitSecret, err := cs.deriveSecret(pathSecret, []byte("path")) if err != nil { return nil, err } return commitSecret, nil }