Files
rdns-go/rdns.go
T

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)
}
}
}