aboutsummaryrefslogtreecommitdiffstats
path: root/oae2.go
diff options
context:
space:
mode:
Diffstat (limited to 'oae2.go')
-rw-r--r--oae2.go130
1 files changed, 109 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")
+ }
+}