diff options
| -rw-r--r-- | internal/cryptoutil/cryptoutil.go | 151 | ||||
| -rw-r--r-- | internal/cryptoutil/cryptoutil_test.go | 33 | ||||
| -rw-r--r-- | internal/cryptoutil/oae2.go | 265 | ||||
| -rw-r--r-- | roseh.moe.go | 58 | ||||
| -rw-r--r-- | tools/notes/notes.go | 34 |
5 files changed, 325 insertions, 216 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 +} diff --git a/roseh.moe.go b/roseh.moe.go index e6db419..f9c5478 100644 --- a/roseh.moe.go +++ b/roseh.moe.go @@ -254,15 +254,16 @@ func login(w http.ResponseWriter, r *http.Request) { } func readNotepad(key cryptoutil.EncryptionKey) (string, error) { - encrypted, err := os.ReadFile(*notepadDir + "/notepad") + f, err := os.Open(*notepadDir + "/notepad") if err != nil { return "", err } - decrypted, err := key.Decrypt(encrypted) + defer f.Close() + notepad, err := io.ReadAll(key.NewReader(f, []byte("notepad"))) if err != nil { return "", err } - return string(decrypted), nil + return string(notepad), nil } var ( @@ -315,16 +316,17 @@ func saveNote(w http.ResponseWriter, r *http.Request) error { if key == nil { return errors.New("not logged in") } - encrypted, err := key.Encrypt([]byte(r.FormValue("content"))) - if err != nil { - return err - } f, err := os.CreateTemp(*notepadDir, "notepad") if err != nil { return err } defer f.Close() - if _, err = f.Write(encrypted); err != nil { + encryptingWriter := key.NewWriter(f, []byte("notepad")) + if _, err := encryptingWriter.Write([]byte(r.FormValue("content"))); err != nil { + os.Remove(f.Name()) + return err + } + if err := encryptingWriter.Close(); err != nil { os.Remove(f.Name()) return err } @@ -434,10 +436,7 @@ func createNote(stream *gob.Decoder) (*api.CreateNoteResponse, error) { break } defer f.Close() - encryptingWriter, err := key.Writer(f) - if err != nil { - return nil, err - } + encryptingWriter := key.NewWriter(f, []byte("notes/"+name)) for { req := new(api.CreateNoteRequestStream) if err := stream.Decode(req); err != nil { @@ -450,7 +449,7 @@ func createNote(stream *gob.Decoder) (*api.CreateNoteResponse, error) { return nil, fmt.Errorf("create note: write note to file: %s", err) } } - if err := encryptingWriter.Flush(); err != nil { + if err := encryptingWriter.Close(); err != nil { return nil, err } if err := f.Close(); err != nil { @@ -459,6 +458,17 @@ func createNote(stream *gob.Decoder) (*api.CreateNoteResponse, error) { return &api.CreateNoteResponse{Name: name}, nil } +type readNoteResponseWriter struct { + stream *encoder +} + +func (w *readNoteResponseWriter) Write(buf []byte) (int, error) { + if err := w.stream.send(&api.ReadNoteResponseStream{Chunk: buf}); err != nil { + return 0, err + } + return len(buf), nil +} + func readNote(stream *encoder, req *api.ReadNoteRequest) error { f, err := os.OpenInRoot(*notepadDir+"/notes", req.Note) if err != nil { @@ -474,26 +484,8 @@ func readNote(stream *encoder, req *api.ReadNoteRequest) error { if key == nil { return fmt.Errorf("%w: need login", api.PermissionDenied) } - decryptingReader, err := key.Reader(f) - if err != nil { - return err - } - buf := make([]byte, 4*1024*1024) - for { - n, err := decryptingReader.Read(buf) - if n > 0 { - if err := stream.send(&api.ReadNoteResponseStream{Chunk: buf[:n]}); err != nil { - return err - } - } - if err != nil { - if errors.Is(err, io.EOF) { - break - } - return err - } - } - return nil + _, err = io.Copy(&readNoteResponseWriter{stream: stream}, key.NewReader(f, []byte("notes/"+req.Note))) + return err } func tokenAuth(w http.ResponseWriter, r *http.Request) bool { diff --git a/tools/notes/notes.go b/tools/notes/notes.go index 03070f7..adc7e6a 100644 --- a/tools/notes/notes.go +++ b/tools/notes/notes.go @@ -183,6 +183,17 @@ func (*newCommand) Usage() string { func (*newCommand) SetFlags(*flag.FlagSet) {} +type createNoteRequestStreamWriter struct { + w *gob.Encoder +} + +func (w *createNoteRequestStreamWriter) Write(buf []byte) (int, error) { + if err := w.w.Encode(&api.CreateNoteRequestStream{Chunk: buf}); err != nil { + return 0, err + } + return len(buf), nil +} + func (*newCommand) new(ctx context.Context, fileName string) error { f, err := os.Open(fileName) if err != nil { @@ -200,27 +211,8 @@ func (*newCommand) new(ctx context.Context, fileName string) error { } req.Header.Set("Roseh-Token", token) go func() { - defer w.Close() - w := gob.NewEncoder(w) - buf := make([]byte, 4*1024*1024) - for { - n, err := f.Read(buf) - if n > 0 { - if err := w.Encode(&api.CreateNoteRequestStream{Chunk: buf[:n]}); err != nil { - if err != io.ErrClosedPipe { - fmt.Fprintf(os.Stderr, "upload file: %s\n", err) - } - return - } - } - if err != nil { - if errors.Is(err, io.EOF) { - break - } - fmt.Fprintf(os.Stderr, "read file: %s\n", err) - return - } - } + io.Copy(&createNoteRequestStreamWriter{gob.NewEncoder(w)}, f) + w.Close() }() resp, err := readGobResp[api.CreateNoteResponse](req) if err != nil { |
