recursor.go (1948B)
1 package main 2 3 import ( 4 "fmt" 5 "golang.org/x/net/dns/dnsmessage" 6 "os" 7 8 "olowe.co/dns" 9 ) 10 11 // okQType returns true if t is a query type that we can resolve by 12 // recursively querying nameservers. 13 func okQType(t dnsmessage.Type) bool { 14 switch t { 15 case dnsmessage.TypeA, dnsmessage.TypeNS, dnsmessage.TypeCNAME, dnsmessage.TypeSOA, dnsmessage.TypePTR, dnsmessage.TypeMX, dnsmessage.TypeTXT, dnsmessage.TypeAAAA, dnsmessage.TypeSRV, dnsmessage.TypeOPT: 16 return true 17 } 18 return false 19 } 20 21 // rejectHandler is a safeguard to prevent queries we don't want (or support) 22 // to be recursively resolved. It returns true if the message was rejected. 23 func rejectHandler(w dns.ResponseWriter, qmsg *dnsmessage.Message) bool { 24 if !qmsg.Header.RecursionDesired { 25 dns.Refuse(w, qmsg) 26 return true 27 } else if qmsg.Header.OpCode != dns.OpCodeQUERY { 28 dns.Refuse(w, qmsg) 29 return true 30 } else if len(qmsg.Questions) != 1 { 31 dns.FormatError(w, qmsg) 32 return true 33 } 34 q := qmsg.Questions[0] 35 if !okQType(q.Type) { 36 dns.NotImplemented(w, qmsg) 37 return true 38 } else if q.Class != dnsmessage.ClassINET { 39 dns.NotImplemented(w, qmsg) 40 return true 41 } 42 return false 43 } 44 45 func handler(w dns.ResponseWriter, qmsg *dnsmessage.Message) { 46 if rejected := rejectHandler(w, qmsg); rejected { 47 return 48 } 49 50 var rmsg dnsmessage.Message 51 rmsg.Header.ID = qmsg.Header.ID 52 rmsg.Header.Response = true 53 rmsg.Header.RecursionAvailable = true 54 rmsg.Questions = qmsg.Questions 55 rmsg.RecursionDesired = true 56 57 q := qmsg.Questions[0] 58 resolved, err := resolveFromRoot(q) 59 if err != nil { 60 fmt.Fprintln(os.Stderr, err) 61 rmsg.Header.RCode = dnsmessage.RCodeServerFailure 62 w.WriteMsg(rmsg) 63 return 64 } 65 rmsg.Header.RCode = resolved.Header.RCode 66 rmsg.Answers = resolved.Answers 67 if len(rmsg.Answers) == 0 { 68 rmsg.Authorities = resolved.Authorities 69 w.WriteMsg(rmsg) 70 return 71 } 72 w.WriteMsg(rmsg) 73 } 74 75 func main() { 76 fmt.Fprintln(os.Stderr, dns.ListenAndServe("udp", "", handler)) 77 }