tree_crypto.mx raw
1 package mls
2
3 // MLS tree crypto operations (RFC 9420 §7).
4
5 import "errors"
6
7 var (
8 errNodeKeyMismatch error
9 )
10
11 func (nd *parentNode) computeParentHash(cs CipherSuite, originalSiblingTreeHash []byte) (out []byte, err2 error) {
12 raw, err := marshalParentHashInput(nd.encryptionKey, nd.parentHash, originalSiblingTreeHash)
13 if err != nil {
14 return nil, err
15 }
16 return cs.hash(raw), nil
17 }
18
19 func marshalParentHashInput(encryptionKey []byte, parentHash, originalSiblingTreeHash []byte) (out []byte, err error) {
20 var w Writer
21 w.writeOpaqueVec([]byte(encryptionKey))
22 w.writeOpaqueVec(parentHash)
23 w.writeOpaqueVec(originalSiblingTreeHash)
24 return w.bytes()
25 }
26
27 func (nd *leafNode) sign(cs CipherSuite, groupID []byte, li leafIndex, signerPriv []byte) (err error) {
28 tbs, err := marshalRaw(&leafNodeTBS{
29 node: nd,
30 groupID: groupID,
31 leafIndex: li,
32 })
33 if err != nil {
34 return err
35 }
36 sig, err := cs.signWithLabel(signerPriv, []byte("LeafNodeTBS"), tbs)
37 if err != nil {
38 return err
39 }
40 nd.signature = sig
41 return nil
42 }
43
44 func (nd *leafNode) verifySignature(cs CipherSuite, groupID []byte, li leafIndex) (ok bool) {
45 tbs, err := marshalRaw(&leafNodeTBS{
46 node: nd,
47 groupID: groupID,
48 leafIndex: li,
49 })
50 if err != nil {
51 return false
52 }
53 return cs.verifyWithLabel(nd.signatureKey, []byte("LeafNodeTBS"), tbs, nd.signature)
54 }
55
56 func decryptPathSecret(cs CipherSuite, nodePriv []byte, ctx *groupContext, ct hpkeCiphertext) (out []byte, err2 error) {
57 rawCtx, err := marshalRaw(ctx)
58 if err != nil {
59 return nil, err
60 }
61 return cs.decryptWithLabel(nodePriv, []byte("UpdatePathNode"), rawCtx, ct.kemOutput, ct.ciphertext)
62 }
63
64 func nodePrivFromPathSecret(cs CipherSuite, pathSecret []byte, nodePub []byte) (v []byte, err2 error) {
65 nodeSecret, err := cs.deriveSecret(pathSecret, []byte("node"))
66 if err != nil {
67 return nil, err
68 }
69 pub, priv, err := cs.deriveEncryptionKeyPair(nodeSecret)
70 if err != nil {
71 return nil, err
72 }
73 if !bytesEqual(pub, nodePub) {
74 return nil, errNodeKeyMismatch
75 }
76 return priv, nil
77 }
78
79 // computeRootTreeHash computes the tree hash for integrity verification.
80 func ratchetTreeComputeRootTreeHash(tree []*node, cs CipherSuite) (out []byte, err error) {
81 n := ratchetTreeNumLeaves(tree)
82 return ratchetTreeComputeTreeHash(tree, cs, n.root(), n, nil)
83 }
84
85 func ratchetTreeComputeTreeHash(tree []*node, cs CipherSuite, x nodeIndex, n numLeaves, exclude map[leafIndex]bool) (out []byte, err error) {
86 if x.isLeaf() {
87 return ratchetTreeComputeLeafTreeHash(tree, cs, x, exclude)
88 }
89 return ratchetTreeComputeParentTreeHash(tree, cs, x, n, exclude)
90 }
91
92 func ratchetTreeComputeLeafTreeHash(tree []*node, cs CipherSuite, x nodeIndex, exclude map[leafIndex]bool) (out []byte, err2 error) {
93 var w Writer
94 li, _ := x.leafIndex()
95 w.addUint32(uint32(li))
96 nd := ratchetTreeGet(tree, x)
97 if exclude[li] {
98 nd = nil
99 }
100 w.writeOptional(nd != nil)
101 if nd != nil {
102 nd.leafNode.marshal(&w)
103 }
104 input, err := w.bytes()
105 if err != nil {
106 return nil, err
107 }
108 return cs.hash(input), nil
109 }
110
111 func ratchetTreeComputeParentTreeHash(tree []*node, cs CipherSuite, x nodeIndex, n numLeaves, exclude map[leafIndex]bool) (out []byte, err2 error) {
112 l, r, ok := x.children()
113 if !ok {
114 return nil, errUnexpectedEOF
115 }
116 leftHash, err := ratchetTreeComputeTreeHash(tree, cs, l, n, exclude)
117 if err != nil {
118 return nil, err
119 }
120 rightHash, err := ratchetTreeComputeTreeHash(tree, cs, r, n, exclude)
121 if err != nil {
122 return nil, err
123 }
124
125 var w Writer
126 nd := ratchetTreeGet(tree, x)
127 if nd != nil && len(exclude) > 0 {
128 // Filter unmerged leaves for parent hash verification
129 pn := nd.parentNode
130 if pn != nil && len(pn.unmergedLeaves) > 0 {
131 filtered := []leafIndex{:0:len(pn.unmergedLeaves)}
132 for _, li := range pn.unmergedLeaves {
133 if !exclude[li] {
134 filtered = push(filtered, li)
135 }
136 }
137 filteredNode := *pn
138 filteredNode.unmergedLeaves = filtered
139 w.writeOptional(true)
140 filteredNode.marshal(&w)
141 w.writeOpaqueVec(leftHash)
142 w.writeOpaqueVec(rightHash)
143 in, e := w.bytes()
144 if e != nil {
145 return nil, e
146 }
147 return cs.hash(in), nil
148 }
149 }
150 w.writeOptional(nd != nil)
151 if nd != nil {
152 nd.parentNode.marshal(&w)
153 }
154 w.writeOpaqueVec(leftHash)
155 w.writeOpaqueVec(rightHash)
156 input, err := w.bytes()
157 if err != nil {
158 return nil, err
159 }
160 return cs.hash(input), nil
161 }
162
163 // verifyParentHashes verifies the parent hash chain (RFC 9420 §7.9.2).
164 func ratchetTreeVerifyParentHashes(tree []*node, cs CipherSuite) (ok bool) {
165 n := ratchetTreeNumLeaves(tree)
166 for i, nd := range tree {
167 if nd == nil {
168 continue
169 }
170 x := nodeIndex(i)
171 l, r, hasChildren := x.children()
172 if !hasChildren {
173 continue
174 }
175
176 pn := nd.parentNode
177 exclude := map[leafIndex]bool{}
178 for _, li := range pn.unmergedLeaves {
179 exclude[li] = true
180 }
181
182 leftTreeHash, err := ratchetTreeComputeTreeHash(tree, cs, l, n, exclude)
183 if err != nil {
184 return false
185 }
186 rightTreeHash, err := ratchetTreeComputeTreeHash(tree, cs, r, n, exclude)
187 if err != nil {
188 return false
189 }
190
191 leftParentHash, err := pn.computeParentHash(cs, rightTreeHash)
192 if err != nil {
193 return false
194 }
195 rightParentHash, err := pn.computeParentHash(cs, leftTreeHash)
196 if err != nil {
197 return false
198 }
199
200 isLeft := ratchetTreeFindParentHash(tree, ratchetTreeResolve(tree, l), leftParentHash)
201 isRight := ratchetTreeFindParentHash(tree, ratchetTreeResolve(tree, r), rightParentHash)
202 if isLeft == isRight {
203 return false
204 }
205 }
206 return true
207 }
208
209 // verifyIntegrity verifies the ratchet tree (RFC 9420 §12.4.3.1).
210 func ratchetTreeVerifyIntegrity(tree []*node, ctx *groupContext, nowUnix int64) (err error) {
211 cs := ctx.cipherSuite
212 n := ratchetTreeNumLeaves(tree)
213
214 h, err := ratchetTreeComputeRootTreeHash(tree, cs)
215 if err != nil {
216 return err
217 }
218 if !bytesEqual(h, ctx.treeHash) {
219 return errors.New("mls: tree hash verification failed")
220 }
221 if !ratchetTreeVerifyParentHashes(tree, cs) {
222 return errors.New("mls: parent hash verification failed")
223 }
224
225 supportedCreds := ratchetTreeSupportedCreds(tree)
226 sigKeys := map[string]bool{}
227 encKeys := map[string]bool{}
228 for li := leafIndex(0); li < leafIndex(n); li++ {
229 nd := ratchetTreeGetLeaf(tree, li)
230 if nd == nil {
231 continue
232 }
233 e := nd.verify(&leafNodeVerifyOptions{
234 cipherSuite: cs,
235 groupID: ctx.groupID,
236 leafIndex: li,
237 supportedCreds: supportedCreds,
238 signatureKeys: sigKeys,
239 encryptionKeys: encKeys,
240 nowUnix: nowUnix,
241 })
242 if e != nil {
243 return e
244 }
245 sigKeys[string(nd.signatureKey)] = true
246 encKeys[string(nd.encryptionKey)] = true
247 }
248
249 // Check unmerged leaf ancestry
250 for i, nd := range tree {
251 if nd == nil || nd.nodeType != nodeTypeParent {
252 continue
253 }
254 p := nodeIndex(i)
255 for _, ul := range nd.parentNode.unmergedLeaves {
256 x := ul.nodeIndex()
257 for {
258 var ok bool
259 x, ok = n.parent(x)
260 if !ok {
261 return errors.New("mls: unmerged leaf not descendant of parent")
262 }
263 if x == p {
264 break
265 }
266 intermediate := ratchetTreeGet(tree, x)
267 if intermediate != nil && !hasUnmergedLeaf(intermediate.parentNode, ul) {
268 return errors.New("mls: intermediate node missing unmerged leaf")
269 }
270 }
271 }
272 if encKeys[string(nd.parentNode.encryptionKey)] {
273 return errors.New("mls: duplicate encryption key in tree")
274 }
275 encKeys[string(nd.parentNode.encryptionKey)] = true
276 }
277 return nil
278 }
279
280 // mergeUpdatePath applies an update path to the tree (RFC 9420 §7.5).
281 func ratchetTreeMergeUpdatePath(tree []*node, cs CipherSuite, senderLI leafIndex, path *updatePath) (err error) {
282 senderNI := senderLI.nodeIndex()
283 n := ratchetTreeNumLeaves(tree)
284
285 directPath := n.directPath(senderNI)
286 for _, ni := range directPath {
287 ratchetTreeSet(tree, ni, nil)
288 }
289
290 filteredDP := ratchetTreeFilteredDirectPath(tree, senderNI)
291 if len(filteredDP) != len(path.nodes) {
292 return errors.New("mls: update path length mismatch")
293 }
294 for i, ni := range filteredDP {
295 ratchetTreeSet(tree, ni, &node{
296 nodeType: nodeTypeParent,
297 parentNode: &parentNode{
298 encryptionKey: path.nodes[i].encryptionKey,
299 },
300 })
301 }
302
303 // Compute parent hashes root-to-leaf
304 var prevParentHash []byte
305 for i := len(filteredDP) - 1; i >= 0; i-- {
306 ni := filteredDP[i]
307 pn := ratchetTreeGet(tree, ni).parentNode
308
309 l, r, ok := ni.children()
310 if !ok {
311 panic("unreachable")
312 }
313 s := l
314 found := false
315 for _, dp := range directPath {
316 if dp == s {
317 found = true
318 break
319 }
320 }
321 if s == senderNI || found {
322 s = r
323 }
324
325 treeHash, e := ratchetTreeComputeTreeHash(tree, cs, s, n, nil)
326 if e != nil {
327 return e
328 }
329
330 pn.parentHash = prevParentHash
331 h, e := pn.computeParentHash(cs, treeHash)
332 if e != nil {
333 return e
334 }
335 prevParentHash = h
336 }
337
338 if !bytesEqual(path.leafNode.parentHash, prevParentHash) {
339 return errors.New("mls: parent hash mismatch for update path leaf node")
340 }
341
342 ratchetTreeSet(tree, senderNI, &node{
343 nodeType: nodeTypeLeaf,
344 leafNode: &path.leafNode,
345 })
346 return nil
347 }
348
349 // decryptPathSecrets decrypts path secrets from an update path (RFC 9420 §7.6).
350 func ratchetTreeDecryptPathSecrets(tree []*node, cs CipherSuite, ctx *groupContext, senderLI, recipientLI leafIndex, path *updatePath, privTree [][]byte) (out []byte, err2 error) {
351 senderNI := senderLI.nodeIndex()
352 recipientNI := recipientLI.nodeIndex()
353
354 senderFDP := ratchetTreeFilteredDirectPath(tree, senderNI)
355 if len(path.nodes) != len(senderFDP) {
356 return nil, errors.New("mls: invalid update path length")
357 }
358
359 // Find the common ancestor in the filtered direct path
360 recipientAncestor := commonAncestor(senderNI, recipientNI)
361 recipientAncestorIdx := -1
362 for i, ni := range senderFDP {
363 if ni == recipientAncestor {
364 recipientAncestorIdx = i
365 break
366 }
367 }
368 if recipientAncestorIdx < 0 {
369 return nil, errors.New("mls: cannot find recipient ancestor")
370 }
371 upNode := path.nodes[recipientAncestorIdx]
372
373 // Find the copath node
374 ancestor := commonAncestor(senderNI, recipientNI)
375 var copathNode nodeIndex
376 var ok bool
377 if recipientNI < senderNI {
378 copathNode, ok = ancestor.left()
379 } else {
380 copathNode, ok = ancestor.right()
381 }
382 if !ok {
383 panic("unreachable")
384 }
385
386 copathRes := ratchetTreeResolve(tree, copathNode)
387 if len(upNode.encryptedPathSecret) != len(copathRes) {
388 return nil, errors.New("mls: invalid encrypted path secret length")
389 }
390
391 // Find a node in the resolution for which we have a private key
392 var nodePriv []byte
393 resIdx := -1
394 for i, ni := range copathRes {
395 if p := privTree[int32(ni)]; p != nil {
396 nodePriv = p
397 resIdx = i
398 break
399 }
400 }
401 if nodePriv == nil {
402 return nil, errors.New("mls: no private key found")
403 }
404
405 pathSecret, err := decryptPathSecret(cs, nodePriv, ctx, upNode.encryptedPathSecret[resIdx])
406 if err != nil {
407 return nil, err
408 }
409 nodePub := ratchetTreeGet(tree, recipientAncestor).encryptionKey()
410 nodePriv, err = nodePrivFromPathSecret(cs, pathSecret, nodePub)
411 if err != nil {
412 return nil, err
413 }
414 privTree[int32(recipientAncestor)] = nodePriv
415
416 // Derive path secrets for remaining ancestors
417 for _, ni := range senderFDP[recipientAncestorIdx+1:] {
418 pathSecret, err = cs.deriveSecret(pathSecret, []byte("path"))
419 if err != nil {
420 return nil, err
421 }
422 nodePriv, err = nodePrivFromPathSecret(cs, pathSecret, ratchetTreeGet(tree, ni).encryptionKey())
423 if err != nil {
424 return nil, err
425 }
426 privTree[int32(ni)] = nodePriv
427 }
428
429 commitSecret, err := cs.deriveSecret(pathSecret, []byte("path"))
430 if err != nil {
431 return nil, err
432 }
433 return commitSecret, nil
434 }
435