From 2aecb1f4f570849ace353c577436f63a9f4de209 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Sat, 3 Jan 2026 08:22:51 -0800 Subject: Add support for seeking within a stream Oh God what have I done --- oae2.go | 130 ++++++++++++++++++++++++++++++++------- oae2_test.go | 194 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 303 insertions(+), 21 deletions(-) diff --git a/oae2.go b/oae2.go index b0868a0..05bc52d 100644 --- a/oae2.go +++ b/oae2.go @@ -15,8 +15,11 @@ import ( "crypto/hkdf" "crypto/rand" "crypto/sha256" + "encoding/binary" "errors" + "fmt" "io" + "math" ) const ( @@ -222,17 +225,25 @@ func (r *bufReader) unreadByte() { r.buffered = true } +func (r *bufReader) seek(offset int64, whence int) (int64, error) { + n, err := r.r.(io.Seeker).Seek(offset, whence) + if err != nil { + return n, err + } + r.buffered = false + return n, nil +} + // A Reader wraps an io.Reader and decrypts the underlying data stream. type Reader struct { initialized bool err error - r bufReader - segmentSize int - decrypter segmentEncrypter - buf []byte - nRead int - readLastChunk bool + r bufReader + segmentSize int + decrypter segmentEncrypter + buf []byte + nRead int } // NewReader returns a Reader that wraps r and decrypts the data using key in @@ -251,14 +262,15 @@ func NewReader(r io.Reader, key []byte, segmentSize int) *Reader { func (r *Reader) initialize() error { r.initialized = true - salt := make([]byte, saltSize) - if _, r.err = io.ReadFull(&r.r, salt); r.err != nil { + buf := make([]byte, saltSize+1) + if _, r.err = io.ReadFull(&r.r, buf); r.err != nil { if r.err == io.EOF { r.err = io.ErrUnexpectedEOF } return r.err } - r.err = r.decrypter.init(salt) + r.r.unreadByte() + r.err = r.decrypter.init(buf[:saltSize]) return r.err } @@ -270,26 +282,19 @@ func (r *Reader) init() error { } func (r *Reader) fillBuf() error { - if r.readLastChunk { - return io.EOF - } n, err := io.ReadFull(&r.r, r.buf[:r.segmentSize+aeadOverhead+1]) if err != nil && err != io.ErrUnexpectedEOF { - if err == io.EOF { - err = io.ErrUnexpectedEOF + if err != io.EOF { + r.err = err } - r.err = err return err } - r.buf = r.buf[:n] + buf := r.buf[:n] if n == r.segmentSize+aeadOverhead+1 { r.r.unreadByte() - r.buf = r.buf[:r.segmentSize+aeadOverhead] - } - if err == io.ErrUnexpectedEOF { - r.readLastChunk = true + buf = buf[:r.segmentSize+aeadOverhead] } - if r.buf, r.err = r.decrypter.decryptSegment(r.buf, r.readLastChunk); r.err != nil { + if r.buf, r.err = r.decrypter.decryptSegment(buf, err == io.ErrUnexpectedEOF); r.err != nil { return r.err } r.nRead = 0 @@ -337,3 +342,86 @@ func (r *Reader) WriteTo(w io.Writer) (int64, error) { } } } + +func (r *Reader) encryptedToPlaintextSize(offset int64) int64 { + segments := (offset - saltSize) / int64(r.segmentSize+aeadOverhead) + if (offset-saltSize)%int64(r.segmentSize+aeadOverhead) != 0 { + segments++ + } + return offset - saltSize - segments*aeadOverhead +} + +func (r *Reader) seek(offset int64) (int64, error) { + if offset < 0 { + return 0, fmt.Errorf("oae2.Reader.Seek: absolute offset is negative or would overflow") + } + segment := offset / int64(r.segmentSize) + segmentOffset := offset % int64(r.segmentSize) + if segment > math.MaxInt64/int64(r.segmentSize+aeadOverhead) { + return 0, fmt.Errorf("oae2.Reader.Seek: seek offset would overflow int64") + } + streamOffset := saltSize + segment*int64(r.segmentSize+aeadOverhead) + if streamOffset < 0 { + return 0, fmt.Errorf("oae2.Reader.Seek: seek offset would overflow int64") + } + n, err := r.r.seek(streamOffset, io.SeekStart) + if err != nil { + return 0, err + } + r.buf = r.buf[:0] + binary.LittleEndian.PutUint64(r.decrypter.nonce[:], uint64(segment)) + clear(r.decrypter.nonce[8:]) + if err := r.fillBuf(); err != nil && err != io.EOF { + return 0, err + } + r.nRead = min(int(segmentOffset), len(r.buf)) + return r.encryptedToPlaintextSize(n) + int64(r.nRead), nil +} + +// Seek implements io.Seeker. If the underlying reader does not implement +// io.Seeker, Seek returns [errors.ErrUnsupported]. +func (r *Reader) Seek(offset int64, whence int) (int64, error) { + if _, ok := r.r.r.(io.Seeker); !ok { + return 0, errors.ErrUnsupported + } + if err := r.init(); err != nil { + return 0, err + } + switch whence { + case io.SeekStart: + return r.seek(offset) + case io.SeekCurrent: + if r.decrypter.nonce[8]|r.decrypter.nonce[9]|r.decrypter.nonce[10] != 0 { + return 0, fmt.Errorf("oae2.Reader.Seek: internal error: current stream position would overflow int64") + } + currentSegment := binary.LittleEndian.Uint64(r.decrypter.nonce[:8]) + if currentSegment > 0 { + currentSegment-- + } + if currentSegment > math.MaxInt64/uint64(r.segmentSize) { + return 0, fmt.Errorf("oae2.Reader.Seek: internal error: current stream position would overflow int64") + } + currentOffset := int64(currentSegment)*int64(r.segmentSize) + int64(r.nRead) + if currentOffset < 0 { + return 0, fmt.Errorf("oae2.Reader.Seek: internal error: current stream position would overflow int64") + } + return r.seek(currentOffset + offset) + case io.SeekEnd: + // Unfortunately we have to do an additional seek for SeekEnd + // in the general case to determine the total size of the + // underlying stream. Fortunately, the most common case for + // SeekEnd is with offset 0 and we can avoid the additional + // seek in that case. + encryptedSize, err := r.r.seek(0, io.SeekEnd) + if err != nil { + return 0, err + } + plaintextSize := r.encryptedToPlaintextSize(encryptedSize) + if offset == 0 { + return plaintextSize, nil + } + return r.seek(plaintextSize + offset) + default: + return 0, errors.New("oae2.Reader.Seek: invalid whence") + } +} diff --git a/oae2_test.go b/oae2_test.go index 96b83af..e326ef2 100644 --- a/oae2_test.go +++ b/oae2_test.go @@ -55,6 +55,74 @@ func TestReadFromWriteTo(t *testing.T) { } } +func TestReader_Seek(t *testing.T) { + t.Parallel() + const ( + msg = "Hello World!" + key = "asdf" + 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) + 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) + 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) { @@ -118,3 +186,129 @@ func FuzzReadFromWriteTo(f *testing.F) { } }) } + +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" + + if segmentSize <= 0 || offset > len(msg) { + return + } + buf := new(bytes.Buffer) + w := NewWriter(buf, []byte(password), segmentSize) + 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) + 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" + + if segmentSize <= 0 || offset > 0 { + return + } + buf := new(bytes.Buffer) + w := NewWriter(buf, []byte(password), segmentSize) + 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) + 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" + + if segmentSize <= 0 || offset1 < 0 || offset1 > len(msg) || offset1+offset2 > len(msg) { + return + } + buf := new(bytes.Buffer) + w := NewWriter(buf, []byte(password), segmentSize) + 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) + 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:]) + } + }) +} -- cgit v1.3.1