dohproxy.go (3583B)
1 package main 2 3 import ( 4 "fmt" 5 "golang.org/x/crypto/acme/autocert" 6 "golang.org/x/net/dns/dnsmessage" 7 "io" 8 "log" 9 "net/http" 10 "os" 11 "strconv" 12 13 "olowe.co/dns" 14 ) 15 16 type metrics struct { 17 httpOK int 18 httpError int 19 httpBadReq int 20 } 21 22 func dnsHandler(w http.ResponseWriter, req *http.Request) { 23 if v, ok := req.Header["Content-Type"]; ok { 24 for _, s := range v { 25 if s != dns.MediaType { 26 err := fmt.Errorf("unsupported media type %s", s) 27 log.Println(err.Error()) 28 http.Error(w, err.Error(), http.StatusUnsupportedMediaType) 29 counter.httpBadReq++ 30 return 31 } 32 } 33 } 34 35 if v, ok := req.Header["Content-Length"]; ok { 36 for _, s := range v { 37 length, err := strconv.Atoi(s) 38 if err != nil { 39 err = fmt.Errorf("parse Content-Length: %v", err) 40 log.Println(err.Error()) 41 http.Error(w, err.Error(), http.StatusInternalServerError) 42 counter.httpError++ 43 return 44 } 45 if length > dns.MaxMsgSize { 46 err = fmt.Errorf("content length %d larger than permitted %d", length, dns.MaxMsgSize) 47 log.Println(err.Error()) 48 http.Error(w, err.Error(), http.StatusRequestEntityTooLarge) 49 counter.httpBadReq++ 50 return 51 } 52 } 53 } 54 55 if req.Method != http.MethodPost && req.Method != http.MethodGet { 56 err := fmt.Errorf("invalid HTTP method %s, must be GET or POST", req.Method) 57 log.Println(err.Error()) 58 http.Error(w, err.Error(), http.StatusNotImplemented) 59 counter.httpBadReq++ 60 return 61 } 62 63 buf := make([]byte, dns.MaxMsgSize) 64 var n int 65 var err error 66 switch req.Method { 67 case http.MethodPost: 68 n, err = req.Body.Read(buf) 69 if err != nil && err != io.EOF { 70 http.Error(w, err.Error(), http.StatusInternalServerError) 71 counter.httpError++ 72 return 73 } 74 req.Body.Close() 75 case http.MethodGet: 76 log.Println("got a GET request but that's not implemented") 77 http.Error(w, "in progress!", http.StatusNotImplemented) 78 return 79 } 80 81 var msg dnsmessage.Message 82 if err := msg.Unpack(buf[:n]); err != nil { 83 log.Println("unpack query:", err) 84 http.Error(w, "unpack query: "+err.Error(), http.StatusInternalServerError) 85 counter.httpError++ 86 return 87 } 88 89 var resolved dnsmessage.Message 90 if conf.usetls { 91 resolved, err = dns.ExchangeTLS(msg, conf.forwardaddr) 92 } else { 93 resolved, err = dns.Exchange(msg, conf.forwardaddr) 94 } 95 if err != nil { 96 log.Println(err.Error()) 97 http.Error(w, err.Error(), http.StatusInternalServerError) 98 counter.httpError++ 99 return 100 } 101 packed, err := resolved.Pack() 102 if err != nil { 103 log.Println("pack resolved query:", err.Error()) 104 http.Error(w, err.Error(), http.StatusInternalServerError) 105 counter.httpError++ 106 return 107 } 108 w.Header().Add("Content-Type", dns.MediaType) 109 if _, err := w.Write(packed); err != nil { 110 log.Fatalln(err) 111 } 112 counter.httpOK++ 113 } 114 115 var conf config 116 var counter metrics 117 118 func metricsHandler (w http.ResponseWriter, req *http.Request) { 119 w.Header().Add("Content-Type", "text/plain") 120 w.Write([]byte("# TYPE http_requests_total counter\n")) 121 w.Write([]byte(fmt.Sprintf("http_requests_total{code=\"%d\"} %d\n", http.StatusOK, counter.httpOK))) 122 w.Write([]byte(fmt.Sprintf("http_requests_total{code=\"%d\"} %d\n", http.StatusInternalServerError, counter.httpError))) 123 w.Write([]byte(fmt.Sprintf("http_requests_total{code=\"4xx\"} %d\n", counter.httpBadReq))) 124 } 125 126 func main() { 127 var err error 128 conf, err = configFromFile("dohproxy.conf") 129 if err != nil { 130 fmt.Fprintln(os.Stderr, "read configuration:", err) 131 os.Exit(1) 132 } 133 http.HandleFunc("/dns-query", dnsHandler) 134 http.HandleFunc("/metrics", metricsHandler) 135 log.Fatalln(http.Serve(autocert.NewListener(conf.listenaddr), nil)) 136 }