aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--oae2.go130
-rw-r--r--oae2_test.go194
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:])
+ }
+ })
+}