From 35076fdc4ecdfdac27d6e48c0c563716507fb6b9 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Mon, 26 Jan 2026 07:09:07 -0800 Subject: Make sure to reject a 32 byte message --- oae2.go | 9 ++++++++- oae2_test.go | 45 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 53 insertions(+), 1 deletion(-) diff --git a/oae2.go b/oae2.go index e2bb1c1..ff59b76 100644 --- a/oae2.go +++ b/oae2.go @@ -306,7 +306,14 @@ func (r *Reader) initialize() error { } // Decrypt the first segment now to make sure we validate the // additional data - return r.fillBuf() + if err := r.fillBuf(); err != nil { + if err == io.EOF { + r.err = io.ErrUnexpectedEOF + return r.err + } + return err + } + return nil } func (r *Reader) init() error { diff --git a/oae2_test.go b/oae2_test.go index 4099aa7..f0ed37f 100644 --- a/oae2_test.go +++ b/oae2_test.go @@ -2,6 +2,7 @@ package oae2 import ( "bytes" + "crypto/rand" "io" "strings" "testing" @@ -57,6 +58,19 @@ func TestReadFromWriteTo(t *testing.T) { } } +func TestInvalid(t *testing.T) { + t.Parallel() + + const ( + msg = "01234567890123456789012345678901" + key = "asdf" + ) + _, err := io.ReadAll(NewReader(strings.NewReader(msg), []byte(key), 1, nil)) + if err == nil { + t.Errorf("Message %q passed validation, want error", msg) + } +} + func TestReader_Seek(t *testing.T) { t.Parallel() const ( @@ -196,6 +210,37 @@ func FuzzReadFromWriteTo(f *testing.F) { }) } +func FuzzReadInvalid(f *testing.F) { + f.Add(1, "", "") + f.Add(1, "01234567890123456789012345678901", "") + f.Fuzz(func(t *testing.T, segmentSize int, msg string, additionalData string) { + if segmentSize <= 0 { + return + } + password := make([]byte, 32) + // Get a random password so that we can be sure msg is invalid + rand.Read(password) + if _, err := io.ReadAll(NewReader(strings.NewReader(msg), password, segmentSize, []byte(additionalData))); err == nil { + t.Errorf("Reader.Read: message passed validation") + } + }) +} + +func FuzzWriteToInvalid(f *testing.F) { + f.Add(1, "", "") + f.Add(1, "01234567890123456789012345678901", "") + f.Fuzz(func(t *testing.T, segmentSize int, msg string, additionalData string) { + if segmentSize <= 0 { + return + } + password := make([]byte, 32) + rand.Read(password) + if _, err := NewReader(strings.NewReader(msg), password, segmentSize, []byte(additionalData)).WriteTo(io.Discard); err == nil { + t.Errorf("Reader.WriteTo: message passed validation") + } + }) +} + func FuzzSeekStart(f *testing.F) { f.Add(1, []byte("Hello World!"), 6) f.Add(30, []byte("000000"), 6) -- cgit v1.3.1