aboutsummaryrefslogtreecommitdiffstats
path: root/oae2.go
diff options
context:
space:
mode:
Diffstat (limited to 'oae2.go')
-rw-r--r--oae2.go328
1 files changed, 328 insertions, 0 deletions
diff --git a/oae2.go b/oae2.go
new file mode 100644
index 0000000..7b94825
--- /dev/null
+++ b/oae2.go
@@ -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
+ }
+ }
+}