aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--oae2.go57
-rw-r--r--oae2_test.go84
2 files changed, 86 insertions, 55 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
diff --git a/oae2_test.go b/oae2_test.go
index b60be24..4099aa7 100644
--- a/oae2_test.go
+++ b/oae2_test.go
@@ -10,19 +10,20 @@ import (
func TestRoundTrip(t *testing.T) {
t.Parallel()
const (
- msg = "Hello World!"
- key = "password123"
- blockSize = 1
+ msg = "Hello World!"
+ key = "password123"
+ additionalData = "additional data"
+ blockSize = 1
)
buf := new(bytes.Buffer)
- w := NewWriter(buf, []byte(key), blockSize)
+ w := NewWriter(buf, []byte(key), blockSize, []byte(additionalData))
if _, err := io.WriteString(w, msg); err != nil {
t.Fatalf("Writer.Write failed: %s", err)
}
if err := w.Close(); err != nil {
t.Fatalf("Writer.Close failed: %s", err)
}
- got, err := io.ReadAll(NewReader(buf, []byte(key), blockSize))
+ got, err := io.ReadAll(NewReader(buf, []byte(key), blockSize, []byte(additionalData)))
if err != nil {
t.Fatalf("Reader.Read failed: %s", err)
}
@@ -34,12 +35,13 @@ func TestRoundTrip(t *testing.T) {
func TestReadFromWriteTo(t *testing.T) {
t.Parallel()
const (
- msg = "Hello World!"
- key = "asdf"
- blockSize = 2
+ msg = "Hello World!"
+ key = "asdf"
+ additionalData = "additional data"
+ blockSize = 2
)
buf := new(bytes.Buffer)
- w := NewWriter(buf, []byte(key), blockSize)
+ w := NewWriter(buf, []byte(key), blockSize, []byte(additionalData))
if _, err := io.Copy(w, struct{ io.Reader }{strings.NewReader(msg)}); err != nil {
t.Fatalf("Writer.WriteTo failed: %s", err)
}
@@ -47,7 +49,7 @@ func TestReadFromWriteTo(t *testing.T) {
t.Fatalf("Writer.Close failed: %s", err)
}
got := new(strings.Builder)
- if _, err := io.Copy(got, NewReader(buf, []byte(key), blockSize)); err != nil {
+ if _, err := io.Copy(got, NewReader(buf, []byte(key), blockSize, []byte(additionalData))); err != nil {
t.Fatalf("Reader.Read failed: %s", err)
}
if got.String() != msg {
@@ -58,9 +60,10 @@ func TestReadFromWriteTo(t *testing.T) {
func TestReader_Seek(t *testing.T) {
t.Parallel()
const (
- msg = "Hello World!"
- key = "asdf"
- blockSize = 1
+ msg = "Hello World!"
+ key = "asdf"
+ additionalData = "additional data"
+ blockSize = 1
)
for _, tc := range []struct {
desc string
@@ -97,14 +100,14 @@ func TestReader_Seek(t *testing.T) {
t.Parallel()
buf := new(bytes.Buffer)
- w := NewWriter(buf, []byte(key), blockSize)
+ w := NewWriter(buf, []byte(key), blockSize, []byte(additionalData))
if _, err := io.WriteString(w, msg); err != nil {
t.Fatalf("Writer.Write failed: %s", err)
}
if err := w.Close(); err != nil {
t.Fatalf("Writer.Close failed: %s", err)
}
- r := NewReader(bytes.NewReader(buf.Bytes()), []byte(key), blockSize)
+ r := NewReader(bytes.NewReader(buf.Bytes()), []byte(key), blockSize, []byte(additionalData))
n, err := r.Seek(tc.offset, tc.whence)
if err != nil {
t.Fatalf("Reader.Seek(%d, %d) failed: %s", tc.offset, tc.whence, err)
@@ -126,13 +129,16 @@ func TestReader_Seek(t *testing.T) {
func FuzzRoundTrip(f *testing.F) {
f.Add(1, []byte("Hello World!"))
f.Fuzz(func(t *testing.T, segmentSize int, msg []byte) {
- const password = "asdf"
+ const (
+ password = "asdf"
+ additionalData = "additional data"
+ )
if segmentSize <= 0 {
return
}
buf := new(bytes.Buffer)
- w := NewWriter(buf, []byte(password), segmentSize)
+ w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData))
n, err := w.Write(msg)
if err != nil {
t.Fatalf("Writer.Write(%q) failed: %s", msg, err)
@@ -143,7 +149,7 @@ func FuzzRoundTrip(f *testing.F) {
if err := w.Close(); err != nil {
t.Fatalf("Writer.Close failed: %s", err)
}
- got, err := io.ReadAll(NewReader(buf, []byte(password), segmentSize))
+ got, err := io.ReadAll(NewReader(buf, []byte(password), segmentSize, []byte(additionalData)))
if err != nil {
t.Fatalf("Reader.Read failed: %s", err)
}
@@ -156,13 +162,16 @@ func FuzzRoundTrip(f *testing.F) {
func FuzzReadFromWriteTo(f *testing.F) {
f.Add(1, "Hello World!")
f.Fuzz(func(t *testing.T, segmentSize int, msg string) {
- const password = "asdf"
+ const (
+ password = "asdf"
+ additionalData = "additional data"
+ )
if segmentSize <= 0 {
return
}
buf := new(bytes.Buffer)
- w := NewWriter(buf, []byte(password), segmentSize)
+ w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData))
n, err := w.ReadFrom(strings.NewReader(msg))
if err != nil {
t.Fatalf("Writer.ReadFrom(%q) failed: %s", msg, err)
@@ -174,7 +183,7 @@ func FuzzReadFromWriteTo(f *testing.F) {
t.Fatalf("Writer.Close failed: %s", err)
}
got := new(strings.Builder)
- n, err = NewReader(buf, []byte(password), segmentSize).WriteTo(got)
+ n, err = NewReader(buf, []byte(password), segmentSize, []byte(additionalData)).WriteTo(got)
if err != nil {
t.Fatalf("Reader.WriteTo failed: %s", err)
}
@@ -191,20 +200,23 @@ func FuzzSeekStart(f *testing.F) {
f.Add(1, []byte("Hello World!"), 6)
f.Add(30, []byte("000000"), 6)
f.Fuzz(func(t *testing.T, segmentSize int, msg []byte, offset int) {
- const password = "asdf"
+ const (
+ password = "asdf"
+ additionalData = "additional data"
+ )
if segmentSize <= 0 || offset > len(msg) {
return
}
buf := new(bytes.Buffer)
- w := NewWriter(buf, []byte(password), segmentSize)
+ w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData))
if _, err := w.Write(msg); err != nil {
t.Fatalf("Writer.Write(%q) failed: %s", msg, err)
}
if err := w.Close(); err != nil {
t.Fatalf("Writer.Close failed: %s", err)
}
- r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize)
+ r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize, []byte(additionalData))
n, err := r.Seek(int64(offset), io.SeekStart)
if offset < 0 {
if err == nil {
@@ -231,20 +243,23 @@ func FuzzSeekStart(f *testing.F) {
func FuzzSeekEnd(f *testing.F) {
f.Add(1, []byte("Hello World!"), -6)
f.Fuzz(func(t *testing.T, segmentSize int, msg []byte, offset int) {
- const password = "asdf"
+ const (
+ password = "asdf"
+ additionalData = "additionalData"
+ )
if segmentSize <= 0 || offset > 0 {
return
}
buf := new(bytes.Buffer)
- w := NewWriter(buf, []byte(password), segmentSize)
+ w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData))
if _, err := w.Write(msg); err != nil {
t.Fatalf("Writer.Write(%q) failed: %s", msg, err)
}
if err := w.Close(); err != nil {
t.Fatalf("Writer.Close failed: %s", err)
}
- r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize)
+ r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize, []byte(additionalData))
n, err := r.Seek(int64(offset), io.SeekEnd)
if len(msg)+offset < 0 {
if err == nil {
@@ -273,20 +288,23 @@ func FuzzSeekCurrent(f *testing.F) {
f.Add(1, []byte("000000"), 3, 3)
f.Add(1, []byte("000"), 3, -3)
f.Fuzz(func(t *testing.T, segmentSize int, msg []byte, offset1, offset2 int) {
- const password = "asdf"
+ const (
+ password = "asdf"
+ additionalData = "additional data"
+ )
if segmentSize <= 0 || offset1 < 0 || offset1 > len(msg) || offset1+offset2 > len(msg) {
return
}
buf := new(bytes.Buffer)
- w := NewWriter(buf, []byte(password), segmentSize)
+ w := NewWriter(buf, []byte(password), segmentSize, []byte(additionalData))
if _, err := w.Write(msg); err != nil {
t.Fatalf("Writer.Write(%q) failed: %s", msg, err)
}
if err := w.Close(); err != nil {
t.Fatalf("Writer.Close failed: %s", err)
}
- r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize)
+ r := NewReader(bytes.NewReader(buf.Bytes()), []byte(password), segmentSize, []byte(additionalData))
if _, err := io.ReadFull(r, make([]byte, offset1)); err != nil {
t.Fatalf("Reader.Read failed: %s", err)
}
@@ -314,7 +332,7 @@ func FuzzSeekCurrent(f *testing.F) {
}
func encrypt(buf *bytes.Buffer, data []byte) error {
- w := NewWriter(buf, []byte("asdf"), 4*1024*1024)
+ w := NewWriter(buf, []byte("asdf"), 4*1024*1024, nil)
if _, err := w.Write(data); err != nil {
return err
}
@@ -333,7 +351,7 @@ func BenchmarkWriter(b *testing.B) {
}
func decrypt(out []byte, data []byte) error {
- r := NewReader(bytes.NewReader(data), []byte("asdf"), 4*1024*1024)
+ r := NewReader(bytes.NewReader(data), []byte("asdf"), 4*1024*1024, nil)
_, err := io.ReadFull(r, out)
return err
}
@@ -345,7 +363,7 @@ func BenchmarkReader(b *testing.B) {
)
data := make([]byte, 10*1024*1024)
encryptedBuf := new(bytes.Buffer)
- w := NewWriter(encryptedBuf, []byte(password), segmentSize)
+ w := NewWriter(encryptedBuf, []byte(password), segmentSize, nil)
if _, err := w.Write(data); err != nil {
b.Fatal(err)
}