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.
This commit is contained in:
David Fifield
2020-04-19 16:52:48 -06:00
parent a6af2f1df1
commit 83a67e3874
2 changed files with 70 additions and 3 deletions
+5 -3
View File
@@ -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 {
+65
View File
@@ -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