diff options
| -rw-r--r-- | enc/enc.go | 2 | ||||
| -rw-r--r-- | internal/sym/dec.go | 13 | ||||
| -rw-r--r-- | internal/sym/enc.go | 12 | ||||
| -rw-r--r-- | internal/sym/magic.go | 3 | ||||
| -rw-r--r-- | internal/sym/oae.go | 16 |
5 files changed, 30 insertions, 16 deletions
@@ -53,7 +53,7 @@ func enc() error { if *asciiOutput { return sym.EncryptBase64(os.Stdout, os.Stdin, password) } - return sym.Encrypt(os.Stdout, os.Stdin, password) + return sym.EncryptBinary(os.Stdout, os.Stdin, password) } for _, fileName := range args { if err := sym.EncryptFile(fileName, password, *asciiOutput); err != nil { diff --git a/internal/sym/dec.go b/internal/sym/dec.go index 7362db4..31c4ed0 100644 --- a/internal/sym/dec.go +++ b/internal/sym/dec.go @@ -33,12 +33,14 @@ func (r *lineReader) Read(buf []byte) (int, error) { } func decryptBinary(w io.Writer, r io.Reader, password string) error { - header := make([]byte, 1) + header := make([]byte, len(magic)) if _, err := io.ReadFull(r, header); err != nil { return err } - reader := newDecryptingReader(r, password) - _, err := io.Copy(w, reader) + if !bytes.Equal(header, []byte(magic)) { + return fmt.Errorf("bad file format") + } + _, err := io.Copy(w, newDecryptingReader(r, password)) return err } @@ -51,13 +53,14 @@ func Decrypt(w io.Writer, r io.Reader, password string) error { } return err } - if b[0] == 0 { + if b[0] == 0x80 { return decryptBinary(w, bufReader, password) } if b[0] != '-' { return errors.New("invalid input") } - return decryptBinary(w, base64.NewDecoder(base64.StdEncoding, &lineReader{r: bufReader}), password) + _, err = io.Copy(w, newDecryptingReader(base64.NewDecoder(base64.StdEncoding, &lineReader{r: bufReader}), password)) + return err } func DecryptFile(fileName string, password string) (err error) { diff --git a/internal/sym/enc.go b/internal/sym/enc.go index e945142..33ec7ca 100644 --- a/internal/sym/enc.go +++ b/internal/sym/enc.go @@ -33,8 +33,8 @@ func (w *newlineWriter) Write(buf []byte) (int, error) { return nn, nil } -func Encrypt(w io.Writer, r io.Reader, password string) error { - if _, err := w.Write([]byte{0}); err != nil { +func EncryptBinary(w io.Writer, r io.Reader, password string) error { + if _, err := io.WriteString(w, magic); err != nil { return err } writer := newEncryptingWriter(w, password) @@ -52,7 +52,11 @@ func EncryptBase64(w io.Writer, r io.Reader, password string) error { return err } base64Writer := base64.NewEncoder(base64.StdEncoding, &newlineWriter{w: bufWriter}) - if err := Encrypt(base64Writer, r, password); err != nil { + encryptingWriter := newEncryptingWriter(base64Writer, password) + if _, err := io.Copy(encryptingWriter, r); err != nil { + return err + } + if err := encryptingWriter.Close(); err != nil { return err } if err := base64Writer.Close(); err != nil { @@ -87,7 +91,7 @@ func EncryptFile(fileName string, password string, asciiOutput bool) (err error) if asciiOutput { err = EncryptBase64(fOut, f, password) } else { - err = Encrypt(fOut, f, password) + err = EncryptBinary(fOut, f, password) } if err != nil { return err diff --git a/internal/sym/magic.go b/internal/sym/magic.go new file mode 100644 index 0000000..198fe86 --- /dev/null +++ b/internal/sym/magic.go @@ -0,0 +1,3 @@ +package sym + +const magic = "\x80sym" diff --git a/internal/sym/oae.go b/internal/sym/oae.go index 865159f..0ff9540 100644 --- a/internal/sym/oae.go +++ b/internal/sym/oae.go @@ -6,6 +6,7 @@ import ( "crypto/aes" "crypto/cipher" "crypto/rand" + "errors" "io" ) @@ -13,10 +14,10 @@ const ( nonceSize = 12 aeadOverhead = 16 - segmentSize = 4 * 1024 * 1024 - encryptedSegmentSize = segmentSize + aeadOverhead + encryptedSegmentSize = 1024 * 1024 + plaintextSegmentSize = encryptedSegmentSize - aeadOverhead - saltSize = 16 + saltSize = 32 ) type segmentEncrypter struct { @@ -41,7 +42,10 @@ func (se *segmentEncrypter) initialize(salt []byte) error { func (se *segmentEncrypter) ad(lastSegment bool) ([]byte, error) { // Increment counter - for i := range se.nonce { + for i := 0; ; i++ { + if i == len(se.nonce) { + return nil, errors.New("counter overflowed") + } se.nonce[i]++ if se.nonce[i] != 0 { break @@ -122,12 +126,12 @@ func (w *encryptingWriter) Write(buf []byte) (int, error) { } nn := 0 for len(buf) > 0 { - if len(w.buf) == segmentSize { + if len(w.buf) == plaintextSegmentSize { if err := w.writeBuf(false); err != nil { return nn, err } } - n := copy(w.buf[len(w.buf):segmentSize], buf) + n := copy(w.buf[len(w.buf):plaintextSegmentSize], buf) nn += n w.buf = w.buf[:len(w.buf)+n] buf = buf[n:] |
