dns.go (6216B)
1 /* 2 Package dns provides small DNS client and server implementations built 3 around the Go project's dnsmessage package. It supports both UDP and 4 TCP (including TLS). 5 6 The package deliberately does not implement all features of the DNS 7 specifications. Notably EDNS and DNSSEC are unsupported. 8 9 The most basic operation is creating a question, asking the DNS server 10 the question, then handling the response using Ask: 11 12 q := dnsmessage.Question{ 13 Name: dnsmessage.MustNewName("www.example.com."), 14 Type: dnsmessage.TypeAAAA, 15 } 16 resp, err := dns.Ask(q, "192.0.2.1:domain") 17 // ... 18 19 Queries to a recursive resolver via DNS over TLS (DoT) can be made with ExchangeTLS: 20 21 name, err := dnsmessage.NewName("www.example.com.") 22 if err != nil { 23 // handle error 24 } 25 qmsg := dnsmessage.Message{ 26 Header: dnsmessage.Header{ID: 69, RecursionDesired: true}, 27 Questions: []dnsmessage.Question{ 28 dnsmessage.Question{ 29 Name: name, 30 Type: dnsmessage.TypeA, 31 Class: dnsmessage.ClassINET, 32 }, 33 }, 34 } 35 rmsg, err := dns.ExchangeTLS(qmsg, "192.0.2.1:853") 36 37 ListenAndServe starts a DNS server listening on the given network and 38 address. Received messages are managed with the given Handler in a new 39 goroutine. Handler may be nil, in which case all messages are 40 gracefully refused. 41 42 log.Fatal(dns.ListenAndServe("udp", ":domain", nil)) 43 44 Handlers are just functions to which a DNS message from the server is 45 passed. Responses are written to ResponseWriter. 46 47 func myHandler(w dns.ResponseWriter, qmsg *dnsmessage.Message) { 48 var rmsg dnsmessage.Message 49 rmsg.Header.ID = qmsg.Header.ID 50 if rmsg.Header.RecursionDesired { 51 rmsg.Header.RCode = dnsmessage.RCodeRefused 52 w.WriteMsg(rmsg) 53 return 54 } 55 // answer questions... 56 } 57 58 A Server may be created with a custom net.Listener: 59 60 l, err := tls.Listen(network, addr, config) 61 srv := &dns.Server{Handler: myHandler} 62 log.Fatal(srv.Serve(l)) 63 64 */ 65 package dns 66 67 import ( 68 "crypto/tls" 69 "errors" 70 "fmt" 71 "io" 72 "math/rand" 73 "net" 74 "time" 75 76 "golang.org/x/net/dns/dnsmessage" 77 ) 78 79 // https://datatracker.ietf.org/doc/html/rfc8484 80 const MediaType string = "application/dns-message" 81 const MaxMsgSize int = 65535 // max size of a message in bytes 82 83 const OpCodeQUERY dnsmessage.OpCode = 0 84 85 var errMismatchedID = errors.New("mismatched message id") 86 87 var randomsrc *rand.Rand = rand.New(rand.NewSource(time.Now().UnixNano())) 88 89 func newID() uint16 { 90 return uint16(randomsrc.Intn(65535)) 91 } 92 93 // Ask sends a message with q to addr and returns its response. 94 // The exchange is unencrypted using UDP. 95 func Ask(q dnsmessage.Question, addr string) (dnsmessage.Message, error) { 96 qmsg := dnsmessage.Message{ 97 Header: dnsmessage.Header{ID: newID()}, 98 Questions: []dnsmessage.Question{q}, 99 } 100 return Exchange(qmsg, addr) 101 } 102 103 // Ask sends a message with q to addr and returns its response. 104 // The exchange is unencrypted using TCP. 105 func AskTCP(q dnsmessage.Question, addr string) (dnsmessage.Message, error) { 106 qmsg := dnsmessage.Message{ 107 Header: dnsmessage.Header{ID: newID()}, 108 Questions: []dnsmessage.Question{q}, 109 } 110 return ExchangeTCP(qmsg, addr) 111 } 112 113 // Ask sends a message with q to addr and returns its response. 114 // The exchange is encrypted using DNS over TLS. 115 func AskTLS(q dnsmessage.Question, addr string) (dnsmessage.Message, error) { 116 qmsg := dnsmessage.Message{ 117 Header: dnsmessage.Header{ID: newID()}, 118 Questions: []dnsmessage.Question{q}, 119 } 120 return ExchangeTLS(qmsg, addr) 121 } 122 123 // Exchange performs a synchronous, unencrypted UDP DNS exchange with addr and returns its 124 // reply to msg. 125 func Exchange(msg dnsmessage.Message, addr string) (dnsmessage.Message, error) { 126 conn, err := net.Dial("udp", addr) 127 if err != nil { 128 return dnsmessage.Message{}, err 129 } 130 defer conn.Close() 131 return exchange(msg, conn) 132 } 133 134 func ExchangeTCP(msg dnsmessage.Message, addr string) (dnsmessage.Message, error) { 135 conn, err := net.Dial("tcp", addr) 136 if err != nil { 137 return dnsmessage.Message{}, err 138 } 139 defer conn.Close() 140 return exchange(msg, conn) 141 } 142 143 // ExchangeTLS performs a synchronous DNS-over-TLS exchange with addr and returns its 144 // reply to msg. 145 func ExchangeTLS(msg dnsmessage.Message, addr string) (dnsmessage.Message, error) { 146 conn, err := tls.Dial("tcp", addr, nil) 147 if err != nil { 148 return dnsmessage.Message{}, err 149 } 150 defer conn.Close() 151 return exchange(msg, conn) 152 } 153 154 func exchange(msg dnsmessage.Message, conn net.Conn) (dnsmessage.Message, error) { 155 if err := sendMsg(msg, conn); err != nil { 156 return dnsmessage.Message{}, err 157 } 158 rmsg, err := receive(conn) 159 if err != nil { 160 return dnsmessage.Message{}, err 161 } 162 if rmsg.Header.ID != msg.Header.ID { 163 return rmsg, errMismatchedID 164 } else if rmsg.Questions[0] != msg.Questions[0] { 165 return rmsg, fmt.Errorf("mismatched response to question") 166 } 167 return rmsg, nil 168 } 169 170 func sendMsg(msg dnsmessage.Message, conn net.Conn) error { 171 packed, err := msg.Pack() 172 if err != nil { 173 return err 174 } 175 _, err = send(packed, conn) 176 return err 177 } 178 179 func send(p []byte, conn net.Conn) (int, error) { 180 if _, ok := conn.(net.PacketConn); ok { 181 return conn.Write(p) 182 } 183 // DNS over TCP requires you to prepend the message with a 184 // 2-octet length field. 185 l := len(p) 186 m := make([]byte, 2+l) 187 m[0] = byte(l >> 8) 188 m[1] = byte(l) 189 copy(m[2:], p) 190 return conn.Write(m) 191 } 192 193 func sendMsgTo(msg dnsmessage.Message, conn net.PacketConn, addr net.Addr) error { 194 packed, err := msg.Pack() 195 if err != nil { 196 return err 197 } 198 _, err = conn.WriteTo(packed, addr) 199 return err 200 } 201 202 func receive(conn net.Conn) (dnsmessage.Message, error) { 203 var buf []byte 204 var n int 205 var err error 206 if _, ok := conn.(net.PacketConn); ok { 207 buf = make([]byte, 512) 208 n, err = conn.Read(buf) 209 if err != nil { 210 return dnsmessage.Message{}, err 211 } 212 } else { 213 buf = make([]byte, 1280) 214 if _, err := io.ReadFull(conn, buf[:2]); err != nil { 215 return dnsmessage.Message{}, fmt.Errorf("read length: %w", err) 216 } 217 l := int(buf[0])<<8 | int(buf[1]) 218 if l > len(buf) { 219 buf = make([]byte, l) 220 } 221 n, err = io.ReadFull(conn, buf[:l]) 222 if err != nil { 223 return dnsmessage.Message{}, fmt.Errorf("read after length: %w", err) 224 } 225 } 226 var msg dnsmessage.Message 227 if err := msg.Unpack(buf[:n]); err != nil { 228 return dnsmessage.Message{}, err 229 } 230 return msg, nil 231 }