aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--internal/sym/dec.go38
-rw-r--r--internal/sym/enc.go74
-rw-r--r--internal/sym/enc_test.go2
-rw-r--r--internal/sym/sym_test.go2
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)
}