dns

DNS client and server implementations using the Go project's dnsmessage package
Log | Files | Refs | README | LICENSE

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 }