mirror of
https://github.com/tladesignz/dnstt.git
synced 2026-10-02 20:08:06 +03:00
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:
+5
-3
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user