diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2026-01-25 07:47:19 -0800 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2026-01-25 08:12:23 -0800 |
| commit | 6412b3912b17cd22ba8bdbf07a12f8b005bf058f (patch) | |
| tree | dbe5d21e300a20d0617e3a6b8bd21f59a4b1adbf /oae2.go | |
| parent | fafe0213c9022f298ec391c7907ed92762fbcd1b (diff) | |
| download | oae2-6412b3912b17cd22ba8bdbf07a12f8b005bf058f.tar.zst | |
Add support for additional data
We can't really call it "streaming AEAD" and then not support
additional data 🤦♀️
Diffstat (limited to 'oae2.go')
| -rw-r--r-- | oae2.go | 57 |
1 files changed, 35 insertions, 22 deletions
@@ -40,7 +40,8 @@ const ( ) type segmentEncrypter struct { - key []byte + key []byte + additionalData []byte aead cipher.AEAD nonce [nonceSize]byte @@ -59,6 +60,16 @@ func (e *segmentEncrypter) init(salt []byte) error { return err } +func (e *segmentEncrypter) firstSegment() bool { + n := e.nonce + var or byte + // First 11 bytes of the nonce are the segment index + for _, b := range n[:nonceSize-1] { + or |= b + } + return or != 0 +} + func (e *segmentEncrypter) nextNonce(lastSegment bool) ([]byte, error) { for i := 0; ; i++ { if i == nonceSize-1 { @@ -76,19 +87,27 @@ func (e *segmentEncrypter) nextNonce(lastSegment bool) ([]byte, error) { } func (e *segmentEncrypter) encryptSegment(segment []byte, lastSegment bool) ([]byte, error) { + var ad []byte + if e.firstSegment() { + ad = e.additionalData + } nonce, err := e.nextNonce(lastSegment) if err != nil { return nil, err } - return e.aead.Seal(segment[:len(segment)+aeadOverhead][:0], nonce, segment, nil), nil + return e.aead.Seal(segment[:len(segment)+aeadOverhead][:0], nonce, segment, ad), nil } func (e *segmentEncrypter) decryptSegment(segment []byte, lastSegment bool) ([]byte, error) { + var ad []byte + if e.firstSegment() { + ad = e.additionalData + } nonce, err := e.nextNonce(lastSegment) if err != nil { return nil, err } - return e.aead.Open(segment[:0], nonce, segment, nil) + return e.aead.Open(segment[:0], nonce, segment, ad) } // A Writer wraps an io.Writer and encrypts the data in segments. Make sure to @@ -108,14 +127,14 @@ type Writer struct { // panics if segmentSize <= 0. // // Make sure to call [Writer.Close] to flush the final segment. -func NewWriter(w io.Writer, key []byte, segmentSize int) *Writer { +func NewWriter(w io.Writer, key []byte, segmentSize int, additionalData []byte) *Writer { if segmentSize <= 0 { panic("oae2.NewWriter: segmentSize must be strictly greater than 0") } return &Writer{ w: w, segmentSize: segmentSize, - encrypter: segmentEncrypter{key: key}, + encrypter: segmentEncrypter{key: key, additionalData: additionalData}, buf: make([]byte, 0, segmentSize+aeadOverhead), } } @@ -261,40 +280,33 @@ type Reader struct { // NewReader returns a Reader that wraps r and decrypts the data using key in // chunks of size segmentSize. segmentSize must match the segment size that was // used to write the encrypted stream. NewReader panics if segmentSize <= 0. -func NewReader(r io.Reader, key []byte, segmentSize int) *Reader { +func NewReader(r io.Reader, key []byte, segmentSize int, additionalData []byte) *Reader { if segmentSize <= 0 { panic("oae2.NewReader: segmentSize must be strictly greater than 0") } return &Reader{ r: bufReader{r: r}, segmentSize: segmentSize, - decrypter: segmentEncrypter{key: key}, + decrypter: segmentEncrypter{key: key, additionalData: additionalData}, buf: make([]byte, 0, segmentSize+aeadOverhead+1), } } func (r *Reader) initialize() error { r.initialized = true - // N.B. there's a subtle bug lurking here: for an input stream of - // exactly 32 bytes, it's important that we reject this stream since it - // doesn't have an authentication tag. So if we naively read exactly 32 - // bytes inside initialize, then Read will see that the underlying - // reader is at EOF, and won't be able to distinguish this from a valid - // EOF case. Fortunately we can easily work around this by reading one - // extra byte here: since the shortest possible encrypted stream is - // 32 + 16 bytes, every valid stream will have an extra byte for us - // here, and other than on the first segment, Read always knows - // precisely whether the stream is at EOF. - buf := make([]byte, saltSize+1) + buf := make([]byte, saltSize) if _, r.err = io.ReadFull(&r.r, buf); r.err != nil { if r.err == io.EOF { r.err = io.ErrUnexpectedEOF } return r.err } - r.r.unreadByte() - r.err = r.decrypter.init(buf[:saltSize]) - return r.err + if r.err = r.decrypter.init(buf); r.err != nil { + return r.err + } + // Decrypt the first segment now to make sure we validate the + // additional data + return r.fillBuf() } func (r *Reader) init() error { @@ -335,7 +347,7 @@ func (r *Reader) Read(buf []byte) (int, error) { if err := r.init(); err != nil { return 0, err } - if r.nRead == len(r.buf) { + if r.nRead >= len(r.buf) { if err := r.fillBuf(); err != nil { return 0, err } @@ -445,6 +457,7 @@ func (r *Reader) Seek(offset int64, whence int) (int64, error) { if err != nil { return 0, err } + r.buf = r.buf[:0] plaintextSize := r.encryptedToPlaintextSize(encryptedSize) if offset == 0 { return plaintextSize, nil |
