From 69060292c43fec008d2f7f084625703b42a82fb9 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Sat, 11 Oct 2025 19:30:46 -0700 Subject: Move some tests around --- dec/dec.go | 50 +---------- dec/dec_test.go | 96 -------------------- enc/enc.go | 80 +---------------- enc/enc_test.go | 147 ------------------------------- internal/sym/dec.go | 55 +++++++++--- internal/sym/dec_test.go | 195 +++++++++++++++++++++++++++++++++++++++++ internal/sym/enc.go | 100 +++++++++++++++------ internal/sym/enc_test.go | 190 +++++++++++++++++++++++++++++++++++++++ internal/sym/pwhash.go | 14 +++ internal/sym/shared_options.go | 20 ----- internal/sym/sym_test.go | 146 +----------------------------- 11 files changed, 526 insertions(+), 567 deletions(-) delete mode 100644 dec/dec_test.go delete mode 100644 enc/enc_test.go create mode 100644 internal/sym/dec_test.go create mode 100644 internal/sym/enc_test.go delete mode 100644 internal/sym/shared_options.go diff --git a/dec/dec.go b/dec/dec.go index 049bc74..2414bfb 100644 --- a/dec/dec.go +++ b/dec/dec.go @@ -3,60 +3,16 @@ package main import ( "flag" "fmt" - "io" "os" - "golang.org/x/term" "roseh.moe/cmd/sym/internal/sym" ) -type options struct { - password string - force bool - - stdin io.Reader - stdout io.Writer -} - -func (o *options) dec(args ...string) error { - if o.stdin == nil { - o.stdin = os.Stdin - } - if o.stdout == nil { - o.stdout = os.Stdout - } - if len(args) == 0 && o.password == "" { - return fmt.Errorf("-p is required when reading from stdin") - } - var password string - if o.password != "" { - password = o.password - } else { - fmt.Fprint(os.Stderr, "Enter password: ") - pw, err := term.ReadPassword(int(os.Stdin.Fd())) - fmt.Fprintln(os.Stderr) - if err != nil { - return err - } - password = string(pw) - } - if len(args) == 0 { - return sym.Decrypt(o.stdout, o.stdin, password) - } - for _, fileName := range args { - if err := sym.DecryptFile(fileName, password, sym.Force(o.force)); err != nil { - return err - } - } - return nil -} - func main() { - o := new(options) - flag.StringVar(&o.password, "p", "", "use the specified password; if not provided, dec will prompt for a password") - flag.BoolVar(&o.force, "f", false, "overwrite output files even if they already exist") + o := sym.DefaultDecryptOptions + o.RegisterFlags(flag.CommandLine) flag.Parse() - if err := o.dec(flag.Args()...); err != nil { + if err := o.Run(flag.Args()...); err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) } diff --git a/dec/dec_test.go b/dec/dec_test.go deleted file mode 100644 index 823ab68..0000000 --- a/dec/dec_test.go +++ /dev/null @@ -1,96 +0,0 @@ -package main - -import ( - "bytes" - "os" - "path/filepath" - "testing" - - "roseh.moe/cmd/sym/internal/sym" -) - -func mustWriteFile(t *testing.T, path string, content []byte) { - t.Helper() - if err := os.WriteFile(path, content, 0600); err != nil { - t.Fatalf("Failed to write test file: %s", err) - } -} - -func mustReadFile(t *testing.T, path string) []byte { - t.Helper() - content, err := os.ReadFile(path) - if err != nil { - t.Fatalf("Failed to read file: %s", err) - } - return content -} - -func mustRemove(t *testing.T, path string) { - t.Helper() - if err := os.Remove(path); err != nil { - t.Fatalf("Failed to remove file: %s", err) - } -} - -func TestDec(t *testing.T) { - t.Parallel() - - const password = "asdf" - fileContent := []byte("test file content") - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, fileContent) - if err := sym.EncryptFile(fileName, password); err != nil { - t.Errorf("EncryptFile failed: %s", err) - } - mustRemove(t, fileName) - err := (&options{password: password}).dec(fileName + ".enc") - if err != nil { - t.Errorf("dec failed: %s", err) - } - gotFileContents := mustReadFile(t, fileName) - if !bytes.Equal(gotFileContents, fileContent) { - t.Errorf("dec returned incorrect contents %q, want %q", gotFileContents, fileContent) - } -} - -func TestDec_UsageError(t *testing.T) { - t.Parallel() - - err := (&options{}).dec() - if err == nil { - t.Errorf("dec without -p when reading from stdin, want error") - } -} - -func TestDec_NotFound(t *testing.T) { - t.Parallel() - - err := (&options{password: "asdf"}).dec("my-nonexistent-file-name.txt") - if err == nil { - t.Errorf("dec succeeded with nonexistent file, want error") - } -} - -func TestDec_Stdin(t *testing.T) { - t.Parallel() - - const password = "asdf" - content := []byte("test contents") - encrypted := new(bytes.Buffer) - if err := sym.EncryptBinary(encrypted, bytes.NewReader(content), password); err != nil { - t.Fatalf("Failed to encrypt: %s", err) - } - gotContentBuf := new(bytes.Buffer) - opts := &options{ - password: password, - stdin: bytes.NewReader(encrypted.Bytes()), - stdout: gotContentBuf, - } - if err := opts.dec(); err != nil { - t.Fatalf("dec failed: %s", err) - } - gotContent := gotContentBuf.Bytes() - if !bytes.Equal(gotContent, content) { - t.Errorf("dec returned incorrect contents %q, want %q", gotContent, content) - } -} diff --git a/enc/enc.go b/enc/enc.go index 5fb382e..c98875d 100644 --- a/enc/enc.go +++ b/enc/enc.go @@ -1,92 +1,18 @@ package main import ( - "crypto/rand" - "encoding/binary" "flag" "fmt" - "io" "os" - "strings" - "golang.org/x/term" "roseh.moe/cmd/sym/internal/sym" - "roseh.moe/pkg/wordlist" ) -type options struct { - generatePassword bool - password string - asciiOutput bool - force bool - - passwordOut io.Writer - stdin io.Reader - stdout io.Writer -} - -func (o *options) enc(args ...string) error { - if o.passwordOut == nil { - o.passwordOut = os.Stderr - } - if o.stdin == nil { - o.stdin = os.Stdin - } - if o.stdout == nil { - o.stdout = os.Stdout - } - if o.generatePassword && o.password != "" { - return fmt.Errorf("-g and -p cannot be used together") - } - if len(args) == 0 && !o.generatePassword && o.password == "" { - return fmt.Errorf("must use -g or -p when reading from stdin") - } - var password string - if o.password != "" { - password = o.password - } else if o.generatePassword { - const nWords = 10 - buf := make([]byte, 2*nWords) - rand.Read(buf) - words := make([]string, nWords) - for i := range words { - words[i] = wordlist.Words[binary.NativeEndian.Uint16(buf[2*i:])&0x1fff] - } - password = strings.Join(words, " ") - fmt.Fprint(os.Stderr, "Your password: ") - fmt.Fprint(o.passwordOut, password) - fmt.Fprintln(os.Stderr) - } else { - fmt.Fprint(os.Stderr, "Enter password: ") - pw, err := term.ReadPassword(int(os.Stdin.Fd())) - fmt.Fprintln(os.Stderr) - if err != nil { - return err - } - password = string(pw) - } - if len(args) == 0 { - if o.asciiOutput { - return sym.EncryptBase64(o.stdout, o.stdin, password) - } - return sym.EncryptBinary(o.stdout, o.stdin, password) - } - for _, fileName := range args { - if err := sym.EncryptFile(fileName, password, sym.WithASCIIOutput(o.asciiOutput), sym.Force(o.force)); err != nil { - return err - } - } - return nil -} - func main() { - o := new(options) - flag.BoolVar(&o.generatePassword, "g", false, "generate a secure password automatically (password will be printed to stderr)") - flag.StringVar(&o.password, "p", "", "use the specified password; if not provided, enc will prompt for a password") - flag.BoolVar(&o.asciiOutput, "a", false, "output in base64, default is binary output") - flag.BoolVar(&o.force, "f", false, "overwrite output files even if they already exist") + o := sym.DefaultEncryptOptions + o.RegisterFlags(flag.CommandLine) flag.Parse() - if err := o.enc(flag.Args()...); err != nil { + if err := o.Run(flag.Args()...); err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) } diff --git a/enc/enc_test.go b/enc/enc_test.go deleted file mode 100644 index e133b04..0000000 --- a/enc/enc_test.go +++ /dev/null @@ -1,147 +0,0 @@ -package main - -import ( - "bytes" - "os" - "path/filepath" - "strings" - "testing" - - "roseh.moe/cmd/sym/internal/sym" -) - -func mustWriteFile(t *testing.T, path string, content []byte) { - t.Helper() - if err := os.WriteFile(path, content, 0600); err != nil { - t.Fatalf("Failed to write test file: %s", err) - } -} - -func mustReadFile(t *testing.T, path string) []byte { - t.Helper() - content, err := os.ReadFile(path) - if err != nil { - t.Fatalf("Failed to read file: %s", err) - } - return content -} - -func TestEnc(t *testing.T) { - t.Parallel() - - const password = "asdf" - fileContent := []byte("test file content") - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, fileContent) - if err := (&options{password: password}).enc(fileName); err != nil { - t.Fatalf("enc failed: %s", err) - } - if err := sym.DecryptFile(fileName+".enc", password, sym.Force(true)); err != nil { - t.Fatalf("Failed to decrypt encrypted file: %s", err) - } - gotFileContents := mustReadFile(t, fileName) - if !bytes.Equal(gotFileContents, fileContent) { - t.Errorf("encrypt round trip returned incorrect contents %q, want %q", gotFileContents, fileContent) - } -} - -func TestEnc_UsageError(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - generatePassword bool - password string - files []string - }{{ - desc: "GeneratePasswordAndPassword", - generatePassword: true, - password: "asdf", - }, { - desc: "MissingPasswordStdin", - generatePassword: false, - password: "", - }, { - desc: "NonexistentFile", - password: "asdf", - files: []string{"my-nonexistent-file.txt"}, - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - opts := &options{ - generatePassword: tc.generatePassword, - password: tc.password, - } - if err := opts.enc(tc.files...); err == nil { - t.Errorf("enc(%+v) succeeded, want error", opts) - } - }) - } -} - -func TestEnc_GeneratePassword(t *testing.T) { - t.Parallel() - - fileContent := []byte("test file content") - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, fileContent) - - password := new(strings.Builder) - opts := &options{ - generatePassword: true, - passwordOut: password, - } - if err := opts.enc(fileName); err != nil { - t.Fatalf("enc(%+v) failed: %s", opts, err) - } - pw := password.String() - if err := sym.DecryptFile(fileName+".enc", pw, sym.Force(true)); err != nil { - t.Fatalf("Failed to decrypt encrypted file with generated password %q: %s", pw, err) - } - gotFileContents := mustReadFile(t, fileName) - if !bytes.Equal(gotFileContents, fileContent) { - t.Errorf("encrypt round trip returned incorrect contents %q, want %q", gotFileContents, fileContent) - } -} - -func TestEnc_Stdin(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - ascii bool - }{{ - desc: "Binary", - ascii: false, - }, { - desc: "ASCII", - ascii: true, - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - const ( - input = "test input" - password = "asdf" - ) - stdout := new(strings.Builder) - opts := &options{ - password: password, - asciiOutput: tc.ascii, - stdin: strings.NewReader(input), - stdout: stdout, - } - if err := opts.enc(); err != nil { - t.Errorf("enc(+%v) failed: %s", opts, err) - } - got := new(strings.Builder) - if err := sym.Decrypt(got, strings.NewReader(stdout.String()), password); err != nil { - t.Errorf("Failed to decrypt stdout content: %s", err) - } - if got, want := got.String(), input; got != want { - t.Errorf("Encrypt round-trip to stdout returned incorrect contents: %q, want %q", got, want) - } - }) - } -} diff --git a/internal/sym/dec.go b/internal/sym/dec.go index 086bd44..a54d5e5 100644 --- a/internal/sym/dec.go +++ b/internal/sym/dec.go @@ -5,6 +5,7 @@ import ( "bytes" "encoding/base64" "errors" + "flag" "fmt" "io" "os" @@ -44,7 +45,7 @@ func decryptBinary(w io.Writer, r io.Reader, password string) error { return err } -func Decrypt(w io.Writer, r io.Reader, password string) error { +func decrypt(w io.Writer, r io.Reader, password string) error { bufReader := bufio.NewReaderSize(r, 81) b, err := bufReader.Peek(1) if err != nil { @@ -63,20 +64,25 @@ func Decrypt(w io.Writer, r io.Reader, password string) error { return err } -type decryptOptions struct { - force bool +type DecryptOptions struct { + password string + force bool + + stdin io.Reader + stdout io.Writer } -type DecryptFileOption interface { - decryptOpt(*decryptOptions) +var DefaultDecryptOptions = DecryptOptions{ + stdin: os.Stdin, + stdout: os.Stdout, } -func DecryptFile(fileName string, password string, options ...DecryptFileOption) (err error) { - opts := new(decryptOptions) - for _, o := range options { - o.decryptOpt(opts) - } +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") +} +func (o *DecryptOptions) decryptFile(fileName string, password string) (err error) { var outFileName string if name, ok := strings.CutSuffix(fileName, ".enc"); ok { outFileName = name @@ -91,7 +97,7 @@ func DecryptFile(fileName string, password string, options ...DecryptFileOption) } defer fIn.Close() fileOpts := os.O_CREATE | os.O_WRONLY - if opts.force { + if o.force { fileOpts |= os.O_TRUNC } else { fileOpts |= os.O_EXCL @@ -106,8 +112,33 @@ func DecryptFile(fileName string, password string, options ...DecryptFileOption) os.Remove(fOut.Name()) } }() - if err := Decrypt(fOut, fIn, password); err != nil { + if err := decrypt(fOut, fIn, password); err != nil { return fmt.Errorf("decrypt %q: %s", fileName, err) } return fOut.Close() } + +func (o *DecryptOptions) Run(args ...string) error { + if len(args) == 0 && o.password == "" { + return fmt.Errorf("-p is required when reading from stdin") + } + var password string + if o.password != "" { + password = o.password + } else { + var err error + password, err = readPassword() + if err != nil { + return err + } + } + if len(args) == 0 { + return decrypt(o.stdout, o.stdin, password) + } + for _, fileName := range args { + if err := o.decryptFile(fileName, password); err != nil { + return err + } + } + return nil +} diff --git a/internal/sym/dec_test.go b/internal/sym/dec_test.go new file mode 100644 index 0000000..cbe0829 --- /dev/null +++ b/internal/sym/dec_test.go @@ -0,0 +1,195 @@ +package sym + +import ( + "bytes" + "flag" + "path/filepath" + "slices" + "strings" + "testing" +) + +func TestDecryptFile_Force(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + force bool + wantErr bool + }{{ + desc: "OutputExists", + force: false, + wantErr: true, + }, { + desc: "Force", + force: true, + wantErr: false, + }} { + 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.Fatalf("Failed to encrypt file: %s", err) + } + decOpts := DefaultDecryptOptions + decOpts.force = tc.force + err := decOpts.decryptFile(fileName+".enc", password) + if gotErr := err != nil; gotErr != tc.wantErr { + t.Errorf("decryptFile(force=%t) returned returned error %v when output file exists, want error? %t", tc.force, err, tc.wantErr) + } + }) + } +} + +func TestDecrypt_BadFileFormat(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + fileContent []byte + }{{ + desc: "Empty", + fileContent: nil, + }, { + desc: "Short", + fileContent: []byte{0x80}, + }, { + desc: "BadHeader", + fileContent: []byte{0x80, 'a', 's', 'd', 'f'}, + }, { + desc: "BadFormat", + fileContent: []byte("bad file format"), + }, { + desc: "BadContent", + fileContent: []byte("\x80symasdfasdf"), + }, { + desc: "BadContentLong", + fileContent: slices.Concat([]byte("\x80sym"), bytes.Repeat([]byte("asdf"), 100)), + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, tc.fileContent) + err := DefaultDecryptOptions.decryptFile(fileName, "asdf") + if err == nil { + t.Errorf("DecryptFile succeeded for incorrect file format, want error") + } + }) + } +} + +func TestDecryptFile_WeirdName(t *testing.T) { + t.Parallel() + + const password = "asdf" + fileContent := []byte("file content") + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, fileContent) + if err := DefaultEncryptOptions.encryptFile(fileName, password); err != nil { + t.Fatalf("EncryptFile failed: %s", err) + } + mustRename(t, fileName+".enc", fileName+".encrypted") + if err := DefaultDecryptOptions.decryptFile(fileName+".encrypted", password); err != nil { + t.Fatalf("DecryptFile failed: %s", err) + } + gotContents := mustReadFile(t, fileName+".encrypted.dec") + if !bytes.Equal(gotContents, fileContent) { + t.Errorf("contents differ") + } +} + +func TestDecryptFile_NotFound(t *testing.T) { + t.Parallel() + + err := DefaultDecryptOptions.decryptFile("my-nonexistent-file.txt", "asdf") + if err == nil { + t.Fatal("DecryptFile succeeded for nonexistent file, want error") + } +} + +func TestDecryptOptions_RegisterFlags(t *testing.T) { + t.Parallel() + + var o DecryptOptions + fs := flag.NewFlagSet("test", flag.ContinueOnError) + o.RegisterFlags(fs) + const cmd = "-p asdf -f" + fs.Parse(strings.Split(cmd, " ")) + want := DecryptOptions{ + password: "asdf", + force: true, + } + if o != want { + t.Errorf("Command line %q parsed incorrect DecryptOptions, got %+v, want %+v", cmd, o, want) + } +} + +func TestDecryptOptions_Run(t *testing.T) { + t.Parallel() + + const password = "asdf" + fileContent := []byte("test file content") + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, fileContent) + if err := DefaultEncryptOptions.encryptFile(fileName, password); err != nil { + t.Errorf("EncryptFile failed: %s", err) + } + mustRemove(t, fileName) + opts := DefaultDecryptOptions + opts.password = password + err := opts.Run(fileName + ".enc") + if err != nil { + t.Errorf("dec failed: %s", err) + } + gotFileContents := mustReadFile(t, fileName) + if !bytes.Equal(gotFileContents, fileContent) { + t.Errorf("Run returned incorrect contents %q, want %q", gotFileContents, fileContent) + } +} + +func TestDecryptOptions_Run_UsageError(t *testing.T) { + t.Parallel() + + err := DefaultDecryptOptions.Run() + if err == nil { + t.Errorf("Run without -p when reading from stdin, want error") + } +} + +func TestDecryptOptions_Run_NotFound(t *testing.T) { + t.Parallel() + + opts := DefaultDecryptOptions + opts.password = "asdf" + err := opts.Run("my-nonexistent-file-name.txt") + if err == nil { + t.Errorf("Run succeeded with nonexistent file, want error") + } +} + +func TestDecryptOptions_Run_Stdin(t *testing.T) { + t.Parallel() + + const password = "asdf" + content := []byte("test contents") + encrypted := new(bytes.Buffer) + if err := encryptBinary(encrypted, bytes.NewReader(content), password); err != nil { + t.Fatalf("Failed to encrypt: %s", err) + } + gotContentBuf := new(bytes.Buffer) + opts := DefaultDecryptOptions + opts.password = password + opts.stdin = bytes.NewReader(encrypted.Bytes()) + opts.stdout = gotContentBuf + if err := opts.Run(); err != nil { + t.Fatalf("Run failed: %s", err) + } + gotContent := gotContentBuf.Bytes() + if !bytes.Equal(gotContent, content) { + t.Errorf("dec returned incorrect contents %q, want %q", gotContent, content) + } +} diff --git a/internal/sym/enc.go b/internal/sym/enc.go index 18b7de7..87e1954 100644 --- a/internal/sym/enc.go +++ b/internal/sym/enc.go @@ -2,10 +2,16 @@ package sym import ( "bufio" + "crypto/rand" "encoding/base64" + "encoding/binary" + "flag" "fmt" "io" "os" + "strings" + + "roseh.moe/pkg/wordlist" ) type newlineWriter struct { @@ -34,7 +40,7 @@ func (w *newlineWriter) Write(buf []byte) (int, error) { return nn, nil } -func EncryptBinary(w io.Writer, r io.Reader, password string) error { +func encryptBinary(w io.Writer, r io.Reader, password string) error { if _, err := io.WriteString(w, magic); err != nil { return err } @@ -45,7 +51,7 @@ func EncryptBinary(w io.Writer, r io.Reader, password string) error { return writer.Close() } -func EncryptBase64(w io.Writer, r io.Reader, password string) 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? --------------------------------- @@ -69,42 +75,42 @@ func EncryptBase64(w io.Writer, r io.Reader, password string) error { return bufWriter.Flush() } -type encryptOptions struct { - asciiOutput bool - force bool -} +type EncryptOptions struct { + generatePassword bool + password string + asciiOutput bool + force bool -type EncryptFileOption interface { - encryptOpt(*encryptOptions) + passwordOut io.Writer + stdin io.Reader + stdout io.Writer } -type encryptFileOptionFunc func(*encryptOptions) - -func (f encryptFileOptionFunc) encryptOpt(opts *encryptOptions) { f(opts) } - -func WithASCIIOutput(asciiOutput bool) EncryptFileOption { - return encryptFileOptionFunc(func(opts *encryptOptions) { - opts.asciiOutput = asciiOutput - }) +var DefaultEncryptOptions = EncryptOptions{ + passwordOut: os.Stderr, + stdin: os.Stdin, + stdout: os.Stdout, } -func EncryptFile(fileName string, password string, options ...EncryptFileOption) (err error) { - opts := new(encryptOptions) - for _, o := range options { - o.encryptOpt(opts) - } +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 { return err } defer f.Close() ext := ".enc" - if opts.asciiOutput { + if o.asciiOutput { ext = ".enc.txt" } fileOpts := os.O_CREATE | os.O_WRONLY - if opts.force { + if o.force { fileOpts |= os.O_TRUNC } else { fileOpts |= os.O_EXCL @@ -119,13 +125,55 @@ func EncryptFile(fileName string, password string, options ...EncryptFileOption) os.Remove(fOut.Name()) } }() - if opts.asciiOutput { - err = EncryptBase64(fOut, f, password) + if o.asciiOutput { + err = encryptBase64(fOut, f, password) } else { - err = EncryptBinary(fOut, f, password) + err = encryptBinary(fOut, f, password) } if err != nil { return fmt.Errorf("encrypt %q: %s", fileName, err) } return fOut.Close() } + +func (o *EncryptOptions) Run(args ...string) error { + if o.generatePassword && o.password != "" { + return fmt.Errorf("-g and -p cannot be used together") + } + if len(args) == 0 && !o.generatePassword && o.password == "" { + return fmt.Errorf("must use -g or -p when reading from stdin") + } + var password string + if o.password != "" { + password = o.password + } else if o.generatePassword { + const nWords = 10 + buf := make([]byte, 2*nWords) + rand.Read(buf) + words := make([]string, nWords) + for i := range words { + words[i] = wordlist.Words[binary.NativeEndian.Uint16(buf[2*i:])&0x1fff] + } + password = strings.Join(words, " ") + fmt.Fprint(os.Stderr, "Your password: ") + fmt.Fprint(o.passwordOut, password) + fmt.Fprintln(os.Stderr) + } else { + var err error + if password, err = readPassword(); err != nil { + return err + } + } + if len(args) == 0 { + if o.asciiOutput { + return encryptBase64(o.stdout, o.stdin, password) + } + return encryptBinary(o.stdout, o.stdin, password) + } + for _, fileName := range args { + if err := o.encryptFile(fileName, password); err != nil { + return err + } + } + return nil +} diff --git a/internal/sym/enc_test.go b/internal/sym/enc_test.go new file mode 100644 index 0000000..c726dea --- /dev/null +++ b/internal/sym/enc_test.go @@ -0,0 +1,190 @@ +package sym + +import ( + "bytes" + "flag" + "path/filepath" + "strings" + "testing" +) + +func TestEncryptFile_Force(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + force bool + wantErr bool + }{{ + desc: "OutputExists", + force: false, + wantErr: true, + }, { + desc: "Force", + force: true, + wantErr: false, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, []byte("test file content")) + mustWriteFile(t, fileName+".enc", []byte("file already exists")) + encOpts := DefaultEncryptOptions + encOpts.force = tc.force + err := encOpts.encryptFile(fileName, "asdf") + if gotErr := err != nil; gotErr != tc.wantErr { + t.Errorf("EncryptFile(force=%t) returned returned error %v when output file exists, want error? %t", tc.force, err, tc.wantErr) + } + }) + } +} + +func TestEncryptFile_NotFound(t *testing.T) { + t.Parallel() + + err := DefaultEncryptOptions.encryptFile("my-nonexistent-file.txt", "asdf") + if err == nil { + t.Fatal("EncryptFile succeeded for nonexistent file, want error") + } +} + +func TestEncryptOptions_RegisterFlags(t *testing.T) { + t.Parallel() + + var o EncryptOptions + fs := flag.NewFlagSet("test", flag.ContinueOnError) + o.RegisterFlags(fs) + const cmd = "-g -p asdf -a -f" + fs.Parse(strings.Split(cmd, " ")) + want := EncryptOptions{ + generatePassword: true, + password: "asdf", + asciiOutput: true, + force: true, + } + if o != want { + t.Errorf("Command line %q parsed incorrect EncryptOptions, got %+v, want %+v", cmd, o, want) + } +} + +func TestEncryptOptions_Run(t *testing.T) { + t.Parallel() + + const password = "asdf" + fileContent := []byte("test file content") + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, fileContent) + opts := DefaultEncryptOptions + opts.password = password + if err := opts.Run(fileName); err != nil { + t.Fatalf("enc failed: %s", err) + } + mustRemove(t, fileName) + if err := DefaultDecryptOptions.decryptFile(fileName+".enc", password); err != nil { + t.Fatalf("Failed to decrypt encrypted file: %s", err) + } + gotFileContents := mustReadFile(t, fileName) + if !bytes.Equal(gotFileContents, fileContent) { + t.Errorf("encrypt round trip returned incorrect contents %q, want %q", gotFileContents, fileContent) + } +} + +func TestEncryptOptions_Run_UsageError(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + generatePassword bool + password string + files []string + }{{ + desc: "GeneratePasswordAndPassword", + generatePassword: true, + password: "asdf", + }, { + desc: "MissingPasswordStdin", + generatePassword: false, + password: "", + }, { + desc: "NonexistentFile", + password: "asdf", + files: []string{"my-nonexistent-file.txt"}, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + opts := DefaultEncryptOptions + opts.generatePassword = tc.generatePassword + opts.password = tc.password + if err := opts.Run(tc.files...); err == nil { + t.Errorf("Run(%+v) succeeded, want error", opts) + } + }) + } +} + +func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { + t.Parallel() + + fileContent := []byte("test file content") + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, fileContent) + + password := new(strings.Builder) + opts := DefaultEncryptOptions + opts.generatePassword = true + opts.passwordOut = password + if err := opts.Run(fileName); err != nil { + t.Fatalf("enc(%+v) failed: %s", opts, err) + } + pw := password.String() + mustRemove(t, fileName) + if err := DefaultDecryptOptions.decryptFile(fileName+".enc", pw); err != nil { + t.Fatalf("Failed to decrypt encrypted file with generated password %q: %s", pw, err) + } + gotFileContents := mustReadFile(t, fileName) + if !bytes.Equal(gotFileContents, fileContent) { + t.Errorf("encrypt round trip returned incorrect contents %q, want %q", gotFileContents, fileContent) + } +} + +func TestEncryptOptions_Run_Stdin(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + ascii bool + }{{ + desc: "Binary", + ascii: false, + }, { + desc: "ASCII", + ascii: true, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + const ( + input = "test input" + password = "asdf" + ) + stdout := new(strings.Builder) + opts := DefaultEncryptOptions + opts.password = password + opts.asciiOutput = tc.ascii + opts.stdin = strings.NewReader(input) + opts.stdout = stdout + if err := opts.Run(); err != nil { + t.Errorf("enc(+%v) failed: %s", opts, err) + } + got := new(strings.Builder) + if err := decrypt(got, strings.NewReader(stdout.String()), password); err != nil { + t.Errorf("Failed to decrypt stdout content: %s", err) + } + if got, want := got.String(), input; got != want { + t.Errorf("Encrypt round-trip to stdout returned incorrect contents: %q, want %q", got, want) + } + }) + } +} diff --git a/internal/sym/pwhash.go b/internal/sym/pwhash.go index 263e2bc..886554e 100644 --- a/internal/sym/pwhash.go +++ b/internal/sym/pwhash.go @@ -3,8 +3,22 @@ package sym import ( "crypto/pbkdf2" "crypto/sha256" + "fmt" + "os" + + "golang.org/x/term" ) 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: ") + pw, err := term.ReadPassword(int(os.Stdin.Fd())) + fmt.Fprintln(os.Stderr) + if err != nil { + return "", err + } + return string(pw), nil +} diff --git a/internal/sym/shared_options.go b/internal/sym/shared_options.go deleted file mode 100644 index 74c3210..0000000 --- a/internal/sym/shared_options.go +++ /dev/null @@ -1,20 +0,0 @@ -package sym - -type Option interface { - EncryptFileOption - DecryptFileOption -} - -type forceOption bool - -func (force forceOption) encryptOpt(opts *encryptOptions) { - opts.force = bool(force) -} - -func (force forceOption) decryptOpt(opts *decryptOptions) { - opts.force = bool(force) -} - -func Force(force bool) Option { - return forceOption(force) -} diff --git a/internal/sym/sym_test.go b/internal/sym/sym_test.go index 43f29fd..1eef363 100644 --- a/internal/sym/sym_test.go +++ b/internal/sym/sym_test.go @@ -4,7 +4,6 @@ import ( "bytes" "os" "path/filepath" - "slices" "testing" ) @@ -61,7 +60,9 @@ func TestEncryptDecrypt(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, buf) const password = "karp cache tidal mars fed rajah uses graze pobox flew" - if err := EncryptFile(fileName, password, WithASCIIOutput(tc.ascii)); err != nil { + encOpts := DefaultEncryptOptions + encOpts.asciiOutput = tc.ascii + if err := encOpts.encryptFile(fileName, password); err != nil { t.Fatalf("EncryptFile failed: %s", err) } mustRemove(t, fileName) @@ -69,7 +70,7 @@ func TestEncryptDecrypt(t *testing.T) { if tc.ascii { ext = ".enc.txt" } - if err := DecryptFile(fileName+ext, password); err != nil { + if err := DefaultDecryptOptions.decryptFile(fileName+ext, password); err != nil { t.Fatalf("DecryptFile failed: %s", err) } gotContents := mustReadFile(t, fileName) @@ -79,142 +80,3 @@ func TestEncryptDecrypt(t *testing.T) { }) } } - -func TestEncryptFile_Force(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - force bool - wantErr bool - }{{ - desc: "OutputExists", - force: false, - wantErr: true, - }, { - desc: "Force", - force: true, - wantErr: false, - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, []byte("test file content")) - mustWriteFile(t, fileName+".enc", []byte("file already exists")) - err := EncryptFile(fileName, "asdf", Force(tc.force)) - if gotErr := err != nil; gotErr != tc.wantErr { - t.Errorf("EncryptFile(force=%t) returned returned error %v when output file exists, want error? %t", tc.force, err, tc.wantErr) - } - }) - } -} - -func TestDecryptFile_Force(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - force bool - wantErr bool - }{{ - desc: "OutputExists", - force: false, - wantErr: true, - }, { - desc: "Force", - force: true, - wantErr: false, - }} { - 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 := EncryptFile(fileName, password); err != nil { - t.Fatalf("Failed to encrypt file: %s", err) - } - err := DecryptFile(fileName+".enc", password, Force(tc.force)) - if gotErr := err != nil; gotErr != tc.wantErr { - t.Errorf("DecryptFile(force=%t) returned returned error %v when output file exists, want error? %t", tc.force, err, tc.wantErr) - } - }) - } -} - -func TestDecrypt_BadFileFormat(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - fileContent []byte - }{{ - desc: "Empty", - fileContent: nil, - }, { - desc: "Short", - fileContent: []byte{0x80}, - }, { - desc: "BadHeader", - fileContent: []byte{0x80, 'a', 's', 'd', 'f'}, - }, { - desc: "BadFormat", - fileContent: []byte("bad file format"), - }, { - desc: "BadContent", - fileContent: []byte("\x80symasdfasdf"), - }, { - desc: "BadContentLong", - fileContent: slices.Concat([]byte("\x80sym"), bytes.Repeat([]byte("asdf"), 100)), - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, tc.fileContent) - err := DecryptFile(fileName, "asdf") - if err == nil { - t.Errorf("DecryptFile succeeded for incorrect file format, want error") - } - }) - } -} - -func TestDecryptFile_WeirdName(t *testing.T) { - t.Parallel() - - const password = "asdf" - fileContent := []byte("file content") - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, fileContent) - if err := EncryptFile(fileName, password); err != nil { - t.Fatalf("EncryptFile failed: %s", err) - } - mustRename(t, fileName+".enc", fileName+".encrypted") - if err := DecryptFile(fileName+".encrypted", password); err != nil { - t.Fatalf("DecryptFile failed: %s", err) - } - gotContents := mustReadFile(t, fileName+".encrypted.dec") - if !bytes.Equal(gotContents, fileContent) { - t.Errorf("contents differ") - } -} - -func TestEncryptFile_NotFound(t *testing.T) { - t.Parallel() - - err := EncryptFile("my-nonexistent-file.txt", "asdf") - if err == nil { - t.Fatal("EncryptFile succeeded for nonexistent file, want error") - } -} - -func TestDecryptFile_NotFound(t *testing.T) { - t.Parallel() - - err := DecryptFile("my-nonexistent-file.txt", "asdf") - if err == nil { - t.Fatal("DecryptFile succeeded for nonexistent file, want error") - } -} -- cgit v1.3.1