summaryrefslogtreecommitdiffstats
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/cryptoutil/cryptoutil.go151
-rw-r--r--internal/cryptoutil/cryptoutil_test.go33
-rw-r--r--internal/cryptoutil/oae2.go265
3 files changed, 287 insertions, 162 deletions
diff --git a/internal/cryptoutil/cryptoutil.go b/internal/cryptoutil/cryptoutil.go
index 09ed996..b269803 100644
--- a/internal/cryptoutil/cryptoutil.go
+++ b/internal/cryptoutil/cryptoutil.go
@@ -3,8 +3,6 @@
package cryptoutil
import (
- "crypto/aes"
- "crypto/cipher"
"crypto/hkdf"
"crypto/hmac"
"crypto/pbkdf2"
@@ -13,10 +11,6 @@ import (
"crypto/subtle"
"encoding/hex"
"errors"
- "fmt"
- "hash"
- "io"
- "slices"
)
const (
@@ -25,8 +19,6 @@ const (
oneSaltSize = 64
hashSize = 64
certSize = sha512.Size
- aesKeySize = 32
- nonceSize = aes.BlockSize
)
// SaltSize is the expected salt length for HashIter.
@@ -59,13 +51,6 @@ type HMACKey []byte
// with HMAC-SHA512.
type SignedMessage []byte
-// An EncryptionKey must be EncryptionKeySize bytes.
-const EncryptionKeySize = aesKeySize + HMACKeySize
-
-// An EncryptionKey is used for encrypting and decrypting data. An EncryptionKey
-// must be EncryptionKeySize bytes.
-type EncryptionKey []byte
-
// An EncryptedMessage is encrypted with AES-256-CTR-HMAC-SHA512.
type EncryptedMessage []byte
@@ -140,139 +125,5 @@ func (k HMACKey) Verify(msg SignedMessage) ([]byte, bool) {
// EncryptionKey derives an EncryptionKey.
func (k RawKey) EncryptionKey() (EncryptionKey, error) {
- return hkdf.Expand(sha512.New, k.key, "encrypt", EncryptionKeySize)
-}
-
-// Encrypt encrypts a message with AES-256-CTR-HMAC-SHA512.
-func (k EncryptionKey) Encrypt(msg []byte) (EncryptedMessage, error) {
- aesKey, hmacKey := k[:aesKeySize], HMACKey(k[aesKeySize:])
- block, err := aes.NewCipher(aesKey)
- if err != nil {
- return nil, err
- }
- cipherText := make([]byte, len(msg)+nonceSize, len(msg)+nonceSize+certSize)
- nonce := cipherText[len(msg) : len(msg)+nonceSize]
- rand.Read(nonce)
- cipher.NewCTR(block, nonce).XORKeyStream(cipherText, msg)
- return EncryptedMessage(hmacKey.Sign(cipherText)), nil
-}
-
-// Decrypt verifies and decrypts an encrypted message. It overwrites msg with
-// the resulting plain text.
-func (k EncryptionKey) Decrypt(msg EncryptedMessage) ([]byte, error) {
- aesKey, hmacKey := k[:aesKeySize], HMACKey(k[aesKeySize:])
- msg, ok := hmacKey.Verify(SignedMessage(msg))
- if !ok {
- return nil, errors.New("bad signature")
- }
- block, err := aes.NewCipher(aesKey)
- if err != nil {
- return nil, err
- }
- msg, nonce := msg[:len(msg)-nonceSize], msg[len(msg)-nonceSize:]
- cipher.NewCTR(block, nonce).XORKeyStream(msg, msg)
- return msg, nil
-}
-
-// An EncryptingWriter encrypts the output and writes it to the underlying
-// writer. It's very important to call .Flush() to write the MAC after the data
-// has been written.
-type EncryptingWriter struct {
- w io.Writer
- stream cipher.Stream
- mac hash.Hash
- buf []byte
-}
-
-// Writer returns a new writer that encrypts its output. Don't forget to call
-// .Flush() to write the MAC.
-func (k EncryptionKey) Writer(w io.Writer) (*EncryptingWriter, error) {
- aesKey, hmacKey := k[:aesKeySize], k[aesKeySize:]
- block, err := aes.NewCipher(aesKey)
- if err != nil {
- return nil, err
- }
- mac := hmac.New(sha512.New, hmacKey)
- nonce := make([]byte, nonceSize)
- rand.Read(nonce)
- mac.Write(nonce)
- if _, err := w.Write(nonce); err != nil {
- return nil, err
- }
- return &EncryptingWriter{
- w: w,
- stream: cipher.NewCTR(block, nonce),
- mac: mac,
- }, nil
-}
-
-// Write encrypts buf and writes it to the underlying writer.
-func (w *EncryptingWriter) Write(buf []byte) (int, error) {
- if len(w.buf) < len(buf) {
- w.buf = slices.Grow(w.buf, len(buf)-len(w.buf))
- }
- w.buf = w.buf[:len(buf)]
- // Encrypt-then-MAC
- w.stream.XORKeyStream(w.buf, buf)
- w.mac.Write(w.buf)
- return w.w.Write(w.buf)
-}
-
-// Flush writes the MAC for the encrypted message. Flush must be called after
-// all data has been written.
-func (w *EncryptingWriter) Flush() error {
- _, err := w.w.Write(w.mac.Sum(nil))
- return err
-}
-
-// A DecryptingReader decrypts data from an underlying reader.
-type DecryptingReader struct {
- r cipher.StreamReader
-}
-
-// Reader returns a new DecryptingReader. The data is processed in two passes,
-// first to verify the MAC, then the io.Seeker interface is used to reset the
-// reader for decryption.
-func (k EncryptionKey) Reader(r io.ReadSeeker) (*DecryptingReader, error) {
- aesKey, hmacKey := k[:aesKeySize], k[aesKeySize:]
- block, err := aes.NewCipher(aesKey)
- if err != nil {
- return nil, err
- }
- totalLen, err := r.Seek(0, io.SeekEnd)
- if err != nil {
- return nil, err
- }
- if totalLen < certSize {
- return nil, errors.New("file too short")
- }
- dataLen := totalLen - certSize
- if _, err := r.Seek(0, io.SeekStart); err != nil {
- return nil, err
- }
- mac := hmac.New(sha512.New, hmacKey)
- if _, err := io.Copy(mac, &io.LimitedReader{R: r, N: dataLen}); err != nil {
- return nil, err
- }
- expectedMAC := mac.Sum(nil)
- sig := make([]byte, certSize)
- if _, err := io.ReadFull(r, sig); err != nil {
- return nil, err
- }
- if !hmac.Equal(sig, expectedMAC) {
- return nil, fmt.Errorf("invalid mac (got %x, want %x)", sig, expectedMAC)
- }
- if _, err := r.Seek(0, io.SeekStart); err != nil {
- return nil, err
- }
- nonce := make([]byte, nonceSize)
- if _, err := io.ReadFull(r, nonce); err != nil {
- return nil, err
- }
- return &DecryptingReader{cipher.StreamReader{S: cipher.NewCTR(block, nonce), R: &io.LimitedReader{R: r, N: dataLen - nonceSize}}}, nil
-}
-
-// Read reads and decrypts data from the underlying reader.
-func (r *DecryptingReader) Read(buf []byte) (int, error) {
- return r.r.Read(buf)
+ return hkdf.Expand(sha512.New, k.key, "encrypt", sha512.Size)
}
diff --git a/internal/cryptoutil/cryptoutil_test.go b/internal/cryptoutil/cryptoutil_test.go
index 79a53a9..dc9074d 100644
--- a/internal/cryptoutil/cryptoutil_test.go
+++ b/internal/cryptoutil/cryptoutil_test.go
@@ -3,6 +3,7 @@ package cryptoutil
import (
"bytes"
"encoding/hex"
+ "io"
"testing"
)
@@ -43,16 +44,20 @@ func TestSignature(t *testing.T) {
func TestEncrypt(t *testing.T) {
key := EncryptionKey(mustHex(t, "b6de26860e0a39aa134732e58055c06ba028675453f73a6912bc932e58d5743d0e423217f06487d3e96a59980301fcc97dfcc4c6b19765f2947de3c33ab7ef9e"))
msg := []byte("test message")
- encryptedMsg, err := key.Encrypt(msg)
- if err != nil {
- t.Fatalf("Encrypt(%q) failed: %s", msg, err)
+ encryptedMsg := new(bytes.Buffer)
+ w := key.NewWriter(encryptedMsg, nil)
+ if _, err := w.Write(msg); err != nil {
+ t.Fatalf("EncryptingWriter.Write(%q) failed: %s", msg, err)
+ }
+ if err := w.Close(); err != nil {
+ t.Fatalf("EncryptingWriter.Close() failed: %s", err)
}
- got, err := key.Decrypt(encryptedMsg)
+ got, err := io.ReadAll(key.NewReader(bytes.NewReader(encryptedMsg.Bytes()), nil))
if err != nil {
- t.Fatalf("Decrypt(%x) failed: %s", encryptedMsg, err)
+ t.Fatalf("DecryptingReader.Read(%x) failed: %s", encryptedMsg, err)
}
if !bytes.Equal(got, msg) {
- t.Errorf("Decrypt(%x) = %x, want %x", encryptedMsg, got, msg)
+ t.Errorf("DecryptingReader.Read(%x) = %x, want %x", encryptedMsg, got, msg)
}
}
@@ -71,15 +76,19 @@ func TestDeriveKey(t *testing.T) {
t.Fatalf("EncryptionKey() failed: %s", err)
}
msg := []byte("test message")
- encryptedMsg, err := key.Encrypt(msg)
- if err != nil {
- t.Fatalf("Encrypt(%q) failed: %s", msg, err)
+ encryptedMsg := new(bytes.Buffer)
+ w := key.NewWriter(encryptedMsg, nil)
+ if _, err := w.Write(msg); err != nil {
+ t.Fatalf("EncryptingWriter.Write(%q) failed: %s", msg, err)
+ }
+ if err := w.Close(); err != nil {
+ t.Fatalf("EncryptingWriter.Close() failed: %s", err)
}
- got, err := key.Decrypt(encryptedMsg)
+ got, err := io.ReadAll(key.NewReader(encryptedMsg, nil))
if err != nil {
- t.Fatalf("Decrypt(%x) failed: %s", encryptedMsg, err)
+ t.Fatalf("DecryptingReader.Read(%x) failed: %s", encryptedMsg, err)
}
if !bytes.Equal(got, msg) {
- t.Errorf("Decrypt(%x) = %x, want %x", encryptedMsg, got, msg)
+ t.Errorf("DecryptingReader.Read(%x) = %x, want %x", encryptedMsg, got, msg)
}
}
diff --git a/internal/cryptoutil/oae2.go b/internal/cryptoutil/oae2.go
new file mode 100644
index 0000000..28ec9d7
--- /dev/null
+++ b/internal/cryptoutil/oae2.go
@@ -0,0 +1,265 @@
+package cryptoutil
+
+import (
+ "crypto/aes"
+ "crypto/cipher"
+ "crypto/hkdf"
+ "crypto/rand"
+ "crypto/sha512"
+ "encoding/binary"
+ "errors"
+ "io"
+)
+
+// Online Authenticated Encryption from https://eprint.iacr.org/2015/189.pdf
+
+const (
+ aeadOverhead = 16
+ aesKeySize = 32
+ noncePrefixSize = 3
+ gcmNonceSize = 12
+ headerSize = aesKeySize + noncePrefixSize
+ cacheSize = 192 * 1024
+ encryptedSegmentSize = cacheSize - 1
+ segmentSize = encryptedSegmentSize - aeadOverhead
+)
+
+// An EncryptionKey is used for encrypting and decrypting data. A key can be any
+// length, but at least 64 bytes would be recommended.
+type EncryptionKey []byte
+
+func (k EncryptionKey) deriveAESKey(salt []byte) (cipher.AEAD, error) {
+ derivedKey, err := hkdf.Key(sha512.New, k, salt, "", aesKeySize)
+ if err != nil {
+ return nil, err
+ }
+ block, err := aes.NewCipher(derivedKey)
+ if err != nil {
+ return nil, err
+ }
+ return cipher.NewGCM(block)
+}
+
+func makeNonce(nonce []byte, noncePrefix []byte, i *uint64) error {
+ copy(nonce, noncePrefix)
+ binary.BigEndian.PutUint64(nonce[noncePrefixSize:], *i)
+ *i++
+ if *i == 0 {
+ return errors.New("counter overflowed (64 bits??)")
+ }
+ return nil
+}
+
+// An EncryptingWriter encrypts data in segments using the STREAM construction
+// described in https://eprint.iacr.org/2015/189.pdf. The writer buffers data up
+// to the segment size, so it's important to call Close to flush the
+// final segment.
+type EncryptingWriter struct {
+ w io.Writer
+ key EncryptionKey
+ aead cipher.AEAD
+ additionalData []byte
+ noncePrefix [noncePrefixSize]byte
+ i uint64
+ initialized bool
+ err error
+ bufN int
+ buf [encryptedSegmentSize]byte
+}
+
+// NewWriter returns a new EncryptingWriter that writes to w. The additionalData
+// will be authenticated with the first segment, but is not written to w. The
+// same additional data must be provided when decrypting.
+func (k EncryptionKey) NewWriter(w io.Writer, additionalData []byte) *EncryptingWriter {
+ return &EncryptingWriter{
+ w: w,
+ key: k,
+ additionalData: additionalData,
+ i: 1,
+ }
+}
+
+func (w *EncryptingWriter) initialize() error {
+ w.initialized = true
+ header := make([]byte, headerSize)
+ rand.Read(header)
+ copy(w.noncePrefix[:], header[aesKeySize:])
+ if w.aead, w.err = w.key.deriveAESKey(header[:aesKeySize]); w.err != nil {
+ return w.err
+ }
+ _, w.err = w.w.Write(header)
+ return w.err
+}
+
+func (w *EncryptingWriter) nonce(nonce []byte) error {
+ w.err = makeNonce(nonce, w.noncePrefix[:], &w.i)
+ return w.err
+}
+
+func (w *EncryptingWriter) writeBuf() error {
+ nonce := make([]byte, gcmNonceSize)
+ if err := w.nonce(nonce); err != nil {
+ return err
+ }
+ if _, w.err = w.w.Write(w.aead.Seal(w.buf[:0], nonce, w.buf[:w.bufN], w.additionalData)); w.err != nil {
+ return w.err
+ }
+ w.bufN = 0
+ w.additionalData = nil
+ return nil
+}
+
+func (w *EncryptingWriter) Write(buf []byte) (int, error) {
+ if !w.initialized {
+ if err := w.initialize(); err != nil {
+ return 0, err
+ }
+ }
+ if w.err != nil {
+ return 0, w.err
+ }
+ nn := 0
+ for len(buf) > 0 {
+ if w.bufN == segmentSize {
+ if err := w.writeBuf(); err != nil {
+ return nn, err
+ }
+ }
+ n := copy(w.buf[w.bufN:segmentSize], buf)
+ w.bufN += n
+ nn += n
+ buf = buf[n:]
+ }
+ return nn, nil
+}
+
+var errClosed = errors.New("closed")
+
+// Close encrypts writes the final segment. It does not close the
+// underlying writer.
+func (w *EncryptingWriter) Close() error {
+ if !w.initialized {
+ if err := w.initialize(); err != nil {
+ return err
+ }
+ }
+ if w.err == errClosed {
+ return nil
+ }
+ if w.err != nil {
+ return w.err
+ }
+ if w.bufN > 0 {
+ nonce := make([]byte, gcmNonceSize)
+ if err := w.nonce(nonce); err != nil {
+ return err
+ }
+ nonce[gcmNonceSize-1] = 1
+ if _, w.err = w.w.Write(w.aead.Seal(w.buf[:0], nonce, w.buf[:w.bufN], w.additionalData)); w.err != nil {
+ return w.err
+ }
+ }
+ w.err = errClosed
+ return nil
+}
+
+// A DecryptingReader decrypts data using the STREAM construction.
+type DecryptingReader struct {
+ r io.Reader
+ key EncryptionKey
+ aead cipher.AEAD
+ additionalData []byte
+ noncePrefix [noncePrefixSize]byte
+ i uint64
+ initialized bool
+ err error
+ bufRead, bufN int
+ peekedByte bool
+ buf [encryptedSegmentSize + 1]byte
+}
+
+// NewReader returns a new DecryptingWriter that decrypts data from r. The
+// additionaData must be the same that was provided when encrypting.
+func (k EncryptionKey) NewReader(r io.Reader, additionalData []byte) *DecryptingReader {
+ return &DecryptingReader{
+ r: r,
+ key: k,
+ additionalData: additionalData,
+ i: 1,
+ }
+}
+
+func (r *DecryptingReader) initialize() error {
+ r.initialized = true
+ header := make([]byte, headerSize)
+ if _, r.err = io.ReadFull(r.r, header); r.err != nil {
+ return r.err
+ }
+ copy(r.noncePrefix[:], header[aesKeySize:])
+ r.aead, r.err = r.key.deriveAESKey(header[:aesKeySize])
+ return r.err
+}
+
+func (r *DecryptingReader) nonce(nonce []byte) error {
+ if err := makeNonce(nonce, r.noncePrefix[:], &r.i); err != nil {
+ r.err = err
+ return err
+ }
+ return nil
+}
+
+func (r *DecryptingReader) fillBuf() error {
+ n := 0
+ if r.peekedByte {
+ n = 1
+ r.buf[0] = r.buf[encryptedSegmentSize]
+ r.peekedByte = false
+ }
+ for n < len(r.buf) && r.err == nil {
+ var m int
+ m, r.err = r.r.Read(r.buf[n:])
+ n += m
+ }
+ if n == 0 {
+ return r.err
+ }
+ if n == encryptedSegmentSize+1 {
+ r.peekedByte = true
+ n = encryptedSegmentSize
+ }
+ nonce := make([]byte, gcmNonceSize)
+ if err := r.nonce(nonce); err != nil {
+ return err
+ }
+ if r.err == io.EOF {
+ nonce[gcmNonceSize-1] = 1
+ }
+ result, err := r.aead.Open(r.buf[:0], nonce, r.buf[:n], r.additionalData)
+ if err != nil {
+ r.err = err
+ return err
+ }
+ r.bufRead = 0
+ r.bufN = len(result)
+ r.additionalData = nil
+ return nil
+}
+
+func (r *DecryptingReader) Read(buf []byte) (int, error) {
+ if !r.initialized {
+ if err := r.initialize(); err != nil {
+ return 0, err
+ }
+ }
+ if r.bufRead == r.bufN {
+ if err := r.fillBuf(); err != nil {
+ return 0, err
+ }
+ }
+ n := copy(buf, r.buf[r.bufRead:r.bufN])
+ r.bufRead += n
+ if r.bufRead == r.bufN {
+ return n, r.err
+ }
+ return n, nil
+}