From 9a08b03b83ff28a26b0c20ccebc178df29cf7312 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Thu, 23 Oct 2025 21:19:00 -0700 Subject: Use one binary with subcommands --- dec.go | 14 +++++++------- dec/dec.go | 19 ------------------- dec_test.go | 44 ++++++++++++++++++++++---------------------- enc.go | 18 +++++++++--------- enc/enc.go | 19 ------------------- enc_test.go | 26 +++++++++++++------------- encryptionalg_string.go | 2 +- metadata.go | 2 +- oae.go | 4 ++-- oae_test.go | 4 ++-- pwhash.go | 2 +- pwhash_string.go | 2 +- sym.go | 48 ++++++++++++++++++++++++++++++++++++++++++++++++ sym_test.go | 8 ++++---- 14 files changed, 111 insertions(+), 101 deletions(-) delete mode 100644 dec/dec.go delete mode 100644 enc/enc.go create mode 100644 sym.go diff --git a/dec.go b/dec.go index 657da35..fdea3fe 100644 --- a/dec.go +++ b/dec.go @@ -1,4 +1,4 @@ -package sym +package main import ( "encoding/binary" @@ -34,12 +34,12 @@ type decryptFlags struct { force bool } -func (f *decryptFlags) RegisterFlags(fs *flag.FlagSet) { +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") } -type DecryptOptions struct { +type decryptOptions struct { decryptFlags passwordIn func() (string, error) @@ -47,13 +47,13 @@ type DecryptOptions struct { stdout io.Writer } -var DefaultDecryptOptions = DecryptOptions{ +var defaultDecryptOptions = decryptOptions{ passwordIn: termReadPassword, stdin: os.Stdin, stdout: os.Stdout, } -func (o *DecryptOptions) decryptFile(fileName string, password string) (err error) { +func (o *decryptOptions) decryptFile(fileName string, password string) (err error) { var outFileName string if name, ok := strings.CutSuffix(fileName, ".enc"); ok { outFileName = name @@ -92,14 +92,14 @@ func (o *DecryptOptions) decryptFile(fileName string, password string) (err erro return fOut.Close() } -func (o *DecryptOptions) readPassword() (string, error) { +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 { +func (o *decryptOptions) run(args ...string) error { if len(args) == 0 && o.password == "" { return fmt.Errorf("-p is required when reading from stdin") } diff --git a/dec/dec.go b/dec/dec.go deleted file mode 100644 index 2414bfb..0000000 --- a/dec/dec.go +++ /dev/null @@ -1,19 +0,0 @@ -package main - -import ( - "flag" - "fmt" - "os" - - "roseh.moe/cmd/sym/internal/sym" -) - -func main() { - o := sym.DefaultDecryptOptions - o.RegisterFlags(flag.CommandLine) - flag.Parse() - if err := o.Run(flag.Args()...); err != nil { - fmt.Fprintln(os.Stderr, err) - os.Exit(1) - } -} diff --git a/dec_test.go b/dec_test.go index e1d6d9b..46bb588 100644 --- a/dec_test.go +++ b/dec_test.go @@ -1,4 +1,4 @@ -package sym +package main import ( "bytes" @@ -36,7 +36,7 @@ func TestDecryptFile_Force(t *testing.T) { if err := testEncryptOptions.encryptFile(fileName, password); err != nil { t.Fatalf("Failed to encrypt file: %s", err) } - decOpts := DefaultDecryptOptions + decOpts := defaultDecryptOptions decOpts.force = tc.force err := decOpts.decryptFile(fileName+".enc", password) if gotErr := err != nil; gotErr != tc.wantErr { @@ -152,7 +152,7 @@ func TestDecrypt_BadFileFormat(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, tc.fileContent) - err := DefaultDecryptOptions.decryptFile(fileName, "asdf") + err := defaultDecryptOptions.decryptFile(fileName, "asdf") if err == nil { t.Errorf("DecryptFile succeeded for incorrect file format, want error") } @@ -171,7 +171,7 @@ func TestDecryptFile_WeirdName(t *testing.T) { t.Fatalf("EncryptFile failed: %s", err) } mustRename(t, fileName+".enc", fileName+".encrypted") - if err := DefaultDecryptOptions.decryptFile(fileName+".encrypted", password); err != nil { + if err := defaultDecryptOptions.decryptFile(fileName+".encrypted", password); err != nil { t.Fatalf("DecryptFile failed: %s", err) } gotContents := mustReadFile(t, fileName+".encrypted.dec") @@ -183,7 +183,7 @@ func TestDecryptFile_WeirdName(t *testing.T) { func TestDecryptFile_NotFound(t *testing.T) { t.Parallel() - err := DefaultDecryptOptions.decryptFile("my-nonexistent-file.txt", "asdf") + err := defaultDecryptOptions.decryptFile("my-nonexistent-file.txt", "asdf") if err == nil { t.Fatal("decryptFile succeeded for nonexistent file, want error") } @@ -196,7 +196,7 @@ 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) - opts := DefaultDecryptOptions + opts := defaultDecryptOptions opts.force = true err := opts.decryptFile(fileName, "asdf") if err == nil { @@ -207,9 +207,9 @@ func TestDecryptFile_NoPermission(t *testing.T) { func TestDecryptOptions_RegisterFlags(t *testing.T) { t.Parallel() - var o DecryptOptions + var o decryptOptions fs := flag.NewFlagSet("test", flag.ContinueOnError) - o.RegisterFlags(fs) + o.registerFlags(fs) const cmd = "-p asdf -f" fs.Parse(strings.Split(cmd, " ")) want := decryptFlags{ @@ -217,7 +217,7 @@ func TestDecryptOptions_RegisterFlags(t *testing.T) { force: true, } if o.decryptFlags != want { - t.Errorf("Command line %q parsed incorrect DecryptOptions, got %+v, want %+v", cmd, o, want) + t.Errorf("Command line %q parsed incorrect decryptOptions, got %+v, want %+v", cmd, o, want) } } @@ -232,22 +232,22 @@ func TestDecryptOptions_Run(t *testing.T) { t.Errorf("EncryptFile failed: %s", err) } mustRemove(t, fileName) - opts := DefaultDecryptOptions + opts := defaultDecryptOptions opts.password = password - err := opts.Run(fileName + ".enc") + 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) + t.Errorf("run returned incorrect contents %q, want %q", gotFileContents, fileContent) } } func TestDecryptOptions_Run_UsageError(t *testing.T) { t.Parallel() - err := DefaultDecryptOptions.Run() + err := defaultDecryptOptions.run() if err == nil { t.Errorf("Run without -p when reading from stdin, want error") } @@ -256,11 +256,11 @@ func TestDecryptOptions_Run_UsageError(t *testing.T) { func TestDecryptOptions_Run_NotFound(t *testing.T) { t.Parallel() - opts := DefaultDecryptOptions + opts := defaultDecryptOptions opts.password = "asdf" - err := opts.Run("my-nonexistent-file-name.txt") + err := opts.run("my-nonexistent-file-name.txt") if err == nil { - t.Errorf("Run succeeded with nonexistent file, want error") + t.Errorf("run succeeded with nonexistent file, want error") } } @@ -274,12 +274,12 @@ func TestDecryptOptions_Run_Stdin(t *testing.T) { t.Fatalf("Failed to encrypt: %s", err) } gotContentBuf := new(bytes.Buffer) - opts := DefaultDecryptOptions + 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) + if err := opts.run(); err != nil { + t.Fatalf("run failed: %s", err) } gotContent := gotContentBuf.Bytes() if !bytes.Equal(gotContent, content) { @@ -313,13 +313,13 @@ func TestDecryptOptions_Run_ReadPassword(t *testing.T) { } mustRemove(t, fileName) - opts := DefaultDecryptOptions + opts := defaultDecryptOptions opts.passwordIn = func() (string, error) { return password, tc.err } - err := opts.Run(fileName + ".enc") + 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) + t.Errorf("decryptOptions.run returned error %v reading password from stdin, want error? %t", err, tc.wantErr) } }) } diff --git a/enc.go b/enc.go index 623174c..6327f25 100644 --- a/enc.go +++ b/enc.go @@ -1,4 +1,4 @@ -package sym +package main import ( "crypto/rand" @@ -19,13 +19,13 @@ type encryptFlags struct { force bool } -func (f *encryptFlags) RegisterFlags(fs *flag.FlagSet) { +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") } -type EncryptOptions struct { +type encryptOptions struct { encryptFlags iterations int @@ -35,7 +35,7 @@ type EncryptOptions struct { stdout io.Writer } -var DefaultEncryptOptions = EncryptOptions{ +var defaultEncryptOptions = encryptOptions{ iterations: defaultPBKDF2Iters, passwordIn: termReadPassword, passwordOut: os.Stderr, @@ -43,7 +43,7 @@ var DefaultEncryptOptions = EncryptOptions{ stdout: os.Stdout, } -func (o *EncryptOptions) encrypt(w io.Writer, r io.Reader, password string) error { +func (o *encryptOptions) encrypt(w io.Writer, r io.Reader, password string) error { if _, err := io.WriteString(w, magic); err != nil { return err } @@ -66,10 +66,10 @@ func (o *EncryptOptions) encrypt(w io.Writer, r io.Reader, password string) erro if _, err := io.Copy(writer, r); err != nil { return err } - return writer.Close() + return writer.close() } -func (o *EncryptOptions) encryptFile(fileName string, password string) (err error) { +func (o *encryptOptions) encryptFile(fileName string, password string) (err error) { f, err := os.Open(fileName) if err != nil { return err @@ -100,7 +100,7 @@ func (o *EncryptOptions) encryptFile(fileName string, password string) (err erro return fOut.Close() } -func (o *EncryptOptions) readPassword() (string, error) { +func (o *encryptOptions) readPassword() (string, error) { const maxAttempts = 3 for i := 1; i <= maxAttempts; i++ { fmt.Fprint(os.Stderr, "Enter password") @@ -132,7 +132,7 @@ func (o *EncryptOptions) readPassword() (string, error) { return "", fmt.Errorf("too many attempts") } -func (o *EncryptOptions) Run(args ...string) error { +func (o *encryptOptions) run(args ...string) error { if o.generatePassword && o.password != "" { return fmt.Errorf("-g and -p cannot be used together") } diff --git a/enc/enc.go b/enc/enc.go deleted file mode 100644 index c98875d..0000000 --- a/enc/enc.go +++ /dev/null @@ -1,19 +0,0 @@ -package main - -import ( - "flag" - "fmt" - "os" - - "roseh.moe/cmd/sym/internal/sym" -) - -func main() { - o := sym.DefaultEncryptOptions - o.RegisterFlags(flag.CommandLine) - flag.Parse() - if err := o.Run(flag.Args()...); err != nil { - fmt.Fprintln(os.Stderr, err) - os.Exit(1) - } -} diff --git a/enc_test.go b/enc_test.go index 3a0c493..d2c6690 100644 --- a/enc_test.go +++ b/enc_test.go @@ -1,4 +1,4 @@ -package sym +package main import ( "bytes" @@ -68,9 +68,9 @@ func TestEncryptFile_NoPermission(t *testing.T) { func TestEncryptOptions_RegisterFlags(t *testing.T) { t.Parallel() - var o EncryptOptions + var o encryptOptions fs := flag.NewFlagSet("test", flag.ContinueOnError) - o.RegisterFlags(fs) + o.registerFlags(fs) const cmd = "-g -p asdf -f" fs.Parse(strings.Split(cmd, " ")) want := encryptFlags{ @@ -79,7 +79,7 @@ func TestEncryptOptions_RegisterFlags(t *testing.T) { force: true, } if o.encryptFlags != want { - t.Errorf("Command line %q parsed incorrect EncryptOptions, got %+v, want %+v", cmd, o, want) + t.Errorf("Command line %q parsed incorrect encryptOptions, got %+v, want %+v", cmd, o, want) } } @@ -92,11 +92,11 @@ func TestEncryptOptions_Run(t *testing.T) { mustWriteFile(t, fileName, fileContent) opts := testEncryptOptions opts.password = password - if err := opts.Run(fileName); err != nil { + 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 { + if err := defaultDecryptOptions.decryptFile(fileName+".enc", password); err != nil { t.Fatalf("Failed to decrypt encrypted file: %s", err) } gotFileContents := mustReadFile(t, fileName) @@ -132,8 +132,8 @@ func TestEncryptOptions_Run_UsageError(t *testing.T) { opts := testEncryptOptions opts.generatePassword = tc.generatePassword opts.password = tc.password - if err := opts.Run(tc.files...); err == nil { - t.Errorf("Run(%+v) succeeded, want error", opts) + if err := opts.run(tc.files...); err == nil { + t.Errorf("run(%+v) succeeded, want error", opts) } }) } @@ -150,12 +150,12 @@ func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { opts := testEncryptOptions opts.generatePassword = true opts.passwordOut = password - if err := opts.Run(fileName); err != nil { + 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 { + 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) @@ -176,7 +176,7 @@ func TestEncryptOptions_Run_Stdin(t *testing.T) { opts.password = password opts.stdin = strings.NewReader(input) opts.stdout = stdout - if err := opts.Run(); err != nil { + if err := opts.run(); err != nil { t.Errorf("enc(+%v) failed: %s", opts, err) } got := new(strings.Builder) @@ -233,9 +233,9 @@ func TestEncryptOptions_Run_ReadPassword(t *testing.T) { passwordI++ return pw, nil } - err := opts.Run(fileName) + err := opts.run(fileName) if gotErr := err != nil; gotErr != tc.wantErr { - t.Errorf("EncryptOptions.Run returned error %v, want error? %t", err, tc.wantErr) + t.Errorf("encryptOptions.run returned error %v, want error? %t", err, tc.wantErr) } }) } diff --git a/encryptionalg_string.go b/encryptionalg_string.go index b7384ee..9860f20 100644 --- a/encryptionalg_string.go +++ b/encryptionalg_string.go @@ -1,6 +1,6 @@ // Code generated by "stringer -type=encryptionAlg -linecomment"; DO NOT EDIT. -package sym +package main import "strconv" diff --git a/metadata.go b/metadata.go index 3875a63..8c789c0 100644 --- a/metadata.go +++ b/metadata.go @@ -1,4 +1,4 @@ -package sym +package main import "fmt" diff --git a/oae.go b/oae.go index e40cb60..a462b05 100644 --- a/oae.go +++ b/oae.go @@ -1,4 +1,4 @@ -package sym +package main import ( "bufio" @@ -170,7 +170,7 @@ func (w *encryptingWriter) Write(buf []byte) (int, error) { return nn, nil } -func (w *encryptingWriter) Close() error { +func (w *encryptingWriter) close() error { if err := w.initialize(); err != nil { return err } diff --git a/oae_test.go b/oae_test.go index 16f9f96..8a0a639 100644 --- a/oae_test.go +++ b/oae_test.go @@ -1,4 +1,4 @@ -package sym +package main import ( "bytes" @@ -28,7 +28,7 @@ func TestOAEReadWrite(t *testing.T) { if _, err := io.WriteString(writer, input); err != nil { t.Fatalf("Failed to write: %s", err) } - if err := writer.Close(); err != nil { + if err := writer.close(); err != nil { t.Fatalf("writer.Close() failed: %s", err) } got, err := io.ReadAll(testEncryptionMetadata.newDecryptingReader(bytes.NewReader(out.Bytes()), password, &testHashMetadata)) diff --git a/pwhash.go b/pwhash.go index d05f6b5..2d7a70f 100644 --- a/pwhash.go +++ b/pwhash.go @@ -1,4 +1,4 @@ -package sym +package main import ( "crypto/pbkdf2" diff --git a/pwhash_string.go b/pwhash_string.go index ceb4d1f..37796a8 100644 --- a/pwhash_string.go +++ b/pwhash_string.go @@ -1,6 +1,6 @@ // Code generated by "stringer -type=pwHash -linecomment"; DO NOT EDIT. -package sym +package main import "strconv" diff --git a/sym.go b/sym.go new file mode 100644 index 0000000..2405feb --- /dev/null +++ b/sym.go @@ -0,0 +1,48 @@ +package main + +import ( + "flag" + "fmt" + "os" + "path/filepath" +) + +type subcommand interface { + registerFlags(*flag.FlagSet) + run(...string) error +} + +func whichSubcommand(name string) (subcommand, bool) { + switch filepath.Base(name) { + case "enc": + return &defaultEncryptOptions, true + case "dec": + return &defaultDecryptOptions, true + default: + return nil, false + } +} + +func run(args []string) error { + cmd, ok := whichSubcommand(args[0]) + if !ok { + if len(args) < 2 { + return fmt.Errorf("missing subcommand") + } + cmd, ok = whichSubcommand(args[1]) + if !ok { + return fmt.Errorf("invalid subcommand %q", args[1]) + } + args = args[1:] + } + cmd.registerFlags(flag.CommandLine) + flag.CommandLine.Parse(args[1:]) + return cmd.run(flag.Args()...) +} + +func main() { + if err := run(os.Args); err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } +} diff --git a/sym_test.go b/sym_test.go index 704d477..7a223e3 100644 --- a/sym_test.go +++ b/sym_test.go @@ -1,4 +1,4 @@ -package sym +package main import ( "bytes" @@ -7,8 +7,8 @@ import ( "testing" ) -var testEncryptOptions = func() EncryptOptions { - opts := DefaultEncryptOptions +var testEncryptOptions = func() encryptOptions { + opts := defaultEncryptOptions opts.iterations = 10 return opts }() @@ -65,7 +65,7 @@ func TestEncryptDecrypt(t *testing.T) { t.Fatalf("EncryptFile failed: %s", err) } mustRemove(t, fileName) - if err := DefaultDecryptOptions.decryptFile(fileName+".enc", password); err != nil { + if err := defaultDecryptOptions.decryptFile(fileName+".enc", password); err != nil { t.Fatalf("DecryptFile failed: %s", err) } gotContents := mustReadFile(t, fileName) -- cgit v1.3.1