From 6ff4e9c0622dee2663d70059594e58b3c96f135a Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Sat, 4 Oct 2025 21:17:49 -0700 Subject: Add some extra methods to oae2 --- internal/cryptoutil/oae2.go | 134 +++++++++++++++++++++++++++++++++++++++++--- 1 file 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 +} -- cgit v1.3.1