summaryrefslogtreecommitdiffstats
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/cryptoutil/oae2.go100
1 files changed, 42 insertions, 58 deletions
diff --git a/internal/cryptoutil/oae2.go b/internal/cryptoutil/oae2.go
index 018eb74..2819852 100644
--- a/internal/cryptoutil/oae2.go
+++ b/internal/cryptoutil/oae2.go
@@ -1,6 +1,8 @@
package cryptoutil
import (
+ "bufio"
+ "bytes"
"crypto/aes"
"crypto/cipher"
"crypto/hkdf"
@@ -38,6 +40,7 @@ type oae2 struct {
}
func (o *oae2) initialize(header []byte) error {
+ o.i = 1
copy(o.noncePrefix[:], header[aesKeySize:])
derivedKey, err := hkdf.Key(sha512.New, o.key, header[:aesKeySize], "", aesKeySize)
if err != nil {
@@ -64,22 +67,22 @@ func (o *oae2) nonce(nonce []byte, lastBlock bool) error {
return nil
}
-func (o *oae2) encryptBlock(block []byte, lastBlock bool) ([]byte, error) {
+func (o *oae2) encryptBlock(out, block []byte, lastBlock bool) ([]byte, error) {
nonce := make([]byte, gcmNonceSize)
if err := o.nonce(nonce, lastBlock); err != nil {
return nil, err
}
- encrypted := o.aead.Seal(block[:len(block)+aeadOverhead][:0], nonce, block, o.additionalData)
+ encrypted := o.aead.Seal(out, nonce, block, o.additionalData)
o.additionalData = nil
return encrypted, nil
}
-func (o *oae2) decryptBlock(block []byte, lastBlock bool) ([]byte, error) {
+func (o *oae2) decryptBlock(out, block []byte, lastBlock bool) ([]byte, error) {
nonce := make([]byte, gcmNonceSize)
if err := o.nonce(nonce, lastBlock); err != nil {
return nil, err
}
- decrypted, err := o.aead.Open(block[:0], nonce, block, o.additionalData)
+ decrypted, err := o.aead.Open(out, nonce, block, o.additionalData)
o.additionalData = nil
return decrypted, err
}
@@ -94,9 +97,7 @@ type EncryptingWriter struct {
initialized bool
err error
-
- bufN int
- buf [encryptedSegmentSize]byte
+ buf []byte
}
// NewWriter returns a new EncryptingWriter that writes to w. The additionalData
@@ -108,13 +109,13 @@ func (k EncryptionKey) NewWriter(w io.Writer, additionalData []byte) *Encrypting
oae2: oae2{
key: k,
additionalData: additionalData,
- i: 1,
},
}
}
func (w *EncryptingWriter) initialize() error {
w.initialized = true
+ w.buf = make([]byte, 0, encryptedSegmentSize)
header := make([]byte, headerSize)
rand.Read(header)
if w.err = w.oae2.initialize(header); w.err != nil {
@@ -133,22 +134,22 @@ func (w *EncryptingWriter) Write(buf []byte) (int, error) {
if w.err != nil {
return 0, w.err
}
+ r := bytes.NewReader(buf)
nn := 0
- for len(buf) > 0 {
- if w.bufN == segmentSize {
+ for r.Len() > 0 {
+ if len(w.buf) == segmentSize {
var encrypted []byte
- if encrypted, w.err = w.oae2.encryptBlock(w.buf[:w.bufN], false); w.err != nil {
+ if encrypted, w.err = w.oae2.encryptBlock(w.buf[:0], w.buf, false); w.err != nil {
return nn, w.err
}
if _, w.err = w.w.Write(encrypted); w.err != nil {
return nn, w.err
}
- w.bufN = 0
+ w.buf = w.buf[:0]
}
- n := copy(w.buf[w.bufN:segmentSize], buf)
- w.bufN += n
+ n, _ := r.Read(w.buf[len(w.buf):segmentSize])
+ w.buf = w.buf[:len(w.buf)+n]
nn += n
- buf = buf[n:]
}
return nn, nil
}
@@ -170,7 +171,7 @@ func (w *EncryptingWriter) Close() error {
return w.err
}
var encrypted []byte
- if encrypted, w.err = w.oae2.encryptBlock(w.buf[:w.bufN], true); w.err != nil {
+ if encrypted, w.err = w.oae2.encryptBlock(w.buf[:0], w.buf, true); w.err != nil {
return w.err
}
if _, w.err = w.w.Write(encrypted); w.err != nil {
@@ -182,66 +183,53 @@ func (w *EncryptingWriter) Close() error {
// A DecryptingReader decrypts data using the STREAM construction.
type DecryptingReader struct {
- r io.Reader
+ r *bufio.Reader
oae2 oae2
- initialized bool
- err error
-
- bufRead, bufN int
- peekedByte bool
- buf [encryptedSegmentSize + 1]byte
+ initialized bool
+ decryptedBuf bytes.Buffer
}
// NewReader returns a new DecryptingWriter that decrypts data from r. The
// additionaData must be the same that was provided when encrypting.
func (k EncryptionKey) NewReader(r io.Reader, additionalData []byte) *DecryptingReader {
return &DecryptingReader{
- r: r,
+ r: bufio.NewReaderSize(r, encryptedSegmentSize+1),
oae2: oae2{
key: k,
additionalData: additionalData,
- i: 1,
},
}
}
func (r *DecryptingReader) initialize() error {
- r.initialized = true
- header := make([]byte, headerSize)
- if _, r.err = io.ReadFull(r.r, header); r.err != nil {
- return r.err
+ header, err := r.r.Peek(headerSize)
+ if err != nil {
+ return err
}
- r.err = r.oae2.initialize(header)
- return r.err
+ if err := r.oae2.initialize(header); err != nil {
+ return err
+ }
+ r.r.Discard(len(header))
+ r.decryptedBuf = *bytes.NewBuffer(make([]byte, 0, segmentSize))
+ r.initialized = true
+ return nil
}
func (r *DecryptingReader) fillBuf() error {
- n := 0
- if r.peekedByte {
- n = 1
- r.buf[0] = r.buf[encryptedSegmentSize]
- r.peekedByte = false
+ // 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
}
- for n < len(r.buf) && r.err == nil {
- var m int
- m, r.err = r.r.Read(r.buf[n:])
- n += m
- }
- if n == 0 {
- return r.err
- }
- if n == encryptedSegmentSize+1 {
- r.peekedByte = true
- n = encryptedSegmentSize
- }
- result, err := r.oae2.decryptBlock(r.buf[:n], r.err == io.EOF)
+ 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.err = err
return err
}
- r.bufRead = 0
- r.bufN = len(result)
+ r.decryptedBuf.Write(result)
+ r.r.Discard(len(block))
return nil
}
@@ -251,15 +239,11 @@ func (r *DecryptingReader) Read(buf []byte) (int, error) {
return 0, err
}
}
- if r.bufRead == r.bufN {
+ if r.decryptedBuf.Len() == 0 {
if err := r.fillBuf(); err != nil {
return 0, err
}
}
- n := copy(buf, r.buf[r.bufRead:r.bufN])
- r.bufRead += n
- if r.bufRead == r.bufN {
- return n, r.err
- }
+ n, _ := r.decryptedBuf.Read(buf)
return n, nil
}