proposal.mx raw

   1  package mls
   2  
   3  // MLS proposals (RFC 9420 §12).
   4  
   5  var (
   6  	errInvalidProposalType    error
   7  	errDupAddKey              error
   8  	errUpdateByCommitter      error
   9  	errDupUpdateOrRemove      error
  10  	errRemoveCommitter        error
  11  	errDupPSK                 error
  12  	errDupGroupContextExts    error
  13  	errReinitWithOther        error
  14  	errExternalInitNotAllowed error
  15  )
  16  
  17  // --- ProposalType ---
  18  
  19  type proposalType uint16
  20  
  21  const (
  22  	proposalTypeAdd                    proposalType = 0x0001
  23  	proposalTypeUpdate                 proposalType = 0x0002
  24  	proposalTypeRemove                 proposalType = 0x0003
  25  	proposalTypePSK                    proposalType = 0x0004
  26  	proposalTypeReinit                 proposalType = 0x0005
  27  	proposalTypeExternalInit           proposalType = 0x0006
  28  	proposalTypeGroupContextExtensions proposalType = 0x0007
  29  )
  30  
  31  func (t *proposalType) unmarshal(r *Reader) (err error) {
  32  	v, ok := r.readUint16()
  33  	if !ok {
  34  		return errUnexpectedEOF
  35  	}
  36  	*t = proposalType(v)
  37  	switch *t {
  38  	case proposalTypeAdd, proposalTypeUpdate, proposalTypeRemove,
  39  		proposalTypePSK, proposalTypeReinit, proposalTypeExternalInit,
  40  		proposalTypeGroupContextExtensions:
  41  		return nil
  42  	default:
  43  		return errInvalidProposalType
  44  	}
  45  }
  46  
  47  func (t *proposalType) marshal(w *Writer) {
  48  	w.addUint16(uint16(t))
  49  }
  50  
  51  // --- Proposal ---
  52  
  53  type proposal struct {
  54  	proposalType           proposalType
  55  	add                    *add
  56  	update                 *update
  57  	remove                 *remove
  58  	preSharedKey           *preSharedKey
  59  	reInit                 *reInit
  60  	externalInit           *externalInit
  61  	groupContextExtensions *groupContextExtensions
  62  }
  63  
  64  func (prop *proposal) unmarshal(r *Reader) (err error) {
  65  	*prop = proposal{}
  66  	if e := prop.proposalType.unmarshal(r); e != nil {
  67  		return e
  68  	}
  69  	switch prop.proposalType {
  70  	case proposalTypeAdd:
  71  		prop.add = &add{}
  72  		return prop.add.unmarshal(r)
  73  	case proposalTypeUpdate:
  74  		prop.update = &update{}
  75  		return prop.update.unmarshal(r)
  76  	case proposalTypeRemove:
  77  		prop.remove = &remove{}
  78  		return prop.remove.unmarshal(r)
  79  	case proposalTypePSK:
  80  		prop.preSharedKey = &preSharedKey{}
  81  		return prop.preSharedKey.unmarshal(r)
  82  	case proposalTypeReinit:
  83  		prop.reInit = &reInit{}
  84  		return prop.reInit.unmarshal(r)
  85  	case proposalTypeExternalInit:
  86  		prop.externalInit = &externalInit{}
  87  		return prop.externalInit.unmarshal(r)
  88  	case proposalTypeGroupContextExtensions:
  89  		prop.groupContextExtensions = &groupContextExtensions{}
  90  		return prop.groupContextExtensions.unmarshal(r)
  91  	default:
  92  		panic("unreachable")
  93  	}
  94  }
  95  
  96  func (prop *proposal) marshal(w *Writer) {
  97  	prop.proposalType.marshal(w)
  98  	switch prop.proposalType {
  99  	case proposalTypeAdd:
 100  		prop.add.marshal(w)
 101  	case proposalTypeUpdate:
 102  		prop.update.marshal(w)
 103  	case proposalTypeRemove:
 104  		prop.remove.marshal(w)
 105  	case proposalTypePSK:
 106  		prop.preSharedKey.marshal(w)
 107  	case proposalTypeReinit:
 108  		prop.reInit.marshal(w)
 109  	case proposalTypeExternalInit:
 110  		prop.externalInit.marshal(w)
 111  	case proposalTypeGroupContextExtensions:
 112  		prop.groupContextExtensions.marshal(w)
 113  	default:
 114  		panic("unreachable")
 115  	}
 116  }
 117  
 118  // --- Proposal sub-types ---
 119  
 120  type add struct {
 121  	keyPackage KeyPackage
 122  }
 123  
 124  func (a *add) unmarshal(r *Reader) (err error) {
 125  	*a = add{}
 126  	return a.keyPackage.unmarshal(r)
 127  }
 128  
 129  func (a *add) marshal(w *Writer) {
 130  	a.keyPackage.marshal(w)
 131  }
 132  
 133  type update struct {
 134  	leafNode leafNode
 135  }
 136  
 137  func (upd *update) unmarshal(r *Reader) (err error) {
 138  	*upd = update{}
 139  	return upd.leafNode.unmarshal(r)
 140  }
 141  
 142  func (upd *update) marshal(w *Writer) {
 143  	upd.leafNode.marshal(w)
 144  }
 145  
 146  type remove struct {
 147  	removed leafIndex
 148  }
 149  
 150  func (rm *remove) unmarshal(r *Reader) (err error) {
 151  	*rm = remove{}
 152  	v, ok := r.readUint32()
 153  	if !ok {
 154  		return errUnexpectedEOF
 155  	}
 156  	rm.removed = leafIndex(v)
 157  	return nil
 158  }
 159  
 160  func (rm *remove) marshal(w *Writer) {
 161  	w.addUint32(uint32(rm.removed))
 162  }
 163  
 164  type preSharedKey struct {
 165  	psk preSharedKeyID
 166  }
 167  
 168  func (psk *preSharedKey) unmarshal(r *Reader) (err error) {
 169  	*psk = preSharedKey{}
 170  	return psk.psk.unmarshal(r)
 171  }
 172  
 173  func (psk *preSharedKey) marshal(w *Writer) {
 174  	psk.psk.marshal(w)
 175  }
 176  
 177  type reInit struct {
 178  	groupID     []byte
 179  	version     protocolVersion
 180  	cipherSuite CipherSuite
 181  	extensions  []extension
 182  }
 183  
 184  func (ri *reInit) unmarshal(r *Reader) (err error) {
 185  	*ri = reInit{}
 186  	var ok bool
 187  	ri.groupID, ok = r.readOpaqueVec()
 188  	if !ok {
 189  		return errUnexpectedEOF
 190  	}
 191  	v, ok := r.readUint16()
 192  	if !ok {
 193  		return errUnexpectedEOF
 194  	}
 195  	ri.version = protocolVersion(v)
 196  	v, ok = r.readUint16()
 197  	if !ok {
 198  		return errUnexpectedEOF
 199  	}
 200  	ri.cipherSuite = CipherSuite(v)
 201  	exts, err := unmarshalExtensionVec(r)
 202  	if err != nil {
 203  		return err
 204  	}
 205  	ri.extensions = exts
 206  	return nil
 207  }
 208  
 209  func (ri *reInit) marshal(w *Writer) {
 210  	w.writeOpaqueVec([]byte(ri.groupID))
 211  	w.addUint16(uint16(ri.version))
 212  	w.addUint16(uint16(ri.cipherSuite))
 213  	marshalExtensionVec(w, ri.extensions)
 214  }
 215  
 216  type externalInit struct {
 217  	kemOutput []byte
 218  }
 219  
 220  func (ei *externalInit) unmarshal(r *Reader) (err error) {
 221  	*ei = externalInit{}
 222  	var ok bool
 223  	ei.kemOutput, ok = r.readOpaqueVec()
 224  	if !ok {
 225  		return errUnexpectedEOF
 226  	}
 227  	return nil
 228  }
 229  
 230  func (ei *externalInit) marshal(w *Writer) {
 231  	w.writeOpaqueVec(ei.kemOutput)
 232  }
 233  
 234  type groupContextExtensions struct {
 235  	extensions []extension
 236  }
 237  
 238  func (exts *groupContextExtensions) unmarshal(r *Reader) (err error) {
 239  	*exts = groupContextExtensions{}
 240  	l, err := unmarshalExtensionVec(r)
 241  	if err != nil {
 242  		return err
 243  	}
 244  	exts.extensions = l
 245  	return nil
 246  }
 247  
 248  func (exts *groupContextExtensions) marshal(w *Writer) {
 249  	marshalExtensionVec(w, exts.extensions)
 250  }
 251  
 252  // --- ProposalOrRef ---
 253  
 254  type proposalOrRefType uint8
 255  
 256  const (
 257  	proposalOrRefTypeProposal  proposalOrRefType = 1
 258  	proposalOrRefTypeReference proposalOrRefType = 2
 259  )
 260  
 261  func (t *proposalOrRefType) unmarshal(r *Reader) (err error) {
 262  	b, ok := r.readByte()
 263  	if !ok {
 264  		return errUnexpectedEOF
 265  	}
 266  	*t = proposalOrRefType(b)
 267  	switch *t {
 268  	case proposalOrRefTypeProposal, proposalOrRefTypeReference:
 269  		return nil
 270  	default:
 271  		return errInvalidProposalOrRefType
 272  	}
 273  }
 274  
 275  func (t *proposalOrRefType) marshal(w *Writer) {
 276  	w.addByte(byte(t))
 277  }
 278  
 279  // proposalRefEqual reports whether two proposal references are equal.
 280  func proposalRefEqual(ref []byte, other []byte) (ok bool) {
 281  	return bytesEqual(ref, other)
 282  }
 283  
 284  type proposalOrRef struct {
 285  	typ       proposalOrRefType
 286  	proposal  *proposal
 287  	reference []byte
 288  }
 289  
 290  func (por *proposalOrRef) unmarshal(r *Reader) (err error) {
 291  	*por = proposalOrRef{}
 292  	if e := por.typ.unmarshal(r); e != nil {
 293  		return e
 294  	}
 295  	switch por.typ {
 296  	case proposalOrRefTypeProposal:
 297  		por.proposal = &proposal{}
 298  		return por.proposal.unmarshal(r)
 299  	case proposalOrRefTypeReference:
 300  		var ok bool
 301  		por.reference, ok = r.readOpaqueVec()
 302  		if !ok {
 303  			return errUnexpectedEOF
 304  		}
 305  		return nil
 306  	default:
 307  		panic("unreachable")
 308  	}
 309  }
 310  
 311  func (por *proposalOrRef) marshal(w *Writer) {
 312  	por.typ.marshal(w)
 313  	switch por.typ {
 314  	case proposalOrRefTypeProposal:
 315  		por.proposal.marshal(w)
 316  	case proposalOrRefTypeReference:
 317  		w.writeOpaqueVec([]byte(por.reference))
 318  	default:
 319  		panic("unreachable")
 320  	}
 321  }
 322  
 323  // --- Proposal list validation (RFC 9420 §12.2) ---
 324  
 325  func verifyProposalList(proposals []proposal, senders []leafIndex, committer leafIndex) (err error) {
 326  	if len(proposals) != len(senders) {
 327  		panic("unreachable")
 328  	}
 329  
 330  	addKeys := map[string]struct{}{}
 331  	updateOrRemove := map[leafIndex]struct{}{}
 332  	pskKeys := map[string]struct{}{}
 333  	hasGroupCtxExts := false
 334  
 335  	for i, prop := range proposals {
 336  		snd := senders[i]
 337  		switch prop.proposalType {
 338  		case proposalTypeAdd:
 339  			k := string(prop.add.keyPackage.leafNode.signatureKey)
 340  			if _, dup := addKeys[k]; dup {
 341  				return errDupAddKey
 342  			}
 343  			addKeys[k] = struct{}{}
 344  		case proposalTypeUpdate:
 345  			if snd == committer {
 346  				return errUpdateByCommitter
 347  			}
 348  			if _, dup := updateOrRemove[snd]; dup {
 349  				return errDupUpdateOrRemove
 350  			}
 351  			updateOrRemove[snd] = struct{}{}
 352  		case proposalTypeRemove:
 353  			if prop.remove.removed == committer {
 354  				return errRemoveCommitter
 355  			}
 356  			if _, dup := updateOrRemove[prop.remove.removed]; dup {
 357  				return errDupUpdateOrRemove
 358  			}
 359  			updateOrRemove[prop.remove.removed] = struct{}{}
 360  		case proposalTypePSK:
 361  			raw, e := marshalRaw(&prop.preSharedKey.psk)
 362  			if e != nil {
 363  				return e
 364  			}
 365  			k := string(raw)
 366  			if _, dup := pskKeys[k]; dup {
 367  				return errDupPSK
 368  			}
 369  			pskKeys[k] = struct{}{}
 370  		case proposalTypeGroupContextExtensions:
 371  			if hasGroupCtxExts {
 372  				return errDupGroupContextExts
 373  			}
 374  			hasGroupCtxExts = true
 375  		case proposalTypeReinit:
 376  			if len(proposals) > 1 {
 377  				return errReinitWithOther
 378  			}
 379  		case proposalTypeExternalInit:
 380  			return errExternalInitNotAllowed
 381  		}
 382  	}
 383  	return nil
 384  }
 385  
 386  func proposalListNeedsPath(proposals []proposal) (ok bool) {
 387  	if len(proposals) == 0 {
 388  		return true
 389  	}
 390  	for _, prop := range proposals {
 391  		switch prop.proposalType {
 392  		case proposalTypeUpdate, proposalTypeRemove,
 393  			proposalTypeExternalInit, proposalTypeGroupContextExtensions:
 394  			return true
 395  		}
 396  	}
 397  	return false
 398  }
 399