aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--dec.go12
-rw-r--r--dec_test.go56
-rw-r--r--enc.go11
-rw-r--r--enc_test.go70
-rw-r--r--oae.go12
-rw-r--r--oae_test.go6
-rw-r--r--pwhash.go6
-rw-r--r--sym.go13
-rw-r--r--sym_test.go18
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)