aboutsummaryrefslogtreecommitdiffstats
path: root/oae2.go
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2026-01-25 07:47:19 -0800
committerRose Hogenson <rosehogenson@posteo.net>2026-01-25 08:12:23 -0800
commit6412b3912b17cd22ba8bdbf07a12f8b005bf058f (patch)
treedbe5d21e300a20d0617e3a6b8bd21f59a4b1adbf /oae2.go
parentfafe0213c9022f298ec391c7907ed92762fbcd1b (diff)
downloadoae2-6412b3912b17cd22ba8bdbf07a12f8b005bf058f.tar.zst
Add support for additional data
We can't really call it "streaming AEAD" and then not support additional data 🤦‍♀️
Diffstat (limited to 'oae2.go')
-rw-r--r--oae2.go57
1 files changed, 35 insertions, 22 deletions
diff --git a/oae2.go b/oae2.go
index adef4f7..e2bb1c1 100644
--- a/oae2.go
+++ b/oae2.go
@@ -40,7 +40,8 @@ const (
)
type segmentEncrypter struct {
- key []byte
+ key []byte
+ additionalData []byte
aead cipher.AEAD
nonce [nonceSize]byte
@@ -59,6 +60,16 @@ func (e *segmentEncrypter) init(salt []byte) error {
return err
}
+func (e *segmentEncrypter) firstSegment() bool {
+ n := e.nonce
+ var or byte
+ // First 11 bytes of the nonce are the segment index
+ for _, b := range n[:nonceSize-1] {
+ or |= b
+ }
+ return or != 0
+}
+
func (e *segmentEncrypter) nextNonce(lastSegment bool) ([]byte, error) {
for i := 0; ; i++ {
if i == nonceSize-1 {
@@ -76,19 +87,27 @@ func (e *segmentEncrypter) nextNonce(lastSegment bool) ([]byte, error) {
}
func (e *segmentEncrypter) encryptSegment(segment []byte, lastSegment bool) ([]byte, error) {
+ var ad []byte
+ if e.firstSegment() {
+ ad = e.additionalData
+ }
nonce, err := e.nextNonce(lastSegment)
if err != nil {
return nil, err
}
- return e.aead.Seal(segment[:len(segment)+aeadOverhead][:0], nonce, segment, nil), nil
+ return e.aead.Seal(segment[:len(segment)+aeadOverhead][:0], nonce, segment, ad), nil
}
func (e *segmentEncrypter) decryptSegment(segment []byte, lastSegment bool) ([]byte, error) {
+ var ad []byte
+ if e.firstSegment() {
+ ad = e.additionalData
+ }
nonce, err := e.nextNonce(lastSegment)
if err != nil {
return nil, err
}
- return e.aead.Open(segment[:0], nonce, segment, nil)
+ return e.aead.Open(segment[:0], nonce, segment, ad)
}
// A Writer wraps an io.Writer and encrypts the data in segments. Make sure to
@@ -108,14 +127,14 @@ type Writer struct {
// panics if segmentSize <= 0.
//
// Make sure to call [Writer.Close] to flush the final segment.
-func NewWriter(w io.Writer, key []byte, segmentSize int) *Writer {
+func NewWriter(w io.Writer, key []byte, segmentSize int, additionalData []byte) *Writer {
if segmentSize <= 0 {
panic("oae2.NewWriter: segmentSize must be strictly greater than 0")
}
return &Writer{
w: w,
segmentSize: segmentSize,
- encrypter: segmentEncrypter{key: key},
+ encrypter: segmentEncrypter{key: key, additionalData: additionalData},
buf: make([]byte, 0, segmentSize+aeadOverhead),
}
}
@@ -261,40 +280,33 @@ type Reader struct {
// NewReader returns a Reader that wraps r and decrypts the data using key in
// chunks of size segmentSize. segmentSize must match the segment size that was
// used to write the encrypted stream. NewReader panics if segmentSize <= 0.
-func NewReader(r io.Reader, key []byte, segmentSize int) *Reader {
+func NewReader(r io.Reader, key []byte, segmentSize int, additionalData []byte) *Reader {
if segmentSize <= 0 {
panic("oae2.NewReader: segmentSize must be strictly greater than 0")
}
return &Reader{
r: bufReader{r: r},
segmentSize: segmentSize,
- decrypter: segmentEncrypter{key: key},
+ decrypter: segmentEncrypter{key: key, additionalData: additionalData},
buf: make([]byte, 0, segmentSize+aeadOverhead+1),
}
}
func (r *Reader) initialize() error {
r.initialized = true
- // N.B. there's a subtle bug lurking here: for an input stream of
- // exactly 32 bytes, it's important that we reject this stream since it
- // doesn't have an authentication tag. So if we naively read exactly 32
- // bytes inside initialize, then Read will see that the underlying
- // reader is at EOF, and won't be able to distinguish this from a valid
- // EOF case. Fortunately we can easily work around this by reading one
- // extra byte here: since the shortest possible encrypted stream is
- // 32 + 16 bytes, every valid stream will have an extra byte for us
- // here, and other than on the first segment, Read always knows
- // precisely whether the stream is at EOF.
- buf := make([]byte, saltSize+1)
+ buf := make([]byte, saltSize)
if _, r.err = io.ReadFull(&r.r, buf); r.err != nil {
if r.err == io.EOF {
r.err = io.ErrUnexpectedEOF
}
return r.err
}
- r.r.unreadByte()
- r.err = r.decrypter.init(buf[:saltSize])
- return r.err
+ if r.err = r.decrypter.init(buf); r.err != nil {
+ return r.err
+ }
+ // Decrypt the first segment now to make sure we validate the
+ // additional data
+ return r.fillBuf()
}
func (r *Reader) init() error {
@@ -335,7 +347,7 @@ func (r *Reader) Read(buf []byte) (int, error) {
if err := r.init(); err != nil {
return 0, err
}
- if r.nRead == len(r.buf) {
+ if r.nRead >= len(r.buf) {
if err := r.fillBuf(); err != nil {
return 0, err
}
@@ -445,6 +457,7 @@ func (r *Reader) Seek(offset int64, whence int) (int64, error) {
if err != nil {
return 0, err
}
+ r.buf = r.buf[:0]
plaintextSize := r.encryptedToPlaintextSize(encryptedSize)
if offset == 0 {
return plaintextSize, nil