summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2025-10-04 21:17:49 -0700
committerRose Hogenson <rosehogenson@posteo.net>2025-10-04 21:20:22 -0700
commit6ff4e9c0622dee2663d70059594e58b3c96f135a (patch)
treeabc42b8d2fc4d162f52fd977924260beef73725f
parent9833d792339b7401e395bf0fa35f1f0486ec80f4 (diff)
downloadroseh.moe-6ff4e9c0622dee2663d70059594e58b3c96f135a.tar.zst
Add some extra methods to oae2
-rw-r--r--internal/cryptoutil/oae2.go134
1 files changed, 127 insertions, 7 deletions
diff --git a/internal/cryptoutil/oae2.go b/internal/cryptoutil/oae2.go
index 4997051..be4f1a9 100644
--- a/internal/cryptoutil/oae2.go
+++ b/internal/cryptoutil/oae2.go
@@ -11,6 +11,7 @@ import (
"encoding/binary"
"errors"
"io"
+ "strings"
)
// Online Authenticated Encryption from https://eprint.iacr.org/2015/189.pdf
@@ -153,6 +154,18 @@ func (w *EncryptingWriter) initialize() error {
return w.err
}
+func (w *EncryptingWriter) writeBuf() error {
+ var encrypted []byte
+ if encrypted, w.err = w.oae2.encryptBlock(w.buf[:0], w.buf, false); w.err != nil {
+ return w.err
+ }
+ if _, w.err = w.w.Write(encrypted); w.err != nil {
+ return w.err
+ }
+ w.buf = w.buf[:0]
+ return nil
+}
+
func (w *EncryptingWriter) Write(buf []byte) (int, error) {
if !w.initialized {
if err := w.initialize(); err != nil {
@@ -166,14 +179,9 @@ func (w *EncryptingWriter) Write(buf []byte) (int, error) {
nn := 0
for r.Len() > 0 {
if len(w.buf) == w.segmentSize {
- var encrypted []byte
- if encrypted, w.err = w.oae2.encryptBlock(w.buf[:0], w.buf, false); w.err != nil {
- return nn, w.err
+ if err := w.writeBuf(); err != nil {
+ return nn, err
}
- if _, w.err = w.w.Write(encrypted); w.err != nil {
- return nn, w.err
- }
- w.buf = w.buf[:0]
}
n, _ := r.Read(w.buf[len(w.buf):w.segmentSize])
w.buf = w.buf[:len(w.buf)+n]
@@ -209,6 +217,77 @@ func (w *EncryptingWriter) Close() error {
return nil
}
+func (w *EncryptingWriter) ReadFrom(r io.Reader) (int64, error) {
+ if !w.initialized {
+ if err := w.initialize(); err != nil {
+ return 0, err
+ }
+ }
+ if w.err != nil {
+ return 0, w.err
+ }
+ var nn int64
+ bufReader := bufio.NewReaderSize(r, 1)
+ for {
+ if _, err := bufReader.Peek(1); err != nil {
+ if err == io.EOF {
+ return nn, nil
+ }
+ return nn, err
+ }
+ if len(w.buf) == w.segmentSize {
+ if err := w.writeBuf(); err != nil {
+ return nn, err
+ }
+ }
+ n, _ := io.ReadFull(bufReader, w.buf[len(w.buf):w.segmentSize])
+ w.buf = w.buf[:len(w.buf)+n]
+ nn += int64(n)
+ }
+}
+
+func (w *EncryptingWriter) WriteByte(c byte) error {
+ if !w.initialized {
+ if err := w.initialize(); err != nil {
+ return err
+ }
+ }
+ if w.err != nil {
+ return w.err
+ }
+ if len(w.buf) == w.segmentSize {
+ if err := w.writeBuf(); err != nil {
+ return err
+ }
+ }
+ w.buf = append(w.buf, c)
+ return nil
+}
+
+func (w *EncryptingWriter) WriteString(s string) (int, error) {
+ if !w.initialized {
+ if err := w.initialize(); err != nil {
+ return 0, err
+ }
+ }
+ if w.err != nil {
+ return 0, w.err
+ }
+ r := strings.NewReader(s)
+ nn := 0
+ for r.Len() > 0 {
+ if len(w.buf) == w.segmentSize {
+ if err := w.writeBuf(); err != nil {
+ return nn, err
+ }
+ }
+ n, _ := r.Read(w.buf[len(w.buf):w.segmentSize])
+ w.buf = w.buf[:len(w.buf)+n]
+ nn += n
+ }
+ return nn, nil
+}
+
// A DecryptingReader decrypts data using the STREAM construction.
type DecryptingReader struct {
r *bufio.Reader
@@ -287,3 +366,44 @@ func (r *DecryptingReader) Read(buf []byte) (int, error) {
n, _ := r.buf.Read(buf)
return n, nil
}
+
+func (r *DecryptingReader) WriteTo(w io.Writer) (int64, error) {
+ if !r.initialized {
+ if err := r.initialize(); err != nil {
+ return 0, err
+ }
+ }
+ var nn int64
+ if r.buf.Len() > 0 {
+ n, err := w.Write(r.buf.Bytes())
+ nn += int64(n)
+ if err != nil {
+ return nn, err
+ }
+ }
+ for {
+ if err := r.fillBuf(); err != nil {
+ return nn, err
+ }
+ n, err := w.Write(r.buf.Bytes())
+ nn += int64(n)
+ if err != nil {
+ return nn, err
+ }
+ }
+}
+
+func (r *DecryptingReader) ReadByte() (byte, error) {
+ if !r.initialized {
+ if err := r.initialize(); err != nil {
+ return 0, err
+ }
+ }
+ if r.buf.Len() == 0 {
+ if err := r.fillBuf(); err != nil {
+ return 0, err
+ }
+ }
+ b, _ := r.buf.ReadByte()
+ return b, nil
+}