diff options
| -rw-r--r-- | oae2.go | 28 | ||||
| -rw-r--r-- | oae2_test.go | 2 |
2 files changed, 22 insertions, 8 deletions
@@ -45,10 +45,10 @@ func (e *segmentEncrypter) init(salt []byte) error { return err } -func (e *segmentEncrypter) nextNonce(lastSegment bool) []byte { +func (e *segmentEncrypter) nextNonce(lastSegment bool) ([]byte, error) { for i := 0; ; i++ { if i == nonceSize-1 { - panic("counter overflowed") // impossible, 11 bytes + return nil, errors.New("oae2: nonce counter overflowed") } e.nonce[i]++ if e.nonce[i] != 0 { @@ -58,15 +58,23 @@ func (e *segmentEncrypter) nextNonce(lastSegment bool) []byte { if lastSegment { e.nonce[nonceSize-1] = 1 } - return e.nonce[:] + return e.nonce[:], nil } -func (e *segmentEncrypter) encryptSegment(segment []byte, lastSegment bool) []byte { - return e.aead.Seal(segment[:len(segment)+aeadOverhead][:0], e.nextNonce(lastSegment), segment, nil) +func (e *segmentEncrypter) encryptSegment(segment []byte, lastSegment bool) ([]byte, error) { + nonce, err := e.nextNonce(lastSegment) + if err != nil { + return nil, err + } + return e.aead.Seal(segment[:len(segment)+aeadOverhead][:0], nonce, segment, nil), nil } func (e *segmentEncrypter) decryptSegment(segment []byte, lastSegment bool) ([]byte, error) { - return e.aead.Open(segment[:0], e.nextNonce(lastSegment), segment, nil) + nonce, err := e.nextNonce(lastSegment) + if err != nil { + return nil, err + } + return e.aead.Open(segment[:0], nonce, segment, nil) } // A Writer wraps an io.Writer and encrypts the data in segments. Make sure to @@ -119,7 +127,11 @@ func (w *Writer) init() error { } func (w *Writer) writeBuf(lastSegment bool) error { - if _, w.err = w.w.Write(w.encrypter.encryptSegment(w.buf, lastSegment)); w.err != nil { + var encrypted []byte + if encrypted, w.err = w.encrypter.encryptSegment(w.buf, lastSegment); w.err != nil { + return w.err + } + if _, w.err = w.w.Write(encrypted); w.err != nil { return w.err } w.buf = w.buf[:0] @@ -278,7 +290,7 @@ func (r *Reader) fillBuf() error { if err == io.ErrUnexpectedEOF { r.readLastChunk = true } - if r.buf, err = r.decrypter.decryptSegment(r.buf, r.readLastChunk); err != nil { + if r.buf, r.err = r.decrypter.decryptSegment(r.buf, r.readLastChunk); r.err != nil { return r.err } r.nRead = 0 diff --git a/oae2_test.go b/oae2_test.go index 91194bf..90cf84d 100644 --- a/oae2_test.go +++ b/oae2_test.go @@ -8,6 +8,7 @@ import ( ) func TestRoundTrip(t *testing.T) { + t.Parallel() const ( msg = "Hello World!" key = "password123" @@ -31,6 +32,7 @@ func TestRoundTrip(t *testing.T) { } func TestReadFromWriteTo(t *testing.T) { + t.Parallel() const ( msg = "Hello World!" key = "asdf" |
