// Pacakage oae2 implements “Online Authenticated-Encryption” (also known as // streaming AEAD) 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. // // # Encrypted stream format // // The encrypted streams written by this package start with a 32 byte random // salt, followed by a number of encrypted segments each of length // segmentSize + 16. That is, each encrypted segment has an overhead of 16 // bytes for the authentication tag. Each segment is encrypted using // AES-256-GCM, using a key derived from the given key and the random salt // using HKDF-HMAC-SHA256. // // The encrypted stream format may change in the future. // // [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" "encoding/binary" "errors" "fmt" "io" "math" ) 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}, buf: make([]byte, 0, segmentSize+aeadOverhead), } } 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 } _, w.err = w.w.Write(salt) return w.err } 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: writer 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 } // We avoid bufio.Reader to minimize extra allocations. 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 } 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 } // NewReader returns a Reader that wraps r and decrypts the data using key in // chunks of size segmentSize. segmentSize must match the segment size that was // used to write the encrypted stream. 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}, buf: make([]byte, 0, segmentSize+aeadOverhead+1), } } func (r *Reader) initialize() error { r.initialized = true // N.B. there's a subtle bug lurking here: for an input stream of // exactly 32 bytes, it's important that we reject this stream since it // doesn't have an authentication tag. So if we naively read exactly 32 // bytes inside initialize, then Read will see that the underlying // reader is at EOF, and won't be able to distinguish this from a valid // EOF case. Fortunately we can easily work around this by reading one // extra byte here: since the shortest possible encrypted stream is // 32 + 16 bytes, every valid stream will have an extra byte for us // here, and other than on the first segment, Read always knows // precisely whether the stream is at EOF. 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.r.unreadByte() r.err = r.decrypter.init(buf[:saltSize]) return r.err } func (r *Reader) init() error { if !r.initialized { return r.initialize() } return r.err } func (r *Reader) fillBuf() error { // It's important that we know whether this is the last segment, to set // the appropriate byte in the IV. Some other implementations just try // both options, but I think we can be a little bit more precise and // read an extra byte here to be sure if we're at the end. n, err := io.ReadFull(&r.r, r.buf[:r.segmentSize+aeadOverhead+1]) if err != nil && err != io.ErrUnexpectedEOF { // Don't set r.err to EOF here, since we might Seek and reset // the stream, but r.err is unrecoverable. if err != io.EOF { r.err = err } return err } buf := r.buf[:n] if n == r.segmentSize+aeadOverhead+1 { r.r.unreadByte() buf = buf[:r.segmentSize+aeadOverhead] } if r.buf, r.err = r.decrypter.decryptSegment(buf, err == io.ErrUnexpectedEOF); 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 } } } 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, fmt.Errorf("%w", 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") } }