From 83a67e3874bc58bf15d1d8064c58dc32742e9431 Mon Sep 17 00:00:00 2001 From: David Fifield Date: Sun, 19 Apr 2020 16:33:20 -0600 Subject: [PATCH] Fix a bug in noise.readMessage. Was returning nil in case of an error from io.ReadFull, and even then it should have been io.ErrUnexpectedEOF, not io.EOF. --- noise/noise.go | 8 +++--- noise/noise_test.go | 65 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 70 insertions(+), 3 deletions(-) diff --git a/noise/noise.go b/noise/noise.go index a9622ed..eb1dc41 100644 --- a/noise/noise.go +++ b/noise/noise.go @@ -21,14 +21,16 @@ func readMessage(r io.Reader) ([]byte, error) { var length uint16 err := binary.Read(r, binary.BigEndian, &length) if err != nil { + // We may return a real io.EOF only here. return nil, err } msg := make([]byte, int(length)) _, err = io.ReadFull(r, msg) - if err != nil { - return nil, nil + // Here we must change io.EOF to io.ErrUnexpectedEOF. + if err == io.EOF { + err = io.ErrUnexpectedEOF } - return msg, nil + return msg, err } func writeMessage(w io.Writer, msg []byte) error { diff --git a/noise/noise_test.go b/noise/noise_test.go index dac100d..66a2e4f 100644 --- a/noise/noise_test.go +++ b/noise/noise_test.go @@ -2,9 +2,74 @@ package noise import ( "bytes" + "io" "testing" ) +func allMessages(buf []byte) ([][]byte, error) { + var messages [][]byte + r := bytes.NewReader(buf) + for { + msg, err := readMessage(r) + if err != nil { + return messages, err + } + messages = append(messages, msg) + } +} + +func messagesEqual(a, b [][]byte) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if !bytes.Equal(a[i], b[i]) { + return false + } + } + return true +} + +func TestReadMessage(t *testing.T) { + for _, test := range []struct { + input string + messages [][]byte + err error + }{ + {"", [][]byte{}, io.EOF}, + {"\x00", [][]byte{}, io.ErrUnexpectedEOF}, + {"\x00\x00", [][]byte{{}}, io.EOF}, + {"\x00\x00\x00", [][]byte{{}}, io.ErrUnexpectedEOF}, + {"\x00\x01", [][]byte{}, io.ErrUnexpectedEOF}, + {"\x00\x05hello\x00\x05world", [][]byte{[]byte("hello"), []byte("world")}, io.EOF}, + } { + packets, err := allMessages([]byte(test.input)) + if !messagesEqual(packets, test.messages) || err != test.err { + t.Errorf("%x\nreturned %x %v\nexpected %x %v", + test.input, packets, err, test.messages, test.err) + } + } +} + +func TestMessageRoundTrip(t *testing.T) { + for _, messages := range [][][]byte{ + {}, + } { + var buf bytes.Buffer + for _, msg := range messages { + err := writeMessage(&buf, msg) + if err != nil { + panic(err) + } + } + output, err := allMessages(buf.Bytes()) + if !messagesEqual(output, messages) || err != io.EOF { + t.Errorf("%x roundtripped to %x %v", + messages, output, err) + } + } +} + func TestReadKey(t *testing.T) { for _, test := range []struct { input string