diff --git a/dns/dns.go b/dns/dns.go index eb7b638..7b7458a 100644 --- a/dns/dns.go +++ b/dns/dns.go @@ -5,7 +5,9 @@ import ( "bytes" "encoding/binary" "errors" + "fmt" "io" + "strings" ) // The maximum number of DNS name compression pointers we are willing to follow. @@ -103,14 +105,31 @@ func ParseName(s string) (Name, error) { } } -// String returns a string representation of name, with labels separated by -// dots. +// String returns a reversible string representation of name. Labels are +// separated by dots, and any bytes in a label that are outside the set +// [0-9A-Za-z-] are replaced with a \xXX hex escape sequence. func (name Name) String() string { if len(name) == 0 { return "." - } else { - return string(bytes.Join(name, []byte("."))) } + + var buf strings.Builder + for i, label := range name { + if i > 0 { + buf.WriteByte('.') + } + for _, b := range label { + if b == '-' || + ('0' <= b && b <= '9') || + ('A' <= b && b <= 'Z') || + ('a' <= b && b <= 'z') { + buf.WriteByte(b) + } else { + fmt.Fprintf(&buf, "\\x%02x", b) + } + } + } + return buf.String() } // TrimSuffix returns a Name with the given suffix removed, if it was present. diff --git a/dns/dns_test.go b/dns/dns_test.go index 4603e63..c7ae490 100644 --- a/dns/dns_test.go +++ b/dns/dns_test.go @@ -2,7 +2,9 @@ package dns import ( "bytes" + "fmt" "io" + "strconv" "strings" "testing" ) @@ -19,15 +21,6 @@ func namesEqual(a, b Name) bool { return true } -func anyLabelContainsDot(labels [][]byte) bool { - for _, label := range labels { - if bytes.Contains(label, []byte(".")) { - return true - } - } - return false -} - func TestName(t *testing.T) { for _, test := range []struct { labels [][]byte @@ -85,10 +78,6 @@ func TestName(t *testing.T) { {'0'}, {'1'}, {'2'}, {'3'}, {'4'}, {'5'}, {'6'}, {'7'}, {'8'}, {'9'}, {'a'}, {'b'}, {'c'}, {'d'}, {'e'}, {'f'}, {'0'}, {'1'}, {'2'}, {'3'}, {'4'}, {'5'}, {'6'}, {'7'}, {'8'}, {'9'}, {'A'}, {'B'}, {'C'}, {'D'}, {'E'}, {'F'}, }, ErrNameTooLong, ""}, - - // Labels may contain any octets, though ones containing dots - // cannot be losslessly roundtripped through a string. - {[][]byte{[]byte("\x00"), []byte("a.b")}, nil, "\x00.a.b"}, } { // Test that NewName returns proper error codes, and otherwise // returns an equal slice of labels. @@ -112,22 +101,20 @@ func TestName(t *testing.T) { // Test that parsing from a string back to a Name results in the // original slice of labels. - if !anyLabelContainsDot(test.labels) { - name, err := ParseName(s) - if err != nil || !namesEqual(name, test.labels) { + name, err = ParseName(s) + if err != nil || !namesEqual(name, test.labels) { + t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)", + test.labels, s, name, err, test.labels, nil) + continue + } + // A trailing dot should be ignored. + if !strings.HasSuffix(s, ".") { + dotName, dotErr := ParseName(s + ".") + if dotErr != err || !namesEqual(dotName, name) { t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)", - test.labels, s, name, err, test.labels, nil) + test.labels, s+".", dotName, dotErr, name, err) continue } - // A trailing dot should be ignored. - if !strings.HasSuffix(s, ".") { - dotName, dotErr := ParseName(s + ".") - if dotErr != err || !namesEqual(dotName, name) { - t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)", - test.labels, s+".", dotName, dotErr, name, err) - continue - } - } } } } @@ -151,6 +138,65 @@ func TestParseName(t *testing.T) { } } +func unescapeString(s string) ([][]byte, error) { + if s == "." { + return [][]byte{}, nil + } + + var result [][]byte + for _, label := range strings.Split(s, ".") { + var buf bytes.Buffer + i := 0 + for i < len(label) { + switch label[i] { + case '\\': + if i+3 >= len(label) { + return nil, fmt.Errorf("truncated escape sequence at index %v", i) + } + if label[i+1] != 'x' { + return nil, fmt.Errorf("malformed escape sequence at index %v", i) + } + b, err := strconv.ParseInt(string(label[i+2:i+4]), 16, 8) + if err != nil { + return nil, fmt.Errorf("malformed hex sequence at index %v", i+2) + } + buf.WriteByte(byte(b)) + i += 4 + default: + buf.WriteByte(label[i]) + i++ + } + } + result = append(result, buf.Bytes()) + } + return result, nil +} + +func TestNameString(t *testing.T) { + for _, test := range []struct { + name Name + s string + }{ + {[][]byte{}, "."}, + {[][]byte{[]byte("\x00"), []byte("a.b"), []byte("c\nd\\")}, "\\x00.a\\x2eb.c\\x0ad\\x5c"}, + } { + s := test.name.String() + if s != test.s { + t.Errorf("%+q escaped to %+q, expected %+q", test.name, s, test.s) + continue + } + unescaped, err := unescapeString(s) + if err != nil { + t.Errorf("%+q unescaping %+q resulted in error %v", test.name, s, err) + continue + } + if !namesEqual(Name(unescaped), test.name) { + t.Errorf("%+q roundtripped through %+q to %+q", test.name, s, unescaped) + continue + } + } +} + func TestNameTrimSuffix(t *testing.T) { for _, test := range []struct { name, suffix string diff --git a/dnstt-server/main.go b/dnstt-server/main.go index 753fdff..cf40cb7 100644 --- a/dnstt-server/main.go +++ b/dnstt-server/main.go @@ -692,9 +692,17 @@ func computeMaxEncodedPayload(limit int) int { if err != nil { panic(err) } - if len(maxLengthName.String())+2 != 255 { - panic(fmt.Sprintf("max-length name is %d octets, should be %d %s", - len(maxLengthName.String())+2, 255, maxLengthName)) + { + // Compute the encoded length of maxLengthName and that its + // length is actually at the maximum of 255 octets. + n := 0 + for _, label := range maxLengthName { + n += len(label) + 1 + } + n += 1 // For the terminating null label. + if n != 255 { + panic(fmt.Sprintf("max-length name is %d octets, should be %d %s", n, 255, maxLengthName)) + } } queryLimit := uint16(limit)