From 2bea9cd57509e87c0b7583e045e1d1545cf64875 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Tue, 14 Oct 2025 18:59:02 -0700 Subject: Implement ReadFrom and WriteTo for less copying --- internal/sym/dec_test.go | 12 ++++---- internal/sym/enc.go | 78 +++++++++++++++++++++++++----------------------- internal/sym/enc_test.go | 16 +++++----- internal/sym/oae.go | 45 ++++++++++++++++++++++++++++ internal/sym/oae_test.go | 41 +++++++++++++++++++++++++ internal/sym/sym_test.go | 8 ++++- 6 files changed, 147 insertions(+), 53 deletions(-) create mode 100644 internal/sym/oae_test.go diff --git a/internal/sym/dec_test.go b/internal/sym/dec_test.go index 727fa51..b38497f 100644 --- a/internal/sym/dec_test.go +++ b/internal/sym/dec_test.go @@ -33,7 +33,7 @@ func TestDecryptFile_Force(t *testing.T) { const password = "asdf" fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, []byte("test file content")) - if err := DefaultEncryptOptions.encryptFile(fileName, password); err != nil { + if err := testEncryptOptions.encryptFile(fileName, password); err != nil { t.Fatalf("Failed to encrypt file: %s", err) } decOpts := DefaultDecryptOptions @@ -53,7 +53,7 @@ func encodeHeader(t *testing.T, f *fileMetadata) []byte { f.HashMetadata.PasswordHashType = pwHashPBKDF2_HMAC_SHA256 } if f.HashMetadata.Iterations == 0 { - f.HashMetadata.Iterations = defaultPBKDF2Iters + f.HashMetadata.Iterations = 10 } if f.HashMetadata.SaltSize == 0 { f.HashMetadata.SaltSize = defaultSaltSize @@ -167,7 +167,7 @@ func TestDecryptFile_WeirdName(t *testing.T) { fileContent := []byte("file content") fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, fileContent) - if err := DefaultEncryptOptions.encryptFile(fileName, password); err != nil { + if err := testEncryptOptions.encryptFile(fileName, password); err != nil { t.Fatalf("EncryptFile failed: %s", err) } mustRename(t, fileName+".enc", fileName+".encrypted") @@ -228,7 +228,7 @@ func TestDecryptOptions_Run(t *testing.T) { fileContent := []byte("test file content") fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, fileContent) - if err := DefaultEncryptOptions.encryptFile(fileName, password); err != nil { + if err := testEncryptOptions.encryptFile(fileName, password); err != nil { t.Errorf("EncryptFile failed: %s", err) } mustRemove(t, fileName) @@ -270,7 +270,7 @@ func TestDecryptOptions_Run_Stdin(t *testing.T) { const password = "asdf" content := []byte("test contents") encrypted := new(bytes.Buffer) - if err := encryptBinary(encrypted, bytes.NewReader(content), password); err != nil { + if err := testEncryptOptions.encryptBinary(encrypted, bytes.NewReader(content), password); err != nil { t.Fatalf("Failed to encrypt: %s", err) } gotContentBuf := new(bytes.Buffer) @@ -308,7 +308,7 @@ func TestDecryptOptions_Run_ReadPassword(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, []byte("test file content")) - if err := DefaultEncryptOptions.encryptFile(fileName, password); err != nil { + if err := testEncryptOptions.encryptFile(fileName, password); err != nil { t.Errorf("EncryptFile failed: %s", err) } mustRemove(t, fileName) diff --git a/internal/sym/enc.go b/internal/sym/enc.go index 274954a..be4f839 100644 --- a/internal/sym/enc.go +++ b/internal/sym/enc.go @@ -41,7 +41,39 @@ func (w *newlineWriter) Write(buf []byte) (int, error) { return nn, nil } -func encryptBinary(w io.Writer, r io.Reader, password string) error { +type encryptFlags struct { + generatePassword bool + password string + asciiOutput bool + 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.asciiOutput, "a", false, "output in base64, default is binary output") + 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) encryptBinary(w io.Writer, r io.Reader, password string) error { if _, err := io.WriteString(w, magic); err != nil { return err } @@ -49,7 +81,7 @@ func encryptBinary(w io.Writer, r io.Reader, password string) error { Version: 0, HashMetadata: hashMetadata{ PasswordHashType: pwHashPBKDF2_HMAC_SHA256, - Iterations: defaultPBKDF2Iters, + Iterations: int32(o.iterations), SaltSize: defaultSaltSize, }, EncryptionMetadata: encryptionMetadata{ @@ -67,7 +99,7 @@ func encryptBinary(w io.Writer, r io.Reader, password string) error { return writer.Close() } -func encryptBase64(w io.Writer, r io.Reader, password string) error { +func (o *EncryptOptions) encryptBase64(w io.Writer, r io.Reader, password string) error { bufWriter := bufio.NewWriter(w) if _, err := bufWriter.WriteString(`-------------------------- Begin encrypted text block -------------------------- -------------------------- am i cool like gpg? --------------------------------- @@ -75,7 +107,7 @@ func encryptBase64(w io.Writer, r io.Reader, password string) error { return err } base64Writer := base64.NewEncoder(base64.StdEncoding, &newlineWriter{w: bufWriter}) - if err := encryptBinary(base64Writer, r, password); err != nil { + if err := o.encryptBinary(base64Writer, r, password); err != nil { return err } if err := base64Writer.Close(); err != nil { @@ -87,36 +119,6 @@ func encryptBase64(w io.Writer, r io.Reader, password string) error { return bufWriter.Flush() } -type encryptFlags struct { - generatePassword bool - password string - asciiOutput bool - 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.asciiOutput, "a", false, "output in base64, default is binary output") - fs.BoolVar(&f.force, "f", false, "overwrite output files even if they already exist") -} - -type EncryptOptions struct { - encryptFlags - - passwordIn func() (string, error) - passwordOut io.Writer - stdin io.Reader - stdout io.Writer -} - -var DefaultEncryptOptions = EncryptOptions{ - passwordIn: termReadPassword, - passwordOut: os.Stderr, - stdin: os.Stdin, - stdout: os.Stdout, -} - func (o *EncryptOptions) encryptFile(fileName string, password string) (err error) { f, err := os.Open(fileName) if err != nil { @@ -147,9 +149,9 @@ func (o *EncryptOptions) encryptFile(fileName string, password string) (err erro } }() if o.asciiOutput { - err = encryptBase64(fOut, f, password) + err = o.encryptBase64(fOut, f, password) } else { - err = encryptBinary(fOut, f, password) + err = o.encryptBinary(fOut, f, password) } if err != nil { return fmt.Errorf("encrypt %q: %s", fileName, err) @@ -219,9 +221,9 @@ func (o *EncryptOptions) Run(args ...string) error { } if len(args) == 0 { if o.asciiOutput { - return encryptBase64(o.stdout, o.stdin, password) + return o.encryptBase64(o.stdout, o.stdin, password) } - return encryptBinary(o.stdout, o.stdin, password) + return o.encryptBinary(o.stdout, o.stdin, password) } for _, fileName := range args { if err := o.encryptFile(fileName, password); err != nil { diff --git a/internal/sym/enc_test.go b/internal/sym/enc_test.go index c6f697c..1aa2e06 100644 --- a/internal/sym/enc_test.go +++ b/internal/sym/enc_test.go @@ -31,7 +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 := DefaultEncryptOptions + encOpts := testEncryptOptions encOpts.force = tc.force err := encOpts.encryptFile(fileName, "asdf") if gotErr := err != nil; gotErr != tc.wantErr { @@ -44,7 +44,7 @@ func TestEncryptFile_Force(t *testing.T) { func TestEncryptFile_NotFound(t *testing.T) { t.Parallel() - err := DefaultEncryptOptions.encryptFile("my-nonexistent-file.txt", "asdf") + err := testEncryptOptions.encryptFile("my-nonexistent-file.txt", "asdf") if err == nil { t.Fatal("encryptFile succeeded for nonexistent file, want error") } @@ -57,7 +57,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 := DefaultEncryptOptions + opts := testEncryptOptions opts.force = true err := opts.encryptFile(fileName, "asdf") if err == nil { @@ -91,7 +91,7 @@ func TestEncryptOptions_Run(t *testing.T) { fileContent := []byte("test file content") fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, fileContent) - opts := DefaultEncryptOptions + opts := testEncryptOptions opts.password = password if err := opts.Run(fileName); err != nil { t.Fatalf("enc failed: %s", err) @@ -130,7 +130,7 @@ func TestEncryptOptions_Run_UsageError(t *testing.T) { t.Run(tc.desc, func(t *testing.T) { t.Parallel() - opts := DefaultEncryptOptions + opts := testEncryptOptions opts.generatePassword = tc.generatePassword opts.password = tc.password if err := opts.Run(tc.files...); err == nil { @@ -148,7 +148,7 @@ func TestEncryptOptions_Run_GeneratePassword(t *testing.T) { mustWriteFile(t, fileName, fileContent) password := new(strings.Builder) - opts := DefaultEncryptOptions + opts := testEncryptOptions opts.generatePassword = true opts.passwordOut = password if err := opts.Run(fileName); err != nil { @@ -186,7 +186,7 @@ func TestEncryptOptions_Run_Stdin(t *testing.T) { password = "asdf" ) stdout := new(strings.Builder) - opts := DefaultEncryptOptions + opts := testEncryptOptions opts.password = password opts.asciiOutput = tc.ascii opts.stdin = strings.NewReader(input) @@ -240,7 +240,7 @@ func TestEncryptOptions_Run_ReadPassword(t *testing.T) { fileName := filepath.Join(t.TempDir(), "file") mustWriteFile(t, fileName, []byte("test file content")) - opts := DefaultEncryptOptions + opts := testEncryptOptions passwordI := 0 opts.passwordIn = func() (string, error) { if passwordI == len(tc.passwords) && tc.err != nil { diff --git a/internal/sym/oae.go b/internal/sym/oae.go index 8fbbeea..e40cb60 100644 --- a/internal/sym/oae.go +++ b/internal/sym/oae.go @@ -177,6 +177,31 @@ func (w *encryptingWriter) Close() error { 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 @@ -246,3 +271,23 @@ func (r *decryptingReader) Read(buf []byte) (int, error) { } 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 new file mode 100644 index 0000000..16f9f96 --- /dev/null +++ b/internal/sym/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/internal/sym/sym_test.go b/internal/sym/sym_test.go index a277ae6..c9e396a 100644 --- a/internal/sym/sym_test.go +++ b/internal/sym/sym_test.go @@ -7,6 +7,12 @@ import ( "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 { @@ -67,7 +73,7 @@ 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" - encOpts := DefaultEncryptOptions + encOpts := testEncryptOptions encOpts.asciiOutput = tc.ascii if err := encOpts.encryptFile(fileName, password); err != nil { t.Fatalf("EncryptFile failed: %s", err) -- cgit v1.3.1