key_schedule.mx raw

   1  package mls
   2  
   3  // MLS key schedule types (RFC 9420 ยง8).
   4  // Data types and serialization - crypto operations go in key_schedule_crypto.mx.
   5  
   6  var (
   7  	errInvalidPSKType           error
   8  	errInvalidPSKUsage          error
   9  	errInvalidProposalOrRefType error
  10  )
  11  
  12  // --- CipherSuite ---
  13  
  14  type CipherSuite uint16
  15  
  16  const (
  17  	// MLS_128_DHKEMP256_AES128GCM_SHA256_P256
  18  	CipherSuite0x0001 CipherSuite = 0x0001
  19  	// MLS_128_DHKEMX25519_CHACHA20POLY1305_SHA256_Ed25519
  20  	CipherSuite0x0003 CipherSuite = 0x0003
  21  )
  22  
  23  // --- GroupContext ---
  24  
  25  type groupContext struct {
  26  	version                 protocolVersion
  27  	cipherSuite             CipherSuite
  28  	groupID                 []byte
  29  	epoch                   uint64
  30  	treeHash                []byte
  31  	confirmedTranscriptHash []byte
  32  	extensions              []extension
  33  }
  34  
  35  func (ctx *groupContext) unmarshal(r *Reader) (err error) {
  36  	*ctx = groupContext{}
  37  
  38  	v, ok := r.readUint16()
  39  	if !ok {
  40  		return errUnexpectedEOF
  41  	}
  42  	ctx.version = protocolVersion(v)
  43  	if ctx.version != protocolVersionMLS10 {
  44  		return errInvalidVersion
  45  	}
  46  
  47  	v, ok = r.readUint16()
  48  	if !ok {
  49  		return errUnexpectedEOF
  50  	}
  51  	ctx.cipherSuite = CipherSuite(v)
  52  
  53  	ctx.groupID, ok = r.readOpaqueVec()
  54  	if !ok {
  55  		return errUnexpectedEOF
  56  	}
  57  	ctx.epoch, ok = r.readUint64()
  58  	if !ok {
  59  		return errUnexpectedEOF
  60  	}
  61  	ctx.treeHash, ok = r.readOpaqueVec()
  62  	if !ok {
  63  		return errUnexpectedEOF
  64  	}
  65  	ctx.confirmedTranscriptHash, ok = r.readOpaqueVec()
  66  	if !ok {
  67  		return errUnexpectedEOF
  68  	}
  69  
  70  	exts, err := unmarshalExtensionVec(r)
  71  	if err != nil {
  72  		return err
  73  	}
  74  	ctx.extensions = exts
  75  	return nil
  76  }
  77  
  78  func (ctx *groupContext) marshal(w *Writer) {
  79  	w.addUint16(uint16(ctx.version))
  80  	w.addUint16(uint16(ctx.cipherSuite))
  81  	w.writeOpaqueVec([]byte(ctx.groupID))
  82  	w.addUint64(ctx.epoch)
  83  	w.writeOpaqueVec(ctx.treeHash)
  84  	w.writeOpaqueVec(ctx.confirmedTranscriptHash)
  85  	marshalExtensionVec(w, ctx.extensions)
  86  }
  87  
  88  // --- Secret labels ---
  89  
  90  var (
  91  	secretLabelInit           []byte
  92  	secretLabelSenderData     []byte
  93  	secretLabelEncryption     []byte
  94  	secretLabelExporter       []byte
  95  	secretLabelExternal       []byte
  96  	secretLabelConfirm        []byte
  97  	secretLabelMembership     []byte
  98  	secretLabelResumption     []byte
  99  	secretLabelAuthentication []byte
 100  )
 101  
 102  // --- ConfirmedTranscriptHashInput ---
 103  
 104  type confirmedTranscriptHashInput struct {
 105  	wireFormat wireFormat
 106  	content    framedContent
 107  	signature  []byte
 108  }
 109  
 110  func (input *confirmedTranscriptHashInput) marshal(w *Writer) {
 111  	input.wireFormat.marshal(w)
 112  	input.content.marshal(w)
 113  	w.writeOpaqueVec(input.signature)
 114  }
 115  
 116  // --- PSK types ---
 117  
 118  type pskType uint8
 119  
 120  const (
 121  	pskTypeExternal   pskType = 1
 122  	pskTypeResumption pskType = 2
 123  )
 124  
 125  func (t *pskType) unmarshal(r *Reader) (err error) {
 126  	b, ok := r.readByte()
 127  	if !ok {
 128  		return errUnexpectedEOF
 129  	}
 130  	*t = pskType(b)
 131  	switch *t {
 132  	case pskTypeExternal, pskTypeResumption:
 133  		return nil
 134  	default:
 135  		return errInvalidPSKType
 136  	}
 137  }
 138  
 139  func (t *pskType) marshal(w *Writer) {
 140  	w.addByte(byte(t))
 141  }
 142  
 143  type resumptionPSKUsage uint8
 144  
 145  const (
 146  	resumptionPSKUsageApplication resumptionPSKUsage = 1
 147  	resumptionPSKUsageReinit      resumptionPSKUsage = 2
 148  	resumptionPSKUsageBranch      resumptionPSKUsage = 3
 149  )
 150  
 151  func (usage *resumptionPSKUsage) unmarshal(r *Reader) (err error) {
 152  	b, ok := r.readByte()
 153  	if !ok {
 154  		return errUnexpectedEOF
 155  	}
 156  	*usage = resumptionPSKUsage(b)
 157  	switch *usage {
 158  	case resumptionPSKUsageApplication, resumptionPSKUsageReinit, resumptionPSKUsageBranch:
 159  		return nil
 160  	default:
 161  		return errInvalidPSKUsage
 162  	}
 163  }
 164  
 165  func (usage *resumptionPSKUsage) marshal(w *Writer) {
 166  	w.addByte(byte(usage))
 167  }
 168  
 169  // --- PreSharedKeyID ---
 170  
 171  type preSharedKeyID struct {
 172  	pskType pskType
 173  
 174  	pskID []byte // for pskTypeExternal
 175  
 176  	usage      resumptionPSKUsage // for pskTypeResumption
 177  	pskGroupID []byte            // for pskTypeResumption
 178  	pskEpoch   uint64             // for pskTypeResumption
 179  
 180  	pskNonce []byte
 181  }
 182  
 183  func (id *preSharedKeyID) unmarshal(r *Reader) (err error) {
 184  	*id = preSharedKeyID{}
 185  	if e := id.pskType.unmarshal(r); e != nil {
 186  		return e
 187  	}
 188  
 189  	switch id.pskType {
 190  	case pskTypeExternal:
 191  		var okID bool
 192  		id.pskID, okID = r.readOpaqueVec()
 193  		if !okID {
 194  			return errUnexpectedEOF
 195  		}
 196  	case pskTypeResumption:
 197  		if e := id.usage.unmarshal(r); e != nil {
 198  			return e
 199  		}
 200  		var okGrp bool
 201  		id.pskGroupID, okGrp = r.readOpaqueVec()
 202  		if !okGrp {
 203  			return errUnexpectedEOF
 204  		}
 205  		id.pskEpoch, okGrp = r.readUint64()
 206  		if !okGrp {
 207  			return errUnexpectedEOF
 208  		}
 209  	default:
 210  		panic("unreachable")
 211  	}
 212  
 213  	var ok bool
 214  	id.pskNonce, ok = r.readOpaqueVec()
 215  	if !ok {
 216  		return errUnexpectedEOF
 217  	}
 218  	return nil
 219  }
 220  
 221  func (id *preSharedKeyID) marshal(w *Writer) {
 222  	id.pskType.marshal(w)
 223  	switch id.pskType {
 224  	case pskTypeExternal:
 225  		w.writeOpaqueVec(id.pskID)
 226  	case pskTypeResumption:
 227  		id.usage.marshal(w)
 228  		w.writeOpaqueVec([]byte(id.pskGroupID))
 229  		w.addUint64(id.pskEpoch)
 230  	default:
 231  		panic("unreachable")
 232  	}
 233  	w.writeOpaqueVec(id.pskNonce)
 234  }
 235