aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2025-11-11 11:27:17 -0800
committerRose Hogenson <rosehogenson@posteo.net>2025-11-11 11:27:17 -0800
commit128d5c513b7a5c50a65f03480fc61b06fb5ef8e8 (patch)
tree4128102ee2d5fd31f7a4760aefb9eddce478422f
parentf57b3ba29eff4e5fcd02fa8923dc17f5eea1efb7 (diff)
downloadsym-128d5c513b7a5c50a65f03480fc61b06fb5ef8e8.tar.zst
Use subcommands package
This reduces a little bit of complexity. (does it?)
-rw-r--r--README.md2
-rw-r--r--dec.go84
-rw-r--r--dec_test.go64
-rw-r--r--enc.go132
-rw-r--r--enc_test.go65
-rw-r--r--go.mod1
-rw-r--r--go.sum2
-rw-r--r--oae.go2
-rw-r--r--sym.go86
-rw-r--r--sym_test.go86
10 files changed, 284 insertions, 240 deletions
diff --git a/README.md b/README.md
index 3b2c070..b55ff4f 100644
--- a/README.md
+++ b/README.md
@@ -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".
diff --git a/dec.go b/dec.go
index d1404b7..357f979 100644
--- a/dec.go
+++ b/dec.go
@@ -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)
}
})
}
diff --git a/enc.go b/enc.go
index fc54756..898fc10 100644
--- a/enc.go
+++ b/enc.go
@@ -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)
}
})
}
diff --git a/go.mod b/go.mod
index db929f6..59b5c03 100644
--- a/go.mod
+++ b/go.mod
@@ -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
diff --git a/go.sum b/go.sum
index 56eda83..4393630 100644
--- a/go.sum
+++ b/go.sum
@@ -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=
diff --git a/oae.go b/oae.go
index a3f14a5..0389a91 100644
--- a/oae.go
+++ b/oae.go
@@ -18,7 +18,7 @@ const (
segmentSize = 1024 * 1024
plaintextSegmentSize = segmentSize - aeadOverhead
- saltSize = 16
+ saltSize = 32
)
type segmentEncrypter struct {
diff --git a/sym.go b/sym.go
index b726562..b08657d 100644
--- a/sym.go
+++ b/sym.go
@@ -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")
+}