diff options
| -rw-r--r-- | dec/dec.go | 5 | ||||
| -rw-r--r-- | enc/enc.go | 14 | ||||
| -rw-r--r-- | internal/sym/dec.go | 43 | ||||
| -rw-r--r-- | internal/sym/enc.go | 20 | ||||
| -rw-r--r-- | internal/sym/oae.go | 76 | ||||
| -rw-r--r-- | internal/sym/pwhash.go | 4 | ||||
| -rw-r--r-- | internal/sym/sym_test.go | 19 |
7 files changed, 70 insertions, 111 deletions
@@ -28,12 +28,11 @@ func dec() error { } password = string(pw) } - pwCache := make(sym.PasswordCache) if len(args) == 0 { - return sym.Decrypt(os.Stdout, os.Stdin, password, pwCache) + return sym.Decrypt(os.Stdout, os.Stdin, password) } for _, fileName := range args { - if err := sym.DecryptFile(fileName, password, pwCache); err != nil { + if err := sym.DecryptFile(fileName, password); err != nil { return err } } @@ -49,20 +49,14 @@ func enc() error { } password = string(pw) } - salt := make([]byte, sym.SaltSize) - rand.Read(salt) - key, err := sym.HashPassword(password, salt) - if err != nil { - return err - } if len(args) == 0 { if *asciiOutput { - return sym.EncryptBase64(os.Stdout, os.Stdin, key, salt, 0) + return sym.EncryptBase64(os.Stdout, os.Stdin, password) } - return sym.Encrypt(os.Stdout, os.Stdin, key, salt, 0) + return sym.Encrypt(os.Stdout, os.Stdin, password) } - for i, fileName := range args { - if err := sym.EncryptFile(fileName, key, salt, i, *asciiOutput); err != nil { + for _, fileName := range args { + if err := sym.EncryptFile(fileName, password, *asciiOutput); err != nil { return err } } diff --git a/internal/sym/dec.go b/internal/sym/dec.go index 7cbf5d4..7362db4 100644 --- a/internal/sym/dec.go +++ b/internal/sym/dec.go @@ -12,53 +12,38 @@ import ( ) type lineReader struct { - r bufio.Scanner + r *bufio.Reader line []byte } func (r *lineReader) Read(buf []byte) (int, error) { for len(r.line) == 0 { - if !r.r.Scan() { - if err := r.r.Err(); err != nil { - return 0, err - } - return 0, io.EOF + line, err := r.r.ReadBytes('\n') + if len(line) == 0 { + return 0, err } - line := r.r.Bytes() if bytes.HasPrefix(line, []byte("-")) { continue } - r.line = line + r.line = bytes.TrimSuffix(line, []byte("\n")) } n := copy(buf, r.line) r.line = r.line[n:] return n, nil } -type PasswordCache map[[SaltSize]byte][]byte - -func decryptBinary(w io.Writer, r io.Reader, password string, pwCache PasswordCache) error { - header := make([]byte, 1+SaltSize) +func decryptBinary(w io.Writer, r io.Reader, password string) error { + header := make([]byte, 1) if _, err := io.ReadFull(r, header); err != nil { return err } - salt := header[1:] - key, ok := pwCache[[SaltSize]byte(salt)] - if !ok { - var err error - key, err = HashPassword(password, salt) - if err != nil { - return err - } - pwCache[[SaltSize]byte(salt)] = key - } - reader := newDecryptingReader(r, key) + reader := newDecryptingReader(r, password) _, err := io.Copy(w, reader) return err } -func Decrypt(w io.Writer, r io.Reader, password string, pwCache PasswordCache) error { - bufReader := bufio.NewReaderSize(r, 1) +func Decrypt(w io.Writer, r io.Reader, password string) error { + bufReader := bufio.NewReaderSize(r, 81) b, err := bufReader.Peek(1) if err != nil { if err == io.EOF { @@ -67,15 +52,15 @@ func Decrypt(w io.Writer, r io.Reader, password string, pwCache PasswordCache) e return err } if b[0] == 0 { - return decryptBinary(w, bufReader, password, pwCache) + return decryptBinary(w, bufReader, password) } if b[0] != '-' { return errors.New("invalid input") } - return decryptBinary(w, base64.NewDecoder(base64.StdEncoding, &lineReader{r: *bufio.NewScanner(bufReader)}), password, pwCache) + return decryptBinary(w, base64.NewDecoder(base64.StdEncoding, &lineReader{r: bufReader}), password) } -func DecryptFile(fileName, password string, pwCache PasswordCache) (err error) { +func DecryptFile(fileName string, password string) (err error) { var outFileName string if name, ok := strings.CutSuffix(fileName, ".enc"); ok { outFileName = name @@ -99,7 +84,7 @@ func DecryptFile(fileName, password string, pwCache PasswordCache) (err error) { os.Remove(fOut.Name()) } }() - if err := Decrypt(fOut, fIn, password, pwCache); err != nil { + if err := Decrypt(fOut, fIn, password); err != nil { return err } return fOut.Close() diff --git a/internal/sym/enc.go b/internal/sym/enc.go index 32ec307..e945142 100644 --- a/internal/sym/enc.go +++ b/internal/sym/enc.go @@ -3,7 +3,6 @@ package sym import ( "bufio" "encoding/base64" - "encoding/binary" "io" "os" ) @@ -34,23 +33,18 @@ func (w *newlineWriter) Write(buf []byte) (int, error) { return nn, nil } -func Encrypt(w io.Writer, r io.Reader, key, salt []byte, counter int) error { +func Encrypt(w io.Writer, r io.Reader, password string) error { if _, err := w.Write([]byte{0}); err != nil { return err } - if _, err := w.Write(salt); err != nil { - return err - } - var noncePrefix [noncePrefixSize]byte - binary.BigEndian.PutUint32(noncePrefix[:], uint32(counter)) - writer := newEncryptingWriter(w, key, noncePrefix) + writer := newEncryptingWriter(w, password) if _, err := io.Copy(writer, r); err != nil { return err } return writer.Close() } -func EncryptBase64(w io.Writer, r io.Reader, key, salt []byte, count int) error { +func EncryptBase64(w io.Writer, r io.Reader, password string) error { bufWriter := bufio.NewWriter(w) if _, err := bufWriter.WriteString(`-------------------------- Begin encrypted text block -------------------------- -------------------------- am i cool like gpg? --------------------------------- @@ -58,7 +52,7 @@ func EncryptBase64(w io.Writer, r io.Reader, key, salt []byte, count int) error return err } base64Writer := base64.NewEncoder(base64.StdEncoding, &newlineWriter{w: bufWriter}) - if err := Encrypt(base64Writer, r, key, salt, count); err != nil { + if err := Encrypt(base64Writer, r, password); err != nil { return err } if err := base64Writer.Close(); err != nil { @@ -70,7 +64,7 @@ func EncryptBase64(w io.Writer, r io.Reader, key, salt []byte, count int) error return bufWriter.Flush() } -func EncryptFile(fileName string, key, salt []byte, count int, asciiOutput bool) (err error) { +func EncryptFile(fileName string, password string, asciiOutput bool) (err error) { f, err := os.Open(fileName) if err != nil { return err @@ -91,9 +85,9 @@ func EncryptFile(fileName string, key, salt []byte, count int, asciiOutput bool) } }() if asciiOutput { - err = EncryptBase64(fOut, f, key, salt, count) + err = EncryptBase64(fOut, f, password) } else { - err = Encrypt(fOut, f, key, salt, count) + err = Encrypt(fOut, f, password) } if err != nil { return err diff --git a/internal/sym/oae.go b/internal/sym/oae.go index e71d5c3..865159f 100644 --- a/internal/sym/oae.go +++ b/internal/sym/oae.go @@ -5,31 +5,33 @@ import ( "bytes" "crypto/aes" "crypto/cipher" - "encoding/binary" - "errors" + "crypto/rand" "io" - "math" ) const ( - nonceSize = 12 - aeadOverhead = 16 - noncePrefixSize = nonceSize - 8 + nonceSize = 12 + aeadOverhead = 16 segmentSize = 4 * 1024 * 1024 encryptedSegmentSize = segmentSize + aeadOverhead + + saltSize = 16 ) type segmentEncrypter struct { - key []byte - noncePrefix [noncePrefixSize]byte + password string - aead cipher.AEAD - i uint64 + aead cipher.AEAD + nonce [nonceSize]byte } -func (se *segmentEncrypter) initialize() error { - block, err := aes.NewCipher(se.key) +func (se *segmentEncrypter) initialize(salt []byte) error { + key, err := hashPassword(se.password, salt) + if err != nil { + return err + } + block, err := aes.NewCipher(key) if err != nil { return err } @@ -37,36 +39,36 @@ func (se *segmentEncrypter) initialize() error { return err } -func (se *segmentEncrypter) nonce(lastSegment bool) ([]byte, []byte, error) { - if se.i == math.MaxUint64 { - return nil, nil, errors.New("counter overflowed") +func (se *segmentEncrypter) ad(lastSegment bool) ([]byte, error) { + // Increment counter + for i := range se.nonce { + se.nonce[i]++ + if se.nonce[i] != 0 { + break + } } - nonce := make([]byte, nonceSize) - copy(nonce, se.noncePrefix[:]) - binary.BigEndian.PutUint64(nonce[noncePrefixSize:], se.i) - ad := make([]byte, 9) - binary.BigEndian.PutUint64(ad, se.i) + ad := make([]byte, len(se.nonce)+1) + copy(ad, se.nonce[:]) if lastSegment { - ad[8] = 1 + ad[len(ad)-1] = 1 } - se.i++ - return nonce, ad, nil + return ad, nil } func (se *segmentEncrypter) encrypt(out, buf []byte, lastSegment bool) ([]byte, error) { - nonce, ad, err := se.nonce(lastSegment) + ad, err := se.ad(lastSegment) if err != nil { return nil, err } - return se.aead.Seal(out, nonce, buf, ad), nil + return se.aead.Seal(out, se.nonce[:], buf, ad), nil } func (se *segmentEncrypter) decrypt(out, buf []byte, lastSegment bool) ([]byte, error) { - nonce, ad, err := se.nonce(lastSegment) + ad, err := se.ad(lastSegment) if err != nil { return nil, err } - return se.aead.Open(out, nonce, buf, ad) + return se.aead.Open(out, se.nonce[:], buf, ad) } type encryptingWriter struct { @@ -76,12 +78,11 @@ type encryptingWriter struct { initialized bool } -func newEncryptingWriter(w io.Writer, key []byte, noncePrefix [noncePrefixSize]byte) *encryptingWriter { +func newEncryptingWriter(w io.Writer, password string) *encryptingWriter { return &encryptingWriter{ w: w, encrypter: segmentEncrypter{ - key: key, - noncePrefix: noncePrefix, + password: password, }, } } @@ -90,10 +91,12 @@ func (w *encryptingWriter) initialize() error { if w.initialized { return nil } - if err := w.encrypter.initialize(); err != nil { + header := make([]byte, saltSize) + rand.Read(header) + if err := w.encrypter.initialize(header); err != nil { return err } - if _, err := w.w.Write(w.encrypter.noncePrefix[:]); err != nil { + if _, err := w.w.Write(header); err != nil { return err } w.buf = make([]byte, 0, encryptedSegmentSize) @@ -146,11 +149,11 @@ type decryptingReader struct { initialized bool } -func newDecryptingReader(r io.Reader, key []byte) *decryptingReader { +func newDecryptingReader(r io.Reader, password string) *decryptingReader { return &decryptingReader{ r: bufio.NewReaderSize(r, 1), decrypter: segmentEncrypter{ - key: key, + password: password, }, } } @@ -159,10 +162,11 @@ func (r *decryptingReader) initialize() error { if r.initialized { return nil } - if _, err := io.ReadFull(r.r, r.decrypter.noncePrefix[:]); err != nil { + header := make([]byte, saltSize) + if _, err := io.ReadFull(r.r, header); err != nil { return err } - if err := r.decrypter.initialize(); err != nil { + if err := r.decrypter.initialize(header); err != nil { return err } r.buf = *bytes.NewBuffer(make([]byte, 0, encryptedSegmentSize)) diff --git a/internal/sym/pwhash.go b/internal/sym/pwhash.go index ef17182..263e2bc 100644 --- a/internal/sym/pwhash.go +++ b/internal/sym/pwhash.go @@ -5,8 +5,6 @@ import ( "crypto/sha256" ) -const SaltSize = 16 - -func HashPassword(password string, salt []byte) ([]byte, error) { +func hashPassword(password string, salt []byte) ([]byte, error) { return pbkdf2.Key(sha256.New, password, salt, 35_000_000, 32) } diff --git a/internal/sym/sym_test.go b/internal/sym/sym_test.go index 0ab0144..c5f136c 100644 --- a/internal/sym/sym_test.go +++ b/internal/sym/sym_test.go @@ -2,30 +2,15 @@ package sym import ( "bytes" - "encoding/hex" "os" "path/filepath" "testing" ) -func mustHex(t *testing.T, s string) []byte { - t.Helper() - b, err := hex.DecodeString(s) - if err != nil { - t.Fatalf("Bad hex %q: %s", s, err) - } - return b -} - func TestEncryptDecrypt(t *testing.T) { t.Parallel() const password = "karp cache tidal mars fed rajah uses graze pobox flew" - salt := mustHex(t, "9aa7d8bb6d19f794162f4062c789b230") - key, err := HashPassword(password, salt) - if err != nil { - t.Fatalf("HashPassword failed: %s", err) - } buf := make([]byte, 10*1024*1024) for i := range buf { buf[i] = byte(i) @@ -47,14 +32,14 @@ func TestEncryptDecrypt(t *testing.T) { if err := os.WriteFile(fileName, buf, 0600); err != nil { t.Fatalf("Failed to write test file: %s", err) } - if err := EncryptFile(fileName, key, salt, 0, tc.ascii); err != nil { + if err := EncryptFile(fileName, password, tc.ascii); err != nil { t.Fatalf("EncryptFile failed: %s", err) } ext := ".enc" if tc.ascii { ext = ".enc.txt" } - if err := DecryptFile(fileName+ext, password, make(PasswordCache)); err != nil { + if err := DecryptFile(fileName+ext, password); err != nil { t.Fatalf("DecryptFile failed: %s", err) } gotContents, err := os.ReadFile(fileName) |
