diff options
Diffstat (limited to 'internal/cryptoutil')
| -rw-r--r-- | internal/cryptoutil/oae2.go | 47 |
1 files changed, 26 insertions, 21 deletions
diff --git a/internal/cryptoutil/oae2.go b/internal/cryptoutil/oae2.go index 2310039..4997051 100644 --- a/internal/cryptoutil/oae2.go +++ b/internal/cryptoutil/oae2.go @@ -215,8 +215,8 @@ type DecryptingReader struct { oae2 oae2 segmentSize int - initialized bool - decryptedBuf bytes.Buffer + initialized bool + buf bytes.Buffer } // NewReader returns a new DecryptingWriter that decrypts data from r. The @@ -227,7 +227,7 @@ func (k EncryptionKey) NewReader(r io.Reader, options ...Option) *DecryptingRead o(&opts) } return &DecryptingReader{ - r: bufio.NewReaderSize(r, opts.segmentSize+aeadOverhead+1), + r: bufio.NewReaderSize(r, 1), oae2: oae2{ key: k, additionalData: opts.additionalData, @@ -237,34 +237,39 @@ func (k EncryptionKey) NewReader(r io.Reader, options ...Option) *DecryptingRead } func (r *DecryptingReader) initialize() error { - header, err := r.r.Peek(headerSize) - if err != nil { + header := make([]byte, headerSize) + if _, err := io.ReadFull(r.r, header); err != nil { return err } if err := r.oae2.initialize(header); err != nil { return err } - r.r.Discard(len(header)) - r.decryptedBuf = *bytes.NewBuffer(make([]byte, 0, r.segmentSize)) + r.buf = *bytes.NewBuffer(make([]byte, 0, r.segmentSize+aeadOverhead)) r.initialized = true return nil } func (r *DecryptingReader) fillBuf() error { - encryptedSegmentSize := r.segmentSize + aeadOverhead - // Peek one extra byte to make sure if this is the last segment - block, readErr := r.r.Peek(encryptedSegmentSize + 1) - if len(block) == 0 { - return readErr - } - block = block[:min(len(block), encryptedSegmentSize)] - r.decryptedBuf.Reset() - result, err := r.oae2.decryptBlock(r.decryptedBuf.AvailableBuffer(), block, readErr == io.EOF) - if err != nil { + r.buf.Reset() + buf := r.buf.AvailableBuffer() + buf = buf[:cap(buf)] + n, err := io.ReadFull(r.r, buf) + if n == 0 { + if err == io.ErrUnexpectedEOF { + return io.EOF + } return err } - r.decryptedBuf.Write(result) - r.r.Discard(len(block)) + buf = buf[:n] + if n > 0 { + // Peek one extra byte to check if this is the last segment + _, readErr := r.r.Peek(1) + result, err := r.oae2.decryptBlock(buf[:0], buf, readErr == io.EOF) + if err != nil { + return err + } + r.buf.Write(result) + } return nil } @@ -274,11 +279,11 @@ func (r *DecryptingReader) Read(buf []byte) (int, error) { return 0, err } } - if r.decryptedBuf.Len() == 0 { + if r.buf.Len() == 0 { if err := r.fillBuf(); err != nil { return 0, err } } - n, _ := r.decryptedBuf.Read(buf) + n, _ := r.buf.Read(buf) return n, nil } |
