tree.mx raw

   1  package mls
   2  
   3  // MLS ratchet tree types (RFC 9420 ยง7).
   4  // Data types and serialization only - crypto operations (sign, verify,
   5  // parentHash) go in tree_crypto.mx with the cipher suite.
   6  
   7  import "errors"
   8  
   9  var (
  10  	errInvalidLeafNodeSource error
  11  	errInvalidNodeType       error
  12  )
  13  
  14  func bytesEqual(a, b []byte) (ok bool) {
  15  	if len(a) != len(b) {
  16  		return false
  17  	}
  18  	for i := range a {
  19  		if a[i] != b[i] {
  20  			return false
  21  		}
  22  	}
  23  	return true
  24  }
  25  
  26  // --- ParentNode ---
  27  
  28  type parentNode struct {
  29  	encryptionKey  []byte
  30  	parentHash     []byte
  31  	unmergedLeaves []leafIndex
  32  }
  33  
  34  func (nd *parentNode) unmarshal(r *Reader) (err error) {
  35  	*nd = parentNode{}
  36  	var ok bool
  37  	nd.encryptionKey, ok = r.readOpaqueVec()
  38  	if !ok {
  39  		return errUnexpectedEOF
  40  	}
  41  	nd.parentHash, ok = r.readOpaqueVec()
  42  	if !ok {
  43  		return errUnexpectedEOF
  44  	}
  45  	return r.readVector(func(r *Reader) error {
  46  		v, hasV := r.readUint32()
  47  		if !hasV {
  48  			return errUnexpectedEOF
  49  		}
  50  		nd.unmergedLeaves = push(nd.unmergedLeaves, leafIndex(v))
  51  		return nil
  52  	})
  53  }
  54  
  55  func (nd *parentNode) marshal(w *Writer) {
  56  	w.writeOpaqueVec([]byte(nd.encryptionKey))
  57  	w.writeOpaqueVec(nd.parentHash)
  58  	w.writeVector(len(nd.unmergedLeaves), func(w *Writer, i int32) {
  59  		w.addUint32(uint32(nd.unmergedLeaves[i]))
  60  	})
  61  }
  62  
  63  // --- LeafNodeSource ---
  64  
  65  type leafNodeSource uint8
  66  
  67  const (
  68  	leafNodeSourceKeyPackage leafNodeSource = 1
  69  	leafNodeSourceUpdate     leafNodeSource = 2
  70  	leafNodeSourceCommit     leafNodeSource = 3
  71  )
  72  
  73  func (src *leafNodeSource) unmarshal(r *Reader) (err error) {
  74  	b, ok := r.readByte()
  75  	if !ok {
  76  		return errUnexpectedEOF
  77  	}
  78  	*src = leafNodeSource(b)
  79  	switch *src {
  80  	case leafNodeSourceKeyPackage, leafNodeSourceUpdate, leafNodeSourceCommit:
  81  		return nil
  82  	default:
  83  		return errInvalidLeafNodeSource
  84  	}
  85  }
  86  
  87  func (src *leafNodeSource) marshal(w *Writer) {
  88  	w.addByte(byte(src))
  89  }
  90  
  91  // --- Capabilities ---
  92  
  93  type capabilities struct {
  94  	versions     []protocolVersion
  95  	cipherSuites []CipherSuite
  96  	extensions   []extensionType
  97  	proposals    []proposalType
  98  	credentials  []credentialType
  99  }
 100  
 101  func (caps *capabilities) unmarshal(r *Reader) (err error) {
 102  	*caps = capabilities{}
 103  
 104  	err := r.readVector(func(r *Reader) error {
 105  		v, ok := r.readUint16()
 106  		if !ok {
 107  			return errUnexpectedEOF
 108  		}
 109  		caps.versions = push(caps.versions, protocolVersion(v))
 110  		return nil
 111  	})
 112  	if err != nil {
 113  		return err
 114  	}
 115  
 116  	err = r.readVector(func(r *Reader) error {
 117  		v, ok := r.readUint16()
 118  		if !ok {
 119  			return errUnexpectedEOF
 120  		}
 121  		caps.cipherSuites = push(caps.cipherSuites, CipherSuite(v))
 122  		return nil
 123  	})
 124  	if err != nil {
 125  		return err
 126  	}
 127  
 128  	err = r.readVector(func(r *Reader) error {
 129  		v, ok := r.readUint16()
 130  		if !ok {
 131  			return errUnexpectedEOF
 132  		}
 133  		caps.extensions = push(caps.extensions, extensionType(v))
 134  		return nil
 135  	})
 136  	if err != nil {
 137  		return err
 138  	}
 139  
 140  	err = r.readVector(func(r *Reader) error {
 141  		v, ok := r.readUint16()
 142  		if !ok {
 143  			return errUnexpectedEOF
 144  		}
 145  		caps.proposals = push(caps.proposals, proposalType(v))
 146  		return nil
 147  	})
 148  	if err != nil {
 149  		return err
 150  	}
 151  
 152  	return r.readVector(func(r *Reader) error {
 153  		v, ok := r.readUint16()
 154  		if !ok {
 155  			return errUnexpectedEOF
 156  		}
 157  		caps.credentials = push(caps.credentials, credentialType(v))
 158  		return nil
 159  	})
 160  }
 161  
 162  func (caps *capabilities) marshal(w *Writer) {
 163  	w.writeVector(len(caps.versions), func(w *Writer, i int32) {
 164  		w.addUint16(uint16(caps.versions[i]))
 165  	})
 166  	w.writeVector(len(caps.cipherSuites), func(w *Writer, i int32) {
 167  		w.addUint16(uint16(caps.cipherSuites[i]))
 168  	})
 169  	w.writeVector(len(caps.extensions), func(w *Writer, i int32) {
 170  		w.addUint16(uint16(caps.extensions[i]))
 171  	})
 172  	w.writeVector(len(caps.proposals), func(w *Writer, i int32) {
 173  		w.addUint16(uint16(caps.proposals[i]))
 174  	})
 175  	w.writeVector(len(caps.credentials), func(w *Writer, i int32) {
 176  		w.addUint16(uint16(caps.credentials[i]))
 177  	})
 178  }
 179  
 180  // --- Lifetime ---
 181  
 182  type lifetime struct {
 183  	notBefore, notAfter uint64
 184  }
 185  
 186  func (lt *lifetime) unmarshal(r *Reader) (err error) {
 187  	*lt = lifetime{}
 188  	var ok bool
 189  	lt.notBefore, ok = r.readUint64()
 190  	if !ok {
 191  		return errUnexpectedEOF
 192  	}
 193  	lt.notAfter, ok = r.readUint64()
 194  	if !ok {
 195  		return errUnexpectedEOF
 196  	}
 197  	return nil
 198  }
 199  
 200  func (lt *lifetime) marshal(w *Writer) {
 201  	w.addUint64(lt.notBefore)
 202  	w.addUint64(lt.notAfter)
 203  }
 204  
 205  // --- Extension ---
 206  
 207  type extensionType uint16
 208  
 209  const (
 210  	extensionTypeApplicationID        extensionType = 0x0001
 211  	extensionTypeRatchetTree          extensionType = 0x0002
 212  	extensionTypeRequiredCapabilities extensionType = 0x0003
 213  	extensionTypeExternalPub          extensionType = 0x0004
 214  	extensionTypeExternalSenders      extensionType = 0x0005
 215  	// Marmot: KeyPackage reusable for multiple Welcomes (MIP-00)
 216  	ExtensionTypeLastResort extensionType = 0x000a
 217  	// Marmot: Nostr group metadata (MIP-01)
 218  	ExtensionTypeNostrGroupData extensionType = 0xf2ee
 219  )
 220  
 221  type Extension = extension
 222  
 223  type extension struct {
 224  	extensionType extensionType
 225  	extensionData []byte
 226  }
 227  
 228  func NewExtension(t extensionType, data []byte) (v extension) {
 229  	return extension{extensionType: t, extensionData: data}
 230  }
 231  
 232  type ExtensionType = extensionType
 233  
 234  func unmarshalExtensionVec(r *Reader) (list []extension, err2 error) {
 235  	var exts []extension
 236  	err := r.readVector(func(r *Reader) error {
 237  		var ext extension
 238  		v, ok := r.readUint16()
 239  		if !ok {
 240  			return errUnexpectedEOF
 241  		}
 242  		ext.extensionType = extensionType(v)
 243  		ext.extensionData, ok = r.readOpaqueVec()
 244  		if !ok {
 245  			return errUnexpectedEOF
 246  		}
 247  		exts = push(exts, ext)
 248  		return nil
 249  	})
 250  	return exts, err
 251  }
 252  
 253  func marshalExtensionVec(w *Writer, exts []extension) {
 254  	w.writeVector(len(exts), func(w *Writer, i int32) {
 255  		ext := exts[i]
 256  		w.addUint16(uint16(ext.extensionType))
 257  		w.writeOpaqueVec(ext.extensionData)
 258  	})
 259  }
 260  
 261  func findExtensionData(exts []extension, t extensionType) (buf []byte) {
 262  	for _, ext := range exts {
 263  		if ext.extensionType == t {
 264  			return ext.extensionData
 265  		}
 266  	}
 267  	return nil
 268  }
 269  
 270  // --- LeafNode ---
 271  
 272  type leafNode struct {
 273  	encryptionKey []byte
 274  	signatureKey  []byte
 275  	credential    Credential
 276  	capabilities  capabilities
 277  
 278  	leafNodeSource leafNodeSource
 279  	lifetime       *lifetime // for leafNodeSourceKeyPackage
 280  	parentHash     []byte    // for leafNodeSourceCommit
 281  
 282  	extensions []extension
 283  	signature  []byte
 284  }
 285  
 286  func (nd *leafNode) unmarshal(r *Reader) (err error) {
 287  	*nd = leafNode{}
 288  
 289  	var ok bool
 290  	nd.encryptionKey, ok = r.readOpaqueVec()
 291  	if !ok {
 292  		return errUnexpectedEOF
 293  	}
 294  	nd.signatureKey, ok = r.readOpaqueVec()
 295  	if !ok {
 296  		return errUnexpectedEOF
 297  	}
 298  
 299  	if e := nd.credential.unmarshal(r); e != nil {
 300  		return e
 301  	}
 302  	if e := nd.capabilities.unmarshal(r); e != nil {
 303  		return e
 304  	}
 305  	if e := nd.leafNodeSource.unmarshal(r); e != nil {
 306  		return e
 307  	}
 308  
 309  	var err error
 310  	switch nd.leafNodeSource {
 311  	case leafNodeSourceKeyPackage:
 312  		nd.lifetime = &lifetime{}
 313  		err = nd.lifetime.unmarshal(r)
 314  	case leafNodeSourceCommit:
 315  		nd.parentHash, ok = r.readOpaqueVec()
 316  		if !ok {
 317  			err = errUnexpectedEOF
 318  		}
 319  	}
 320  	if err != nil {
 321  		return err
 322  	}
 323  
 324  	exts, err := unmarshalExtensionVec(r)
 325  	if err != nil {
 326  		return err
 327  	}
 328  	nd.extensions = exts
 329  
 330  	nd.signature, ok = r.readOpaqueVec()
 331  	if !ok {
 332  		return errUnexpectedEOF
 333  	}
 334  	return nil
 335  }
 336  
 337  func (nd *leafNode) marshalBase(w *Writer) {
 338  	w.writeOpaqueVec([]byte(nd.encryptionKey))
 339  	w.writeOpaqueVec([]byte(nd.signatureKey))
 340  	nd.credential.marshal(w)
 341  	nd.capabilities.marshal(w)
 342  	nd.leafNodeSource.marshal(w)
 343  	switch nd.leafNodeSource {
 344  	case leafNodeSourceKeyPackage:
 345  		nd.lifetime.marshal(w)
 346  	case leafNodeSourceCommit:
 347  		w.writeOpaqueVec(nd.parentHash)
 348  	}
 349  	marshalExtensionVec(w, nd.extensions)
 350  }
 351  
 352  func (nd *leafNode) marshal(w *Writer) {
 353  	nd.marshalBase(w)
 354  	w.writeOpaqueVec(nd.signature)
 355  }
 356  
 357  // --- LeafNodeTBS ---
 358  
 359  type leafNodeTBS struct {
 360  	node *leafNode
 361  	// for leafNodeSourceUpdate and leafNodeSourceCommit
 362  	groupID   []byte
 363  	leafIndex leafIndex
 364  }
 365  
 366  func (tbs *leafNodeTBS) marshal(w *Writer) {
 367  	tbs.node.marshalBase(w)
 368  	switch tbs.node.leafNodeSource {
 369  	case leafNodeSourceUpdate, leafNodeSourceCommit:
 370  		w.writeOpaqueVec([]byte(tbs.groupID))
 371  		w.addUint32(uint32(tbs.leafIndex))
 372  	}
 373  }
 374  
 375  // --- UpdatePathNode ---
 376  
 377  type updatePathNode struct {
 378  	encryptionKey       []byte
 379  	encryptedPathSecret []hpkeCiphertext
 380  }
 381  
 382  func (upn *updatePathNode) unmarshal(r *Reader) (err error) {
 383  	*upn = updatePathNode{}
 384  	var ok bool
 385  	upn.encryptionKey, ok = r.readOpaqueVec()
 386  	if !ok {
 387  		return errUnexpectedEOF
 388  	}
 389  	return r.readVector(func(r *Reader) error {
 390  		var ct hpkeCiphertext
 391  		if e := ct.unmarshal(r); e != nil {
 392  			return e
 393  		}
 394  		upn.encryptedPathSecret = push(upn.encryptedPathSecret, ct)
 395  		return nil
 396  	})
 397  }
 398  
 399  func (upn *updatePathNode) marshal(w *Writer) {
 400  	w.writeOpaqueVec([]byte(upn.encryptionKey))
 401  	w.writeVector(len(upn.encryptedPathSecret), func(w *Writer, i int32) {
 402  		upn.encryptedPathSecret[i].marshal(w)
 403  	})
 404  }
 405  
 406  // --- UpdatePath ---
 407  
 408  type updatePath struct {
 409  	leafNode leafNode
 410  	nodes    []updatePathNode
 411  }
 412  
 413  func (up *updatePath) unmarshal(r *Reader) (err error) {
 414  	*up = updatePath{}
 415  	if e := up.leafNode.unmarshal(r); e != nil {
 416  		return e
 417  	}
 418  	return r.readVector(func(r *Reader) error {
 419  		var nd updatePathNode
 420  		if e := nd.unmarshal(r); e != nil {
 421  			return e
 422  		}
 423  		up.nodes = push(up.nodes, nd)
 424  		return nil
 425  	})
 426  }
 427  
 428  func (up *updatePath) marshal(w *Writer) {
 429  	up.leafNode.marshal(w)
 430  	w.writeVector(len(up.nodes), func(w *Writer, i int32) {
 431  		up.nodes[i].marshal(w)
 432  	})
 433  }
 434  
 435  // --- NodeType ---
 436  
 437  type nodeType uint8
 438  
 439  const (
 440  	nodeTypeLeaf   nodeType = 1
 441  	nodeTypeParent nodeType = 2
 442  )
 443  
 444  func (t *nodeType) unmarshal(r *Reader) (err error) {
 445  	b, ok := r.readByte()
 446  	if !ok {
 447  		return errUnexpectedEOF
 448  	}
 449  	*t = nodeType(b)
 450  	switch *t {
 451  	case nodeTypeLeaf, nodeTypeParent:
 452  		return nil
 453  	default:
 454  		return errInvalidNodeType
 455  	}
 456  }
 457  
 458  func (t *nodeType) marshal(w *Writer) {
 459  	w.addByte(byte(t))
 460  }
 461  
 462  // --- Node ---
 463  
 464  type node struct {
 465  	nodeType   nodeType
 466  	leafNode   *leafNode   // for nodeTypeLeaf
 467  	parentNode *parentNode // for nodeTypeParent
 468  }
 469  
 470  func (n *node) unmarshal(r *Reader) (err error) {
 471  	*n = node{}
 472  	if e := n.nodeType.unmarshal(r); e != nil {
 473  		return e
 474  	}
 475  	switch n.nodeType {
 476  	case nodeTypeLeaf:
 477  		n.leafNode = &leafNode{}
 478  		return n.leafNode.unmarshal(r)
 479  	case nodeTypeParent:
 480  		n.parentNode = &parentNode{}
 481  		return n.parentNode.unmarshal(r)
 482  	default:
 483  		panic("unreachable")
 484  	}
 485  }
 486  
 487  func (n *node) marshal(w *Writer) {
 488  	n.nodeType.marshal(w)
 489  	switch n.nodeType {
 490  	case nodeTypeLeaf:
 491  		n.leafNode.marshal(w)
 492  	case nodeTypeParent:
 493  		n.parentNode.marshal(w)
 494  	default:
 495  		panic("unreachable")
 496  	}
 497  }
 498  
 499  func (n *node) encryptionKey() (v []byte) {
 500  	switch n.nodeType {
 501  	case nodeTypeLeaf:
 502  		return n.leafNode.encryptionKey
 503  	case nodeTypeParent:
 504  		return n.parentNode.encryptionKey
 505  	default:
 506  		panic("unreachable")
 507  	}
 508  }
 509  
 510  // --- RatchetTree ---
 511  //
 512  // A ratchet tree is a []*node. Moxie forbids named slice types, so the old
 513  // methods are free functions whose first parameter is the tree.
 514  
 515  func ratchetTreeUnmarshal(tree *[]*node, r *Reader) (err error) {
 516  	*tree = []*node{}
 517  	err = r.readVector(func(r *Reader) error {
 518  		present, ok := r.readOptional()
 519  		if !ok {
 520  			return errUnexpectedEOF
 521  		}
 522  		if present {
 523  			n := &node{}
 524  			if e := n.unmarshal(r); e != nil {
 525  				return e
 526  			}
 527  			*tree = push(*tree, n)
 528  		} else {
 529  			*tree = push(*tree, nil)
 530  		}
 531  		return nil
 532  	})
 533  	if err != nil {
 534  		return err
 535  	}
 536  	// Pad to next power of 2 (width + 1 must be power of 2)
 537  	for !isPowerOf2(uint32(len(*tree) + 1)) {
 538  		*tree = push(*tree, nil)
 539  	}
 540  	return nil
 541  }
 542  
 543  func ratchetTreeMarshal(tree []*node, w *Writer) {
 544  	end := len(tree)
 545  	for end > 0 && tree[end-1] == nil {
 546  		end--
 547  	}
 548  	w.writeVector(len(tree[:end]), func(w *Writer, i int32) {
 549  		n := tree[i]
 550  		w.writeOptional(n != nil)
 551  		if n != nil {
 552  			n.marshal(w)
 553  		}
 554  	})
 555  }
 556  
 557  // ratchetTreeMarshalRaw serializes a ratchet tree to bare TLS bytes; a slice
 558  // cannot satisfy the marshaler interface, so the raw envelope is explicit.
 559  func ratchetTreeMarshalRaw(tree []*node) (out []byte, err error) {
 560  	var w Writer
 561  	ratchetTreeMarshal(tree, &w)
 562  	return w.bytes()
 563  }
 564  
 565  // ratchetTreeUnmarshalRaw reads a ratchet tree from bare TLS bytes.
 566  func ratchetTreeUnmarshalRaw(raw []byte, tree *[]*node) (err error) {
 567  	r := newReader(raw)
 568  	if e := ratchetTreeUnmarshal(tree, &r); e != nil {
 569  		return e
 570  	}
 571  	if !r.empty() {
 572  		return errExcessBytes
 573  	}
 574  	return nil
 575  }
 576  
 577  func ratchetTreeNumLeaves(tree []*node) (v numLeaves) {
 578  	return numLeavesFromWidth(uint32(len(tree)))
 579  }
 580  
 581  func ratchetTreeGet(tree []*node, i nodeIndex) (p *node) {
 582  	return tree[int32(i)]
 583  }
 584  
 585  func ratchetTreeSet(tree []*node, i nodeIndex, nd *node) {
 586  	tree[int32(i)] = nd
 587  }
 588  
 589  func ratchetTreeGetLeaf(tree []*node, li leafIndex) (p *leafNode) {
 590  	nd := ratchetTreeGet(tree, li.nodeIndex())
 591  	if nd == nil {
 592  		return nil
 593  	}
 594  	return nd.leafNode
 595  }
 596  
 597  func ratchetTreeResolve(tree []*node, x nodeIndex) (ss []nodeIndex) {
 598  	n := ratchetTreeGet(tree, x)
 599  	if n == nil {
 600  		l, r, ok := x.children()
 601  		if !ok {
 602  			return nil
 603  		}
 604  		return ratchetTreeResolve(tree, l) | ratchetTreeResolve(tree, r)
 605  	}
 606  	res := []nodeIndex{x}
 607  	if n.nodeType == nodeTypeParent {
 608  		for _, li := range n.parentNode.unmergedLeaves {
 609  			res = push(res, li.nodeIndex())
 610  		}
 611  	}
 612  	return res
 613  }
 614  
 615  func ratchetTreeCopy(tree []*node) (v []*node) {
 616  	newTree := []*node{:len(tree)}
 617  	for i, nd := range tree {
 618  		newTree[i] = nd
 619  	}
 620  	return newTree
 621  }
 622  
 623  func ratchetTreeAdd(tree *[]*node, ln *leafNode) {
 624  	li := leafIndex(0)
 625  	var ni nodeIndex
 626  	found := false
 627  	for {
 628  		ni = li.nodeIndex()
 629  		if int32(ni) >= len(*tree) {
 630  			break
 631  		}
 632  		if ratchetTreeGet(*tree, ni) == nil {
 633  			found = true
 634  			break
 635  		}
 636  		li++
 637  	}
 638  	if !found {
 639  		newLen := ((len(*tree) + 1) * 2) - 1
 640  		for len(*tree) < newLen {
 641  			*tree = push(*tree, nil)
 642  		}
 643  	}
 644  
 645  	n := ratchetTreeNumLeaves(*tree)
 646  	p := ni
 647  	var ok bool
 648  	var nd *node
 649  	for {
 650  		p, ok = n.parent(p)
 651  		if !ok {
 652  			break
 653  		}
 654  		nd = ratchetTreeGet(*tree, p)
 655  		if nd != nil {
 656  			nd.parentNode.unmergedLeaves = push(nd.parentNode.unmergedLeaves, li)
 657  		}
 658  	}
 659  
 660  	ratchetTreeSet(*tree, ni, &node{
 661  		nodeType: nodeTypeLeaf,
 662  		leafNode: ln,
 663  	})
 664  }
 665  
 666  func ratchetTreeUpdate(tree []*node, li leafIndex, ln *leafNode) {
 667  	ni := li.nodeIndex()
 668  	ratchetTreeSet(tree, ni, &node{
 669  		nodeType: nodeTypeLeaf,
 670  		leafNode: ln,
 671  	})
 672  	n := ratchetTreeNumLeaves(tree)
 673  	for {
 674  		var ok bool
 675  		ni, ok = n.parent(ni)
 676  		if !ok {
 677  			break
 678  		}
 679  		ratchetTreeSet(tree, ni, nil)
 680  	}
 681  }
 682  
 683  func ratchetTreeRemove(tree *[]*node, li leafIndex) {
 684  	ni := li.nodeIndex()
 685  	n := ratchetTreeNumLeaves(*tree)
 686  	var ok bool
 687  	for {
 688  		ratchetTreeSet(*tree, ni, nil)
 689  		ni, ok = n.parent(ni)
 690  		if !ok {
 691  			break
 692  		}
 693  	}
 694  
 695  	li = leafIndex(n - 1)
 696  	lastPowerOf2 := len(*tree) + 1
 697  	for {
 698  		ni = li.nodeIndex()
 699  		if ratchetTreeGet(*tree, ni) != nil {
 700  			break
 701  		}
 702  		if isPowerOf2(uint32(ni)) {
 703  			lastPowerOf2 = int32(ni)
 704  		}
 705  		if li == 0 {
 706  			*tree = nil
 707  			return
 708  		}
 709  		li--
 710  	}
 711  	if lastPowerOf2 < len(*tree)+1 {
 712  		*tree = (*tree)[:lastPowerOf2-1]
 713  	}
 714  }
 715  
 716  func ratchetTreeApply(tree *[]*node, proposals []proposal, senders []leafIndex) {
 717  	for i, prop := range proposals {
 718  		if prop.proposalType == proposalTypeUpdate {
 719  			ratchetTreeUpdate(*tree, senders[i], &prop.update.leafNode)
 720  		}
 721  	}
 722  	for _, prop := range proposals {
 723  		if prop.proposalType == proposalTypeRemove {
 724  			ratchetTreeRemove(tree, prop.remove.removed)
 725  		}
 726  	}
 727  	for _, prop := range proposals {
 728  		if prop.proposalType == proposalTypeAdd {
 729  			ratchetTreeAdd(tree, &prop.add.keyPackage.leafNode)
 730  		}
 731  	}
 732  }
 733  
 734  func ratchetTreeFindLeaf(tree []*node, ln *leafNode) (v leafIndex, ok bool) {
 735  	for li := leafIndex(0); li < leafIndex(ratchetTreeNumLeaves(tree)); li++ {
 736  		nd := ratchetTreeGetLeaf(tree, li)
 737  		if nd == nil {
 738  			continue
 739  		}
 740  		if !bytesEqual(nd.encryptionKey, ln.encryptionKey) {
 741  			continue
 742  		}
 743  		raw1, err1 := marshalRaw(ln)
 744  		raw2, err2 := marshalRaw(nd)
 745  		return li, err1 == nil && err2 == nil && bytesEqual(raw1, raw2)
 746  	}
 747  	return 0, false
 748  }
 749  
 750  func ratchetTreeKeys(tree []*node) (sigKeys, encKeys map[string]bool) {
 751  	sigKeys = map[string]bool{}
 752  	encKeys = map[string]bool{}
 753  	for li := leafIndex(0); li < leafIndex(ratchetTreeNumLeaves(tree)); li++ {
 754  		nd := ratchetTreeGetLeaf(tree, li)
 755  		if nd == nil {
 756  			continue
 757  		}
 758  		sigKeys[string(nd.signatureKey)] = true
 759  		encKeys[string(nd.encryptionKey)] = true
 760  	}
 761  	return sigKeys, encKeys
 762  }
 763  
 764  func ratchetTreeSupportedCreds(tree []*node) (m map[credentialType]bool) {
 765  	numMembers := int32(0)
 766  	counts := map[credentialType]int32{}
 767  	for li := leafIndex(0); li < leafIndex(ratchetTreeNumLeaves(tree)); li++ {
 768  		nd := ratchetTreeGetLeaf(tree, li)
 769  		if nd == nil {
 770  			continue
 771  		}
 772  		numMembers++
 773  		for _, ct := range nd.capabilities.credentials {
 774  			counts[ct]++
 775  		}
 776  	}
 777  	result := map[credentialType]bool{}
 778  	for ct, c := range counts {
 779  		if c == numMembers {
 780  			result[ct] = true
 781  		}
 782  	}
 783  	return result
 784  }
 785  
 786  func ratchetTreeFilteredDirectPath(tree []*node, x nodeIndex) (ss []nodeIndex) {
 787  	n := ratchetTreeNumLeaves(tree)
 788  	var path []nodeIndex
 789  	for {
 790  		p, ok := n.parent(x)
 791  		if !ok {
 792  			break
 793  		}
 794  		s, ok := n.sibling(x)
 795  		if !ok {
 796  			panic("unreachable")
 797  		}
 798  		if len(ratchetTreeResolve(tree, s)) > 0 {
 799  			path = push(path, p)
 800  		}
 801  		x = p
 802  	}
 803  	return path
 804  }
 805  
 806  func hasUnmergedLeaf(pn *parentNode, target leafIndex) (ok bool) {
 807  	for _, li := range pn.unmergedLeaves {
 808  		if li == target {
 809  			return true
 810  		}
 811  	}
 812  	return false
 813  }
 814  
 815  func ratchetTreeFindParentHash(tree []*node, nodeIndices []nodeIndex, ph []byte) (ok bool) {
 816  	for _, x := range nodeIndices {
 817  		nd := ratchetTreeGet(tree, x)
 818  		if nd == nil {
 819  			continue
 820  		}
 821  		var h []byte
 822  		switch nd.nodeType {
 823  		case nodeTypeLeaf:
 824  			h = nd.leafNode.parentHash
 825  		case nodeTypeParent:
 826  			h = nd.parentNode.parentHash
 827  		}
 828  		if bytesEqual(h, ph) {
 829  			return true
 830  		}
 831  	}
 832  	return false
 833  }
 834  
 835  // --- LeafNode verification ---
 836  
 837  type leafNodeVerifyOptions struct {
 838  	cipherSuite    CipherSuite
 839  	groupID        []byte
 840  	leafIndex      leafIndex
 841  	supportedCreds map[credentialType]bool
 842  	signatureKeys  map[string]bool
 843  	encryptionKeys map[string]bool
 844  	nowUnix        int64 // 0 = skip lifetime check
 845  }
 846  
 847  func (ln *leafNode) verify(opts *leafNodeVerifyOptions) (err error) {
 848  	if !ln.verifySignature(opts.cipherSuite, opts.groupID, opts.leafIndex) {
 849  		return errors.New("mls: leaf node signature verification failed")
 850  	}
 851  	if !opts.supportedCreds[ln.credential.credentialType] {
 852  		return errors.New("mls: credential type not supported by all members")
 853  	}
 854  	if ln.lifetime != nil && opts.nowUnix != 0 {
 855  		if !ln.lifetime.verifyAt(opts.nowUnix) {
 856  			return errors.New("mls: lifetime verification failed")
 857  		}
 858  	}
 859  	supportedExts := map[extensionType]bool{}
 860  	for _, et := range ln.capabilities.extensions {
 861  		supportedExts[et] = true
 862  	}
 863  	for _, ext := range ln.extensions {
 864  		if !supportedExts[ext.extensionType] {
 865  			return errors.New("mls: extension type not supported by leaf node")
 866  		}
 867  	}
 868  	if opts.signatureKeys[string(ln.signatureKey)] {
 869  		return errors.New("mls: duplicate signature key")
 870  	}
 871  	if opts.encryptionKeys[string(ln.encryptionKey)] {
 872  		return errors.New("mls: duplicate encryption key")
 873  	}
 874  	return nil
 875  }
 876  
 877  const maxLeafNodeLifetime = 90 * 24 * 3600 // 90 days in seconds
 878  
 879  func (lt *lifetime) verifyAt(nowUnix int64) (ok bool) {
 880  	notBefore := int64(lt.notBefore)
 881  	notAfter := int64(lt.notAfter)
 882  	duration := notAfter - notBefore
 883  	if duration <= 0 || duration > maxLeafNodeLifetime {
 884  		return false
 885  	}
 886  	return nowUnix > notBefore && notAfter > nowUnix
 887  }
 888