From 57fa3c98dc8586a3773bd79066f53e678eb49f61 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Tue, 30 Sep 2025 16:41:43 -0700 Subject: Make the oae2 implementation a bit easier to read --- internal/cryptoutil/oae2.go | 178 ++++++++++++++++++++++---------------------- 1 file changed, 89 insertions(+), 89 deletions(-) (limited to 'internal') diff --git a/internal/cryptoutil/oae2.go b/internal/cryptoutil/oae2.go index 28ec9d7..018eb74 100644 --- a/internal/cryptoutil/oae2.go +++ b/internal/cryptoutil/oae2.go @@ -28,43 +28,75 @@ const ( // length, but at least 64 bytes would be recommended. type EncryptionKey []byte -func (k EncryptionKey) deriveAESKey(salt []byte) (cipher.AEAD, error) { - derivedKey, err := hkdf.Key(sha512.New, k, salt, "", aesKeySize) +type oae2 struct { + key EncryptionKey + additionalData []byte + + aead cipher.AEAD + i uint64 + noncePrefix [noncePrefixSize]byte +} + +func (o *oae2) initialize(header []byte) error { + copy(o.noncePrefix[:], header[aesKeySize:]) + derivedKey, err := hkdf.Key(sha512.New, o.key, header[:aesKeySize], "", aesKeySize) if err != nil { - return nil, err + return err } block, err := aes.NewCipher(derivedKey) if err != nil { - return nil, err + return err } - return cipher.NewGCM(block) + o.aead, err = cipher.NewGCM(block) + return err } -func makeNonce(nonce []byte, noncePrefix []byte, i *uint64) error { - copy(nonce, noncePrefix) - binary.BigEndian.PutUint64(nonce[noncePrefixSize:], *i) - *i++ - if *i == 0 { +func (o *oae2) nonce(nonce []byte, lastBlock bool) error { + copy(nonce, o.noncePrefix[:]) + if o.i == 0 { return errors.New("counter overflowed (64 bits??)") } + binary.BigEndian.PutUint64(nonce[noncePrefixSize:], o.i) + o.i++ + if lastBlock { + nonce[gcmNonceSize-1] = 1 + } return nil } +func (o *oae2) encryptBlock(block []byte, lastBlock bool) ([]byte, error) { + nonce := make([]byte, gcmNonceSize) + if err := o.nonce(nonce, lastBlock); err != nil { + return nil, err + } + encrypted := o.aead.Seal(block[:len(block)+aeadOverhead][:0], nonce, block, o.additionalData) + o.additionalData = nil + return encrypted, nil +} + +func (o *oae2) decryptBlock(block []byte, lastBlock bool) ([]byte, error) { + nonce := make([]byte, gcmNonceSize) + if err := o.nonce(nonce, lastBlock); err != nil { + return nil, err + } + decrypted, err := o.aead.Open(block[:0], nonce, block, o.additionalData) + o.additionalData = nil + return decrypted, err +} + // An EncryptingWriter encrypts data in segments using the STREAM construction // described in https://eprint.iacr.org/2015/189.pdf. The writer buffers data up // to the segment size, so it's important to call Close to flush the // final segment. type EncryptingWriter struct { - w io.Writer - key EncryptionKey - aead cipher.AEAD - additionalData []byte - noncePrefix [noncePrefixSize]byte - i uint64 - initialized bool - err error - bufN int - buf [encryptedSegmentSize]byte + w io.Writer + oae2 oae2 + + initialized bool + err error + + bufN int + buf [encryptedSegmentSize]byte } // NewWriter returns a new EncryptingWriter that writes to w. The additionalData @@ -72,10 +104,12 @@ type EncryptingWriter struct { // same additional data must be provided when decrypting. func (k EncryptionKey) NewWriter(w io.Writer, additionalData []byte) *EncryptingWriter { return &EncryptingWriter{ - w: w, - key: k, - additionalData: additionalData, - i: 1, + w: w, + oae2: oae2{ + key: k, + additionalData: additionalData, + i: 1, + }, } } @@ -83,32 +117,13 @@ func (w *EncryptingWriter) initialize() error { w.initialized = true header := make([]byte, headerSize) rand.Read(header) - copy(w.noncePrefix[:], header[aesKeySize:]) - if w.aead, w.err = w.key.deriveAESKey(header[:aesKeySize]); w.err != nil { + if w.err = w.oae2.initialize(header); w.err != nil { return w.err } _, w.err = w.w.Write(header) return w.err } -func (w *EncryptingWriter) nonce(nonce []byte) error { - w.err = makeNonce(nonce, w.noncePrefix[:], &w.i) - return w.err -} - -func (w *EncryptingWriter) writeBuf() error { - nonce := make([]byte, gcmNonceSize) - if err := w.nonce(nonce); err != nil { - return err - } - if _, w.err = w.w.Write(w.aead.Seal(w.buf[:0], nonce, w.buf[:w.bufN], w.additionalData)); w.err != nil { - return w.err - } - w.bufN = 0 - w.additionalData = nil - return nil -} - func (w *EncryptingWriter) Write(buf []byte) (int, error) { if !w.initialized { if err := w.initialize(); err != nil { @@ -121,9 +136,14 @@ func (w *EncryptingWriter) Write(buf []byte) (int, error) { nn := 0 for len(buf) > 0 { if w.bufN == segmentSize { - if err := w.writeBuf(); err != nil { - return nn, err + var encrypted []byte + if encrypted, w.err = w.oae2.encryptBlock(w.buf[:w.bufN], false); w.err != nil { + return nn, w.err } + if _, w.err = w.w.Write(encrypted); w.err != nil { + return nn, w.err + } + w.bufN = 0 } n := copy(w.buf[w.bufN:segmentSize], buf) w.bufN += n @@ -149,15 +169,12 @@ func (w *EncryptingWriter) Close() error { if w.err != nil { return w.err } - if w.bufN > 0 { - nonce := make([]byte, gcmNonceSize) - if err := w.nonce(nonce); err != nil { - return err - } - nonce[gcmNonceSize-1] = 1 - if _, w.err = w.w.Write(w.aead.Seal(w.buf[:0], nonce, w.buf[:w.bufN], w.additionalData)); w.err != nil { - return w.err - } + var encrypted []byte + if encrypted, w.err = w.oae2.encryptBlock(w.buf[:w.bufN], true); w.err != nil { + return w.err + } + if _, w.err = w.w.Write(encrypted); w.err != nil { + return w.err } w.err = errClosed return nil @@ -165,27 +182,27 @@ func (w *EncryptingWriter) Close() error { // A DecryptingReader decrypts data using the STREAM construction. type DecryptingReader struct { - r io.Reader - key EncryptionKey - aead cipher.AEAD - additionalData []byte - noncePrefix [noncePrefixSize]byte - i uint64 - initialized bool - err error - bufRead, bufN int - peekedByte bool - buf [encryptedSegmentSize + 1]byte + r io.Reader + oae2 oae2 + + initialized bool + err error + + bufRead, bufN int + peekedByte bool + buf [encryptedSegmentSize + 1]byte } // NewReader returns a new DecryptingWriter that decrypts data from r. The // additionaData must be the same that was provided when encrypting. func (k EncryptionKey) NewReader(r io.Reader, additionalData []byte) *DecryptingReader { return &DecryptingReader{ - r: r, - key: k, - additionalData: additionalData, - i: 1, + r: r, + oae2: oae2{ + key: k, + additionalData: additionalData, + i: 1, + }, } } @@ -195,19 +212,10 @@ func (r *DecryptingReader) initialize() error { if _, r.err = io.ReadFull(r.r, header); r.err != nil { return r.err } - copy(r.noncePrefix[:], header[aesKeySize:]) - r.aead, r.err = r.key.deriveAESKey(header[:aesKeySize]) + r.err = r.oae2.initialize(header) return r.err } -func (r *DecryptingReader) nonce(nonce []byte) error { - if err := makeNonce(nonce, r.noncePrefix[:], &r.i); err != nil { - r.err = err - return err - } - return nil -} - func (r *DecryptingReader) fillBuf() error { n := 0 if r.peekedByte { @@ -227,21 +235,13 @@ func (r *DecryptingReader) fillBuf() error { r.peekedByte = true n = encryptedSegmentSize } - nonce := make([]byte, gcmNonceSize) - if err := r.nonce(nonce); err != nil { - return err - } - if r.err == io.EOF { - nonce[gcmNonceSize-1] = 1 - } - result, err := r.aead.Open(r.buf[:0], nonce, r.buf[:n], r.additionalData) + result, err := r.oae2.decryptBlock(r.buf[:n], r.err == io.EOF) if err != nil { r.err = err return err } r.bufRead = 0 r.bufN = len(result) - r.additionalData = nil return nil } -- cgit v1.3.1