From 3c2c07a69daf75d377ba196c28524c0fc6a16db6 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Thu, 23 Oct 2025 21:02:22 -0700 Subject: Move sym files into the root directory --- dec.go | 125 ++++++++++++++ dec_test.go | 326 +++++++++++++++++++++++++++++++++++ enc.go | 172 ++++++++++++++++++ enc_test.go | 242 ++++++++++++++++++++++++++ encryptionalg_string.go | 25 +++ internal/sym/dec.go | 125 -------------- internal/sym/dec_test.go | 326 ----------------------------------- internal/sym/enc.go | 172 ------------------ internal/sym/enc_test.go | 242 -------------------------- internal/sym/encryptionalg_string.go | 25 --- internal/sym/metadata.go | 21 --- internal/sym/oae.go | 293 ------------------------------- internal/sym/oae_test.go | 41 ----- internal/sym/pwhash.go | 51 ------ internal/sym/pwhash_string.go | 25 --- internal/sym/sym_test.go | 75 -------- metadata.go | 21 +++ oae.go | 293 +++++++++++++++++++++++++++++++ oae_test.go | 41 +++++ pwhash.go | 51 ++++++ pwhash_string.go | 25 +++ sym_test.go | 75 ++++++++ 22 files changed, 1396 insertions(+), 1396 deletions(-) create mode 100644 dec.go create mode 100644 dec_test.go create mode 100644 enc.go create mode 100644 enc_test.go create mode 100644 encryptionalg_string.go delete mode 100644 internal/sym/dec.go delete mode 100644 internal/sym/dec_test.go delete mode 100644 internal/sym/enc.go delete mode 100644 internal/sym/enc_test.go delete mode 100644 internal/sym/encryptionalg_string.go delete mode 100644 internal/sym/metadata.go delete mode 100644 internal/sym/oae.go delete mode 100644 internal/sym/oae_test.go delete mode 100644 internal/sym/pwhash.go delete mode 100644 internal/sym/pwhash_string.go delete mode 100644 internal/sym/sym_test.go create mode 100644 metadata.go create mode 100644 oae.go create mode 100644 oae_test.go create mode 100644 pwhash.go create mode 100644 pwhash_string.go create mode 100644 sym_test.go diff --git a/dec.go b/dec.go new file mode 100644 index 0000000..657da35 --- /dev/null +++ b/dec.go @@ -0,0 +1,125 @@ +package sym + +import ( + "encoding/binary" + "errors" + "flag" + "fmt" + "io" + "os" + "strings" +) + +func decrypt(w io.Writer, r io.Reader, password string) error { + fileFormat := make([]byte, 4) + if _, err := io.ReadFull(r, fileFormat); err != nil { + return err + } + if string(fileFormat) != magic { + return fmt.Errorf("bad file format") + } + header := new(fileMetadata) + if err := binary.Read(r, binary.BigEndian, header); err != nil { + return err + } + if err := header.validate(); err != nil { + return err + } + _, err := io.Copy(w, header.EncryptionMetadata.newDecryptingReader(r, password, &header.HashMetadata)) + return err +} + +type decryptFlags struct { + password string + force bool +} + +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 { + decryptFlags + + passwordIn func() (string, error) + stdin io.Reader + stdout io.Writer +} + +var DefaultDecryptOptions = DecryptOptions{ + 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" + } + fIn, err := os.Open(fileName) + if err != nil { + return err + } + defer fIn.Close() + fileOpts := os.O_CREATE | os.O_WRONLY + if o.force { + fileOpts |= os.O_TRUNC + } else { + fileOpts |= os.O_EXCL + } + fOut, err := os.OpenFile(outFileName, fileOpts, 0644) + if err != nil { + if errors.Is(err, os.ErrExist) { + return fmt.Errorf("output file %q exists (use -f to overwrite)", outFileName) + } + return err + } + defer func() { + fOut.Close() + if err != nil { + os.Remove(fOut.Name()) + } + }() + if err := decrypt(fOut, fIn, password); err != nil { + return fmt.Errorf("decrypt %q: %s", fileName, err) + } + return fOut.Close() +} + +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 { + if len(args) == 0 && o.password == "" { + return fmt.Errorf("-p is required when reading from stdin") + } + var password string + if o.password != "" { + password = o.password + } else { + var err error + password, err = o.readPassword() + if err != nil { + return err + } + } + if len(args) == 0 { + return decrypt(o.stdout, o.stdin, password) + } + for _, fileName := range args { + if err := o.decryptFile(fileName, password); err != nil { + return err + } + } + return nil +} diff --git a/dec_test.go b/dec_test.go new file mode 100644 index 0000000..e1d6d9b --- /dev/null +++ b/dec_test.go @@ -0,0 +1,326 @@ +package sym + +import ( + "bytes" + "encoding/binary" + "errors" + "flag" + "path/filepath" + "slices" + "strings" + "testing" +) + +func TestDecryptFile_Force(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + force bool + wantErr bool + }{{ + desc: "OutputExists", + force: false, + wantErr: true, + }, { + desc: "Force", + force: true, + wantErr: false, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + const password = "asdf" + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, []byte("test file content")) + if err := testEncryptOptions.encryptFile(fileName, password); err != nil { + t.Fatalf("Failed to encrypt file: %s", err) + } + decOpts := DefaultDecryptOptions + decOpts.force = tc.force + err := decOpts.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) + } + }) + } +} + +func encodeHeader(t *testing.T, f *fileMetadata) []byte { + t.Helper() + + if f.HashMetadata.PasswordHashType == pwHashInvalid { + f.HashMetadata.PasswordHashType = pwHashPBKDF2_HMAC_SHA256 + } + if f.HashMetadata.Iterations == 0 { + f.HashMetadata.Iterations = 10 + } + if f.HashMetadata.SaltSize == 0 { + f.HashMetadata.SaltSize = defaultSaltSize + } + if f.EncryptionMetadata.EncryptionType == encryptionAlgInvalid { + f.EncryptionMetadata.EncryptionType = encryptionAlgAES256_GCM + } + if f.EncryptionMetadata.SegmentSize == 0 { + f.EncryptionMetadata.SegmentSize = defaultSegmentSize + } + b, err := binary.Append([]byte(magic), binary.BigEndian, f) + if err != nil { + t.Fatalf("Bad file metadata: %s", err) + } + return b +} + +func TestDecrypt_BadFileFormat(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + fileContent []byte + }{{ + desc: "Empty", + fileContent: nil, + }, { + desc: "Short", + fileContent: []byte{0x80}, + }, { + desc: "BadFormat", + fileContent: []byte("bad file format"), + }, { + desc: "BadMagic", + fileContent: []byte("\x80asdf"), + }, { + desc: "BadHeader", + fileContent: []byte("\x80symasdf"), + }, { + desc: "BadVersion", + fileContent: encodeHeader(t, &fileMetadata{ + Version: -1, + }), + }, { + desc: "BadEncryptionAlg", + fileContent: encodeHeader(t, &fileMetadata{ + EncryptionMetadata: encryptionMetadata{ + EncryptionType: -1, + }, + }), + }, { + desc: "BadSegmentSize", + fileContent: encodeHeader(t, &fileMetadata{ + EncryptionMetadata: encryptionMetadata{ + SegmentSize: -1, + }, + }), + }, { + desc: "BadSaltSize", + fileContent: encodeHeader(t, &fileMetadata{ + HashMetadata: hashMetadata{ + SaltSize: -1, + }, + }), + }, { + desc: "BadPasswordHashType", + fileContent: encodeHeader(t, &fileMetadata{ + HashMetadata: hashMetadata{ + PasswordHashType: -1, + }, + }), + }, { + desc: "BadIterations", + fileContent: encodeHeader(t, &fileMetadata{ + HashMetadata: hashMetadata{ + Iterations: -1, + }, + }), + }, { + desc: "NoSalt", + fileContent: encodeHeader(t, &fileMetadata{}), + }, { + desc: "ShortSalt", + fileContent: slices.Concat( + encodeHeader(t, &fileMetadata{}), + []byte("asdf")), + }, { + desc: "BadContent", + fileContent: slices.Concat( + encodeHeader(t, &fileMetadata{}), + bytes.Repeat([]byte{0}, defaultSaltSize), + []byte("bad content")), + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, tc.fileContent) + err := DefaultDecryptOptions.decryptFile(fileName, "asdf") + if err == nil { + t.Errorf("DecryptFile succeeded for incorrect file format, want error") + } + }) + } +} + +func TestDecryptFile_WeirdName(t *testing.T) { + t.Parallel() + + const password = "asdf" + fileContent := []byte("file content") + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, fileContent) + if err := testEncryptOptions.encryptFile(fileName, password); err != nil { + t.Fatalf("EncryptFile failed: %s", err) + } + mustRename(t, fileName+".enc", fileName+".encrypted") + if err := DefaultDecryptOptions.decryptFile(fileName+".encrypted", password); err != nil { + t.Fatalf("DecryptFile failed: %s", err) + } + gotContents := mustReadFile(t, fileName+".encrypted.dec") + if !bytes.Equal(gotContents, fileContent) { + t.Errorf("contents differ") + } +} + +func TestDecryptFile_NotFound(t *testing.T) { + t.Parallel() + + err := DefaultDecryptOptions.decryptFile("my-nonexistent-file.txt", "asdf") + if err == nil { + t.Fatal("decryptFile succeeded for nonexistent file, want error") + } +} + +func TestDecryptFile_NoPermission(t *testing.T) { + t.Parallel() + + fileName := filepath.Join(t.TempDir(), "file.enc") + mustWriteFile(t, fileName, []byte("test file content")) + mustWriteFile(t, strings.TrimSuffix(fileName, ".enc"), nil) + mustChmod(t, strings.TrimSuffix(fileName, ".enc"), 0400) + opts := DefaultDecryptOptions + opts.force = true + err := opts.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) { + t.Parallel() + + const password = "asdf" + fileContent := []byte("test file content") + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, fileContent) + if err := testEncryptOptions.encryptFile(fileName, password); err != nil { + t.Errorf("EncryptFile failed: %s", err) + } + mustRemove(t, fileName) + opts := DefaultDecryptOptions + opts.password = password + 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) + } +} + +func TestDecryptOptions_Run_UsageError(t *testing.T) { + t.Parallel() + + err := DefaultDecryptOptions.Run() + if err == nil { + t.Errorf("Run without -p when reading from stdin, want error") + } +} + +func TestDecryptOptions_Run_NotFound(t *testing.T) { + t.Parallel() + + opts := DefaultDecryptOptions + opts.password = "asdf" + err := opts.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) { + t.Parallel() + + const password = "asdf" + content := []byte("test contents") + encrypted := new(bytes.Buffer) + if err := testEncryptOptions.encrypt(encrypted, bytes.NewReader(content), password); err != nil { + t.Fatalf("Failed to encrypt: %s", err) + } + gotContentBuf := new(bytes.Buffer) + 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) + } + gotContent := gotContentBuf.Bytes() + if !bytes.Equal(gotContent, content) { + t.Errorf("dec returned incorrect contents %q, want %q", gotContent, content) + } +} + +func TestDecryptOptions_Run_ReadPassword(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + err error + wantErr bool + }{{ + desc: "Ok", + }, { + desc: "Err", + err: errors.New("test error"), + wantErr: true, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + const password = "asdf" + + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, []byte("test file content")) + if err := testEncryptOptions.encryptFile(fileName, password); err != nil { + t.Errorf("EncryptFile failed: %s", err) + } + mustRemove(t, fileName) + + opts := DefaultDecryptOptions + opts.passwordIn = func() (string, error) { + return password, tc.err + } + 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) + } + }) + } +} diff --git a/enc.go b/enc.go new file mode 100644 index 0000000..623174c --- /dev/null +++ b/enc.go @@ -0,0 +1,172 @@ +package sym + +import ( + "crypto/rand" + "encoding/binary" + "errors" + "flag" + "fmt" + "io" + "os" + "strings" + + "roseh.moe/pkg/wordlist" +) + +type encryptFlags struct { + generatePassword bool + password string + force bool +} + +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 { + encryptFlags + + iterations int + passwordIn func() (string, error) + passwordOut io.Writer + stdin io.Reader + stdout io.Writer +} + +var DefaultEncryptOptions = EncryptOptions{ + iterations: defaultPBKDF2Iters, + passwordIn: termReadPassword, + passwordOut: os.Stderr, + stdin: os.Stdin, + stdout: os.Stdout, +} + +func (o *EncryptOptions) encrypt(w io.Writer, r io.Reader, password string) error { + if _, err := io.WriteString(w, magic); err != nil { + return err + } + header := &fileMetadata{ + Version: 0, + HashMetadata: hashMetadata{ + PasswordHashType: pwHashPBKDF2_HMAC_SHA256, + Iterations: int32(o.iterations), + SaltSize: defaultSaltSize, + }, + EncryptionMetadata: encryptionMetadata{ + EncryptionType: encryptionAlgAES256_GCM, + SegmentSize: defaultSegmentSize, + }, + } + if err := binary.Write(w, binary.BigEndian, header); err != nil { + return err + } + writer := header.EncryptionMetadata.newEncryptingWriter(w, password, &header.HashMetadata) + if _, err := io.Copy(writer, r); err != nil { + return err + } + return writer.Close() +} + +func (o *EncryptOptions) 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 { + fileOpts |= os.O_TRUNC + } else { + fileOpts |= os.O_EXCL + } + fOut, err := os.OpenFile(fileName+".enc", fileOpts, 0644) + if err != nil { + if errors.Is(err, os.ErrExist) { + return fmt.Errorf("output file %q exists (use -f to overwrite)", fileName+".enc") + } + return err + } + defer func() { + fOut.Close() + if err != nil { + os.Remove(fOut.Name()) + } + }() + if err = o.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 + } + return "", fmt.Errorf("too many attempts") +} + +func (o *EncryptOptions) Run(args ...string) error { + if o.generatePassword && o.password != "" { + return fmt.Errorf("-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") + } + var password string + if o.password != "" { + password = o.password + } else if o.generatePassword { + const nWords = 10 + buf := make([]byte, 2*nWords) + rand.Read(buf) + words := make([]string, nWords) + for i := range words { + words[i] = wordlist.Words[binary.NativeEndian.Uint16(buf[2*i:])&0x1fff] + } + password = strings.Join(words, " ") + fmt.Fprint(os.Stderr, "Your password: ") + fmt.Fprint(o.passwordOut, password) + fmt.Fprintln(os.Stderr) + } else { + var err error + if password, err = o.readPassword(); err != nil { + return err + } + } + if len(args) == 0 { + return o.encrypt(o.stdout, o.stdin, password) + } + for _, fileName := range args { + if err := o.encryptFile(fileName, password); err != nil { + return err + } + } + return nil +} diff --git a/enc_test.go b/enc_test.go new file mode 100644 index 0000000..3a0c493 --- /dev/null +++ b/enc_test.go @@ -0,0 +1,242 @@ +package sym + +import ( + "bytes" + "errors" + "flag" + "path/filepath" + "strings" + "testing" +) + +func TestEncryptFile_Force(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + force bool + wantErr bool + }{{ + desc: "OutputExists", + force: false, + wantErr: true, + }, { + desc: "Force", + force: true, + wantErr: false, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + 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") + 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) + } + }) + } +} + +func TestEncryptFile_NotFound(t *testing.T) { + t.Parallel() + + err := testEncryptOptions.encryptFile("my-nonexistent-file.txt", "asdf") + if err == nil { + t.Fatal("encryptFile succeeded for nonexistent file, want error") + } +} + +func TestEncryptFile_NoPermission(t *testing.T) { + t.Parallel() + + fileName := filepath.Join(t.TempDir(), "file") + 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") + 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) { + t.Parallel() + + const password = "asdf" + 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 { + t.Fatalf("enc failed: %s", err) + } + mustRemove(t, fileName) + if err := DefaultDecryptOptions.decryptFile(fileName+".enc", password); err != nil { + t.Fatalf("Failed to decrypt encrypted file: %s", err) + } + gotFileContents := mustReadFile(t, fileName) + if !bytes.Equal(gotFileContents, fileContent) { + t.Errorf("encrypt round trip returned incorrect contents %q, want %q", gotFileContents, fileContent) + } +} + +func TestEncryptOptions_Run_UsageError(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + generatePassword bool + password string + files []string + }{{ + desc: "GeneratePasswordAndPassword", + generatePassword: true, + password: "asdf", + }, { + desc: "MissingPasswordStdin", + generatePassword: false, + password: "", + }, { + desc: "NonexistentFile", + password: "asdf", + files: []string{"my-nonexistent-file.txt"}, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + 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) + } + }) + } +} + +func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { + t.Parallel() + + fileContent := []byte("test file content") + fileName := filepath.Join(t.TempDir(), "file") + mustWriteFile(t, fileName, fileContent) + + password := new(strings.Builder) + opts := testEncryptOptions + opts.generatePassword = true + opts.passwordOut = password + 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 { + t.Fatalf("Failed to decrypt encrypted file with generated password %q: %s", pw, err) + } + gotFileContents := mustReadFile(t, fileName) + if !bytes.Equal(gotFileContents, fileContent) { + t.Errorf("encrypt round trip returned incorrect contents %q, want %q", gotFileContents, fileContent) + } +} + +func TestEncryptOptions_Run_Stdin(t *testing.T) { + t.Parallel() + + const ( + input = "test input" + 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) + } + got := new(strings.Builder) + if err := 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 { + t.Errorf("Encrypt round-trip to stdout returned incorrect contents: %q, want %q", got, want) + } +} + +func TestEncryptOptions_Run_ReadPassword(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + desc string + passwords []string + err error + wantErr bool + }{{ + desc: "Ok", + passwords: []string{"asdf"}, + }, { + desc: "EmptyPassword", + passwords: []string{""}, + wantErr: true, + }, { + desc: "PasswordsDoNotMatch", + passwords: []string{"asdf", "jkl"}, + wantErr: true, + }, { + desc: "ReadPasswordErr", + err: errors.New("test error"), + wantErr: true, + }, { + desc: "RepeatPasswordErr", + passwords: []string{"asdf"}, + err: errors.New("test error"), + wantErr: true, + }} { + t.Run(tc.desc, func(t *testing.T) { + t.Parallel() + + 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) + if gotErr := err != nil; gotErr != 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 new file mode 100644 index 0000000..b7384ee --- /dev/null +++ b/encryptionalg_string.go @@ -0,0 +1,25 @@ +// Code generated by "stringer -type=encryptionAlg -linecomment"; DO NOT EDIT. + +package sym + +import "strconv" + +func _() { + // An "invalid array index" compiler error signifies that the constant values have changed. + // Re-run the stringer command to generate them again. + var x [1]struct{} + _ = x[encryptionAlgInvalid-0] + _ = x[encryptionAlgAES256_GCM-1] +} + +const _encryptionAlg_name = "encryptionAlgInvalidAES-256-GCM" + +var _encryptionAlg_index = [...]uint8{0, 20, 31} + +func (i encryptionAlg) String() string { + idx := int(i) - 0 + if i < 0 || idx >= len(_encryptionAlg_index)-1 { + return "encryptionAlg(" + strconv.FormatInt(int64(i), 10) + ")" + } + return _encryptionAlg_name[_encryptionAlg_index[idx]:_encryptionAlg_index[idx+1]] +} diff --git a/internal/sym/dec.go b/internal/sym/dec.go deleted file mode 100644 index 657da35..0000000 --- a/internal/sym/dec.go +++ /dev/null @@ -1,125 +0,0 @@ -package sym - -import ( - "encoding/binary" - "errors" - "flag" - "fmt" - "io" - "os" - "strings" -) - -func decrypt(w io.Writer, r io.Reader, password string) error { - fileFormat := make([]byte, 4) - if _, err := io.ReadFull(r, fileFormat); err != nil { - return err - } - if string(fileFormat) != magic { - return fmt.Errorf("bad file format") - } - header := new(fileMetadata) - if err := binary.Read(r, binary.BigEndian, header); err != nil { - return err - } - if err := header.validate(); err != nil { - return err - } - _, err := io.Copy(w, header.EncryptionMetadata.newDecryptingReader(r, password, &header.HashMetadata)) - return err -} - -type decryptFlags struct { - password string - force bool -} - -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 { - decryptFlags - - passwordIn func() (string, error) - stdin io.Reader - stdout io.Writer -} - -var DefaultDecryptOptions = DecryptOptions{ - 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" - } - fIn, err := os.Open(fileName) - if err != nil { - return err - } - defer fIn.Close() - fileOpts := os.O_CREATE | os.O_WRONLY - if o.force { - fileOpts |= os.O_TRUNC - } else { - fileOpts |= os.O_EXCL - } - fOut, err := os.OpenFile(outFileName, fileOpts, 0644) - if err != nil { - if errors.Is(err, os.ErrExist) { - return fmt.Errorf("output file %q exists (use -f to overwrite)", outFileName) - } - return err - } - defer func() { - fOut.Close() - if err != nil { - os.Remove(fOut.Name()) - } - }() - if err := decrypt(fOut, fIn, password); err != nil { - return fmt.Errorf("decrypt %q: %s", fileName, err) - } - return fOut.Close() -} - -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 { - if len(args) == 0 && o.password == "" { - return fmt.Errorf("-p is required when reading from stdin") - } - var password string - if o.password != "" { - password = o.password - } else { - var err error - password, err = o.readPassword() - if err != nil { - return err - } - } - if len(args) == 0 { - return decrypt(o.stdout, o.stdin, password) - } - for _, fileName := range args { - if err := o.decryptFile(fileName, password); err != nil { - return err - } - } - return nil -} diff --git a/internal/sym/dec_test.go b/internal/sym/dec_test.go deleted file mode 100644 index e1d6d9b..0000000 --- a/internal/sym/dec_test.go +++ /dev/null @@ -1,326 +0,0 @@ -package sym - -import ( - "bytes" - "encoding/binary" - "errors" - "flag" - "path/filepath" - "slices" - "strings" - "testing" -) - -func TestDecryptFile_Force(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - force bool - wantErr bool - }{{ - desc: "OutputExists", - force: false, - wantErr: true, - }, { - desc: "Force", - force: true, - wantErr: false, - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - const password = "asdf" - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, []byte("test file content")) - if err := testEncryptOptions.encryptFile(fileName, password); err != nil { - t.Fatalf("Failed to encrypt file: %s", err) - } - decOpts := DefaultDecryptOptions - decOpts.force = tc.force - err := decOpts.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) - } - }) - } -} - -func encodeHeader(t *testing.T, f *fileMetadata) []byte { - t.Helper() - - if f.HashMetadata.PasswordHashType == pwHashInvalid { - f.HashMetadata.PasswordHashType = pwHashPBKDF2_HMAC_SHA256 - } - if f.HashMetadata.Iterations == 0 { - f.HashMetadata.Iterations = 10 - } - if f.HashMetadata.SaltSize == 0 { - f.HashMetadata.SaltSize = defaultSaltSize - } - if f.EncryptionMetadata.EncryptionType == encryptionAlgInvalid { - f.EncryptionMetadata.EncryptionType = encryptionAlgAES256_GCM - } - if f.EncryptionMetadata.SegmentSize == 0 { - f.EncryptionMetadata.SegmentSize = defaultSegmentSize - } - b, err := binary.Append([]byte(magic), binary.BigEndian, f) - if err != nil { - t.Fatalf("Bad file metadata: %s", err) - } - return b -} - -func TestDecrypt_BadFileFormat(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - fileContent []byte - }{{ - desc: "Empty", - fileContent: nil, - }, { - desc: "Short", - fileContent: []byte{0x80}, - }, { - desc: "BadFormat", - fileContent: []byte("bad file format"), - }, { - desc: "BadMagic", - fileContent: []byte("\x80asdf"), - }, { - desc: "BadHeader", - fileContent: []byte("\x80symasdf"), - }, { - desc: "BadVersion", - fileContent: encodeHeader(t, &fileMetadata{ - Version: -1, - }), - }, { - desc: "BadEncryptionAlg", - fileContent: encodeHeader(t, &fileMetadata{ - EncryptionMetadata: encryptionMetadata{ - EncryptionType: -1, - }, - }), - }, { - desc: "BadSegmentSize", - fileContent: encodeHeader(t, &fileMetadata{ - EncryptionMetadata: encryptionMetadata{ - SegmentSize: -1, - }, - }), - }, { - desc: "BadSaltSize", - fileContent: encodeHeader(t, &fileMetadata{ - HashMetadata: hashMetadata{ - SaltSize: -1, - }, - }), - }, { - desc: "BadPasswordHashType", - fileContent: encodeHeader(t, &fileMetadata{ - HashMetadata: hashMetadata{ - PasswordHashType: -1, - }, - }), - }, { - desc: "BadIterations", - fileContent: encodeHeader(t, &fileMetadata{ - HashMetadata: hashMetadata{ - Iterations: -1, - }, - }), - }, { - desc: "NoSalt", - fileContent: encodeHeader(t, &fileMetadata{}), - }, { - desc: "ShortSalt", - fileContent: slices.Concat( - encodeHeader(t, &fileMetadata{}), - []byte("asdf")), - }, { - desc: "BadContent", - fileContent: slices.Concat( - encodeHeader(t, &fileMetadata{}), - bytes.Repeat([]byte{0}, defaultSaltSize), - []byte("bad content")), - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, tc.fileContent) - err := DefaultDecryptOptions.decryptFile(fileName, "asdf") - if err == nil { - t.Errorf("DecryptFile succeeded for incorrect file format, want error") - } - }) - } -} - -func TestDecryptFile_WeirdName(t *testing.T) { - t.Parallel() - - const password = "asdf" - fileContent := []byte("file content") - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, fileContent) - if err := testEncryptOptions.encryptFile(fileName, password); err != nil { - t.Fatalf("EncryptFile failed: %s", err) - } - mustRename(t, fileName+".enc", fileName+".encrypted") - if err := DefaultDecryptOptions.decryptFile(fileName+".encrypted", password); err != nil { - t.Fatalf("DecryptFile failed: %s", err) - } - gotContents := mustReadFile(t, fileName+".encrypted.dec") - if !bytes.Equal(gotContents, fileContent) { - t.Errorf("contents differ") - } -} - -func TestDecryptFile_NotFound(t *testing.T) { - t.Parallel() - - err := DefaultDecryptOptions.decryptFile("my-nonexistent-file.txt", "asdf") - if err == nil { - t.Fatal("decryptFile succeeded for nonexistent file, want error") - } -} - -func TestDecryptFile_NoPermission(t *testing.T) { - t.Parallel() - - fileName := filepath.Join(t.TempDir(), "file.enc") - mustWriteFile(t, fileName, []byte("test file content")) - mustWriteFile(t, strings.TrimSuffix(fileName, ".enc"), nil) - mustChmod(t, strings.TrimSuffix(fileName, ".enc"), 0400) - opts := DefaultDecryptOptions - opts.force = true - err := opts.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) { - t.Parallel() - - const password = "asdf" - fileContent := []byte("test file content") - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, fileContent) - if err := testEncryptOptions.encryptFile(fileName, password); err != nil { - t.Errorf("EncryptFile failed: %s", err) - } - mustRemove(t, fileName) - opts := DefaultDecryptOptions - opts.password = password - 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) - } -} - -func TestDecryptOptions_Run_UsageError(t *testing.T) { - t.Parallel() - - err := DefaultDecryptOptions.Run() - if err == nil { - t.Errorf("Run without -p when reading from stdin, want error") - } -} - -func TestDecryptOptions_Run_NotFound(t *testing.T) { - t.Parallel() - - opts := DefaultDecryptOptions - opts.password = "asdf" - err := opts.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) { - t.Parallel() - - const password = "asdf" - content := []byte("test contents") - encrypted := new(bytes.Buffer) - if err := testEncryptOptions.encrypt(encrypted, bytes.NewReader(content), password); err != nil { - t.Fatalf("Failed to encrypt: %s", err) - } - gotContentBuf := new(bytes.Buffer) - 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) - } - gotContent := gotContentBuf.Bytes() - if !bytes.Equal(gotContent, content) { - t.Errorf("dec returned incorrect contents %q, want %q", gotContent, content) - } -} - -func TestDecryptOptions_Run_ReadPassword(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - err error - wantErr bool - }{{ - desc: "Ok", - }, { - desc: "Err", - err: errors.New("test error"), - wantErr: true, - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - const password = "asdf" - - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, []byte("test file content")) - if err := testEncryptOptions.encryptFile(fileName, password); err != nil { - t.Errorf("EncryptFile failed: %s", err) - } - mustRemove(t, fileName) - - opts := DefaultDecryptOptions - opts.passwordIn = func() (string, error) { - return password, tc.err - } - 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) - } - }) - } -} diff --git a/internal/sym/enc.go b/internal/sym/enc.go deleted file mode 100644 index 623174c..0000000 --- a/internal/sym/enc.go +++ /dev/null @@ -1,172 +0,0 @@ -package sym - -import ( - "crypto/rand" - "encoding/binary" - "errors" - "flag" - "fmt" - "io" - "os" - "strings" - - "roseh.moe/pkg/wordlist" -) - -type encryptFlags struct { - generatePassword bool - password string - force bool -} - -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 { - encryptFlags - - iterations int - passwordIn func() (string, error) - passwordOut io.Writer - stdin io.Reader - stdout io.Writer -} - -var DefaultEncryptOptions = EncryptOptions{ - iterations: defaultPBKDF2Iters, - passwordIn: termReadPassword, - passwordOut: os.Stderr, - stdin: os.Stdin, - stdout: os.Stdout, -} - -func (o *EncryptOptions) encrypt(w io.Writer, r io.Reader, password string) error { - if _, err := io.WriteString(w, magic); err != nil { - return err - } - header := &fileMetadata{ - Version: 0, - HashMetadata: hashMetadata{ - PasswordHashType: pwHashPBKDF2_HMAC_SHA256, - Iterations: int32(o.iterations), - SaltSize: defaultSaltSize, - }, - EncryptionMetadata: encryptionMetadata{ - EncryptionType: encryptionAlgAES256_GCM, - SegmentSize: defaultSegmentSize, - }, - } - if err := binary.Write(w, binary.BigEndian, header); err != nil { - return err - } - writer := header.EncryptionMetadata.newEncryptingWriter(w, password, &header.HashMetadata) - if _, err := io.Copy(writer, r); err != nil { - return err - } - return writer.Close() -} - -func (o *EncryptOptions) 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 { - fileOpts |= os.O_TRUNC - } else { - fileOpts |= os.O_EXCL - } - fOut, err := os.OpenFile(fileName+".enc", fileOpts, 0644) - if err != nil { - if errors.Is(err, os.ErrExist) { - return fmt.Errorf("output file %q exists (use -f to overwrite)", fileName+".enc") - } - return err - } - defer func() { - fOut.Close() - if err != nil { - os.Remove(fOut.Name()) - } - }() - if err = o.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 - } - return "", fmt.Errorf("too many attempts") -} - -func (o *EncryptOptions) Run(args ...string) error { - if o.generatePassword && o.password != "" { - return fmt.Errorf("-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") - } - var password string - if o.password != "" { - password = o.password - } else if o.generatePassword { - const nWords = 10 - buf := make([]byte, 2*nWords) - rand.Read(buf) - words := make([]string, nWords) - for i := range words { - words[i] = wordlist.Words[binary.NativeEndian.Uint16(buf[2*i:])&0x1fff] - } - password = strings.Join(words, " ") - fmt.Fprint(os.Stderr, "Your password: ") - fmt.Fprint(o.passwordOut, password) - fmt.Fprintln(os.Stderr) - } else { - var err error - if password, err = o.readPassword(); err != nil { - return err - } - } - if len(args) == 0 { - return o.encrypt(o.stdout, o.stdin, password) - } - for _, fileName := range args { - if err := o.encryptFile(fileName, password); err != nil { - return err - } - } - return nil -} diff --git a/internal/sym/enc_test.go b/internal/sym/enc_test.go deleted file mode 100644 index 3a0c493..0000000 --- a/internal/sym/enc_test.go +++ /dev/null @@ -1,242 +0,0 @@ -package sym - -import ( - "bytes" - "errors" - "flag" - "path/filepath" - "strings" - "testing" -) - -func TestEncryptFile_Force(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - force bool - wantErr bool - }{{ - desc: "OutputExists", - force: false, - wantErr: true, - }, { - desc: "Force", - force: true, - wantErr: false, - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - 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") - 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) - } - }) - } -} - -func TestEncryptFile_NotFound(t *testing.T) { - t.Parallel() - - err := testEncryptOptions.encryptFile("my-nonexistent-file.txt", "asdf") - if err == nil { - t.Fatal("encryptFile succeeded for nonexistent file, want error") - } -} - -func TestEncryptFile_NoPermission(t *testing.T) { - t.Parallel() - - fileName := filepath.Join(t.TempDir(), "file") - 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") - 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) { - t.Parallel() - - const password = "asdf" - 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 { - t.Fatalf("enc failed: %s", err) - } - mustRemove(t, fileName) - if err := DefaultDecryptOptions.decryptFile(fileName+".enc", password); err != nil { - t.Fatalf("Failed to decrypt encrypted file: %s", err) - } - gotFileContents := mustReadFile(t, fileName) - if !bytes.Equal(gotFileContents, fileContent) { - t.Errorf("encrypt round trip returned incorrect contents %q, want %q", gotFileContents, fileContent) - } -} - -func TestEncryptOptions_Run_UsageError(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - generatePassword bool - password string - files []string - }{{ - desc: "GeneratePasswordAndPassword", - generatePassword: true, - password: "asdf", - }, { - desc: "MissingPasswordStdin", - generatePassword: false, - password: "", - }, { - desc: "NonexistentFile", - password: "asdf", - files: []string{"my-nonexistent-file.txt"}, - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - 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) - } - }) - } -} - -func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { - t.Parallel() - - fileContent := []byte("test file content") - fileName := filepath.Join(t.TempDir(), "file") - mustWriteFile(t, fileName, fileContent) - - password := new(strings.Builder) - opts := testEncryptOptions - opts.generatePassword = true - opts.passwordOut = password - 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 { - t.Fatalf("Failed to decrypt encrypted file with generated password %q: %s", pw, err) - } - gotFileContents := mustReadFile(t, fileName) - if !bytes.Equal(gotFileContents, fileContent) { - t.Errorf("encrypt round trip returned incorrect contents %q, want %q", gotFileContents, fileContent) - } -} - -func TestEncryptOptions_Run_Stdin(t *testing.T) { - t.Parallel() - - const ( - input = "test input" - 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) - } - got := new(strings.Builder) - if err := 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 { - t.Errorf("Encrypt round-trip to stdout returned incorrect contents: %q, want %q", got, want) - } -} - -func TestEncryptOptions_Run_ReadPassword(t *testing.T) { - t.Parallel() - - for _, tc := range []struct { - desc string - passwords []string - err error - wantErr bool - }{{ - desc: "Ok", - passwords: []string{"asdf"}, - }, { - desc: "EmptyPassword", - passwords: []string{""}, - wantErr: true, - }, { - desc: "PasswordsDoNotMatch", - passwords: []string{"asdf", "jkl"}, - wantErr: true, - }, { - desc: "ReadPasswordErr", - err: errors.New("test error"), - wantErr: true, - }, { - desc: "RepeatPasswordErr", - passwords: []string{"asdf"}, - err: errors.New("test error"), - wantErr: true, - }} { - t.Run(tc.desc, func(t *testing.T) { - t.Parallel() - - 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) - if gotErr := err != nil; gotErr != tc.wantErr { - t.Errorf("EncryptOptions.Run returned error %v, want error? %t", err, tc.wantErr) - } - }) - } -} diff --git a/internal/sym/encryptionalg_string.go b/internal/sym/encryptionalg_string.go deleted file mode 100644 index b7384ee..0000000 --- a/internal/sym/encryptionalg_string.go +++ /dev/null @@ -1,25 +0,0 @@ -// Code generated by "stringer -type=encryptionAlg -linecomment"; DO NOT EDIT. - -package sym - -import "strconv" - -func _() { - // An "invalid array index" compiler error signifies that the constant values have changed. - // Re-run the stringer command to generate them again. - var x [1]struct{} - _ = x[encryptionAlgInvalid-0] - _ = x[encryptionAlgAES256_GCM-1] -} - -const _encryptionAlg_name = "encryptionAlgInvalidAES-256-GCM" - -var _encryptionAlg_index = [...]uint8{0, 20, 31} - -func (i encryptionAlg) String() string { - idx := int(i) - 0 - if i < 0 || idx >= len(_encryptionAlg_index)-1 { - return "encryptionAlg(" + strconv.FormatInt(int64(i), 10) + ")" - } - return _encryptionAlg_name[_encryptionAlg_index[idx]:_encryptionAlg_index[idx+1]] -} diff --git a/internal/sym/metadata.go b/internal/sym/metadata.go deleted file mode 100644 index 3875a63..0000000 --- a/internal/sym/metadata.go +++ /dev/null @@ -1,21 +0,0 @@ -package sym - -import "fmt" - -const magic = "\x80sym" - -type fileMetadata struct { - Version int8 - HashMetadata hashMetadata - EncryptionMetadata encryptionMetadata -} - -func (f *fileMetadata) validate() error { - if f.Version != 0 { - return fmt.Errorf("bad version") - } - if err := f.HashMetadata.validate(); err != nil { - return err - } - return f.EncryptionMetadata.validate() -} diff --git a/internal/sym/oae.go b/internal/sym/oae.go deleted file mode 100644 index e40cb60..0000000 --- a/internal/sym/oae.go +++ /dev/null @@ -1,293 +0,0 @@ -package sym - -import ( - "bufio" - "bytes" - "crypto/aes" - "crypto/cipher" - "crypto/rand" - "errors" - "fmt" - "io" -) - -const ( - nonceSize = 12 - aeadOverhead = 16 - - defaultSegmentSize = 1024 * 1024 - - defaultSaltSize = 32 -) - -//go:generate go tool stringer -type=encryptionAlg -linecomment -type encryptionAlg int8 - -const ( - encryptionAlgInvalid encryptionAlg = iota - encryptionAlgAES256_GCM // AES-256-GCM -) - -type encryptionMetadata struct { - EncryptionType encryptionAlg - SegmentSize int32 -} - -func (e *encryptionMetadata) validate() error { - if e.EncryptionType != encryptionAlgAES256_GCM { - return fmt.Errorf("invalid encryption alg %q", e.EncryptionType) - } - if e.SegmentSize <= 0 || e.SegmentSize > defaultSegmentSize { - return fmt.Errorf("segment size too long") - } - return nil -} - -func (e *encryptionMetadata) plaintextSegmentSize() int { - return int(e.SegmentSize) - aeadOverhead -} - -type segmentEncrypter struct { - hashMetadata hashMetadata - encryptionMetadata encryptionMetadata - password string - - aead cipher.AEAD - nonce [nonceSize]byte -} - -func (se *segmentEncrypter) initialize(salt []byte) error { - key, err := se.hashMetadata.hashPassword(se.password, salt) - if err != nil { - return err - } - block, err := aes.NewCipher(key) - if err != nil { - return err - } - se.aead, err = cipher.NewGCM(block) - return err -} - -func (se *segmentEncrypter) ad(lastSegment bool) ([]byte, error) { - // Increment counter - for i := 0; ; i++ { - if i == len(se.nonce) { - return nil, errors.New("counter overflowed") - } - se.nonce[i]++ - if se.nonce[i] != 0 { - break - } - } - ad := make([]byte, len(se.nonce)+1) - copy(ad, se.nonce[:]) - if lastSegment { - ad[len(ad)-1] = 1 - } - return ad, nil -} - -func (se *segmentEncrypter) encrypt(out, buf []byte, lastSegment bool) ([]byte, error) { - ad, err := se.ad(lastSegment) - if err != nil { - return nil, err - } - return se.aead.Seal(out, se.nonce[:], buf, ad), nil -} - -func (se *segmentEncrypter) decrypt(out, buf []byte, lastSegment bool) ([]byte, error) { - ad, err := se.ad(lastSegment) - if err != nil { - return nil, err - } - return se.aead.Open(out, se.nonce[:], buf, ad) -} - -type encryptingWriter struct { - w io.Writer - encrypter segmentEncrypter - buf []byte - initialized bool -} - -func (e *encryptionMetadata) newEncryptingWriter(w io.Writer, password string, passwordMetadata *hashMetadata) *encryptingWriter { - return &encryptingWriter{ - w: w, - encrypter: segmentEncrypter{ - hashMetadata: *passwordMetadata, - encryptionMetadata: *e, - password: password, - }, - } -} - -func (w *encryptingWriter) initialize() error { - if w.initialized { - return nil - } - header := make([]byte, w.encrypter.hashMetadata.SaltSize) - rand.Read(header) - if err := w.encrypter.initialize(header); err != nil { - return err - } - if _, err := w.w.Write(header); err != nil { - return err - } - w.buf = make([]byte, 0, w.encrypter.encryptionMetadata.SegmentSize) - w.initialized = true - return nil -} - -func (w *encryptingWriter) writeBuf(lastSegment bool) error { - var err error - if w.buf, err = w.encrypter.encrypt(w.buf[:0], w.buf, lastSegment); err != nil { - return err - } - if _, err := w.w.Write(w.buf); err != nil { - return err - } - w.buf = w.buf[:0] - return nil -} - -func (w *encryptingWriter) Write(buf []byte) (int, error) { - if err := w.initialize(); err != nil { - return 0, err - } - nn := 0 - for len(buf) > 0 { - if len(w.buf) == w.encrypter.encryptionMetadata.plaintextSegmentSize() { - if err := w.writeBuf(false); err != nil { - return nn, err - } - } - n := copy(w.buf[len(w.buf):w.encrypter.encryptionMetadata.plaintextSegmentSize()], buf) - nn += n - w.buf = w.buf[:len(w.buf)+n] - buf = buf[n:] - } - return nn, nil -} - -func (w *encryptingWriter) Close() error { - if err := w.initialize(); err != nil { - return err - } - return w.writeBuf(true) -} - -func (w *encryptingWriter) ReadFrom(r io.Reader) (int64, error) { - if err := w.initialize(); err != nil { - return 0, err - } - var nn int64 - for { - n, err := io.ReadFull(r, w.buf[len(w.buf):w.encrypter.encryptionMetadata.plaintextSegmentSize()+1]) - nn += int64(n) - w.buf = w.buf[:len(w.buf)+n] - if err != nil { - if err == io.EOF || err == io.ErrUnexpectedEOF { - return nn, nil - } - return nn, err - } - nextByte := w.buf[w.encrypter.encryptionMetadata.plaintextSegmentSize()] - w.buf = w.buf[:w.encrypter.encryptionMetadata.plaintextSegmentSize()] - if err := w.writeBuf(false); err != nil { - return nn, err - } - w.buf = w.buf[:1] - w.buf[0] = nextByte - } -} - -type decryptingReader struct { - r *bufio.Reader - decrypter segmentEncrypter - buf bytes.Buffer - initialized bool -} - -func (e *encryptionMetadata) newDecryptingReader(r io.Reader, password string, passwordMetadata *hashMetadata) *decryptingReader { - return &decryptingReader{ - r: bufio.NewReaderSize(r, 0), // we only need .UnreadByte - decrypter: segmentEncrypter{ - hashMetadata: *passwordMetadata, - encryptionMetadata: *e, - password: password, - }, - } -} - -func (r *decryptingReader) initialize() error { - if r.initialized { - return nil - } - header := make([]byte, r.decrypter.hashMetadata.SaltSize) - if _, err := io.ReadFull(r.r, header); err != nil { - if err == io.EOF { - return io.ErrUnexpectedEOF - } - return err - } - if err := r.decrypter.initialize(header); err != nil { - return err - } - r.buf = *bytes.NewBuffer(make([]byte, 0, r.decrypter.encryptionMetadata.SegmentSize+1)) - r.initialized = true - return nil -} - -func (r *decryptingReader) fillBuf() error { - r.buf.Reset() - // Read 1 extra byte to make sure if we're at EOF. - buf := r.buf.AvailableBuffer()[:r.decrypter.encryptionMetadata.SegmentSize+1] - n, err := io.ReadFull(r.r, buf) - if err != nil && err != io.ErrUnexpectedEOF { - return err - } - buf = buf[:n] - if len(buf) == int(r.decrypter.encryptionMetadata.SegmentSize)+1 { - r.r.UnreadByte() - buf = buf[:r.decrypter.encryptionMetadata.SegmentSize] - } - buf, err = r.decrypter.decrypt(buf[:0], buf, err == io.ErrUnexpectedEOF) - if err != nil { - return err - } - r.buf.Write(buf) - return nil -} - -func (r *decryptingReader) Read(buf []byte) (int, error) { - if err := r.initialize(); err != nil { - return 0, err - } - if r.buf.Len() == 0 { - if err := r.fillBuf(); err != nil { - return 0, err - } - } - return r.buf.Read(buf) -} - -func (r *decryptingReader) WriteTo(w io.Writer) (int64, error) { - if err := r.initialize(); err != nil { - return 0, err - } - var nn int64 - for { - n, err := r.buf.WriteTo(w) - nn += n - if err != nil { - return nn, err - } - if err := r.fillBuf(); err != nil { - if err == io.EOF { - return nn, nil - } - return nn, err - } - } -} diff --git a/internal/sym/oae_test.go b/internal/sym/oae_test.go deleted file mode 100644 index 16f9f96..0000000 --- a/internal/sym/oae_test.go +++ /dev/null @@ -1,41 +0,0 @@ -package sym - -import ( - "bytes" - "io" - "strings" - "testing" -) - -var testEncryptionMetadata = encryptionMetadata{ - EncryptionType: encryptionAlgAES256_GCM, - SegmentSize: 18, -} - -var testHashMetadata = hashMetadata{ - PasswordHashType: pwHashPBKDF2_HMAC_SHA256, - Iterations: 10, - SaltSize: defaultSaltSize, -} - -func TestOAEReadWrite(t *testing.T) { - t.Parallel() - - const password = "asdf" - input := strings.Repeat("test input", 1024) - out := new(bytes.Buffer) - writer := testEncryptionMetadata.newEncryptingWriter(out, password, &testHashMetadata) - 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(testEncryptionMetadata.newDecryptingReader(bytes.NewReader(out.Bytes()), password, &testHashMetadata)) - if err != nil { - t.Fatalf("Failed to decrypt: %s", err) - } - if string(got) != input { - t.Errorf("Input failed to round-trip") - } -} diff --git a/internal/sym/pwhash.go b/internal/sym/pwhash.go deleted file mode 100644 index d05f6b5..0000000 --- a/internal/sym/pwhash.go +++ /dev/null @@ -1,51 +0,0 @@ -package sym - -import ( - "crypto/pbkdf2" - "crypto/sha256" - "fmt" - "os" - - "golang.org/x/term" -) - -const defaultPBKDF2Iters = 35_000_000 - -//go:generate go tool stringer -type=pwHash -linecomment -type pwHash int8 - -const ( - pwHashInvalid pwHash = iota - pwHashPBKDF2_HMAC_SHA256 // PBKDF2-HMAC-SHA256 -) - -type hashMetadata struct { - PasswordHashType pwHash - Iterations int32 - SaltSize int8 -} - -func (h *hashMetadata) validate() error { - if h.PasswordHashType != pwHashPBKDF2_HMAC_SHA256 { - return fmt.Errorf("invalid hash type %q", h.PasswordHashType) - } - if h.Iterations <= 0 || h.Iterations > defaultPBKDF2Iters { - return fmt.Errorf("too many iterations") - } - if h.SaltSize <= 0 || h.SaltSize > defaultSaltSize { - return fmt.Errorf("salt size too long") - } - return nil -} - -func (h *hashMetadata) hashPassword(password string, salt []byte) ([]byte, error) { - return pbkdf2.Key(sha256.New, password, salt, int(h.Iterations), 32) -} - -func termReadPassword() (string, error) { - pw, err := term.ReadPassword(int(os.Stdin.Fd())) - if err != nil { - return "", err - } - return string(pw), nil -} diff --git a/internal/sym/pwhash_string.go b/internal/sym/pwhash_string.go deleted file mode 100644 index ceb4d1f..0000000 --- a/internal/sym/pwhash_string.go +++ /dev/null @@ -1,25 +0,0 @@ -// Code generated by "stringer -type=pwHash -linecomment"; DO NOT EDIT. - -package sym - -import "strconv" - -func _() { - // An "invalid array index" compiler error signifies that the constant values have changed. - // Re-run the stringer command to generate them again. - var x [1]struct{} - _ = x[pwHashInvalid-0] - _ = x[pwHashPBKDF2_HMAC_SHA256-1] -} - -const _pwHash_name = "pwHashInvalidPBKDF2-HMAC-SHA256" - -var _pwHash_index = [...]uint8{0, 13, 31} - -func (i pwHash) String() string { - idx := int(i) - 0 - if i < 0 || idx >= len(_pwHash_index)-1 { - return "pwHash(" + strconv.FormatInt(int64(i), 10) + ")" - } - return _pwHash_name[_pwHash_index[idx]:_pwHash_index[idx+1]] -} diff --git a/internal/sym/sym_test.go b/internal/sym/sym_test.go deleted file mode 100644 index 704d477..0000000 --- a/internal/sym/sym_test.go +++ /dev/null @@ -1,75 +0,0 @@ -package sym - -import ( - "bytes" - "os" - "path/filepath" - "testing" -) - -var testEncryptOptions = func() EncryptOptions { - opts := DefaultEncryptOptions - opts.iterations = 10 - return opts -}() - -func mustWriteFile(t *testing.T, path string, content []byte) { - t.Helper() - if err := os.WriteFile(path, content, 0600); err != nil { - t.Fatalf("Failed to write test file: %s", err) - } -} - -func mustReadFile(t *testing.T, path string) []byte { - t.Helper() - content, err := os.ReadFile(path) - if err != nil { - t.Fatalf("Failed to read file: %s", err) - } - return content -} - -func mustRename(t *testing.T, src, dst string) { - t.Helper() - if err := os.Rename(src, dst); err != nil { - t.Fatalf("Failed to rename: %s", err) - } -} - -func mustRemove(t *testing.T, path string) { - t.Helper() - if err := os.Remove(path); err != nil { - t.Fatalf("Failed to remove file: %s", err) - } -} - -func mustChmod(t *testing.T, path string, mod os.FileMode) { - t.Helper() - if err := os.Chmod(path, mod); err != nil { - t.Fatalf("Failed to chmod: %s", err) - } -} - -func TestEncryptDecrypt(t *testing.T) { - t.Parallel() - - buf := make([]byte, 12*1024*1024) - for i := range buf { - buf[i] = byte(i) - } - - 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 { - t.Fatalf("EncryptFile failed: %s", err) - } - mustRemove(t, fileName) - if err := DefaultDecryptOptions.decryptFile(fileName+".enc", password); err != nil { - t.Fatalf("DecryptFile failed: %s", err) - } - gotContents := mustReadFile(t, fileName) - if !bytes.Equal(gotContents, buf) { - t.Errorf("contents differ") - } -} diff --git a/metadata.go b/metadata.go new file mode 100644 index 0000000..3875a63 --- /dev/null +++ b/metadata.go @@ -0,0 +1,21 @@ +package sym + +import "fmt" + +const magic = "\x80sym" + +type fileMetadata struct { + Version int8 + HashMetadata hashMetadata + EncryptionMetadata encryptionMetadata +} + +func (f *fileMetadata) validate() error { + if f.Version != 0 { + return fmt.Errorf("bad version") + } + if err := f.HashMetadata.validate(); err != nil { + return err + } + return f.EncryptionMetadata.validate() +} diff --git a/oae.go b/oae.go new file mode 100644 index 0000000..e40cb60 --- /dev/null +++ b/oae.go @@ -0,0 +1,293 @@ +package sym + +import ( + "bufio" + "bytes" + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "errors" + "fmt" + "io" +) + +const ( + nonceSize = 12 + aeadOverhead = 16 + + defaultSegmentSize = 1024 * 1024 + + defaultSaltSize = 32 +) + +//go:generate go tool stringer -type=encryptionAlg -linecomment +type encryptionAlg int8 + +const ( + encryptionAlgInvalid encryptionAlg = iota + encryptionAlgAES256_GCM // AES-256-GCM +) + +type encryptionMetadata struct { + EncryptionType encryptionAlg + SegmentSize int32 +} + +func (e *encryptionMetadata) validate() error { + if e.EncryptionType != encryptionAlgAES256_GCM { + return fmt.Errorf("invalid encryption alg %q", e.EncryptionType) + } + if e.SegmentSize <= 0 || e.SegmentSize > defaultSegmentSize { + return fmt.Errorf("segment size too long") + } + return nil +} + +func (e *encryptionMetadata) plaintextSegmentSize() int { + return int(e.SegmentSize) - aeadOverhead +} + +type segmentEncrypter struct { + hashMetadata hashMetadata + encryptionMetadata encryptionMetadata + password string + + aead cipher.AEAD + nonce [nonceSize]byte +} + +func (se *segmentEncrypter) initialize(salt []byte) error { + key, err := se.hashMetadata.hashPassword(se.password, salt) + if err != nil { + return err + } + block, err := aes.NewCipher(key) + if err != nil { + return err + } + se.aead, err = cipher.NewGCM(block) + return err +} + +func (se *segmentEncrypter) ad(lastSegment bool) ([]byte, error) { + // Increment counter + for i := 0; ; i++ { + if i == len(se.nonce) { + return nil, errors.New("counter overflowed") + } + se.nonce[i]++ + if se.nonce[i] != 0 { + break + } + } + ad := make([]byte, len(se.nonce)+1) + copy(ad, se.nonce[:]) + if lastSegment { + ad[len(ad)-1] = 1 + } + return ad, nil +} + +func (se *segmentEncrypter) encrypt(out, buf []byte, lastSegment bool) ([]byte, error) { + ad, err := se.ad(lastSegment) + if err != nil { + return nil, err + } + return se.aead.Seal(out, se.nonce[:], buf, ad), nil +} + +func (se *segmentEncrypter) decrypt(out, buf []byte, lastSegment bool) ([]byte, error) { + ad, err := se.ad(lastSegment) + if err != nil { + return nil, err + } + return se.aead.Open(out, se.nonce[:], buf, ad) +} + +type encryptingWriter struct { + w io.Writer + encrypter segmentEncrypter + buf []byte + initialized bool +} + +func (e *encryptionMetadata) newEncryptingWriter(w io.Writer, password string, passwordMetadata *hashMetadata) *encryptingWriter { + return &encryptingWriter{ + w: w, + encrypter: segmentEncrypter{ + hashMetadata: *passwordMetadata, + encryptionMetadata: *e, + password: password, + }, + } +} + +func (w *encryptingWriter) initialize() error { + if w.initialized { + return nil + } + header := make([]byte, w.encrypter.hashMetadata.SaltSize) + rand.Read(header) + if err := w.encrypter.initialize(header); err != nil { + return err + } + if _, err := w.w.Write(header); err != nil { + return err + } + w.buf = make([]byte, 0, w.encrypter.encryptionMetadata.SegmentSize) + w.initialized = true + return nil +} + +func (w *encryptingWriter) writeBuf(lastSegment bool) error { + var err error + if w.buf, err = w.encrypter.encrypt(w.buf[:0], w.buf, lastSegment); err != nil { + return err + } + if _, err := w.w.Write(w.buf); err != nil { + return err + } + w.buf = w.buf[:0] + return nil +} + +func (w *encryptingWriter) Write(buf []byte) (int, error) { + if err := w.initialize(); err != nil { + return 0, err + } + nn := 0 + for len(buf) > 0 { + if len(w.buf) == w.encrypter.encryptionMetadata.plaintextSegmentSize() { + if err := w.writeBuf(false); err != nil { + return nn, err + } + } + n := copy(w.buf[len(w.buf):w.encrypter.encryptionMetadata.plaintextSegmentSize()], buf) + nn += n + w.buf = w.buf[:len(w.buf)+n] + buf = buf[n:] + } + return nn, nil +} + +func (w *encryptingWriter) Close() error { + if err := w.initialize(); err != nil { + return err + } + return w.writeBuf(true) +} + +func (w *encryptingWriter) ReadFrom(r io.Reader) (int64, error) { + if err := w.initialize(); err != nil { + return 0, err + } + var nn int64 + for { + n, err := io.ReadFull(r, w.buf[len(w.buf):w.encrypter.encryptionMetadata.plaintextSegmentSize()+1]) + nn += int64(n) + w.buf = w.buf[:len(w.buf)+n] + if err != nil { + if err == io.EOF || err == io.ErrUnexpectedEOF { + return nn, nil + } + return nn, err + } + nextByte := w.buf[w.encrypter.encryptionMetadata.plaintextSegmentSize()] + w.buf = w.buf[:w.encrypter.encryptionMetadata.plaintextSegmentSize()] + if err := w.writeBuf(false); err != nil { + return nn, err + } + w.buf = w.buf[:1] + w.buf[0] = nextByte + } +} + +type decryptingReader struct { + r *bufio.Reader + decrypter segmentEncrypter + buf bytes.Buffer + initialized bool +} + +func (e *encryptionMetadata) newDecryptingReader(r io.Reader, password string, passwordMetadata *hashMetadata) *decryptingReader { + return &decryptingReader{ + r: bufio.NewReaderSize(r, 0), // we only need .UnreadByte + decrypter: segmentEncrypter{ + hashMetadata: *passwordMetadata, + encryptionMetadata: *e, + password: password, + }, + } +} + +func (r *decryptingReader) initialize() error { + if r.initialized { + return nil + } + header := make([]byte, r.decrypter.hashMetadata.SaltSize) + if _, err := io.ReadFull(r.r, header); err != nil { + if err == io.EOF { + return io.ErrUnexpectedEOF + } + return err + } + if err := r.decrypter.initialize(header); err != nil { + return err + } + r.buf = *bytes.NewBuffer(make([]byte, 0, r.decrypter.encryptionMetadata.SegmentSize+1)) + r.initialized = true + return nil +} + +func (r *decryptingReader) fillBuf() error { + r.buf.Reset() + // Read 1 extra byte to make sure if we're at EOF. + buf := r.buf.AvailableBuffer()[:r.decrypter.encryptionMetadata.SegmentSize+1] + n, err := io.ReadFull(r.r, buf) + if err != nil && err != io.ErrUnexpectedEOF { + return err + } + buf = buf[:n] + if len(buf) == int(r.decrypter.encryptionMetadata.SegmentSize)+1 { + r.r.UnreadByte() + buf = buf[:r.decrypter.encryptionMetadata.SegmentSize] + } + buf, err = r.decrypter.decrypt(buf[:0], buf, err == io.ErrUnexpectedEOF) + if err != nil { + return err + } + r.buf.Write(buf) + return nil +} + +func (r *decryptingReader) Read(buf []byte) (int, error) { + if err := r.initialize(); err != nil { + return 0, err + } + if r.buf.Len() == 0 { + if err := r.fillBuf(); err != nil { + return 0, err + } + } + return r.buf.Read(buf) +} + +func (r *decryptingReader) WriteTo(w io.Writer) (int64, error) { + if err := r.initialize(); err != nil { + return 0, err + } + var nn int64 + for { + n, err := r.buf.WriteTo(w) + nn += n + if err != nil { + return nn, err + } + if err := r.fillBuf(); err != nil { + if err == io.EOF { + return nn, nil + } + return nn, err + } + } +} diff --git a/oae_test.go b/oae_test.go new file mode 100644 index 0000000..16f9f96 --- /dev/null +++ b/oae_test.go @@ -0,0 +1,41 @@ +package sym + +import ( + "bytes" + "io" + "strings" + "testing" +) + +var testEncryptionMetadata = encryptionMetadata{ + EncryptionType: encryptionAlgAES256_GCM, + SegmentSize: 18, +} + +var testHashMetadata = hashMetadata{ + PasswordHashType: pwHashPBKDF2_HMAC_SHA256, + Iterations: 10, + SaltSize: defaultSaltSize, +} + +func TestOAEReadWrite(t *testing.T) { + t.Parallel() + + const password = "asdf" + input := strings.Repeat("test input", 1024) + out := new(bytes.Buffer) + writer := testEncryptionMetadata.newEncryptingWriter(out, password, &testHashMetadata) + 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(testEncryptionMetadata.newDecryptingReader(bytes.NewReader(out.Bytes()), password, &testHashMetadata)) + if err != nil { + t.Fatalf("Failed to decrypt: %s", err) + } + if string(got) != input { + t.Errorf("Input failed to round-trip") + } +} diff --git a/pwhash.go b/pwhash.go new file mode 100644 index 0000000..d05f6b5 --- /dev/null +++ b/pwhash.go @@ -0,0 +1,51 @@ +package sym + +import ( + "crypto/pbkdf2" + "crypto/sha256" + "fmt" + "os" + + "golang.org/x/term" +) + +const defaultPBKDF2Iters = 35_000_000 + +//go:generate go tool stringer -type=pwHash -linecomment +type pwHash int8 + +const ( + pwHashInvalid pwHash = iota + pwHashPBKDF2_HMAC_SHA256 // PBKDF2-HMAC-SHA256 +) + +type hashMetadata struct { + PasswordHashType pwHash + Iterations int32 + SaltSize int8 +} + +func (h *hashMetadata) validate() error { + if h.PasswordHashType != pwHashPBKDF2_HMAC_SHA256 { + return fmt.Errorf("invalid hash type %q", h.PasswordHashType) + } + if h.Iterations <= 0 || h.Iterations > defaultPBKDF2Iters { + return fmt.Errorf("too many iterations") + } + if h.SaltSize <= 0 || h.SaltSize > defaultSaltSize { + return fmt.Errorf("salt size too long") + } + return nil +} + +func (h *hashMetadata) hashPassword(password string, salt []byte) ([]byte, error) { + return pbkdf2.Key(sha256.New, password, salt, int(h.Iterations), 32) +} + +func termReadPassword() (string, error) { + pw, err := term.ReadPassword(int(os.Stdin.Fd())) + if err != nil { + return "", err + } + return string(pw), nil +} diff --git a/pwhash_string.go b/pwhash_string.go new file mode 100644 index 0000000..ceb4d1f --- /dev/null +++ b/pwhash_string.go @@ -0,0 +1,25 @@ +// Code generated by "stringer -type=pwHash -linecomment"; DO NOT EDIT. + +package sym + +import "strconv" + +func _() { + // An "invalid array index" compiler error signifies that the constant values have changed. + // Re-run the stringer command to generate them again. + var x [1]struct{} + _ = x[pwHashInvalid-0] + _ = x[pwHashPBKDF2_HMAC_SHA256-1] +} + +const _pwHash_name = "pwHashInvalidPBKDF2-HMAC-SHA256" + +var _pwHash_index = [...]uint8{0, 13, 31} + +func (i pwHash) String() string { + idx := int(i) - 0 + if i < 0 || idx >= len(_pwHash_index)-1 { + return "pwHash(" + strconv.FormatInt(int64(i), 10) + ")" + } + return _pwHash_name[_pwHash_index[idx]:_pwHash_index[idx+1]] +} diff --git a/sym_test.go b/sym_test.go new file mode 100644 index 0000000..704d477 --- /dev/null +++ b/sym_test.go @@ -0,0 +1,75 @@ +package sym + +import ( + "bytes" + "os" + "path/filepath" + "testing" +) + +var testEncryptOptions = func() EncryptOptions { + opts := DefaultEncryptOptions + opts.iterations = 10 + return opts +}() + +func mustWriteFile(t *testing.T, path string, content []byte) { + t.Helper() + if err := os.WriteFile(path, content, 0600); err != nil { + t.Fatalf("Failed to write test file: %s", err) + } +} + +func mustReadFile(t *testing.T, path string) []byte { + t.Helper() + content, err := os.ReadFile(path) + if err != nil { + t.Fatalf("Failed to read file: %s", err) + } + return content +} + +func mustRename(t *testing.T, src, dst string) { + t.Helper() + if err := os.Rename(src, dst); err != nil { + t.Fatalf("Failed to rename: %s", err) + } +} + +func mustRemove(t *testing.T, path string) { + t.Helper() + if err := os.Remove(path); err != nil { + t.Fatalf("Failed to remove file: %s", err) + } +} + +func mustChmod(t *testing.T, path string, mod os.FileMode) { + t.Helper() + if err := os.Chmod(path, mod); err != nil { + t.Fatalf("Failed to chmod: %s", err) + } +} + +func TestEncryptDecrypt(t *testing.T) { + t.Parallel() + + buf := make([]byte, 12*1024*1024) + for i := range buf { + buf[i] = byte(i) + } + + 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 { + t.Fatalf("EncryptFile failed: %s", err) + } + mustRemove(t, fileName) + if err := DefaultDecryptOptions.decryptFile(fileName+".enc", password); err != nil { + t.Fatalf("DecryptFile failed: %s", err) + } + gotContents := mustReadFile(t, fileName) + if !bytes.Equal(gotContents, buf) { + t.Errorf("contents differ") + } +} -- cgit v1.3.1