package oae2 import ( "bytes" "io" "strings" "testing" ) func TestRoundTrip(t *testing.T) { t.Parallel() const ( msg = "Hello World!" key = "password123" additionalData = "additional data" blockSize = 1 ) buf := new(bytes.Buffer) 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, []byte(additionalData))) if err != nil { t.Fatalf("Reader.Read failed: %s", err) } if string(got) != msg { t.Errorf("Message failed to round-trip, got %q, want %q", got, msg) } } func TestReadFromWriteTo(t *testing.T) { t.Parallel() const ( msg = "Hello World!" key = "asdf" additionalData = "additional data" blockSize = 2 ) buf := new(bytes.Buffer) 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) } if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } got := new(strings.Builder) 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 { t.Errorf("Message failed to round-trip, got %q, want %q", got, msg) } } func TestReader_Seek(t *testing.T) { t.Parallel() const ( msg = "Hello World!" key = "asdf" additionalData = "additional data" blockSize = 1 ) for _, tc := range []struct { desc string offset int64 whence int wantOffset int64 wantMsg string }{{ desc: "SeekStart", offset: 6, whence: io.SeekStart, wantOffset: 6, wantMsg: "World!", }, { desc: "SeekEnd", offset: -6, whence: io.SeekEnd, wantOffset: 6, wantMsg: "World!", }, { desc: "SeekEndOffsetZero", offset: 0, whence: io.SeekEnd, wantOffset: 12, wantMsg: "", }, { desc: "SeekCurrent", offset: 6, whence: io.SeekCurrent, wantOffset: 6, wantMsg: "World!", }} { t.Run(tc.desc, func(t *testing.T) { t.Parallel() buf := new(bytes.Buffer) 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, []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) } if n != tc.wantOffset { t.Errorf("Reader.Seek(%d, %d) = %d, want %d", tc.offset, tc.whence, n, tc.wantOffset) } got, err := io.ReadAll(r) if err != nil { t.Fatalf("Reader.Read failed: %s", err) } if string(got) != tc.wantMsg { t.Errorf("Incorrect message after Seek: got %q, want %q", got, tc.wantMsg) } }) } } func FuzzRoundTrip(f *testing.F) { f.Add(1, []byte("Hello World!")) f.Fuzz(func(t *testing.T, segmentSize int, msg []byte) { const ( password = "asdf" additionalData = "additional data" ) if segmentSize <= 0 { return } buf := new(bytes.Buffer) 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) } if n != len(msg) { t.Errorf("Writer.Write(%q) = %d, want %d", msg, n, len(msg)) } if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } got, err := io.ReadAll(NewReader(buf, []byte(password), segmentSize, []byte(additionalData))) if err != nil { t.Fatalf("Reader.Read failed: %s", err) } if !bytes.Equal(got, msg) { t.Errorf("Message failed to round-trip, got %q, want %q", got, msg) } }) } func FuzzReadFromWriteTo(f *testing.F) { f.Add(1, "Hello World!") f.Fuzz(func(t *testing.T, segmentSize int, msg string) { const ( password = "asdf" additionalData = "additional data" ) if segmentSize <= 0 { return } buf := new(bytes.Buffer) 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) } if n != int64(len(msg)) { t.Errorf("Writer.Write(%q) = %d, want %d", msg, n, len(msg)) } if err := w.Close(); err != nil { t.Fatalf("Writer.Close failed: %s", err) } got := new(strings.Builder) n, err = NewReader(buf, []byte(password), segmentSize, []byte(additionalData)).WriteTo(got) if err != nil { t.Fatalf("Reader.WriteTo failed: %s", err) } if n != int64(len(msg)) { t.Errorf("Reader.WriteTo returned %d bytes, want %d", n, len(msg)) } if got.String() != msg { t.Errorf("Message failed to round-trip, got %q, want %q", got, msg) } }) } 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" additionalData = "additional data" ) if segmentSize <= 0 || offset > len(msg) { return } buf := new(bytes.Buffer) 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, []byte(additionalData)) n, err := r.Seek(int64(offset), io.SeekStart) if offset < 0 { if err == nil { t.Errorf("Seek(%d, io.SeekStart) succeeded, want error", offset) } return } if err != nil { t.Fatalf("Seek(%d, io.SeekStart) failed: %s", offset, err) } if n != int64(offset) { t.Errorf("Seek(%d, io.SeekStart) = %d, want %d", offset, n, offset) } got, err := io.ReadAll(r) if err != nil { t.Fatalf("Reader.Read failed: %s", err) } if !bytes.Equal(got, msg[offset:]) { t.Errorf("Message after seek is %q, want %q", got, msg[offset:]) } }) } 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" additionalData = "additionalData" ) if segmentSize <= 0 || offset > 0 { return } buf := new(bytes.Buffer) 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, []byte(additionalData)) n, err := r.Seek(int64(offset), io.SeekEnd) if len(msg)+offset < 0 { if err == nil { t.Errorf("Seek(%d, io.SeekEnd) succeeded, want error", offset) } return } if err != nil { t.Fatalf("Seek(%d, io.SeekEnd) failed: %s", offset, err) } if n != int64(len(msg)+offset) { t.Errorf("Seek(%d, io.SeekEnd) = %d, want %d", offset, n, offset) } got, err := io.ReadAll(r) if err != nil { t.Fatalf("Reader.Read failed: %s", err) } if !bytes.Equal(got, msg[len(msg)+offset:]) { t.Errorf("Message after seek is %q, want %q", got, msg[offset:]) } }) } func FuzzSeekCurrent(f *testing.F) { f.Add(1, []byte("Hello World!"), 3, 3) 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" 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, []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, []byte(additionalData)) if _, err := io.ReadFull(r, make([]byte, offset1)); err != nil { t.Fatalf("Reader.Read failed: %s", err) } n, err := r.Seek(int64(offset2), io.SeekCurrent) if offset1+offset2 < 0 { if err == nil { t.Errorf("Seek(%d, io.SeekCurrent) succeeded, want error", offset2) } return } if err != nil { t.Fatalf("Seek(%d, io.SeekCurrent) failed: %s", offset2, err) } if n != int64(offset1+offset2) { t.Errorf("Seek(%d, io.SeekEnd) = %d, want %d", offset2, n, offset1+offset2) } got, err := io.ReadAll(r) if err != nil { t.Fatalf("Reader.Read failed: %s", err) } if !bytes.Equal(got, msg[offset1+offset2:]) { t.Errorf("Message after seek is %q, want %q", got, msg[offset1+offset2:]) } }) } func encrypt(buf *bytes.Buffer, data []byte) error { w := NewWriter(buf, []byte("asdf"), 4*1024*1024, nil) if _, err := w.Write(data); err != nil { return err } return w.Close() } func BenchmarkWriter(b *testing.B) { data := make([]byte, 10*1024*1024) outBuffer := bytes.NewBuffer(make([]byte, 0, 32+3*16+len(data))) for b.Loop() { outBuffer.Reset() if err := encrypt(outBuffer, data); err != nil { b.Fatal(err) } } } func decrypt(out []byte, data []byte) error { r := NewReader(bytes.NewReader(data), []byte("asdf"), 4*1024*1024, nil) _, err := io.ReadFull(r, out) return err } func BenchmarkReader(b *testing.B) { const ( segmentSize = 4 * 1024 * 1024 password = "asdf" ) data := make([]byte, 10*1024*1024) encryptedBuf := new(bytes.Buffer) w := NewWriter(encryptedBuf, []byte(password), segmentSize, nil) if _, err := w.Write(data); err != nil { b.Fatal(err) } if err := w.Close(); err != nil { b.Fatal(err) } encryptedData := encryptedBuf.Bytes() for b.Loop() { if err := decrypt(data, encryptedData); err != nil { b.Fatal(err) } } }