From 6412b3912b17cd22ba8bdbf07a12f8b005bf058f Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Sun, 25 Jan 2026 07:47:19 -0800 Subject: Add support for additional data MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit We can't really call it "streaming AEAD" and then not support additional data 🤦‍♀️ --- oae2_test.go | 84 ++++++++++++++++++++++++++++++++++++------------------------ 1 file changed, 51 insertions(+), 33 deletions(-) (limited to 'oae2_test.go') 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) } -- cgit v1.3.1