From 6569b910521fc8a221bb93b41892e2b762a8f48a Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Sat, 18 Oct 2025 18:22:04 -0700 Subject: Try to reduce in-memory copies --- internal/sym/dec.go | 38 +++++++++++++++++-------- internal/sym/enc.go | 74 ++++++++++++++++++++++++++++++++++-------------- internal/sym/enc_test.go | 2 +- internal/sym/sym_test.go | 2 +- 4 files changed, 81 insertions(+), 35 deletions(-) diff --git a/internal/sym/dec.go b/internal/sym/dec.go index ac811b7..4e290d4 100644 --- a/internal/sym/dec.go +++ b/internal/sym/dec.go @@ -2,6 +2,7 @@ package sym import ( "bufio" + "bytes" "encoding/base64" "encoding/binary" "errors" @@ -12,6 +13,29 @@ import ( "strings" ) +type asciiReader struct { + r *bufio.Reader + line []byte +} + +func (r *asciiReader) Read(b []byte) (int, error) { + for len(r.line) == 0 { + line, err := r.r.ReadSlice('\n') + if len(line) == 0 { + return 0, err + } + if bytes.HasPrefix(line, []byte("-")) { + continue + } + if r.line, err = base64.StdEncoding.AppendDecode(line[:0], bytes.TrimSuffix(line, []byte("\n"))); err != nil { + return 0, err + } + } + n := copy(b, r.line) + r.line = r.line[n:] + return n, nil +} + func decryptBinary(w io.Writer, r io.Reader, password string) error { fileFormat := make([]byte, 4) if _, err := io.ReadFull(r, fileFormat); err != nil { @@ -32,7 +56,7 @@ func decryptBinary(w io.Writer, r io.Reader, password string) error { } func decrypt(w io.Writer, r io.Reader, password string) error { - bufReader := bufio.NewReaderSize(r, 81) + bufReader := bufio.NewReader(r) b, err := bufReader.Peek(1) if err != nil { if err == io.EOF { @@ -46,17 +70,7 @@ func decrypt(w io.Writer, r io.Reader, password string) error { if b[0] != '-' { return errors.New("invalid input") } - for { - for { - if _, isPrefix, err := bufReader.ReadLine(); err != nil || !isPrefix { - break - } - } - if b, err := bufReader.Peek(1); err != nil || b[0] != '-' { - break - } - } - return decryptBinary(w, base64.NewDecoder(base64.StdEncoding, bufReader), password) + return decryptBinary(w, &asciiReader{r: bufReader}, password) } type decryptFlags struct { diff --git a/internal/sym/enc.go b/internal/sym/enc.go index be4f839..689df7b 100644 --- a/internal/sym/enc.go +++ b/internal/sym/enc.go @@ -15,32 +15,70 @@ import ( "roseh.moe/pkg/wordlist" ) -type newlineWriter struct { - w *bufio.Writer - n int +type asciiWriter struct { + w *bufio.Writer + buf [2]byte + nBuf int + n int } -func (w *newlineWriter) Write(buf []byte) (int, error) { - const lineSize = 80 +func (w *asciiWriter) writeBase64(b []byte) (int, error) { + _, err := w.w.Write(base64.StdEncoding.AppendEncode(w.w.AvailableBuffer(), b)) + return len(b), err +} + +func (w *asciiWriter) Write(b []byte) (int, error) { + const lineSizeBytes = 60 + nn := 0 - for len(buf) > 0 { - if w.n == lineSize { - if err := w.w.WriteByte('\n'); err != nil { - return nn, err - } - w.n = 0 + if w.nBuf > 0 && w.nBuf+len(b) >= 3 { + leadingChunk := make([]byte, 3) + n := copy(leadingChunk, w.buf[:w.nBuf]) + w.nBuf = 0 + n = copy(leadingChunk[n:], b) + b = b[n:] + n, err := w.writeBase64(leadingChunk) + nn += n + w.n += n + if err != nil { + return nn, err + } + } + for len(b) >= 3 { + chunkSize := lineSizeBytes - w.n + if chunkSize > len(b) { + chunkSize = len(b) - len(b)%3 } - n, err := w.w.Write(buf[:min(lineSize-w.n, len(buf))]) + n, err := w.writeBase64(b[:chunkSize]) nn += n - buf = buf[n:] w.n += n if err != nil { return nn, err } + b = b[n:] + if w.n == lineSizeBytes { + w.w.WriteByte('\n') + w.n = 0 + } } + n := copy(w.buf[:], b) + nn += n + w.nBuf = n return nn, nil } +func (w *asciiWriter) Close() error { + if w.nBuf > 0 { + if _, err := w.writeBase64(w.buf[:w.nBuf]); err != nil { + return err + } + if err := w.w.WriteByte('\n'); err != nil { + return err + } + } + return w.w.Flush() +} + type encryptFlags struct { generatePassword bool password string @@ -106,17 +144,11 @@ func (o *EncryptOptions) encryptBase64(w io.Writer, r io.Reader, password string `); err != nil { return err } - base64Writer := base64.NewEncoder(base64.StdEncoding, &newlineWriter{w: bufWriter}) + base64Writer := &asciiWriter{w: bufWriter} if err := o.encryptBinary(base64Writer, r, password); err != nil { return err } - if err := base64Writer.Close(); err != nil { - return err - } - if err := bufWriter.WriteByte('\n'); err != nil { - return err - } - return bufWriter.Flush() + return base64Writer.Close() } func (o *EncryptOptions) encryptFile(fileName string, password string) (err error) { diff --git a/internal/sym/enc_test.go b/internal/sym/enc_test.go index 1aa2e06..e2ec6a8 100644 --- a/internal/sym/enc_test.go +++ b/internal/sym/enc_test.go @@ -196,7 +196,7 @@ func TestEncryptOptions_Run_Stdin(t *testing.T) { } got := new(strings.Builder) if err := decrypt(got, strings.NewReader(stdout.String()), password); err != nil { - t.Errorf("Failed to decrypt stdout content: %s", err) + 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) diff --git a/internal/sym/sym_test.go b/internal/sym/sym_test.go index c9e396a..98ae025 100644 --- a/internal/sym/sym_test.go +++ b/internal/sym/sym_test.go @@ -53,7 +53,7 @@ func mustChmod(t *testing.T, path string, mod os.FileMode) { func TestEncryptDecrypt(t *testing.T) { t.Parallel() - buf := make([]byte, 10*1024*1024) + buf := make([]byte, 12*1024*1024) for i := range buf { buf[i] = byte(i) } -- cgit v1.3.1