commit ded9109ce8770cd39b01c1fb5a8afdc1bd1fc335
parent 0be319b95e9008dcc63694c16025f23e90e74300
Author: Oliver Lowe <o@olowe.co>
Date: Thu, 16 Dec 2021 14:22:12 +1100
New Ask functions
These let us specify just a question without needing to worry about
the entire DNS message structure.
Diffstat:
3 files changed, 47 insertions(+), 28 deletions(-)
diff --git a/cmd/recursor/recursor.go b/cmd/recursor/recursor.go
@@ -6,8 +6,6 @@ import (
"net"
"strings"
"sync"
- "time"
- "math/rand"
"golang.org/x/net/dns/dnsmessage"
"olowe.co/dns"
)
@@ -28,10 +26,6 @@ func ip2dial(ip net.IP) string {
return net.JoinHostPort(ip.String(), "domain")
}
-func newID() uint16 {
- return uint16(rand.Intn(65535))
-}
-
func nextServerAddrs(resources []dnsmessage.Resource) []net.IP {
var next []net.IP
for _, r := range resources {
@@ -46,10 +40,6 @@ func nextServerAddrs(resources []dnsmessage.Resource) []net.IP {
}
func resolve(q dnsmessage.Question, next []net.IP) (dnsmessage.Message, error) {
- qmsg := dnsmessage.Message{
- Header: dnsmessage.Header{ID: newID()},
- Questions: []dnsmessage.Question{q},
- }
var rmsg dnsmessage.Message
var err error
for _, ip := range next {
@@ -58,7 +48,7 @@ func resolve(q dnsmessage.Question, next []net.IP) (dnsmessage.Message, error) {
continue
}
fmt.Fprintf(os.Stderr, "asking %s about %s\n", ip, q.Name)
- rmsg, err = dns.Exchange(qmsg, ip2dial(ip))
+ rmsg, err = dns.Ask(q, ip2dial(ip))
if rmsg.Header.Authoritative {
return rmsg, err
} else if rmsg.Header.RCode == dnsmessage.RCodeSuccess && err == nil {
@@ -134,7 +124,7 @@ func handler(w dns.ResponseWriter, qmsg *dnsmessage.Message) {
fmt.Fprintf(os.Stderr, "cache served %s %s\n", q.Name, q.Type)
return
}
- cache.RUnlock()
+ cache.RUnlock()
resolved, err := resolveFromRoot(q)
if err != nil {
@@ -164,6 +154,5 @@ var cache = struct{
}{m: make(map[dnsmessage.Question][]dnsmessage.Resource)}
func main() {
- rand.Seed(time.Now().UnixNano())
fmt.Fprintln(os.Stderr, dns.ListenAndServe("udp", "", handler))
}
diff --git a/dns.go b/dns.go
@@ -71,7 +71,9 @@ import (
"errors"
"fmt"
"io"
+ "math/rand"
"net"
+ "time"
"golang.org/x/net/dns/dnsmessage"
)
@@ -82,6 +84,42 @@ const MaxMsgSize int = 65535 // max size of a message in bytes
var errMismatchedID = errors.New("mismatched message id")
+var randomsrc *rand.Rand = rand.New(rand.NewSource(time.Now().UnixNano()))
+
+func newID() uint16 {
+ return uint16(randomsrc.Intn(65535))
+}
+
+// Ask sends a message with q to addr and returns its response.
+// The exchange is unencrypted using UDP.
+func Ask(q dnsmessage.Question, addr string) (dnsmessage.Message, error) {
+ qmsg := dnsmessage.Message{
+ Header: dnsmessage.Header{ID: newID()},
+ Questions: []dnsmessage.Question{q},
+ }
+ return Exchange(qmsg, addr)
+}
+
+// Ask sends a message with q to addr and returns its response.
+// The exchange is unencrypted using TCP.
+func AskTCP(q dnsmessage.Question, addr string) (dnsmessage.Message, error) {
+ qmsg := dnsmessage.Message{
+ Header: dnsmessage.Header{ID: newID()},
+ Questions: []dnsmessage.Question{q},
+ }
+ return ExchangeTCP(qmsg, addr)
+}
+
+// Ask sends a message with q to addr and returns its response.
+// The exchange is encrypted using DNS over TLS.
+func AskTLS(q dnsmessage.Question, addr string) (dnsmessage.Message, error) {
+ qmsg := dnsmessage.Message{
+ Header: dnsmessage.Header{ID: newID()},
+ Questions: []dnsmessage.Question{q},
+ }
+ return ExchangeTLS(qmsg, addr)
+}
+
// Exchange performs a synchronous, unencrypted UDP DNS exchange with addr and returns its
// reply to msg.
func Exchange(msg dnsmessage.Message, addr string) (dnsmessage.Message, error) {
diff --git a/server_test.go b/server_test.go
@@ -1,6 +1,7 @@
package dns
import (
+ "golang.org/x/net/dns/dnsmessage"
"testing"
)
@@ -8,11 +9,8 @@ func TestServer(t *testing.T) {
go func() {
t.Fatal(ListenAndServe("udp", "127.0.0.1:51111", nil))
}()
- q, err := buildmsg("www.example.com.")
- if err != nil {
- t.Fatalf("create query: %v", err)
- }
- rmsg, err := Exchange(q, "127.0.0.1:51111")
+ q := dnsmessage.Question{Name: dnsmessage.MustNewName("www.example.com."), Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET}
+ rmsg, err := Ask(q, "127.0.0.1:51111")
if err != nil {
t.Errorf("exchange: %v", err)
}
@@ -23,11 +21,8 @@ func TestStreamServer(t *testing.T) {
go func() {
t.Fatal(ListenAndServe("tcp", "127.0.0.1:51112", nil))
}()
- q, err := buildmsg("www.example.com.")
- if err != nil {
- t.Fatal("create query:", err)
- }
- rmsg, err := ExchangeTCP(q, "127.0.0.1:51112")
+ q := dnsmessage.Question{Name: dnsmessage.MustNewName("www.example.com."), Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET}
+ rmsg, err := AskTCP(q, "127.0.0.1:51112")
if err != nil {
t.Errorf("exchange: %v", err)
}
@@ -40,11 +35,8 @@ func TestEmptyServer(t *testing.T) {
t.Fatal(srv.ListenAndServe())
t.Log(srv.addr)
}()
- q, err := buildmsg("www.example.com.")
- if err != nil {
- t.Fatal("create query:", err)
- }
- rmsg, err := Exchange(q, "127.0.0.1:domain")
+ q := dnsmessage.Question{Name: dnsmessage.MustNewName("www.example.com."), Type: dnsmessage.TypeA, Class: dnsmessage.ClassINET}
+ rmsg, err := Ask(q, "127.0.0.1:domain")
if err != nil {
t.Errorf("exchange: %v", err)
}