559 lines
14 KiB
Go
559 lines
14 KiB
Go
package main
|
|
|
|
import (
|
|
"crypto/tls"
|
|
"errors"
|
|
"flag"
|
|
"fmt"
|
|
"io"
|
|
"log"
|
|
"math/rand"
|
|
"net"
|
|
"net/http"
|
|
"os"
|
|
"os/signal"
|
|
"runtime"
|
|
"strings"
|
|
"sync"
|
|
"syscall"
|
|
"time"
|
|
|
|
"github.com/miekg/dns"
|
|
"github.com/patrickmn/go-cache"
|
|
"github.com/prometheus/client_golang/prometheus"
|
|
"github.com/prometheus/client_golang/prometheus/promauto"
|
|
"github.com/prometheus/client_golang/prometheus/promhttp"
|
|
"golang.org/x/sync/singleflight"
|
|
)
|
|
|
|
var (
|
|
addr = flag.String("a", "0.0.0.0", "address to use")
|
|
port = flag.String("p", "53", "port to run on")
|
|
ns = flag.String("n", "9.9.9.9", "nameservers to use, separated by ':' or ','")
|
|
blocklistPath = flag.String("b", "", "blocklist path")
|
|
zonesPath = flag.String("z", "", "add a custom zone file")
|
|
hostsPath = flag.String("h", "", "add a custom hosts file")
|
|
cpu = flag.Int("c", 0, "number of cpu to use")
|
|
ttl = flag.Int("t", 10, "cache TTL for queries, in minutes")
|
|
dot = flag.Bool("s", false, "use tls to contact the resolver")
|
|
logs = flag.Bool("l", false, "log queries")
|
|
metrics = flag.Bool("m", false, "enable prometheus metrics")
|
|
qcache *cache.Cache
|
|
inflight singleflight.Group
|
|
|
|
blocklistMap = make(map[string]bool)
|
|
zonesMap = make(map[string]string)
|
|
hostsMap = make(map[string]string)
|
|
|
|
up = promauto.NewGauge(prometheus.GaugeOpts{
|
|
Name: "rdns_up",
|
|
Help: "Non-null value when the server is ready",
|
|
})
|
|
blSize = promauto.NewGauge(prometheus.GaugeOpts{
|
|
Name: "rdns_blocklist_size",
|
|
Help: "Blocklist size in bytes",
|
|
})
|
|
blCount = promauto.NewGauge(prometheus.GaugeOpts{
|
|
Name: "rdns_blocklist_count",
|
|
Help: "Number of items in the blocklist",
|
|
})
|
|
cacheItems = promauto.NewGauge(prometheus.GaugeOpts{
|
|
Name: "rdns_cache_items",
|
|
Help: "Number of cached queries",
|
|
})
|
|
slowAnswers = promauto.NewSummary(prometheus.SummaryOpts{
|
|
Name: "rdns_answers_slow",
|
|
Help: "Latency of upstream responses exceeding 100 milliseconds",
|
|
})
|
|
cacheHits = promauto.NewCounter(prometheus.CounterOpts{
|
|
Name: "rdns_cache_hits",
|
|
Help: "Number of responses using cache",
|
|
})
|
|
queries = promauto.NewCounter(prometheus.CounterOpts{
|
|
Name: "rdns_queries_total",
|
|
Help: "Total number of queries",
|
|
})
|
|
qtypes = promauto.NewCounterVec(prometheus.CounterOpts{
|
|
Name: "rdns_queries",
|
|
Help: "Number of queries by type",
|
|
}, []string{"type"})
|
|
responses = promauto.NewCounterVec(prometheus.CounterOpts{
|
|
Name: "rdns_responses",
|
|
Help: "Number of responses by type",
|
|
}, []string{"type", "status"})
|
|
nameservers = promauto.NewCounterVec(prometheus.CounterOpts{
|
|
Name: "rdns_nameservers",
|
|
Help: "Number of responses by nameserver",
|
|
}, []string{"ns"})
|
|
)
|
|
|
|
func getRcode(rcode int) string {
|
|
if name, ok := dns.RcodeToString[rcode]; ok {
|
|
return name
|
|
}
|
|
return "SERVFAIL"
|
|
}
|
|
|
|
func getDomain(qn string) string {
|
|
if strings.Count(qn, ".") <= 1 {
|
|
return qn
|
|
}
|
|
qs := strings.Split(qn, ".")
|
|
return strings.Join(qs[len(qs)-2:], ".")
|
|
}
|
|
|
|
func getZoneNS(qdomain string) string {
|
|
if ns, exists := zonesMap[qdomain]; exists {
|
|
return ns
|
|
}
|
|
return "."
|
|
}
|
|
|
|
func getHost(qn string) string {
|
|
if ip, exists := hostsMap[qn]; exists {
|
|
return ip
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func doDot(s string) string {
|
|
if len(s) == 0 {
|
|
return "."
|
|
}
|
|
if s[len(s)-1] == '.' {
|
|
return s
|
|
}
|
|
return s + "."
|
|
}
|
|
|
|
func clientIP(w dns.ResponseWriter) string {
|
|
addr := w.RemoteAddr()
|
|
if addr == nil {
|
|
return "unknown"
|
|
}
|
|
if host, _, err := net.SplitHostPort(addr.String()); err == nil {
|
|
return host
|
|
}
|
|
return addr.String()
|
|
}
|
|
|
|
func cacheKey(qname, qclass, qtype string) string {
|
|
return qname + "." + qclass + "." + qtype
|
|
}
|
|
|
|
func handleCache(q *dns.Msg, key string) (*dns.Msg, int, bool) {
|
|
if cached, found := qcache.Get(key); found {
|
|
msg := cached.(*dns.Msg).Copy()
|
|
msg.Id = q.Id
|
|
return msg, msg.Rcode, true
|
|
}
|
|
return nil, 0, false
|
|
}
|
|
|
|
type lookupResult struct {
|
|
msg *dns.Msg
|
|
proto string
|
|
}
|
|
|
|
func lookup(q *dns.Msg, qname string) (*dns.Msg, string, error) {
|
|
var (
|
|
r *dns.Msg
|
|
rtt time.Duration
|
|
err error
|
|
n []string
|
|
notls bool = false
|
|
proto string = "udp"
|
|
qdomain = getDomain(qname)
|
|
)
|
|
|
|
if zoneNS := getZoneNS(qdomain); zoneNS != "." {
|
|
n = []string{zoneNS}
|
|
notls = true
|
|
} else {
|
|
n = strings.FieldsFunc(*ns, func(r rune) bool { return r == ':' || r == ',' })
|
|
if len(n) > 1 {
|
|
rand.Shuffle(len(n), func(i, j int) {
|
|
n[i], n[j] = n[j], n[i]
|
|
})
|
|
}
|
|
}
|
|
|
|
// Advertise an EDNS0 buffer to avoid truncated UDP responses
|
|
if q.IsEdns0() == nil {
|
|
q.SetEdns0(dns.DefaultMsgSize, false)
|
|
}
|
|
|
|
c := &dns.Client{UDPSize: dns.MaxMsgSize}
|
|
for i := 0; i < len(n); i++ {
|
|
if *dot && !notls {
|
|
c.Net = "tcp-tls"
|
|
c.TLSConfig = &tls.Config{
|
|
ServerName: n[i],
|
|
}
|
|
proto = "tcp-tls"
|
|
r, rtt, err = c.Exchange(q, net.JoinHostPort(n[i], "853"))
|
|
if err != nil {
|
|
fmt.Fprintf(os.Stderr, "TLS connection failed to %s: %v\n", n[i], err)
|
|
}
|
|
} else {
|
|
c.Net = "udp"
|
|
proto = "udp"
|
|
r, rtt, err = c.Exchange(q, net.JoinHostPort(n[i], "53"))
|
|
if err != nil && *logs {
|
|
log.Println("WARN:", err)
|
|
}
|
|
if err == nil && r.Truncated {
|
|
// Response was truncated: retry over TCP to get the full answer
|
|
tc := &dns.Client{Net: "tcp"}
|
|
proto = "tcp"
|
|
r, rtt, err = tc.Exchange(q, net.JoinHostPort(n[i], "53"))
|
|
}
|
|
}
|
|
if err == nil {
|
|
nameservers.WithLabelValues(n[i]).Inc()
|
|
break
|
|
}
|
|
}
|
|
|
|
if rtt/time.Millisecond > 100 {
|
|
slowAnswers.Observe(float64(rtt / time.Millisecond))
|
|
}
|
|
|
|
return r, proto, err
|
|
}
|
|
|
|
func resolve(q *dns.Msg, qname, key string) (*dns.Msg, string, error) {
|
|
v, err, _ := inflight.Do(key, func() (any, error) {
|
|
msg, proto, err := lookup(q, qname)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &lookupResult{msg: msg, proto: proto}, nil
|
|
})
|
|
if err != nil {
|
|
return nil, "", err
|
|
}
|
|
res := v.(*lookupResult)
|
|
msg := res.msg.Copy()
|
|
msg.Id = q.Id
|
|
return msg, res.proto, nil
|
|
}
|
|
|
|
func handleQuery(w dns.ResponseWriter, q *dns.Msg) {
|
|
queries.Inc()
|
|
|
|
if len(q.Question) == 0 {
|
|
r := new(dns.Msg)
|
|
r.SetRcode(q, dns.RcodeFormatError)
|
|
if err := w.WriteMsg(r); err != nil && *logs {
|
|
log.Println("ERROR: Failed to write response:", err)
|
|
}
|
|
return
|
|
}
|
|
|
|
var (
|
|
r *dns.Msg
|
|
err error
|
|
rcode int
|
|
proto string
|
|
qname = strings.ToLower(q.Question[0].Name[:len(q.Question[0].Name)-1])
|
|
qclass = dns.Class(q.Question[0].Qclass).String()
|
|
qtype = dns.Type(q.Question[0].Qtype).String()
|
|
key = cacheKey(qname, qclass, qtype)
|
|
client = clientIP(w)
|
|
)
|
|
|
|
qtypes.WithLabelValues(qtype).Inc()
|
|
|
|
if qclass == "CH" && qname == "version.bind" {
|
|
r = new(dns.Msg)
|
|
rcode = dns.RcodeSuccess
|
|
r.SetRcode(q, rcode)
|
|
r.MsgHdr.RecursionAvailable = true
|
|
r.Answer = append(r.Answer, &dns.TXT{
|
|
Hdr: dns.RR_Header{Name: "version.bind.", Rrtype: dns.TypeTXT, Class: dns.ClassCHAOS, Ttl: 86400},
|
|
Txt: []string{"rdns"},
|
|
})
|
|
qcache.SetDefault(key, r)
|
|
if *logs {
|
|
log.Println(client, qname, qclass, qtype, getRcode(rcode))
|
|
}
|
|
if err := w.WriteMsg(r); err != nil && *logs {
|
|
log.Println("ERROR: Failed to write response:", err)
|
|
}
|
|
responses.WithLabelValues(qtype, getRcode(rcode)).Inc()
|
|
return
|
|
}
|
|
|
|
if r, rcode, ok := handleCache(q, key); ok {
|
|
if *logs {
|
|
log.Println(client, qname, qclass, qtype, getRcode(rcode), "cache")
|
|
}
|
|
responses.WithLabelValues(qtype, getRcode(rcode)).Inc()
|
|
cacheHits.Inc()
|
|
if err := w.WriteMsg(r); err != nil && *logs {
|
|
log.Println("ERROR: Failed to write response:", err)
|
|
}
|
|
return
|
|
}
|
|
|
|
if hostIP := getHost(qname); hostIP != "" {
|
|
ip := net.ParseIP(hostIP)
|
|
if (qtype == "A" && ip != nil && ip.To4() != nil) || (qtype == "AAAA" && ip != nil && ip.To4() == nil) {
|
|
r = new(dns.Msg)
|
|
rcode = dns.RcodeSuccess
|
|
r.SetRcode(q, rcode)
|
|
r.MsgHdr.RecursionAvailable = true
|
|
if qtype == "A" {
|
|
r.Answer = append(r.Answer, &dns.A{
|
|
Hdr: dns.RR_Header{Name: doDot(qname), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
|
A: ip,
|
|
})
|
|
} else {
|
|
r.Answer = append(r.Answer, &dns.AAAA{
|
|
Hdr: dns.RR_Header{Name: doDot(qname), Rrtype: dns.TypeAAAA, Class: dns.ClassINET, Ttl: 300},
|
|
AAAA: ip,
|
|
})
|
|
}
|
|
qcache.SetDefault(key, r)
|
|
if *logs {
|
|
log.Println(client, qname, qclass, qtype, getRcode(rcode))
|
|
}
|
|
if err := w.WriteMsg(r); err != nil && *logs {
|
|
log.Println("ERROR: Failed to write response:", err)
|
|
}
|
|
responses.WithLabelValues(qtype, getRcode(rcode)).Inc()
|
|
return
|
|
}
|
|
}
|
|
|
|
if blocklistMap[qname] {
|
|
r = new(dns.Msg)
|
|
rcode = dns.RcodeRefused
|
|
r.SetRcode(q, rcode)
|
|
r.MsgHdr.RecursionAvailable = true
|
|
qcache.SetDefault(key, r)
|
|
if *logs {
|
|
log.Println(client, qname, qclass, qtype, getRcode(rcode))
|
|
}
|
|
if err := w.WriteMsg(r); err != nil && *logs {
|
|
log.Println("ERROR: Failed to write response:", err)
|
|
}
|
|
responses.WithLabelValues(qtype, getRcode(rcode)).Inc()
|
|
return
|
|
}
|
|
|
|
r, proto, err = resolve(q, qname, key)
|
|
if err != nil {
|
|
if *logs {
|
|
log.Println("ERROR:", err)
|
|
}
|
|
r = new(dns.Msg)
|
|
rcode = dns.RcodeServerFailure
|
|
r.SetRcode(q, rcode)
|
|
if *logs {
|
|
log.Println(client, qname, qclass, qtype, getRcode(rcode))
|
|
}
|
|
} else {
|
|
rcode = r.Rcode
|
|
r.MsgHdr.RecursionAvailable = true
|
|
qcache.SetDefault(key, r)
|
|
if *logs {
|
|
log.Println(client, qname, qclass, qtype, getRcode(rcode), proto)
|
|
}
|
|
}
|
|
if err := w.WriteMsg(r); err != nil && *logs {
|
|
log.Println("ERROR: Failed to write response:", err)
|
|
}
|
|
responses.WithLabelValues(qtype, getRcode(rcode)).Inc()
|
|
}
|
|
|
|
func serve(addr string, port string, network string, started func()) *dns.Server {
|
|
server := &dns.Server{
|
|
Addr: net.JoinHostPort(addr, port),
|
|
Net: network,
|
|
ReusePort: true,
|
|
UDPSize: dns.DefaultMsgSize,
|
|
NotifyStartedFunc: started,
|
|
}
|
|
go func() {
|
|
if err := server.ListenAndServe(); err != nil {
|
|
fmt.Fprintf(os.Stderr, "Failed to setup the %s server on %s: %v\n", network, server.Addr, err)
|
|
os.Exit(1)
|
|
}
|
|
}()
|
|
return server
|
|
}
|
|
|
|
func loadBlocklist(path string) {
|
|
data, err := os.ReadFile(path)
|
|
if err != nil {
|
|
log.Println("Error reading blocklist:", err)
|
|
return
|
|
}
|
|
|
|
lines := strings.Split(string(data), "\n")
|
|
count := 0
|
|
for _, line := range lines {
|
|
line = strings.ToLower(strings.TrimSpace(line))
|
|
if line != "" && !strings.HasPrefix(line, "#") {
|
|
blocklistMap[line] = true
|
|
count++
|
|
}
|
|
}
|
|
blCount.Set(float64(count))
|
|
blSize.Set(float64(len(data)))
|
|
}
|
|
|
|
func loadZones(path string) int {
|
|
data, err := os.ReadFile(path)
|
|
if err != nil {
|
|
log.Println("Error reading zones file:", err)
|
|
return 0
|
|
}
|
|
|
|
lines := strings.Split(string(data), "\n")
|
|
count := 0
|
|
|
|
for _, line := range lines {
|
|
if index := strings.Index(line, "#"); index != -1 {
|
|
line = line[:index]
|
|
}
|
|
line = strings.TrimSpace(line)
|
|
if line == "" {
|
|
continue
|
|
}
|
|
|
|
parts := strings.Fields(line)
|
|
if len(parts) >= 2 {
|
|
zonesMap[strings.ToLower(parts[0])] = parts[1]
|
|
count++
|
|
}
|
|
}
|
|
return count
|
|
}
|
|
|
|
func loadHosts(path string) int {
|
|
data, err := os.ReadFile(path)
|
|
if err != nil {
|
|
log.Println("Error reading hosts file:", err)
|
|
return 0
|
|
}
|
|
|
|
lines := strings.Split(string(data), "\n")
|
|
count := 0
|
|
|
|
for _, line := range lines {
|
|
if index := strings.Index(line, "#"); index != -1 {
|
|
line = line[:index]
|
|
}
|
|
line = strings.TrimSpace(line)
|
|
if line == "" {
|
|
continue
|
|
}
|
|
|
|
parts := strings.Fields(line)
|
|
if len(parts) >= 2 && net.ParseIP(parts[0]) != nil {
|
|
for i := 1; i < len(parts); i++ {
|
|
hostsMap[strings.ToLower(parts[i])] = parts[0]
|
|
count++
|
|
}
|
|
}
|
|
}
|
|
return count
|
|
}
|
|
|
|
func cacheEviction(key string, value any) {
|
|
if *logs {
|
|
log.Println("Cache: Evicted", key)
|
|
}
|
|
cacheItems.Set(float64(qcache.ItemCount()))
|
|
}
|
|
|
|
func main() {
|
|
flag.Parse()
|
|
|
|
qcache = cache.New(time.Duration(*ttl)*time.Minute, 5*time.Minute)
|
|
qcache.OnEvicted(cacheEviction)
|
|
|
|
fmt.Println("Starting Proxy Resolver:", net.JoinHostPort(*addr, *port), "->", *ns, "[UDP/TCP]")
|
|
fmt.Println("Cache TTL:", *ttl, "minutes\nTLS enabled:", *dot)
|
|
|
|
if *cpu != 0 {
|
|
runtime.GOMAXPROCS(*cpu)
|
|
}
|
|
if !*logs {
|
|
log.SetOutput(io.Discard)
|
|
}
|
|
if *blocklistPath != "" {
|
|
loadBlocklist(*blocklistPath)
|
|
fmt.Println("Blocklist loaded:", len(blocklistMap), "entries")
|
|
}
|
|
if *zonesPath != "" {
|
|
count := loadZones(*zonesPath)
|
|
if count == 0 {
|
|
fmt.Println("No zones found")
|
|
} else {
|
|
fmt.Println("Zones loaded:", count, "entries")
|
|
}
|
|
}
|
|
if *hostsPath != "" {
|
|
count := loadHosts(*hostsPath)
|
|
if count == 0 {
|
|
fmt.Println("No hosts found")
|
|
} else {
|
|
fmt.Println("Hosts loaded:", count, "entries")
|
|
}
|
|
}
|
|
|
|
var metricsServer *http.Server
|
|
if *metrics {
|
|
fmt.Println("Starting prometheus exporter on port 9153")
|
|
mux := http.NewServeMux()
|
|
mux.Handle("/metrics", promhttp.Handler())
|
|
metricsServer = &http.Server{
|
|
Addr: ":9153",
|
|
Handler: mux,
|
|
ReadHeaderTimeout: 5 * time.Second,
|
|
}
|
|
go func() {
|
|
if err := metricsServer.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
|
fmt.Fprintf(os.Stderr, "Metrics server error: %v\n", err)
|
|
}
|
|
}()
|
|
|
|
go func() {
|
|
ticker := time.NewTicker(10 * time.Second)
|
|
defer ticker.Stop()
|
|
for range ticker.C {
|
|
cacheItems.Set(float64(qcache.ItemCount()))
|
|
}
|
|
}()
|
|
}
|
|
|
|
dns.HandleFunc(".", handleQuery)
|
|
|
|
var readyOnce sync.Once
|
|
ready := func() { readyOnce.Do(func() { up.Set(1) }) }
|
|
servers := []*dns.Server{
|
|
serve(*addr, *port, "tcp", ready),
|
|
serve(*addr, *port, "udp", ready),
|
|
}
|
|
|
|
sig := make(chan os.Signal, 1)
|
|
signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM)
|
|
<-sig
|
|
fmt.Printf("\033[2K\rTime to sleep, goodbye\n")
|
|
|
|
for _, s := range servers {
|
|
if err := s.Shutdown(); err != nil {
|
|
fmt.Fprintf(os.Stderr, "Error shutting down server: %v\n", err)
|
|
}
|
|
}
|
|
if metricsServer != nil {
|
|
if err := metricsServer.Close(); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
|
fmt.Fprintf(os.Stderr, "Error shutting down metrics server: %v\n", err)
|
|
}
|
|
}
|
|
}
|