diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2025-11-11 11:27:17 -0800 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2025-11-11 11:27:17 -0800 |
| commit | 128d5c513b7a5c50a65f03480fc61b06fb5ef8e8 (patch) | |
| tree | 4128102ee2d5fd31f7a4760aefb9eddce478422f | |
| parent | f57b3ba29eff4e5fcd02fa8923dc17f5eea1efb7 (diff) | |
| download | sym-128d5c513b7a5c50a65f03480fc61b06fb5ef8e8.tar.zst | |
Use subcommands package
This reduces a little bit of complexity. (does it?)
| -rw-r--r-- | README.md | 2 | ||||
| -rw-r--r-- | dec.go | 84 | ||||
| -rw-r--r-- | dec_test.go | 64 | ||||
| -rw-r--r-- | enc.go | 132 | ||||
| -rw-r--r-- | enc_test.go | 65 | ||||
| -rw-r--r-- | go.mod | 1 | ||||
| -rw-r--r-- | go.sum | 2 | ||||
| -rw-r--r-- | oae.go | 2 | ||||
| -rw-r--r-- | sym.go | 86 | ||||
| -rw-r--r-- | sym_test.go | 86 |
10 files changed, 284 insertions, 240 deletions
@@ -1,3 +1,5 @@ # Sym: simple symmetric encryption sym is kinda like `gpg --symmetric` + +DO NOT USE THIS PROGRAM. It's more of a "proof of concept". @@ -1,55 +1,54 @@ package main import ( + "context" "errors" "flag" "fmt" "io" "os" "strings" -) -func (o *decryptOptions) decrypt(w io.Writer, r io.Reader, password string) error { - _, err := io.Copy(w, newDecryptingReader(r, password)) - return err -} + "github.com/google/subcommands" +) -type decryptFlags struct { +type decCmd struct { password string force bool + + passwordIn func() (string, error) + 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") - fs.Usage = func() { - fmt.Fprintf(fs.Output(), `usage: %s [OPTION]... [FILE]... +func (*decCmd) Name() string { return "dec" } +func (*decCmd) Synopsis() string { return "decrypt" } +func (*decCmd) Usage() string { + return `usage: sym dec [OPTION]... [FILE]... Decrypt files, or stdin if no files are provided. -`, - fs.Name()) - fs.PrintDefaults() - fmt.Fprintf(fs.Output(), `-p is required when reading from stdin. +-p is required when reading from stdin. For example, - %s my-encrypted-file.txt.enc + sym dec my-encrypted-file.txt.enc would decrypt my-encrypted-file.txt.enc and write the result to my-encrypted-file.txt. If a filename does not end with .enc, the name will be appended with a .dec extension. -`, - fs.Name()) - } + +` } -type decryptOptions struct { - decryptFlags +func (c *decCmd) SetFlags(fs *flag.FlagSet) { + fs.StringVar(&c.password, "p", "", "use the specified password; if not provided, dec will prompt for a password") + fs.BoolVar(&c.force, "f", false, "overwrite output files even if they already exist") +} - passwordIn func() (string, error) - stdin io.Reader - stdout io.Writer +func (c *decCmd) decrypt(w io.Writer, r io.Reader, password string) error { + _, err := io.Copy(w, newDecryptingReader(r, password)) + return err } -func (o *decryptOptions) decryptFile(fileName string, password string) (err error) { +func (c *decCmd) decryptFile(fileName string, password string) (err error) { var outFileName string if name, ok := strings.CutSuffix(fileName, ".enc"); ok { outFileName = name @@ -62,7 +61,7 @@ func (o *decryptOptions) decryptFile(fileName string, password string) (err erro } defer fIn.Close() fileOpts := os.O_CREATE | os.O_WRONLY - if o.force { + if c.force { fileOpts |= os.O_TRUNC } else { fileOpts |= os.O_EXCL @@ -80,40 +79,51 @@ func (o *decryptOptions) decryptFile(fileName string, password string) (err erro os.Remove(fOut.Name()) } }() - if err := o.decrypt(fOut, fIn, password); err != nil { + if err := c.decrypt(fOut, fIn, password); err != nil { return fmt.Errorf("decrypt %q: %s", fileName, err) } return fOut.Close() } -func (o *decryptOptions) readPassword() (string, error) { +func (c *decCmd) readPassword() (string, error) { fmt.Fprint(os.Stderr, "Enter password: ") - pw, err := o.passwordIn() + pw, err := c.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") +func (c *decCmd) run(args ...string) error { + if len(args) == 0 && c.password == "" { + return usageErr("-p is required when reading from stdin") } var password string - if o.password != "" { - password = o.password + if c.password != "" { + password = c.password } else { var err error - password, err = o.readPassword() + password, err = c.readPassword() if err != nil { return err } } if len(args) == 0 { - return o.decrypt(o.stdout, o.stdin, password) + return c.decrypt(c.stdout, c.stdin, password) } for _, fileName := range args { - if err := o.decryptFile(fileName, password); err != nil { + if err := c.decryptFile(fileName, password); err != nil { return err } } return nil } + +func (c *decCmd) Execute(ctx context.Context, f *flag.FlagSet, _ ...any) subcommands.ExitStatus { + if err := c.run(f.Args()...); err != nil { + fmt.Fprintf(os.Stderr, "sym: %s\n", err) + if errors.Is(err, errUsage) { + return subcommands.ExitUsageError + } + return subcommands.ExitFailure + } + return subcommands.ExitSuccess +} diff --git a/dec_test.go b/dec_test.go index 956d0a6..b06c722 100644 --- a/dec_test.go +++ b/dec_test.go @@ -3,7 +3,6 @@ package main import ( "bytes" "errors" - "flag" "path/filepath" "slices" "strings" @@ -32,10 +31,10 @@ func TestDecryptFile_Force(t *testing.T) { const password = "asdf" fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, []byte("test file content")) - if err := (&encryptOptions{}).encryptFile(fileName, password); err != nil { + if err := (&encCmd{}).encryptFile(fileName, password); err != nil { t.Fatalf("Failed to encrypt file: %s", err) } - err := (&decryptOptions{decryptFlags: decryptFlags{force: tc.force}}).decryptFile(fileName+".enc", password) + err := (&decCmd{force: tc.force}).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) } @@ -69,7 +68,7 @@ func TestDecrypt_BadFileFormat(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, tc.fileContent) - err := (&decryptOptions{}).decryptFile(fileName, "asdf") + err := (&decCmd{}).decryptFile(fileName, "asdf") if err == nil { t.Errorf("DecryptFile succeeded for incorrect file format, want error") } @@ -84,11 +83,11 @@ func TestDecryptFile_WeirdName(t *testing.T) { fileContent := []byte("file content") fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, fileContent) - if err := (&encryptOptions{}).encryptFile(fileName, password); err != nil { + if err := (&encCmd{}).encryptFile(fileName, password); err != nil { t.Fatalf("EncryptFile failed: %s", err) } mustRename(t, fileName+".enc", fileName+".encrypted") - if err := (&decryptOptions{}).decryptFile(fileName+".encrypted", password); err != nil { + if err := (&decCmd{}).decryptFile(fileName+".encrypted", password); err != nil { t.Fatalf("DecryptFile failed: %s", err) } gotContents := mustReadFile(t, fileName+".encrypted.dec") @@ -100,7 +99,7 @@ func TestDecryptFile_WeirdName(t *testing.T) { func TestDecryptFile_NotFound(t *testing.T) { t.Parallel() - err := (&decryptOptions{}).decryptFile("my-nonexistent-file.txt", "asdf") + err := (&decCmd{}).decryptFile("my-nonexistent-file.txt", "asdf") if err == nil { t.Fatal("decryptFile succeeded for nonexistent file, want error") } @@ -113,43 +112,26 @@ func TestDecryptFile_NoPermission(t *testing.T) { mustWriteFile(t, fileName, []byte("test file content")) mustWriteFile(t, strings.TrimSuffix(fileName, ".enc"), nil) mustChmod(t, strings.TrimSuffix(fileName, ".enc"), 0400) - err := (&decryptOptions{decryptFlags: decryptFlags{force: true}}).decryptFile(fileName, "asdf") + err := (&decCmd{force: true}).decryptFile(fileName, "asdf") if err == nil { t.Fatal("decryptFile succeeded for unwritable 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 := decryptFlags{ - password: "asdf", - force: true, - } - if o.decryptFlags != want { - t.Errorf("Command line %q parsed incorrect decryptOptions, got %+v, want %+v", cmd, o, want) - } -} - -func TestDecryptOptions_Run(t *testing.T) { +func TestDecCmd_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 := (&encryptOptions{}).encryptFile(fileName, password); err != nil { + if err := (&encCmd{}).encryptFile(fileName, password); err != nil { t.Errorf("EncryptFile failed: %s", err) } mustRemove(t, fileName) - err := (&decryptOptions{decryptFlags: decryptFlags{password: password}}).run(fileName + ".enc") + err := (&decCmd{password: password}).run(fileName + ".enc") if err != nil { - t.Errorf("decryptOptions.run failed: %s", err) + t.Errorf("decCmd.run failed: %s", err) } gotFileContents := mustReadFile(t, fileName) if !bytes.Equal(gotFileContents, fileContent) { @@ -157,36 +139,36 @@ func TestDecryptOptions_Run(t *testing.T) { } } -func TestDecryptOptions_Run_UsageError(t *testing.T) { +func TestDecCmd_Run_UsageError(t *testing.T) { t.Parallel() - err := (&decryptOptions{}).run() + err := (&decCmd{}).run() if err == nil { t.Errorf("Run without -p when reading from stdin, want error") } } -func TestDecryptOptions_Run_NotFound(t *testing.T) { +func TestDecCmd_Run_NotFound(t *testing.T) { t.Parallel() - err := (&decryptOptions{decryptFlags: decryptFlags{password: "asdf"}}).run("my-nonexistent-file-name.txt") + err := (&decCmd{password: "asdf"}).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) { +func TestDecCmd_Run_Stdin(t *testing.T) { t.Parallel() const password = "asdf" content := []byte("test contents") encrypted := new(bytes.Buffer) - if err := (&encryptOptions{}).encrypt(encrypted, bytes.NewReader(content), password); err != nil { + if err := (&encCmd{}).encrypt(encrypted, bytes.NewReader(content), password); err != nil { t.Fatalf("Failed to encrypt: %s", err) } gotContentBuf := new(bytes.Buffer) - if err := (&decryptOptions{ - decryptFlags: decryptFlags{password: password}, + if err := (&decCmd{ + password: password, stdin: bytes.NewReader(encrypted.Bytes()), stdout: gotContentBuf, }).run(); err != nil { @@ -198,7 +180,7 @@ func TestDecryptOptions_Run_Stdin(t *testing.T) { } } -func TestDecryptOptions_Run_ReadPassword(t *testing.T) { +func TestDecCmd_Run_ReadPassword(t *testing.T) { t.Parallel() for _, tc := range []struct { @@ -219,18 +201,18 @@ func TestDecryptOptions_Run_ReadPassword(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, []byte("test file content")) - if err := (&encryptOptions{}).encryptFile(fileName, password); err != nil { + if err := (&encCmd{}).encryptFile(fileName, password); err != nil { t.Errorf("EncryptFile failed: %s", err) } mustRemove(t, fileName) - err := (&decryptOptions{ + err := (&decCmd{ passwordIn: func() (string, error) { return password, tc.err }, }).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) + t.Errorf("decCmd.run returned error %v reading password from stdin, want error? %t", err, tc.wantErr) } }) } @@ -1,6 +1,7 @@ package main import ( + "context" "crypto/rand" "encoding/binary" "errors" @@ -10,46 +11,42 @@ import ( "os" "strings" + "github.com/google/subcommands" "roseh.moe/pkg/wordlist" ) -type encryptFlags struct { +type encCmd struct { generatePassword bool password string force bool + + passwordIn func() (string, error) + passwordOut io.Writer + stdin io.Reader + stdout io.Writer } -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.force, "f", false, "overwrite output files even if they already exist") - fs.Usage = func() { - fmt.Fprintf(fs.Output(), `usage: %s [OPTION]... [FILE]... +func (*encCmd) Name() string { return "enc" } +func (*encCmd) Synopsis() string { return "encrypt" } +func (*encCmd) Usage() string { + return `usage: sym enc [OPTION]... [FILE]... Encrypt files, or stdin if no files are provided. -`, - fs.Name()) - fs.PrintDefaults() - fmt.Fprintf(fs.Output(), ` One of -g or -p must be used when reading from stdin. When encrypting to stdout, consider redirecting the result since binary output can mess up your terminal. Example: - echo test | %s -p 'my super secure password' | base64 -`, - fs.Name()) - } -} + echo test | sym enc -p 'my super secure password' | base64 -type encryptOptions struct { - encryptFlags +` +} - passwordIn func() (string, error) - passwordOut io.Writer - stdin io.Reader - stdout io.Writer +func (c *encCmd) SetFlags(fs *flag.FlagSet) { + fs.BoolVar(&c.generatePassword, "g", false, "generate a secure password automatically (password will be printed to stderr)") + fs.StringVar(&c.password, "p", "", "use the specified password; if not provided, enc will prompt for a password") + fs.BoolVar(&c.force, "f", false, "overwrite output files even if they already exist") } -func (o *encryptOptions) encrypt(w io.Writer, r io.Reader, password string) error { +func (c *encCmd) encrypt(w io.Writer, r io.Reader, password string) error { writer := newEncryptingWriter(w, password) if _, err := io.Copy(writer, r); err != nil { return err @@ -57,14 +54,14 @@ func (o *encryptOptions) encrypt(w io.Writer, r io.Reader, password string) erro return writer.close() } -func (o *encryptOptions) encryptFile(fileName string, password string) (err error) { +func (c *encCmd) encryptFile(fileName string, password string) (err error) { f, err := os.Open(fileName) if err != nil { return err } defer f.Close() fileOpts := os.O_CREATE | os.O_WRONLY - if o.force { + if c.force { fileOpts |= os.O_TRUNC } else { fileOpts |= os.O_EXCL @@ -82,55 +79,45 @@ func (o *encryptOptions) encryptFile(fileName string, password string) (err erro os.Remove(fOut.Name()) } }() - if err = o.encrypt(fOut, f, password); err != nil { + if err = c.encrypt(fOut, f, password); err != nil { return fmt.Errorf("encrypt %q: %s", fileName, err) } 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 +func (c *encCmd) readPassword() (string, error) { + fmt.Fprint(os.Stderr, "Enter password: ") + password, err := c.passwordIn() + fmt.Fprintln(os.Stderr) + if err != nil { + return "", err + } + if password == "" { + return "", usageErr("password cannot be empty") } - return "", fmt.Errorf("too many attempts") + fmt.Fprint(os.Stderr, "Repeat password: ") + pwConfirm, err := c.passwordIn() + fmt.Fprintln(os.Stderr) + if err != nil { + return "", err + } + if pwConfirm != password { + return "", usageErr("passwords do not match") + } + return password, nil } -func (o *encryptOptions) run(args ...string) error { - if o.generatePassword && o.password != "" { - return fmt.Errorf("-g and -p cannot be used together") +func (c *encCmd) run(args ...string) error { + if c.generatePassword && c.password != "" { + return usageErr("-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") + if len(args) == 0 && !c.generatePassword && c.password == "" { + return usageErr("must use -g or -p when reading from stdin") } var password string - if o.password != "" { - password = o.password - } else if o.generatePassword { + if c.password != "" { + password = c.password + } else if c.generatePassword { const nWords = 10 buf := make([]byte, 2*nWords) rand.Read(buf) @@ -140,21 +127,32 @@ func (o *encryptOptions) run(args ...string) error { } password = strings.Join(words, " ") fmt.Fprint(os.Stderr, "Your password: ") - fmt.Fprint(o.passwordOut, password) + fmt.Fprint(c.passwordOut, password) fmt.Fprintln(os.Stderr) } else { var err error - if password, err = o.readPassword(); err != nil { + if password, err = c.readPassword(); err != nil { return err } } if len(args) == 0 { - return o.encrypt(o.stdout, o.stdin, password) + return c.encrypt(c.stdout, c.stdin, password) } for _, fileName := range args { - if err := o.encryptFile(fileName, password); err != nil { + if err := c.encryptFile(fileName, password); err != nil { return err } } return nil } + +func (c *encCmd) Execute(ctx context.Context, f *flag.FlagSet, _ ...any) subcommands.ExitStatus { + if err := c.run(f.Args()...); err != nil { + fmt.Fprintf(os.Stderr, "sym: %s\n", err) + if errors.Is(err, errUsage) { + return subcommands.ExitUsageError + } + return subcommands.ExitFailure + } + return subcommands.ExitSuccess +} diff --git a/enc_test.go b/enc_test.go index f204283..c800652 100644 --- a/enc_test.go +++ b/enc_test.go @@ -3,7 +3,6 @@ package main import ( "bytes" "errors" - "flag" "path/filepath" "strings" "testing" @@ -31,7 +30,7 @@ func TestEncryptFile_Force(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, []byte("test file content")) mustWriteFile(t, fileName+".enc", []byte("file already exists")) - err := (&encryptOptions{encryptFlags: encryptFlags{force: tc.force}}).encryptFile(fileName, "asdf") + err := (&encCmd{force: tc.force}).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) } @@ -42,7 +41,7 @@ func TestEncryptFile_Force(t *testing.T) { func TestEncryptFile_NotFound(t *testing.T) { t.Parallel() - err := (&encryptOptions{}).encryptFile("my-nonexistent-file.txt", "asdf") + err := (&encCmd{}).encryptFile("my-nonexistent-file.txt", "asdf") if err == nil { t.Fatal("encryptFile succeeded for nonexistent file, want error") } @@ -55,42 +54,24 @@ func TestEncryptFile_NoPermission(t *testing.T) { mustWriteFile(t, fileName, []byte("test file content")) mustWriteFile(t, fileName+".enc", nil) mustChmod(t, fileName+".enc", 0400) - err := (&encryptOptions{encryptFlags: encryptFlags{force: true}}).encryptFile(fileName, "asdf") + err := (&encCmd{force: true}).encryptFile(fileName, "asdf") if err == nil { t.Fatal("encryptFile succeeded for unwritable 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 -f" - fs.Parse(strings.Split(cmd, " ")) - want := encryptFlags{ - generatePassword: true, - password: "asdf", - force: true, - } - if o.encryptFlags != want { - t.Errorf("Command line %q parsed incorrect encryptOptions, got %+v, want %+v", cmd, o, want) - } -} - -func TestEncryptOptions_Run(t *testing.T) { +func TestEncCmd_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 := (&encryptOptions{encryptFlags: encryptFlags{password: password}}).run(fileName); err != nil { + if err := (&encCmd{password: password}).run(fileName); err != nil { t.Fatalf("enc failed: %s", err) } mustRemove(t, fileName) - if err := (&decryptOptions{}).decryptFile(fileName+".enc", password); err != nil { + if err := (&decCmd{}).decryptFile(fileName+".enc", password); err != nil { t.Fatalf("Failed to decrypt encrypted file: %s", err) } gotFileContents := mustReadFile(t, fileName) @@ -99,7 +80,7 @@ func TestEncryptOptions_Run(t *testing.T) { } } -func TestEncryptOptions_Run_UsageError(t *testing.T) { +func TestEncCmd_Run_UsageError(t *testing.T) { t.Parallel() for _, tc := range []struct { @@ -123,18 +104,18 @@ func TestEncryptOptions_Run_UsageError(t *testing.T) { t.Run(tc.desc, func(t *testing.T) { t.Parallel() - opts := &encryptOptions{encryptFlags: encryptFlags{ + opts := &encCmd{ generatePassword: tc.generatePassword, password: tc.password, - }} + } if err := opts.run(tc.files...); err == nil { - t.Errorf("encryptOptions.run(%+v) succeeded, want error", opts) + t.Errorf("encCmd.run(%+v) succeeded, want error", opts) } }) } } -func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { +func TestEncCmd_Run_GeneratePassword(t *testing.T) { t.Parallel() fileContent := []byte("test file content") @@ -142,16 +123,16 @@ func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { mustWriteFile(t, fileName, fileContent) password := new(strings.Builder) - opts := &encryptOptions{ - encryptFlags: encryptFlags{generatePassword: true}, + opts := &encCmd{ + generatePassword: true, passwordOut: password, } if err := opts.run(fileName); err != nil { - t.Fatalf("encryptOptions.run(%+v) failed: %s", opts, err) + t.Fatalf("encCmd.run(%+v) failed: %s", opts, err) } pw := password.String() mustRemove(t, fileName) - if err := (&decryptOptions{}).decryptFile(fileName+".enc", pw); err != nil { + if err := (&decCmd{}).decryptFile(fileName+".enc", pw); err != nil { t.Fatalf("Failed to decrypt encrypted file with generated password %q: %s", pw, err) } gotFileContents := mustReadFile(t, fileName) @@ -160,7 +141,7 @@ func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { } } -func TestEncryptOptions_Run_Stdin(t *testing.T) { +func TestEncCmd_Run_Stdin(t *testing.T) { t.Parallel() const ( @@ -168,15 +149,15 @@ func TestEncryptOptions_Run_Stdin(t *testing.T) { password = "asdf" ) stdout := new(strings.Builder) - if err := (&encryptOptions{ - encryptFlags: encryptFlags{password: password}, + if err := (&encCmd{ + password: password, stdin: strings.NewReader(input), stdout: stdout, }).run(); err != nil { - t.Errorf("encryptOptions.run failed: %s", err) + t.Errorf("encCmd.run failed: %s", err) } got := new(strings.Builder) - if err := (&decryptOptions{}).decrypt(got, strings.NewReader(stdout.String()), password); err != nil { + if err := (&decCmd{}).decrypt(got, strings.NewReader(stdout.String()), password); err != nil { t.Errorf("Failed to decrypt stdout content %q: %s", stdout, err) } if got, want := got.String(), input; got != want { @@ -184,7 +165,7 @@ func TestEncryptOptions_Run_Stdin(t *testing.T) { } } -func TestEncryptOptions_Run_ReadPassword(t *testing.T) { +func TestEncCmd_Run_ReadPassword(t *testing.T) { t.Parallel() for _, tc := range []struct { @@ -220,7 +201,7 @@ func TestEncryptOptions_Run_ReadPassword(t *testing.T) { mustWriteFile(t, fileName, []byte("test file content")) passwordI := 0 - err := (&encryptOptions{ + err := (&encCmd{ passwordIn: func() (string, error) { if passwordI == len(tc.passwords) && tc.err != nil { return "", tc.err @@ -231,7 +212,7 @@ func TestEncryptOptions_Run_ReadPassword(t *testing.T) { }, }).run(fileName) if gotErr := err != nil; gotErr != tc.wantErr { - t.Errorf("encryptOptions.run returned error %v, want error? %t", err, tc.wantErr) + t.Errorf("encCmd.run returned error %v, want error? %t", err, tc.wantErr) } }) } @@ -3,6 +3,7 @@ module roseh.moe/cmd/sym go 1.25.0 require ( + github.com/google/subcommands v1.2.0 golang.org/x/crypto v0.43.0 golang.org/x/term v0.36.0 roseh.moe/pkg/wordlist v1.0.2 @@ -1,5 +1,7 @@ github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/google/subcommands v1.2.0 h1:vWQspBTo2nEqTUFita5/KeEWlUL8kQObDFbub/EN9oE= +github.com/google/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk= golang.org/x/crypto v0.43.0 h1:dduJYIi3A3KOfdGOHX8AVZ/jGiyPa3IbBozJ5kNuE04= golang.org/x/crypto v0.43.0/go.mod h1:BFbav4mRNlXJL4wNeejLpWxB7wMbc79PdRGhWKncxR0= golang.org/x/mod v0.29.0 h1:HV8lRxZC4l2cr3Zq1LvtOsi/ThTgWnUk/y64QSs8GwA= @@ -18,7 +18,7 @@ const ( segmentSize = 1024 * 1024 plaintextSegmentSize = segmentSize - aeadOverhead - saltSize = 16 + saltSize = 32 ) type segmentEncrypter struct { @@ -8,43 +8,47 @@ package main import ( + "context" + "errors" "flag" "fmt" + "io" "os" - "path/filepath" + + "github.com/google/subcommands" ) -type subcommand interface { - registerFlags(*flag.FlagSet) - run(...string) error +var errUsage = errors.New("usage error") + +type usageError struct { + msg string } -func whichSubcommand(name string) (subcommand, bool) { - switch name { - case "enc": - return &encryptOptions{ - passwordIn: termReadPassword, - passwordOut: os.Stderr, - stdin: os.Stdin, - stdout: os.Stdout, - }, true - case "dec": - return &decryptOptions{ - passwordIn: termReadPassword, - stdin: os.Stdin, - stdout: os.Stdout, - }, true - default: - return nil, false - } +func usageErr(format string, args ...any) error { + return &usageError{msg: fmt.Sprintf(format, args...)} +} + +func (e *usageError) Error() string { + return e.msg } -func run(args []string) error { - name := filepath.Base(args[0]) - cmd, ok := whichSubcommand(name) - if !ok { - flag.Usage = func() { - fmt.Fprintf(os.Stderr, `usage: sym <subcommand> [OPTION]... [FILE]... +func (e *usageError) Is(target error) bool { return target == errUsage } + +func registerCommands(commander *subcommands.Commander, passwordIn func() (string, error), passwordOut io.Writer, stdin io.Reader, stdout io.Writer) { + commander.Register(&encCmd{ + passwordIn: passwordIn, + passwordOut: passwordOut, + stdin: stdin, + stdout: stdout, + }, "") + commander.Register(&decCmd{ + passwordIn: passwordIn, + stdin: stdin, + stdout: stdout, + }, "") + commander.Register(commander.HelpCommand(), "") + commander.Explain = func(w io.Writer) { + fmt.Fprintf(w, `usage: sym <subcommand> [OPTION]... [FILE]... Encrypt or decrypt files using a password. Subcommands: @@ -52,31 +56,13 @@ Subcommands: dec decrypt Try sym <subcommand> -h for command-specific help. - -Pro tip: use "ln sym enc" or "ln sym dec" to create shortcuts for each subcommand. `) - } - flag.CommandLine.Parse(args[1:]) - args = flag.Args() - if len(args) == 0 { - return fmt.Errorf("missing subcommand (use sym -h for help)") - } - subcommand := filepath.Base(args[0]) - name = "sym " + subcommand - cmd, ok = whichSubcommand(subcommand) - if !ok { - return fmt.Errorf("invalid subcommand %q", args[0]) - } } - fs := flag.NewFlagSet(name, flag.ExitOnError) - cmd.registerFlags(fs) - fs.Parse(args[1:]) - return cmd.run(fs.Args()...) } func main() { - if err := run(os.Args); err != nil { - fmt.Fprintln(os.Stderr, err) - os.Exit(1) - } + ctx := context.Background() + registerCommands(subcommands.DefaultCommander, termReadPassword, os.Stderr, os.Stdin, os.Stdout) + flag.Parse() + os.Exit(int(subcommands.Execute(ctx))) } diff --git a/sym_test.go b/sym_test.go index 3da9fc1..8ae82cc 100644 --- a/sym_test.go +++ b/sym_test.go @@ -2,9 +2,13 @@ package main import ( "bytes" + "context" + "flag" "os" "path/filepath" "testing" + + "github.com/google/subcommands" ) func init() { @@ -59,11 +63,11 @@ 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 := (&encryptOptions{}).encryptFile(fileName, password); err != nil { + if err := (&encCmd{}).encryptFile(fileName, password); err != nil { t.Fatalf("EncryptFile failed: %s", err) } mustRemove(t, fileName) - if err := (&decryptOptions{}).decryptFile(fileName+".enc", password); err != nil { + if err := (&decCmd{}).decryptFile(fileName+".enc", password); err != nil { t.Fatalf("DecryptFile failed: %s", err) } gotContents := mustReadFile(t, fileName) @@ -71,3 +75,81 @@ func TestEncryptDecrypt(t *testing.T) { t.Errorf("contents differ") } } + +func run(ctx context.Context, t *testing.T, cmd ...string) subcommands.ExitStatus { + t.Helper() + + fs := flag.NewFlagSet("test", flag.ContinueOnError) + commander := subcommands.NewCommander(fs, "test") + registerCommands(commander, nil, nil, nil, nil) + if err := fs.Parse(cmd); err != nil { + t.Fatalf("Failed to parse command %q: %s", cmd, err) + } + return commander.Execute(ctx) +} + +func TestCommander(t *testing.T) { + t.Parallel() + + ctx := t.Context() + fileName := filepath.Join(t.TempDir(), "file.txt") + const fileContent = "test file content" + mustWriteFile(t, fileName, []byte(fileContent)) + const password = "asdf" + if st := run(ctx, t, "enc", "-p="+password, fileName); st != subcommands.ExitSuccess { + t.Fatalf("enc failed: status %d", st) + } + mustRemove(t, fileName) + if st := run(ctx, t, "dec", "-p="+password, fileName+".enc"); st != subcommands.ExitSuccess { + t.Fatalf("dec failed: status %d", st) + } + gotContents := mustReadFile(t, fileName) + if !bytes.Equal(gotContents, []byte(fileContent)) { + t.Errorf("dec returned invalid content, got %q, want %q", gotContents, fileContent) + } +} + +func TestCommander_Errors(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + cmd []string + wantStatus subcommands.ExitStatus + } {{ + desc: "EncUsageError", + cmd: []string{"enc", "-g", "-p=asdf", "file.txt"}, + wantStatus: subcommands.ExitUsageError, + }, { + desc: "EncNoSuchFile", + cmd: []string{"enc", "-p=asdf", "nonexistent-file.txt"}, + wantStatus: subcommands.ExitFailure, + }, { + desc: "DecUsageError", + cmd: []string{"dec"}, + wantStatus: subcommands.ExitUsageError, + }, { + desc: "DecNoSuchFile", + cmd: []string{"dec", "-p=asdf", "nonexistent-file.txt.enc"}, + wantStatus: subcommands.ExitFailure, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + ctx := t.Context() + st := run(ctx, t, tc.cmd...) + if st != tc.wantStatus { + t.Errorf("command %q returned status %d, want %d", tc.cmd, st, tc.wantStatus) + } + }) + } +} + +func TestUsage(t *testing.T) { + t.Parallel() + + ctx := t.Context() + run(ctx, t, "help") + run(ctx, t, "enc", "-h") + run(ctx, t, "dec", "-h") +} |
