diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2025-12-28 16:17:59 -0800 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2025-12-28 16:17:59 -0800 |
| commit | 838988c1f2cf8aed46ce0c028d32a2d2a22d01fb (patch) | |
| tree | 7a074e6693f18ceaf2822ad66737b05d26955b3a /oae2.go | |
| download | oae2-838988c1f2cf8aed46ce0c028d32a2d2a22d01fb.tar.zst | |
Initial commit
Diffstat (limited to 'oae2.go')
| -rw-r--r-- | oae2.go | 328 |
1 files changed, 328 insertions, 0 deletions
@@ -0,0 +1,328 @@ +// 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 { + for i := 0; ; i++ { + if i == nonceSize-1 { + panic("counter overflowed") // impossible, 11 bytes + } + e.nonce[i]++ + if e.nonce[i] != 0 { + break + } + } + if lastSegment { + e.nonce[nonceSize-1] = 1 + } + return e.nonce[:] +} + +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) decryptSegment(segment []byte, lastSegment bool) ([]byte, error) { + return e.aead.Open(segment[:0], e.nextNonce(lastSegment), 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 + blockSize 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 blockSize bytes. NewWriter +// panics if blockSize <= 0. +// +// Make sure to call [Writer.Close] to flush the final segment. +func NewWriter(w io.Writer, key []byte, blockSize int) *Writer { + if blockSize <= 0 { + panic("oae2.NewWriter: blockSize must be strictly greater than 0") + } + return &Writer{ + w: w, + blockSize: blockSize, + 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.blockSize+aeadOverhead) + return nil +} + +func (w *Writer) init() error { + if !w.initialized { + return w.initialize() + } + return w.err +} + +func (w *Writer) writeBuf(lastSegment bool) error { + if _, w.err = w.w.Write(w.encrypter.encryptSegment(w.buf, lastSegment)); 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.blockSize { + if err := w.writeBuf(false); err != nil { + return nn, err + } + } + n := copy(w.buf[len(w.buf):w.blockSize], 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.blockSize+1]) + w.buf = w.buf[:len(w.buf)+n] + nn += int64(n) + if len(w.buf) == w.blockSize+1 { + nextByte := w.buf[w.blockSize] + w.buf = w.buf[:w.blockSize] + 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 + blockSize 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 blockSize. NewReader panics if blockSize <= 0. +func NewReader(r io.Reader, key []byte, blockSize int) *Reader { + if blockSize <= 0 { + panic("oae2.NewReader: blockSize must be strictly greater than 0") + } + return &Reader{ + r: bufReader{r: r}, + blockSize: blockSize, + 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.blockSize+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.blockSize+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.blockSize+aeadOverhead+1 { + r.r.unreadByte() + r.buf = r.buf[:r.blockSize+aeadOverhead] + } + if err == io.ErrUnexpectedEOF { + r.readLastChunk = true + } + if r.buf, err = r.decrypter.decryptSegment(r.buf, r.readLastChunk); 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 + } + } +} |
