aboutsummaryrefslogtreecommitdiffstats
path: root/internal
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2025-10-12 08:59:03 -0700
committerRose Hogenson <rosehogenson@posteo.net>2025-10-12 08:59:03 -0700
commitf809b390a40f5d82c2c56efefc02623c6138a5b2 (patch)
tree2960e0361f655c7428a1a11aa5b91da5906400e4 /internal
parentImprove error messages (diff)
downloadsym-f809b390a40f5d82c2c56efefc02623c6138a5b2.tar.zst
Fix reading password
Diffstat (limited to 'internal')
-rw-r--r--internal/sym/dec.go33
-rw-r--r--internal/sym/dec_test.go43
-rw-r--r--internal/sym/enc.go56
-rw-r--r--internal/sym/enc_test.go58
-rw-r--r--internal/sym/pwhash.go5
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
}