From c6b35440c6ea8f934606ab6f3f47057c77b971bf Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Tue, 30 Sep 2025 18:35:58 -0700 Subject: Use some standard data structures for readability --- internal/cryptoutil/oae2.go | 102 +++++++++++++++++++------------------------- 1 file changed, 43 insertions(+), 59 deletions(-) (limited to 'internal/cryptoutil/oae2.go') diff --git a/internal/cryptoutil/oae2.go b/internal/cryptoutil/oae2.go index 018eb74..2819852 100644 --- a/internal/cryptoutil/oae2.go +++ b/internal/cryptoutil/oae2.go @@ -1,6 +1,8 @@ package cryptoutil import ( + "bufio" + "bytes" "crypto/aes" "crypto/cipher" "crypto/hkdf" @@ -38,6 +40,7 @@ type oae2 struct { } func (o *oae2) initialize(header []byte) error { + o.i = 1 copy(o.noncePrefix[:], header[aesKeySize:]) derivedKey, err := hkdf.Key(sha512.New, o.key, header[:aesKeySize], "", aesKeySize) if err != nil { @@ -64,22 +67,22 @@ func (o *oae2) nonce(nonce []byte, lastBlock bool) error { return nil } -func (o *oae2) encryptBlock(block []byte, lastBlock bool) ([]byte, error) { +func (o *oae2) encryptBlock(out, 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) + encrypted := o.aead.Seal(out, nonce, block, o.additionalData) o.additionalData = nil return encrypted, nil } -func (o *oae2) decryptBlock(block []byte, lastBlock bool) ([]byte, error) { +func (o *oae2) decryptBlock(out, 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) + decrypted, err := o.aead.Open(out, nonce, block, o.additionalData) o.additionalData = nil return decrypted, err } @@ -94,9 +97,7 @@ type EncryptingWriter struct { initialized bool err error - - bufN int - buf [encryptedSegmentSize]byte + buf []byte } // NewWriter returns a new EncryptingWriter that writes to w. The additionalData @@ -108,13 +109,13 @@ func (k EncryptionKey) NewWriter(w io.Writer, additionalData []byte) *Encrypting oae2: oae2{ key: k, additionalData: additionalData, - i: 1, }, } } func (w *EncryptingWriter) initialize() error { w.initialized = true + w.buf = make([]byte, 0, encryptedSegmentSize) header := make([]byte, headerSize) rand.Read(header) if w.err = w.oae2.initialize(header); w.err != nil { @@ -133,22 +134,22 @@ func (w *EncryptingWriter) Write(buf []byte) (int, error) { if w.err != nil { return 0, w.err } + r := bytes.NewReader(buf) nn := 0 - for len(buf) > 0 { - if w.bufN == segmentSize { + for r.Len() > 0 { + if len(w.buf) == segmentSize { var encrypted []byte - if encrypted, w.err = w.oae2.encryptBlock(w.buf[:w.bufN], false); w.err != nil { + if encrypted, w.err = w.oae2.encryptBlock(w.buf[:0], w.buf, 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 + w.buf = w.buf[:0] } - n := copy(w.buf[w.bufN:segmentSize], buf) - w.bufN += n + n, _ := r.Read(w.buf[len(w.buf):segmentSize]) + w.buf = w.buf[:len(w.buf)+n] nn += n - buf = buf[n:] } return nn, nil } @@ -170,7 +171,7 @@ func (w *EncryptingWriter) Close() error { return w.err } var encrypted []byte - if encrypted, w.err = w.oae2.encryptBlock(w.buf[:w.bufN], true); w.err != nil { + if encrypted, w.err = w.oae2.encryptBlock(w.buf[:0], w.buf, true); w.err != nil { return w.err } if _, w.err = w.w.Write(encrypted); w.err != nil { @@ -182,66 +183,53 @@ func (w *EncryptingWriter) Close() error { // A DecryptingReader decrypts data using the STREAM construction. type DecryptingReader struct { - r io.Reader + r *bufio.Reader oae2 oae2 - initialized bool - err error - - bufRead, bufN int - peekedByte bool - buf [encryptedSegmentSize + 1]byte + initialized bool + decryptedBuf bytes.Buffer } // 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, + r: bufio.NewReaderSize(r, encryptedSegmentSize+1), oae2: oae2{ key: k, additionalData: additionalData, - i: 1, }, } } func (r *DecryptingReader) initialize() error { - r.initialized = true - header := make([]byte, headerSize) - if _, r.err = io.ReadFull(r.r, header); r.err != nil { - return r.err + header, err := r.r.Peek(headerSize) + if err != nil { + return err } - r.err = r.oae2.initialize(header) - return r.err + if err := r.oae2.initialize(header); err != nil { + return err + } + r.r.Discard(len(header)) + r.decryptedBuf = *bytes.NewBuffer(make([]byte, 0, segmentSize)) + r.initialized = true + return nil } func (r *DecryptingReader) fillBuf() error { - n := 0 - if r.peekedByte { - n = 1 - r.buf[0] = r.buf[encryptedSegmentSize] - r.peekedByte = false - } - for n < len(r.buf) && r.err == nil { - var m int - m, r.err = r.r.Read(r.buf[n:]) - n += m - } - if n == 0 { - return r.err - } - if n == encryptedSegmentSize+1 { - r.peekedByte = true - n = encryptedSegmentSize - } - result, err := r.oae2.decryptBlock(r.buf[:n], r.err == io.EOF) + // Peek one extra byte to make sure if this is the last segment + block, readErr := r.r.Peek(encryptedSegmentSize + 1) + if len(block) == 0 { + return readErr + } + block = block[:min(len(block), encryptedSegmentSize)] + r.decryptedBuf.Reset() + result, err := r.oae2.decryptBlock(r.decryptedBuf.AvailableBuffer(), block, readErr == io.EOF) if err != nil { - r.err = err return err } - r.bufRead = 0 - r.bufN = len(result) + r.decryptedBuf.Write(result) + r.r.Discard(len(block)) return nil } @@ -251,15 +239,11 @@ func (r *DecryptingReader) Read(buf []byte) (int, error) { return 0, err } } - if r.bufRead == r.bufN { + if r.decryptedBuf.Len() == 0 { if err := r.fillBuf(); err != nil { return 0, err } } - n := copy(buf, r.buf[r.bufRead:r.bufN]) - r.bufRead += n - if r.bufRead == r.bufN { - return n, r.err - } + n, _ := r.decryptedBuf.Read(buf) return n, nil } -- cgit v1.3.1