diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/sym/dec.go | 33 | ||||
| -rw-r--r-- | internal/sym/dec_test.go | 43 | ||||
| -rw-r--r-- | internal/sym/enc.go | 56 | ||||
| -rw-r--r-- | internal/sym/enc_test.go | 58 | ||||
| -rw-r--r-- | internal/sym/pwhash.go | 5 |
5 files changed, 168 insertions, 27 deletions
diff --git a/internal/sym/dec.go b/internal/sym/dec.go index f2801b9..71da93c 100644 --- a/internal/sym/dec.go +++ b/internal/sym/dec.go @@ -64,22 +64,28 @@ func decrypt(w io.Writer, r io.Reader, password string) error { return err } -type DecryptOptions struct { +type decryptFlags struct { password string force bool +} - stdin io.Reader - stdout io.Writer +func (f *decryptFlags) RegisterFlags(fs *flag.FlagSet) { + fs.StringVar(&f.password, "p", "", "use the specified password; if not provided, dec will prompt for a password") + fs.BoolVar(&f.force, "f", false, "overwrite output files even if they already exist") } -var DefaultDecryptOptions = DecryptOptions{ - stdin: os.Stdin, - stdout: os.Stdout, +type DecryptOptions struct { + decryptFlags + + passwordIn func() (string, error) + stdin io.Reader + stdout io.Writer } -func (o *DecryptOptions) RegisterFlags(fs *flag.FlagSet) { - fs.StringVar(&o.password, "p", "", "use the specified password; if not provided, dec will prompt for a password") - fs.BoolVar(&o.force, "f", false, "overwrite output files even if they already exist") +var DefaultDecryptOptions = DecryptOptions{ + passwordIn: termReadPassword, + stdin: os.Stdin, + stdout: os.Stdout, } func (o *DecryptOptions) decryptFile(fileName string, password string) (err error) { @@ -121,6 +127,13 @@ func (o *DecryptOptions) decryptFile(fileName string, password string) (err erro return fOut.Close() } +func (o *DecryptOptions) readPassword() (string, error) { + fmt.Fprint(os.Stderr, "Enter password: ") + pw, err := o.passwordIn() + fmt.Fprintln(os.Stderr) + return pw, err +} + func (o *DecryptOptions) Run(args ...string) error { if len(args) == 0 && o.password == "" { return fmt.Errorf("-p is required when reading from stdin") @@ -130,7 +143,7 @@ func (o *DecryptOptions) Run(args ...string) error { password = o.password } else { var err error - password, err = readPassword() + password, err = o.readPassword() if err != nil { return err } diff --git a/internal/sym/dec_test.go b/internal/sym/dec_test.go index 6ee5900..b70d325 100644 --- a/internal/sym/dec_test.go +++ b/internal/sym/dec_test.go @@ -2,6 +2,7 @@ package sym import ( "bytes" + "errors" "flag" "path/filepath" "slices" @@ -134,11 +135,11 @@ func TestDecryptOptions_RegisterFlags(t *testing.T) { o.RegisterFlags(fs) const cmd = "-p asdf -f" fs.Parse(strings.Split(cmd, " ")) - want := DecryptOptions{ + want := decryptFlags{ password: "asdf", force: true, } - if o != want { + if o.decryptFlags != want { t.Errorf("Command line %q parsed incorrect DecryptOptions, got %+v, want %+v", cmd, o, want) } } @@ -208,3 +209,41 @@ func TestDecryptOptions_Run_Stdin(t *testing.T) { t.Errorf("dec returned incorrect contents %q, want %q", gotContent, content) } } + +func TestDecryptOptions_Run_ReadPassword(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + err error + wantErr bool + }{{ + desc: "Ok", + }, { + desc: "Err", + err: errors.New("test error"), + wantErr: true, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + const password = "asdf" + + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, []byte("test file content")) + if err := DefaultEncryptOptions.encryptFile(fileName, password); err != nil { + t.Errorf("EncryptFile failed: %s", err) + } + mustRemove(t, fileName) + + opts := DefaultDecryptOptions + opts.passwordIn = func() (string, error) { + return password, tc.err + } + err := opts.Run(fileName + ".enc") + if gotErr := err != nil; gotErr != tc.wantErr { + t.Errorf("DecryptOptions.Run returned error %v reading password from stdin, want error? %t", err, tc.wantErr) + } + }) + } +} diff --git a/internal/sym/enc.go b/internal/sym/enc.go index 06ba000..be05e68 100644 --- a/internal/sym/enc.go +++ b/internal/sym/enc.go @@ -76,30 +76,36 @@ func encryptBase64(w io.Writer, r io.Reader, password string) error { return bufWriter.Flush() } -type EncryptOptions struct { +type encryptFlags struct { generatePassword bool password string asciiOutput bool force bool +} + +func (f *encryptFlags) RegisterFlags(fs *flag.FlagSet) { + fs.BoolVar(&f.generatePassword, "g", false, "generate a secure password automatically (password will be printed to stderr)") + fs.StringVar(&f.password, "p", "", "use the specified password; if not provided, enc will prompt for a password") + fs.BoolVar(&f.asciiOutput, "a", false, "output in base64, default is binary output") + fs.BoolVar(&f.force, "f", false, "overwrite output files even if they already exist") +} + +type EncryptOptions struct { + encryptFlags + passwordIn func() (string, error) passwordOut io.Writer stdin io.Reader stdout io.Writer } var DefaultEncryptOptions = EncryptOptions{ + passwordIn: termReadPassword, passwordOut: os.Stderr, stdin: os.Stdin, stdout: os.Stdout, } -func (o *EncryptOptions) RegisterFlags(fs *flag.FlagSet) { - fs.BoolVar(&o.generatePassword, "g", false, "generate a secure password automatically (password will be printed to stderr)") - fs.StringVar(&o.password, "p", "", "use the specified password; if not provided, enc will prompt for a password") - fs.BoolVar(&o.asciiOutput, "a", false, "output in base64, default is binary output") - fs.BoolVar(&o.force, "f", false, "overwrite output files even if they already exist") -} - func (o *EncryptOptions) encryptFile(fileName string, password string) (err error) { f, err := os.Open(fileName) if err != nil { @@ -140,6 +146,38 @@ func (o *EncryptOptions) encryptFile(fileName string, password string) (err erro return fOut.Close() } +func (o *EncryptOptions) readPassword() (string, error) { + const maxAttempts = 3 + for i := 1; i <= maxAttempts; i++ { + fmt.Fprint(os.Stderr, "Enter password") + if i > 1 { + fmt.Fprintf(os.Stderr, " (attempt %d/%d)", i, maxAttempts) + } + fmt.Fprint(os.Stderr, ": ") + password, err := o.passwordIn() + fmt.Fprintln(os.Stderr) + if err != nil { + return "", err + } + if password == "" { + fmt.Fprintln(os.Stderr, "Password cannot be empty") + continue + } + fmt.Fprint(os.Stderr, "Repeat password: ") + pwConfirm, err := o.passwordIn() + fmt.Fprintln(os.Stderr) + if err != nil { + return "", err + } + if pwConfirm != password { + fmt.Fprintln(os.Stderr, "Passwords do not match") + continue + } + return password, nil + } + return "", fmt.Errorf("too many attempts") +} + func (o *EncryptOptions) Run(args ...string) error { if o.generatePassword && o.password != "" { return fmt.Errorf("-g and -p cannot be used together") @@ -164,7 +202,7 @@ func (o *EncryptOptions) Run(args ...string) error { fmt.Fprintln(os.Stderr) } else { var err error - if password, err = readPassword(); err != nil { + if password, err = o.readPassword(); err != nil { return err } } diff --git a/internal/sym/enc_test.go b/internal/sym/enc_test.go index e9450e5..c6f697c 100644 --- a/internal/sym/enc_test.go +++ b/internal/sym/enc_test.go @@ -2,6 +2,7 @@ package sym import ( "bytes" + "errors" "flag" "path/filepath" "strings" @@ -72,13 +73,13 @@ func TestEncryptOptions_RegisterFlags(t *testing.T) { o.RegisterFlags(fs) const cmd = "-g -p asdf -a -f" fs.Parse(strings.Split(cmd, " ")) - want := EncryptOptions{ + want := encryptFlags{ generatePassword: true, password: "asdf", asciiOutput: true, force: true, } - if o != want { + if o.encryptFlags != want { t.Errorf("Command line %q parsed incorrect EncryptOptions, got %+v, want %+v", cmd, o, want) } } @@ -203,3 +204,56 @@ func TestEncryptOptions_Run_Stdin(t *testing.T) { }) } } + +func TestEncryptOptions_Run_ReadPassword(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + passwords []string + err error + wantErr bool + }{{ + desc: "Ok", + passwords: []string{"asdf"}, + }, { + desc: "EmptyPassword", + passwords: []string{""}, + wantErr: true, + }, { + desc: "PasswordsDoNotMatch", + passwords: []string{"asdf", "jkl"}, + wantErr: true, + }, { + desc: "ReadPasswordErr", + err: errors.New("test error"), + wantErr: true, + }, { + desc: "RepeatPasswordErr", + passwords: []string{"asdf"}, + err: errors.New("test error"), + wantErr: true, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, []byte("test file content")) + + opts := DefaultEncryptOptions + passwordI := 0 + opts.passwordIn = func() (string, error) { + if passwordI == len(tc.passwords) && tc.err != nil { + return "", tc.err + } + pw := tc.passwords[passwordI%len(tc.passwords)] + passwordI++ + return pw, nil + } + err := opts.Run(fileName) + if gotErr := err != nil; gotErr != tc.wantErr { + t.Errorf("EncryptOptions.Run returned error %v, want error? %t", err, tc.wantErr) + } + }) + } +} diff --git a/internal/sym/pwhash.go b/internal/sym/pwhash.go index 886554e..86539a0 100644 --- a/internal/sym/pwhash.go +++ b/internal/sym/pwhash.go @@ -3,7 +3,6 @@ package sym import ( "crypto/pbkdf2" "crypto/sha256" - "fmt" "os" "golang.org/x/term" @@ -13,10 +12,8 @@ func hashPassword(password string, salt []byte) ([]byte, error) { return pbkdf2.Key(sha256.New, password, salt, 35_000_000, 32) } -func readPassword() (string, error) { - fmt.Fprint(os.Stderr, "Enter password: ") +func termReadPassword() (string, error) { pw, err := term.ReadPassword(int(os.Stdin.Fd())) - fmt.Fprintln(os.Stderr) if err != nil { return "", err } |
