Fix nameserver listener restarts

This commit is contained in:
Daniel
2022-02-14 13:50:54 +01:00
parent 293af89c41
commit 11893a780a

View File

@@ -6,6 +6,7 @@ import (
"net" "net"
"os" "os"
"strconv" "strconv"
"sync"
"github.com/miekg/dns" "github.com/miekg/dns"
@@ -17,8 +18,12 @@ import (
) )
var ( var (
module *modules.Module module *modules.Module
stopListener func() error
stopListeners bool
stopListener1 func() error
stopListener2 func() error
stopListenersLock sync.Mutex
) )
func init() { func init() {
@@ -42,11 +47,13 @@ func start() error {
return err return err
} }
// Get listen addresses.
ip1, ip2, port, err := getListenAddresses(nameserverAddressConfig()) ip1, ip2, port, err := getListenAddresses(nameserverAddressConfig())
if err != nil { if err != nil {
return fmt.Errorf("failed to parse nameserver listen address: %w", err) return fmt.Errorf("failed to parse nameserver listen address: %w", err)
} }
// Get own hostname.
hostname, err = os.Hostname() hostname, err = os.Hostname()
if err != nil { if err != nil {
log.Warningf("nameserver: failed to get hostname: %s", err) log.Warningf("nameserver: failed to get hostname: %s", err)
@@ -56,8 +63,7 @@ func start() error {
// Start listener(s). // Start listener(s).
if ip2 == nil { if ip2 == nil {
// Start a single listener. // Start a single listener.
dnsServer := startListener(ip1, port) startListener(ip1, port, true)
stopListener = dnsServer.Shutdown
// Set nameserver matcher in firewall to fast-track dns queries. // Set nameserver matcher in firewall to fast-track dns queries.
if ip1.Equal(net.IPv4zero) || ip1.Equal(net.IPv6zero) { if ip1.Equal(net.IPv4zero) || ip1.Equal(net.IPv6zero) {
@@ -73,22 +79,11 @@ func start() error {
return firewall.SetNameserverIPMatcher(func(ip net.IP) bool { return firewall.SetNameserverIPMatcher(func(ip net.IP) bool {
return ip.Equal(ip1) return ip.Equal(ip1)
}) })
} }
// Dual listener. // Dual listener.
dnsServer1 := startListener(ip1, port) startListener(ip1, port, true)
dnsServer2 := startListener(ip2, port) startListener(ip2, port, false)
stopListener = func() error {
// Shutdown both listeners.
err1 := dnsServer1.Shutdown()
err2 := dnsServer2.Shutdown()
// Return first error.
if err1 != nil {
return err1
}
return err2
}
// Fast track dns queries destined for one of the listener IPs. // Fast track dns queries destined for one of the listener IPs.
return firewall.SetNameserverIPMatcher(func(ip net.IP) bool { return firewall.SetNameserverIPMatcher(func(ip net.IP) bool {
@@ -96,20 +91,46 @@ func start() error {
}) })
} }
func startListener(ip net.IP, port uint16) *dns.Server { func startListener(ip net.IP, port uint16, first bool) {
// Create DNS server.
dnsServer := &dns.Server{
Addr: net.JoinHostPort(
ip.String(),
strconv.Itoa(int(port)),
),
Net: "udp",
}
dns.HandleFunc(".", handleRequestAsWorker)
// Start DNS server as service worker. // Start DNS server as service worker.
log.Infof("nameserver: starting to listen on %s", dnsServer.Addr)
module.StartServiceWorker("dns resolver", 0, func(ctx context.Context) error { module.StartServiceWorker("dns resolver", 0, func(ctx context.Context) error {
// Create DNS server.
dnsServer := &dns.Server{
Addr: net.JoinHostPort(
ip.String(),
strconv.Itoa(int(port)),
),
Net: "udp",
Handler: dns.HandlerFunc(handleRequestAsWorker),
}
// Register stop function.
func() {
stopListenersLock.Lock()
defer stopListenersLock.Unlock()
// Check if we should stop
if stopListeners {
_ = dnsServer.Shutdown()
dnsServer = nil
return
}
// Register stop function.
if first {
stopListener1 = dnsServer.Shutdown
} else {
stopListener2 = dnsServer.Shutdown
}
}()
// Check if we should stop.
if dnsServer == nil {
return nil
}
// Start listening.
log.Infof("nameserver: starting to listen on %s", dnsServer.Addr)
err := dnsServer.ListenAndServe() err := dnsServer.ListenAndServe()
if err != nil { if err != nil {
// check if we are shutting down // check if we are shutting down
@@ -124,16 +145,25 @@ func startListener(ip net.IP, port uint16) *dns.Server {
} }
return err return err
}) })
return dnsServer
} }
func stop() error { func stop() error {
if stopListener != nil { stopListenersLock.Lock()
if err := stopListener(); err != nil { defer stopListenersLock.Unlock()
log.Warningf("nameserver: failed to stop: %s", err)
// Stop listeners.
stopListeners = true
if stopListener1 != nil {
if err := stopListener1(); err != nil {
log.Warningf("nameserver: failed to stop listener1: %s", err)
} }
} }
if stopListener2 != nil {
if err := stopListener2(); err != nil {
log.Warningf("nameserver: failed to stop listener2: %s", err)
}
}
return nil return nil
} }