summaryrefslogtreecommitdiffstats
path: root/internal/cryptoutil
diff options
context:
space:
mode:
Diffstat (limited to 'internal/cryptoutil')
-rw-r--r--internal/cryptoutil/cryptoutil.go107
1 files changed, 107 insertions, 0 deletions
diff --git a/internal/cryptoutil/cryptoutil.go b/internal/cryptoutil/cryptoutil.go
index 4ed600d..09ed996 100644
--- a/internal/cryptoutil/cryptoutil.go
+++ b/internal/cryptoutil/cryptoutil.go
@@ -13,6 +13,10 @@ import (
"crypto/subtle"
"encoding/hex"
"errors"
+ "fmt"
+ "hash"
+ "io"
+ "slices"
)
const (
@@ -169,3 +173,106 @@ func (k EncryptionKey) Decrypt(msg EncryptedMessage) ([]byte, error) {
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)
+}