// Pacakage oae2 implements “Online Authenticated-Encryption” based on the // influential paper // “[Online Authenticated-Encryption and its Nonce-Reuse Misuse-Resistance]” by // Hoang et al. // // Please do not use this package; it likely has critical // security vulnerabilities. // // [Online Authenticated-Encryption and its Nonce-Reuse Misuse-Resistance]: https://eprint.iacr.org/2015/189.pdf package oae2 import ( "crypto/aes" "crypto/cipher" "crypto/hkdf" "crypto/rand" "crypto/sha256" "errors" "io" ) const ( nonceSize = 12 aeadOverhead = 16 saltSize = 32 ) type segmentEncrypter struct { key []byte aead cipher.AEAD nonce [nonceSize]byte } func (e *segmentEncrypter) init(salt []byte) error { key, err := hkdf.Key(sha256.New, e.key, salt, "", 32) if err != nil { return err } block, err := aes.NewCipher(key) if err != nil { return err } e.aead, err = cipher.NewGCM(block) return err } func (e *segmentEncrypter) nextNonce(lastSegment bool) ([]byte, error) { for i := 0; ; i++ { if i == nonceSize-1 { return nil, errors.New("oae2: nonce counter overflowed") } e.nonce[i]++ if e.nonce[i] != 0 { break } } if lastSegment { e.nonce[nonceSize-1] = 1 } return e.nonce[:], 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) { 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 // call [Writer.Close] once all the data is written. type Writer struct { initialized bool err error w io.Writer segmentSize int encrypter segmentEncrypter buf []byte } // NewWriter returns a Writer that encrypts data using key and writes the // encrypted data to w in chunks of approximately segmentSize bytes. NewWriter // panics if segmentSize <= 0. // // Make sure to call [Writer.Close] to flush the final segment. func NewWriter(w io.Writer, key []byte, segmentSize int) *Writer { if segmentSize <= 0 { panic("oae2.NewWriter: segmentSize must be strictly greater than 0") } return &Writer{ w: w, segmentSize: segmentSize, encrypter: segmentEncrypter{key: key}, } } func (w *Writer) initialize() error { w.initialized = true salt := make([]byte, saltSize) rand.Read(salt) if w.err = w.encrypter.init(salt); w.err != nil { return w.err } if _, w.err = w.w.Write(salt); w.err != nil { return w.err } w.buf = make([]byte, 0, w.segmentSize+aeadOverhead) return nil } func (w *Writer) init() error { if !w.initialized { return w.initialize() } return w.err } func (w *Writer) writeBuf(lastSegment bool) error { 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] return nil } // Write implements io.Writer. func (w *Writer) Write(buf []byte) (int, error) { if err := w.init(); err != nil { return 0, err } nn := 0 for len(buf) > 0 { if len(w.buf) == w.segmentSize { if err := w.writeBuf(false); err != nil { return nn, err } } n := copy(w.buf[len(w.buf):w.segmentSize], buf) w.buf = w.buf[:len(w.buf)+n] buf = buf[n:] nn += n } return nn, nil } // ReadFrom implements io.ReaderFrom. func (w *Writer) ReadFrom(r io.Reader) (int64, error) { if err := w.init(); err != nil { return 0, err } var nn int64 for { n, err := r.Read(w.buf[len(w.buf) : w.segmentSize+1]) w.buf = w.buf[:len(w.buf)+n] nn += int64(n) if len(w.buf) == w.segmentSize+1 { nextByte := w.buf[w.segmentSize] w.buf = w.buf[:w.segmentSize] if err := w.writeBuf(false); err != nil { return nn, err } w.buf = w.buf[:1] w.buf[0] = nextByte } if err != nil { if err == io.EOF { return nn, nil } return nn, err } } } var errClosed = errors.New("oae2.Writer.Close: already closed") // Close flushes the final segment. Failure to call Close will result in a // truncated stream. func (w *Writer) Close() error { if err := w.init(); err != nil { return err } if err := w.writeBuf(true); err != nil { return err } w.err = errClosed return nil } type bufReader struct { r io.Reader nextByte byte buffered bool } func (r *bufReader) Read(buf []byte) (int, error) { if r.buffered && len(buf) > 0 { buf[0] = r.nextByte r.buffered = false return 1, nil } n, err := r.r.Read(buf) if n > 0 { r.nextByte = buf[n-1] } return n, err } func (r *bufReader) unreadByte() { r.buffered = true } // 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 } // NewReader returns a Reader that wraps r and decrypts the data using key in // chunks of size segmentSize. NewReader panics if segmentSize <= 0. func NewReader(r io.Reader, key []byte, segmentSize int) *Reader { if segmentSize <= 0 { panic("oae2.NewReader: segmentSize must be strictly greater than 0") } return &Reader{ r: bufReader{r: r}, segmentSize: segmentSize, decrypter: segmentEncrypter{key: key}, } } func (r *Reader) initialize() error { r.initialized = true salt := make([]byte, saltSize) if _, r.err = io.ReadFull(&r.r, salt); r.err != nil { if r.err == io.EOF { r.err = io.ErrUnexpectedEOF } return r.err } if r.err = r.decrypter.init(salt); r.err != nil { return r.err } r.buf = make([]byte, 0, r.segmentSize+aeadOverhead+1) return nil } func (r *Reader) init() error { if !r.initialized { return r.initialize() } return r.err } func (r *Reader) fillBuf() error { n, err := io.ReadFull(&r.r, r.buf[:r.segmentSize+aeadOverhead+1]) if err != nil && err != io.ErrUnexpectedEOF { if err == io.EOF && !r.readLastChunk { return io.ErrUnexpectedEOF } r.err = err return err } r.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 } if r.buf, r.err = r.decrypter.decryptSegment(r.buf, r.readLastChunk); r.err != nil { return r.err } r.nRead = 0 return nil } // Read implements io.Reader. func (r *Reader) Read(buf []byte) (int, error) { if err := r.init(); err != nil { return 0, err } if r.nRead == len(r.buf) { if err := r.fillBuf(); err != nil { return 0, err } if len(r.buf) == 0 { return 0, io.EOF } } n := copy(buf, r.buf[r.nRead:]) r.nRead += n return n, nil } // WriteTo implements io.WriterTo. func (r *Reader) WriteTo(w io.Writer) (int64, error) { if err := r.init(); err != nil { return 0, err } var nn int64 for { if r.nRead < len(r.buf) { n, err := w.Write(r.buf[r.nRead:]) r.nRead += n nn += int64(n) if err != nil { return nn, err } } if err := r.fillBuf(); err != nil { if err == io.EOF { return nn, nil } return nn, err } } }