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 | |
| 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 🤦♀️
| -rw-r--r-- | oae2.go | 57 | ||||
| -rw-r--r-- | oae2_test.go | 84 |
2 files changed, 86 insertions, 55 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 diff --git a/oae2_test.go b/oae2_test.go index b60be24..4099aa7 100644 --- a/oae2_test.go +++ b/oae2_test.go @@ -10,19 +10,20 @@ import ( func TestRoundTrip(t *testing.T) { t.Parallel() const ( - msg = "Hello World!" - key = "password123" - blockSize = 1 + msg = "Hello World!" + key = "password123" + additionalData = "additional data" + blockSize = 1 ) buf := new(bytes.Buffer) - w := NewWriter(buf, []byte(key), blockSize) + w := NewWriter(buf, []byte(key), blockSize, []byte(additionalData)) if _, err := io.WriteString(w, msg); err != nil { t.Fatalf("Writer.Write failed: %s", err) } if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } - got, err := io.ReadAll(NewReader(buf, []byte(key), blockSize)) + got, err := io.ReadAll(NewReader(buf, []byte(key), blockSize, []byte(additionalData))) if err != nil { t.Fatalf("Reader.Read failed: %s", err) } @@ -34,12 +35,13 @@ func TestRoundTrip(t *testing.T) { func TestReadFromWriteTo(t *testing.T) { t.Parallel() const ( - msg = "Hello World!" - key = "asdf" - blockSize = 2 + msg = "Hello World!" + key = "asdf" + additionalData = "additional data" + blockSize = 2 ) buf := new(bytes.Buffer) - w := NewWriter(buf, []byte(key), blockSize) + w := NewWriter(buf, []byte(key), blockSize, []byte(additionalData)) if _, err := io.Copy(w, struct{ io.Reader }{strings.NewReader(msg)}); err != nil { t.Fatalf("Writer.WriteTo failed: %s", err) } @@ -47,7 +49,7 @@ func TestReadFromWriteTo(t *testing.T) { t.Fatalf("Writer.Close failed: %s", err) } got := new(strings.Builder) - if _, err := io.Copy(got, NewReader(buf, []byte(key), blockSize)); err != nil { + if _, err := io.Copy(got, NewReader(buf, []byte(key), blockSize, []byte(additionalData))); err != nil { t.Fatalf("Reader.Read failed: %s", err) } if got.String() != msg { @@ -58,9 +60,10 @@ func TestReadFromWriteTo(t *testing.T) { func TestReader_Seek(t *testing.T) { t.Parallel() const ( - msg = "Hello World!" - key = "asdf" - blockSize = 1 + msg = "Hello World!" + key = "asdf" + additionalData = "additional data" + blockSize = 1 ) for _, tc := range []struct { desc string @@ -97,14 +100,14 @@ func TestReader_Seek(t *testing.T) { t.Parallel() buf := new(bytes.Buffer) - w := NewWriter(buf, []byte(key), blockSize) + w := NewWriter(buf, []byte(key), blockSize, []byte(additionalData)) if _, err := io.WriteString(w, msg); err != nil { t.Fatalf("Writer.Write failed: %s", err) } if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } - r := NewReader(bytes.NewReader(buf.Bytes()), []byte(key), blockSize) + r := NewReader(bytes.NewReader(buf.Bytes()), []byte(key), blockSize, []byte(additionalData)) n, err := r.Seek(tc.offset, tc.whence) if err != nil { t.Fatalf("Reader.Seek(%d, %d) failed: %s", tc.offset, tc.whence, err) @@ -126,13 +129,16 @@ func TestReader_Seek(t *testing.T) { func FuzzRoundTrip(f *testing.F) { f.Add(1, []byte("Hello World!")) f.Fuzz(func(t *testing.T, segmentSize int, msg []byte) { - const password = "asdf" + const ( + password = "asdf" + additionalData = "additional data" + ) if segmentSize <= 0 { return } buf := new(bytes.Buffer) - w := NewWriter(buf, []byte(password), segmentSize) + w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData)) n, err := w.Write(msg) if err != nil { t.Fatalf("Writer.Write(%q) failed: %s", msg, err) @@ -143,7 +149,7 @@ func FuzzRoundTrip(f *testing.F) { if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } - got, err := io.ReadAll(NewReader(buf, []byte(password), segmentSize)) + got, err := io.ReadAll(NewReader(buf, []byte(password), segmentSize, []byte(additionalData))) if err != nil { t.Fatalf("Reader.Read failed: %s", err) } @@ -156,13 +162,16 @@ func FuzzRoundTrip(f *testing.F) { func FuzzReadFromWriteTo(f *testing.F) { f.Add(1, "Hello World!") f.Fuzz(func(t *testing.T, segmentSize int, msg string) { - const password = "asdf" + const ( + password = "asdf" + additionalData = "additional data" + ) if segmentSize <= 0 { return } buf := new(bytes.Buffer) - w := NewWriter(buf, []byte(password), segmentSize) + w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData)) n, err := w.ReadFrom(strings.NewReader(msg)) if err != nil { t.Fatalf("Writer.ReadFrom(%q) failed: %s", msg, err) @@ -174,7 +183,7 @@ func FuzzReadFromWriteTo(f *testing.F) { t.Fatalf("Writer.Close failed: %s", err) } got := new(strings.Builder) - n, err = NewReader(buf, []byte(password), segmentSize).WriteTo(got) + n, err = NewReader(buf, []byte(password), segmentSize, []byte(additionalData)).WriteTo(got) if err != nil { t.Fatalf("Reader.WriteTo failed: %s", err) } @@ -191,20 +200,23 @@ func FuzzSeekStart(f *testing.F) { f.Add(1, []byte("Hello World!"), 6) f.Add(30, []byte("000000"), 6) f.Fuzz(func(t *testing.T, segmentSize int, msg []byte, offset int) { - const password = "asdf" + const ( + password = "asdf" + additionalData = "additional data" + ) if segmentSize <= 0 || offset > len(msg) { return } buf := new(bytes.Buffer) - w := NewWriter(buf, []byte(password), segmentSize) + w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData)) if _, err := w.Write(msg); err != nil { t.Fatalf("Writer.Write(%q) failed: %s", msg, err) } if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } - r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize) + r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize, []byte(additionalData)) n, err := r.Seek(int64(offset), io.SeekStart) if offset < 0 { if err == nil { @@ -231,20 +243,23 @@ func FuzzSeekStart(f *testing.F) { func FuzzSeekEnd(f *testing.F) { f.Add(1, []byte("Hello World!"), -6) f.Fuzz(func(t *testing.T, segmentSize int, msg []byte, offset int) { - const password = "asdf" + const ( + password = "asdf" + additionalData = "additionalData" + ) if segmentSize <= 0 || offset > 0 { return } buf := new(bytes.Buffer) - w := NewWriter(buf, []byte(password), segmentSize) + w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData)) if _, err := w.Write(msg); err != nil { t.Fatalf("Writer.Write(%q) failed: %s", msg, err) } if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } - r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize) + r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize, []byte(additionalData)) n, err := r.Seek(int64(offset), io.SeekEnd) if len(msg)+offset < 0 { if err == nil { @@ -273,20 +288,23 @@ func FuzzSeekCurrent(f *testing.F) { f.Add(1, []byte("000000"), 3, 3) f.Add(1, []byte("000"), 3, -3) f.Fuzz(func(t *testing.T, segmentSize int, msg []byte, offset1, offset2 int) { - const password = "asdf" + const ( + password = "asdf" + additionalData = "additional data" + ) if segmentSize <= 0 || offset1 < 0 || offset1 > len(msg) || offset1+offset2 > len(msg) { return } buf := new(bytes.Buffer) - w := NewWriter(buf, []byte(password), segmentSize) + w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData)) if _, err := w.Write(msg); err != nil { t.Fatalf("Writer.Write(%q) failed: %s", msg, err) } if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } - r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize) + r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize, []byte(additionalData)) if _, err := io.ReadFull(r, make([]byte, offset1)); err != nil { t.Fatalf("Reader.Read failed: %s", err) } @@ -314,7 +332,7 @@ func FuzzSeekCurrent(f *testing.F) { } func encrypt(buf *bytes.Buffer, data []byte) error { - w := NewWriter(buf, []byte("asdf"), 4*1024*1024) + w := NewWriter(buf, []byte("asdf"), 4*1024*1024, nil) if _, err := w.Write(data); err != nil { return err } @@ -333,7 +351,7 @@ func BenchmarkWriter(b *testing.B) { } func decrypt(out []byte, data []byte) error { - r := NewReader(bytes.NewReader(data), []byte("asdf"), 4*1024*1024) + r := NewReader(bytes.NewReader(data), []byte("asdf"), 4*1024*1024, nil) _, err := io.ReadFull(r, out) return err } @@ -345,7 +363,7 @@ func BenchmarkReader(b *testing.B) { ) data := make([]byte, 10*1024*1024) encryptedBuf := new(bytes.Buffer) - w := NewWriter(encryptedBuf, []byte(password), segmentSize) + w := NewWriter(encryptedBuf, []byte(password), segmentSize, nil) if _, err := w.Write(data); err != nil { b.Fatal(err) } |
