diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2026-01-03 08:22:51 -0800 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2026-01-03 11:17:07 -0800 |
| commit | 2aecb1f4f570849ace353c577436f63a9f4de209 (patch) | |
| tree | 36c0b7af6ae69578ad65a993056c5c2e36e36ca9 /oae2.go | |
| parent | 32efe0766ed06964e72bfe622a5146c9ff324ab0 (diff) | |
| download | oae2-2aecb1f4f570849ace353c577436f63a9f4de209.tar.zst | |
Add support for seeking within a stream
Oh God what have I done
Diffstat (limited to 'oae2.go')
| -rw-r--r-- | oae2.go | 130 |
1 files changed, 109 insertions, 21 deletions
@@ -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") + } +} |
