From 0cd822566776eddf8e47e61b7279aff4fa5481d7 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Fri, 24 Oct 2025 16:37:31 -0700 Subject: Get rid of defaultOptions --- dec.go | 12 +---------- dec_test.go | 56 +++++++++++++++++++++---------------------------- enc.go | 11 +--------- enc_test.go | 70 +++++++++++++++++++++++++++++-------------------------------- oae.go | 12 +++-------- oae_test.go | 6 ++---- pwhash.go | 6 +++--- sym.go | 13 ++++++++++-- sym_test.go | 18 +++++----------- 9 files changed, 83 insertions(+), 121 deletions(-) diff --git a/dec.go b/dec.go index ae9ec33..e67f4b3 100644 --- a/dec.go +++ b/dec.go @@ -10,7 +10,7 @@ import ( ) func (o *decryptOptions) decrypt(w io.Writer, r io.Reader, password string) error { - _, err := io.Copy(w, newDecryptingReader(r, password, o.memory)) + _, err := io.Copy(w, newDecryptingReader(r, password)) return err } @@ -27,25 +27,15 @@ func (f *decryptFlags) registerFlags(fs *flag.FlagSet) { type decryptOptions struct { decryptFlags - memory int passwordIn func() (string, error) stdin io.Reader stdout io.Writer } -var defaultDecryptOptions = decryptOptions{ - memory: defaultArgon2Memory, - passwordIn: termReadPassword, - stdin: os.Stdin, - stdout: os.Stdout, -} - func (o *decryptOptions) decryptFile(fileName string, password string) (err error) { var outFileName string if name, ok := strings.CutSuffix(fileName, ".enc"); ok { outFileName = name - } else if name, ok := strings.CutSuffix(fileName, ".enc.txt"); ok { - outFileName = name } else { outFileName = fileName + ".dec" } diff --git a/dec_test.go b/dec_test.go index 27fdd96..956d0a6 100644 --- a/dec_test.go +++ b/dec_test.go @@ -32,12 +32,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 := testEncryptOptions.encryptFile(fileName, password); err != nil { + if err := (&encryptOptions{}).encryptFile(fileName, password); err != nil { t.Fatalf("Failed to encrypt file: %s", err) } - decOpts := testDecryptOptions - decOpts.force = tc.force - err := decOpts.decryptFile(fileName+".enc", password) + err := (&decryptOptions{decryptFlags: decryptFlags{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) } @@ -71,7 +69,7 @@ func TestDecrypt_BadFileFormat(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, tc.fileContent) - err := testDecryptOptions.decryptFile(fileName, "asdf") + err := (&decryptOptions{}).decryptFile(fileName, "asdf") if err == nil { t.Errorf("DecryptFile succeeded for incorrect file format, want error") } @@ -86,11 +84,11 @@ func TestDecryptFile_WeirdName(t *testing.T) { fileContent := []byte("file content") fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, fileContent) - if err := testEncryptOptions.encryptFile(fileName, password); err != nil { + if err := (&encryptOptions{}).encryptFile(fileName, password); err != nil { t.Fatalf("EncryptFile failed: %s", err) } mustRename(t, fileName+".enc", fileName+".encrypted") - if err := testDecryptOptions.decryptFile(fileName+".encrypted", password); err != nil { + if err := (&decryptOptions{}).decryptFile(fileName+".encrypted", password); err != nil { t.Fatalf("DecryptFile failed: %s", err) } gotContents := mustReadFile(t, fileName+".encrypted.dec") @@ -102,7 +100,7 @@ func TestDecryptFile_WeirdName(t *testing.T) { func TestDecryptFile_NotFound(t *testing.T) { t.Parallel() - err := testDecryptOptions.decryptFile("my-nonexistent-file.txt", "asdf") + err := (&decryptOptions{}).decryptFile("my-nonexistent-file.txt", "asdf") if err == nil { t.Fatal("decryptFile succeeded for nonexistent file, want error") } @@ -115,9 +113,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 := testDecryptOptions - opts.force = true - err := opts.decryptFile(fileName, "asdf") + err := (&decryptOptions{decryptFlags: decryptFlags{force: true}}).decryptFile(fileName, "asdf") if err == nil { t.Fatal("decryptFile succeeded for unwritable file, want error") } @@ -147,15 +143,13 @@ func TestDecryptOptions_Run(t *testing.T) { fileContent := []byte("test file content") fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, fileContent) - if err := testEncryptOptions.encryptFile(fileName, password); err != nil { + if err := (&encryptOptions{}).encryptFile(fileName, password); err != nil { t.Errorf("EncryptFile failed: %s", err) } mustRemove(t, fileName) - opts := testDecryptOptions - opts.password = password - err := opts.run(fileName + ".enc") + err := (&decryptOptions{decryptFlags: decryptFlags{password: password}}).run(fileName + ".enc") if err != nil { - t.Errorf("dec failed: %s", err) + t.Errorf("decryptOptions.run failed: %s", err) } gotFileContents := mustReadFile(t, fileName) if !bytes.Equal(gotFileContents, fileContent) { @@ -166,7 +160,7 @@ func TestDecryptOptions_Run(t *testing.T) { func TestDecryptOptions_Run_UsageError(t *testing.T) { t.Parallel() - err := testDecryptOptions.run() + err := (&decryptOptions{}).run() if err == nil { t.Errorf("Run without -p when reading from stdin, want error") } @@ -175,9 +169,7 @@ func TestDecryptOptions_Run_UsageError(t *testing.T) { func TestDecryptOptions_Run_NotFound(t *testing.T) { t.Parallel() - opts := testDecryptOptions - opts.password = "asdf" - err := opts.run("my-nonexistent-file-name.txt") + err := (&decryptOptions{decryptFlags: decryptFlags{password: "asdf"}}).run("my-nonexistent-file-name.txt") if err == nil { t.Errorf("run succeeded with nonexistent file, want error") } @@ -189,15 +181,15 @@ func TestDecryptOptions_Run_Stdin(t *testing.T) { const password = "asdf" content := []byte("test contents") encrypted := new(bytes.Buffer) - if err := testEncryptOptions.encrypt(encrypted, bytes.NewReader(content), password); err != nil { + if err := (&encryptOptions{}).encrypt(encrypted, bytes.NewReader(content), password); err != nil { t.Fatalf("Failed to encrypt: %s", err) } gotContentBuf := new(bytes.Buffer) - opts := testDecryptOptions - opts.password = password - opts.stdin = bytes.NewReader(encrypted.Bytes()) - opts.stdout = gotContentBuf - if err := opts.run(); err != nil { + if err := (&decryptOptions{ + decryptFlags: decryptFlags{password: password}, + stdin: bytes.NewReader(encrypted.Bytes()), + stdout: gotContentBuf, + }).run(); err != nil { t.Fatalf("run failed: %s", err) } gotContent := gotContentBuf.Bytes() @@ -227,16 +219,16 @@ func TestDecryptOptions_Run_ReadPassword(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, []byte("test file content")) - if err := testEncryptOptions.encryptFile(fileName, password); err != nil { + if err := (&encryptOptions{}).encryptFile(fileName, password); err != nil { t.Errorf("EncryptFile failed: %s", err) } mustRemove(t, fileName) - opts := testDecryptOptions - opts.passwordIn = func() (string, error) { - return password, tc.err - } - err := opts.run(fileName + ".enc") + err := (&decryptOptions{ + 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) } diff --git a/enc.go b/enc.go index c24cba8..1217606 100644 --- a/enc.go +++ b/enc.go @@ -28,23 +28,14 @@ func (f *encryptFlags) registerFlags(fs *flag.FlagSet) { type encryptOptions struct { encryptFlags - memory int passwordIn func() (string, error) passwordOut io.Writer stdin io.Reader stdout io.Writer } -var defaultEncryptOptions = encryptOptions{ - memory: defaultArgon2Memory, - passwordIn: termReadPassword, - passwordOut: os.Stderr, - stdin: os.Stdin, - stdout: os.Stdout, -} - func (o *encryptOptions) encrypt(w io.Writer, r io.Reader, password string) error { - writer := newEncryptingWriter(w, password, o.memory) + writer := newEncryptingWriter(w, password) if _, err := io.Copy(writer, r); err != nil { return err } diff --git a/enc_test.go b/enc_test.go index 35d9f9e..f204283 100644 --- a/enc_test.go +++ b/enc_test.go @@ -31,9 +31,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")) - encOpts := testEncryptOptions - encOpts.force = tc.force - err := encOpts.encryptFile(fileName, "asdf") + err := (&encryptOptions{encryptFlags: encryptFlags{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) } @@ -44,7 +42,7 @@ func TestEncryptFile_Force(t *testing.T) { func TestEncryptFile_NotFound(t *testing.T) { t.Parallel() - err := testEncryptOptions.encryptFile("my-nonexistent-file.txt", "asdf") + err := (&encryptOptions{}).encryptFile("my-nonexistent-file.txt", "asdf") if err == nil { t.Fatal("encryptFile succeeded for nonexistent file, want error") } @@ -57,9 +55,7 @@ func TestEncryptFile_NoPermission(t *testing.T) { mustWriteFile(t, fileName, []byte("test file content")) mustWriteFile(t, fileName+".enc", nil) mustChmod(t, fileName+".enc", 0400) - opts := testEncryptOptions - opts.force = true - err := opts.encryptFile(fileName, "asdf") + err := (&encryptOptions{encryptFlags: encryptFlags{force: true}}).encryptFile(fileName, "asdf") if err == nil { t.Fatal("encryptFile succeeded for unwritable file, want error") } @@ -90,13 +86,11 @@ func TestEncryptOptions_Run(t *testing.T) { fileContent := []byte("test file content") fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, fileContent) - opts := testEncryptOptions - opts.password = password - if err := opts.run(fileName); err != nil { + if err := (&encryptOptions{encryptFlags: encryptFlags{password: password}}).run(fileName); err != nil { t.Fatalf("enc failed: %s", err) } mustRemove(t, fileName) - if err := testDecryptOptions.decryptFile(fileName+".enc", password); err != nil { + if err := (&decryptOptions{}).decryptFile(fileName+".enc", password); err != nil { t.Fatalf("Failed to decrypt encrypted file: %s", err) } gotFileContents := mustReadFile(t, fileName) @@ -129,11 +123,12 @@ func TestEncryptOptions_Run_UsageError(t *testing.T) { t.Run(tc.desc, func(t *testing.T) { t.Parallel() - opts := testEncryptOptions - opts.generatePassword = tc.generatePassword - opts.password = tc.password + opts := &encryptOptions{encryptFlags: encryptFlags{ + generatePassword: tc.generatePassword, + password: tc.password, + }} if err := opts.run(tc.files...); err == nil { - t.Errorf("run(%+v) succeeded, want error", opts) + t.Errorf("encryptOptions.run(%+v) succeeded, want error", opts) } }) } @@ -147,15 +142,16 @@ func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { mustWriteFile(t, fileName, fileContent) password := new(strings.Builder) - opts := testEncryptOptions - opts.generatePassword = true - opts.passwordOut = password + opts := &encryptOptions{ + encryptFlags: encryptFlags{generatePassword: true}, + passwordOut: password, + } if err := opts.run(fileName); err != nil { - t.Fatalf("enc(%+v) failed: %s", opts, err) + t.Fatalf("encryptOptions.run(%+v) failed: %s", opts, err) } pw := password.String() mustRemove(t, fileName) - if err := testDecryptOptions.decryptFile(fileName+".enc", pw); err != nil { + if err := (&decryptOptions{}).decryptFile(fileName+".enc", pw); err != nil { t.Fatalf("Failed to decrypt encrypted file with generated password %q: %s", pw, err) } gotFileContents := mustReadFile(t, fileName) @@ -172,15 +168,15 @@ func TestEncryptOptions_Run_Stdin(t *testing.T) { password = "asdf" ) stdout := new(strings.Builder) - opts := testEncryptOptions - opts.password = password - opts.stdin = strings.NewReader(input) - opts.stdout = stdout - if err := opts.run(); err != nil { - t.Errorf("enc(+%v) failed: %s", opts, err) + if err := (&encryptOptions{ + encryptFlags: encryptFlags{password: password}, + stdin: strings.NewReader(input), + stdout: stdout, + }).run(); err != nil { + t.Errorf("encryptOptions.run failed: %s", err) } got := new(strings.Builder) - if err := testDecryptOptions.decrypt(got, strings.NewReader(stdout.String()), password); err != nil { + if err := (&decryptOptions{}).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 { @@ -223,17 +219,17 @@ func TestEncryptOptions_Run_ReadPassword(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, []byte("test file content")) - opts := testEncryptOptions passwordI := 0 - opts.passwordIn = func() (string, error) { - if passwordI == len(tc.passwords) && tc.err != nil { - return "", tc.err - } - pw := tc.passwords[passwordI%len(tc.passwords)] - passwordI++ - return pw, nil - } - err := opts.run(fileName) + err := (&encryptOptions{ + passwordIn: func() (string, error) { + if passwordI == len(tc.passwords) && tc.err != nil { + return "", tc.err + } + pw := tc.passwords[passwordI%len(tc.passwords)] + passwordI++ + return pw, nil + }, + }).run(fileName) if gotErr := err != nil; gotErr != tc.wantErr { t.Errorf("encryptOptions.run returned error %v, want error? %t", err, tc.wantErr) } diff --git a/oae.go b/oae.go index a4eb6c8..915b562 100644 --- a/oae.go +++ b/oae.go @@ -22,17 +22,13 @@ const ( type segmentEncrypter struct { password string - memory int aead cipher.AEAD nonce [nonceSize]byte } func (se *segmentEncrypter) initialize(salt []byte) error { - key, err := hashPassword(se.password, salt, se.memory) - if err != nil { - return err - } + key := hashPassword(se.password, salt) block, err := aes.NewCipher(key) if err != nil { return err @@ -83,12 +79,11 @@ type encryptingWriter struct { initialized bool } -func newEncryptingWriter(w io.Writer, password string, memory int) *encryptingWriter { +func newEncryptingWriter(w io.Writer, password string) *encryptingWriter { return &encryptingWriter{ w: w, encrypter: segmentEncrypter{ password: password, - memory: memory, }, } } @@ -181,12 +176,11 @@ type decryptingReader struct { readFinalBlock bool } -func newDecryptingReader(r io.Reader, password string, memory int) *decryptingReader { +func newDecryptingReader(r io.Reader, password string) *decryptingReader { return &decryptingReader{ r: bufio.NewReaderSize(r, 0), // we only need .UnreadByte decrypter: segmentEncrypter{ password: password, - memory: memory, }, } } diff --git a/oae_test.go b/oae_test.go index 0cd5dd3..3602418 100644 --- a/oae_test.go +++ b/oae_test.go @@ -7,22 +7,20 @@ import ( "testing" ) -const testIters = 10 - func TestOAEReadWrite(t *testing.T) { t.Parallel() const password = "asdf" input := strings.Repeat("test input", 1024) out := new(bytes.Buffer) - writer := newEncryptingWriter(out, password, testIters) + writer := newEncryptingWriter(out, password) if _, err := io.WriteString(writer, input); err != nil { t.Fatalf("Failed to write: %s", err) } if err := writer.close(); err != nil { t.Fatalf("writer.Close() failed: %s", err) } - got, err := io.ReadAll(newDecryptingReader(bytes.NewReader(out.Bytes()), password, testIters)) + got, err := io.ReadAll(newDecryptingReader(bytes.NewReader(out.Bytes()), password)) if err != nil { t.Fatalf("Failed to decrypt: %s", err) } diff --git a/pwhash.go b/pwhash.go index dabafa7..f1c285d 100644 --- a/pwhash.go +++ b/pwhash.go @@ -7,10 +7,10 @@ import ( "golang.org/x/term" ) -const defaultArgon2Memory = 64 * 1024 +var argon2Memory = 64 * 1024 -func hashPassword(password string, salt []byte, memory int) ([]byte, error) { - return argon2.IDKey([]byte(password), salt, 1, uint32(memory), 4, 32), nil +func hashPassword(password string, salt []byte) []byte { + return argon2.IDKey([]byte(password), salt, 1, uint32(argon2Memory), 4, 32) } func termReadPassword() (string, error) { diff --git a/sym.go b/sym.go index 2405feb..e9eef94 100644 --- a/sym.go +++ b/sym.go @@ -15,9 +15,18 @@ type subcommand interface { func whichSubcommand(name string) (subcommand, bool) { switch filepath.Base(name) { case "enc": - return &defaultEncryptOptions, true + return &encryptOptions{ + passwordIn: termReadPassword, + passwordOut: os.Stderr, + stdin: os.Stdin, + stdout: os.Stdout, + }, true case "dec": - return &defaultDecryptOptions, true + return &decryptOptions{ + passwordIn: termReadPassword, + stdin: os.Stdin, + stdout: os.Stdout, + }, true default: return nil, false } diff --git a/sym_test.go b/sym_test.go index bb81ab8..3da9fc1 100644 --- a/sym_test.go +++ b/sym_test.go @@ -7,17 +7,9 @@ import ( "testing" ) -var testEncryptOptions = func() encryptOptions { - opts := defaultEncryptOptions - opts.memory = 1 - return opts -}() - -var testDecryptOptions = func() decryptOptions { - opts := defaultDecryptOptions - opts.memory = 1 - return opts -}() +func init() { + argon2Memory = 1 +} func mustWriteFile(t *testing.T, path string, content []byte) { t.Helper() @@ -67,11 +59,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 := testEncryptOptions.encryptFile(fileName, password); err != nil { + if err := (&encryptOptions{}).encryptFile(fileName, password); err != nil { t.Fatalf("EncryptFile failed: %s", err) } mustRemove(t, fileName) - if err := testDecryptOptions.decryptFile(fileName+".enc", password); err != nil { + if err := (&decryptOptions{}).decryptFile(fileName+".enc", password); err != nil { t.Fatalf("DecryptFile failed: %s", err) } gotContents := mustReadFile(t, fileName) -- cgit v1.3.1