mirror of
https://github.com/psviderski/uncloud.git
synced 2026-08-26 19:13:34 +00:00
chore(dns-server): authoritative responses for internal DNS queries and forwarding to ustream servers
This commit is contained in:
@@ -2,14 +2,338 @@ package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math/rand/v2"
|
||||
"net"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
)
|
||||
|
||||
// Server provides a DNS server for service discovery and proxying to upstream DNS servers.
|
||||
type Server struct {
|
||||
store RecordsStore
|
||||
const (
|
||||
// InternalDomain is the cluster internal domain for service discovery. All DNS queries ending with this suffix
|
||||
// will be resolved using the internal DNS server.
|
||||
InternalDomain = "internal."
|
||||
// dnsPort is the standard DNS port.
|
||||
dnsPort = 53
|
||||
// maxConcurrentForwards is the maximum number of concurrent forwarded queries to upstream DNS servers.
|
||||
// 1024 is the default used by the Docker internal DNS server.
|
||||
maxConcurrentForwards = 1024
|
||||
// forwardingTimeout is the timeout for forwarding a DNS query to an upstream server.
|
||||
forwardingTimeout = 3 * time.Second
|
||||
)
|
||||
|
||||
// Resolver is an interface for resolving service names to IP addresses.
|
||||
type Resolver interface {
|
||||
// Resolve returns a list of IP addresses of service containers. An empty list is returned if no service is found.
|
||||
Resolve(serviceName string) []netip.Addr
|
||||
}
|
||||
|
||||
type RecordsStore interface {
|
||||
Resolve(ctx context.Context, name string) ([]netip.Addr, error)
|
||||
// Server is an embedded internal DNS server for service discovery and forwarding external queries
|
||||
// to upstream DNS servers.
|
||||
type Server struct {
|
||||
listenAddr netip.Addr
|
||||
resolver Resolver
|
||||
upstreamServers []netip.AddrPort
|
||||
|
||||
udpServer *dns.Server
|
||||
tcpServer *dns.Server
|
||||
inProgressReqs sync.WaitGroup
|
||||
forwardSemaphore chan struct{}
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
// NewServer creates a new DNS server with the given configuration.
|
||||
// If upstreams is nil, nameservers from /etc/resolv.conf will be used. An empty upstreams list means to only resolve
|
||||
// internal DNS queries and not forward any external queries.
|
||||
func NewServer(listenAddr netip.Addr, resolver Resolver, upstreams []netip.AddrPort) (*Server, error) {
|
||||
if listenAddr.IsValid() {
|
||||
return nil, fmt.Errorf("invalid listen address: %s", listenAddr)
|
||||
}
|
||||
if resolver == nil {
|
||||
return nil, fmt.Errorf("resolver must be provided")
|
||||
}
|
||||
|
||||
// Load upstreams from /etc/resolv.conf and set fallback servers only if upstreams is not provided.
|
||||
if upstreams == nil {
|
||||
if nameservers, err := parseNameserversFromResolvConf(); err != nil {
|
||||
slog.Warn("Failed to parse nameservers from /etc/resolv.conf.", "err", err)
|
||||
} else {
|
||||
for _, ns := range nameservers {
|
||||
if ns.Compare(listenAddr) == 0 {
|
||||
// Skip this internal DNS server address if it has been configured in /etc/resolv.conf
|
||||
// to not create a forwarding loop.
|
||||
continue
|
||||
}
|
||||
upstreams = append(upstreams, netip.AddrPortFrom(ns, dnsPort))
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to common public DNS servers if no nameservers were found in /etc/resolv.conf.
|
||||
if upstreams == nil {
|
||||
upstreams = []netip.AddrPort{
|
||||
netip.AddrPortFrom(netip.MustParseAddr("1.1.1.1"), dnsPort), // Cloudflare DNS
|
||||
netip.AddrPortFrom(netip.MustParseAddr("8.8.8.8"), dnsPort), // Google DNS
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return &Server{
|
||||
listenAddr: listenAddr,
|
||||
resolver: resolver,
|
||||
upstreamServers: upstreams,
|
||||
forwardSemaphore: make(chan struct{}, maxConcurrentForwards),
|
||||
log: slog.With("component", "dns-server"),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Run starts the DNS server listening on both UDP and TCP ports. It will stop the server when the context is canceled.
|
||||
func (s *Server) Run(ctx context.Context) error {
|
||||
addr := net.JoinHostPort(s.listenAddr.String(), strconv.Itoa(dnsPort))
|
||||
s.udpServer = &dns.Server{
|
||||
Addr: addr,
|
||||
Net: "udp",
|
||||
Handler: dns.HandlerFunc(s.handleRequest),
|
||||
}
|
||||
s.tcpServer = &dns.Server{
|
||||
Addr: addr,
|
||||
Net: "tcp",
|
||||
Handler: dns.HandlerFunc(s.handleRequest),
|
||||
}
|
||||
|
||||
errCh := make(chan error, 2) // Buffer size 2 for UDP and TCP errors.
|
||||
|
||||
go func() {
|
||||
s.log.Info("Starting DNS server (UDP).", "addr", addr, "proto", "udp")
|
||||
if err := s.udpServer.ListenAndServe(); err != nil {
|
||||
errCh <- fmt.Errorf("listen and serve on %s/udp: %w", addr, err)
|
||||
}
|
||||
}()
|
||||
|
||||
go func() {
|
||||
s.log.Info("Starting DNS server (TCP).", "addr", addr, "proto", "tcp")
|
||||
if err := s.tcpServer.ListenAndServe(); err != nil {
|
||||
errCh <- fmt.Errorf("listen and serve on %s/udp: %w", addr, err)
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-errCh:
|
||||
// Stop the servers if one of them fails.
|
||||
s.stop()
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
s.log.Info("Stopping DNS server.")
|
||||
return s.stop()
|
||||
}
|
||||
}
|
||||
|
||||
// stop gracefully shuts down the DNS server.
|
||||
func (s *Server) stop() error {
|
||||
var udpErr, tcpErr error
|
||||
|
||||
if s.udpServer != nil {
|
||||
udpErr = s.udpServer.Shutdown()
|
||||
}
|
||||
if s.tcpServer != nil {
|
||||
tcpErr = s.tcpServer.Shutdown()
|
||||
}
|
||||
|
||||
// Wait for all in-progress requests to finish.
|
||||
s.inProgressReqs.Wait()
|
||||
|
||||
if udpErr != nil {
|
||||
return fmt.Errorf("shutdown DNS server (UDP): %w", udpErr)
|
||||
}
|
||||
|
||||
if tcpErr != nil {
|
||||
return fmt.Errorf("shutdown DNS server (TCP): %w", tcpErr)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// handleRequest processes a DNS query and returns an appropriate response.
|
||||
func (s *Server) handleRequest(w dns.ResponseWriter, req *dns.Msg) {
|
||||
s.inProgressReqs.Add(1)
|
||||
defer s.inProgressReqs.Done()
|
||||
|
||||
if len(req.Question) == 0 {
|
||||
resp := new(dns.Msg).SetRcode(req, dns.RcodeFormatError)
|
||||
s.reply(w, req, resp)
|
||||
return
|
||||
}
|
||||
|
||||
// While the original DNS RFCs allow multiple questions, in practice it never works. So handle only the first one.
|
||||
q := req.Question[0]
|
||||
log := s.log.With("name", q.Name, "type", dns.TypeToString[q.Qtype])
|
||||
log.Debug("Received DNS query.")
|
||||
|
||||
if !dns.IsSubDomain(InternalDomain, dns.CanonicalName(q.Name)) {
|
||||
log.Debug("Forwarding non-internal DNS query to upstream DNS servers.")
|
||||
|
||||
// Use the same transport for the forwarded request as the original request.
|
||||
resp, err := s.forwardRequest(req, w.LocalAddr().Network())
|
||||
if err != nil {
|
||||
log.Error("Failed to forward DNS query.", "err", err)
|
||||
resp = new(dns.Msg).SetRcode(req, dns.RcodeServerFailure)
|
||||
}
|
||||
|
||||
s.reply(w, req, resp)
|
||||
return
|
||||
}
|
||||
|
||||
// Handle the query for the internal domain.
|
||||
resp := new(dns.Msg).SetReply(req)
|
||||
resp.Authoritative = true
|
||||
resp.RecursionAvailable = true
|
||||
|
||||
switch q.Qtype {
|
||||
case dns.TypeA:
|
||||
records := s.handleAQuery(q.Name)
|
||||
if len(records) > 0 {
|
||||
log.Debug("Found A records for internal DNS query.", "count", len(records))
|
||||
resp.Answer = append(resp.Answer, records...)
|
||||
} else {
|
||||
log.Debug("No records found for internal DNS query.")
|
||||
resp.SetRcode(req, dns.RcodeNameError)
|
||||
}
|
||||
// TODO: Handle other query types (SRV, TXT, etc.) as needed.
|
||||
}
|
||||
|
||||
// Truncate the response if it exceeds the maximum size for the transport protocol.
|
||||
maxSize := dns.MinMsgSize
|
||||
if w.LocalAddr().Network() == "tcp" {
|
||||
maxSize = dns.MaxMsgSize
|
||||
} else {
|
||||
// Retrieve the UDP buffer size from the EDNS0 record if present.
|
||||
if opt := req.IsEdns0(); opt != nil {
|
||||
if udpSize := int(opt.UDPSize()); udpSize > maxSize {
|
||||
maxSize = udpSize
|
||||
}
|
||||
}
|
||||
}
|
||||
resp.Truncate(maxSize)
|
||||
|
||||
s.reply(w, req, resp)
|
||||
}
|
||||
|
||||
func (s *Server) reply(w dns.ResponseWriter, req *dns.Msg, resp *dns.Msg) error {
|
||||
if err := w.WriteMsg(resp); err != nil {
|
||||
s.log.Error("Failed to write DNS response.", "err", err, "msg", resp)
|
||||
// It may fail due to a malformed response message, e.g. exceeding the maximum size. In that case,
|
||||
// attempt to send a server failure response instead so the client doesn't have to wait for a timeout.
|
||||
if resp.Rcode != dns.RcodeServerFailure {
|
||||
resp = new(dns.Msg).SetRcode(req, dns.RcodeServerFailure)
|
||||
if err2 := w.WriteMsg(resp); err2 != nil {
|
||||
s.log.Error("Failed to write DNS error response.", "err", err2, "msg", resp)
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// forwardRequest forwards a DNS query to system DNS servers
|
||||
func (s *Server) forwardRequest(req *dns.Msg, proto string) (*dns.Msg, error) {
|
||||
if len(s.upstreamServers) == 0 {
|
||||
return nil, errors.New("no upstream DNS servers configured")
|
||||
}
|
||||
|
||||
// Apply concurrency control for forwarded queries.
|
||||
select {
|
||||
case s.forwardSemaphore <- struct{}{}:
|
||||
defer func() {
|
||||
<-s.forwardSemaphore
|
||||
}()
|
||||
default:
|
||||
return nil, fmt.Errorf("too many concurrent forwarded queries (max: %d)", maxConcurrentForwards)
|
||||
}
|
||||
|
||||
// Create DNS client with timeout and same transport as original request
|
||||
client := &dns.Client{
|
||||
Net: proto,
|
||||
Timeout: forwardingTimeout,
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
for _, server := range s.upstreamServers {
|
||||
resp, _, err := client.Exchange(req, server.String())
|
||||
if err == nil {
|
||||
return resp, nil
|
||||
}
|
||||
lastErr = err
|
||||
s.log.Debug("Failed to forward DNS query to upstream server.", "server", server, "err", err)
|
||||
}
|
||||
|
||||
return nil, lastErr
|
||||
}
|
||||
|
||||
// handleAQuery processes an A query for the internal domain and returns A records for the requested name.
|
||||
// The internal domain suffix is already stripped from the name. An empty list is returned if no records are found.
|
||||
func (s *Server) handleAQuery(name string) []dns.RR {
|
||||
serviceName := trimInternalDomain(name)
|
||||
ips := s.resolver.Resolve(serviceName)
|
||||
if len(ips) == 0 {
|
||||
s.log.Debug("Failed to resolve service name.", "service", serviceName)
|
||||
return nil
|
||||
}
|
||||
// TODO: verify the formatting of the IP addresses.
|
||||
s.log.Debug("Resolved service name.", "service", serviceName, "ips", ips)
|
||||
|
||||
if len(ips) > 1 {
|
||||
// TODO: sort by proximity to the requesting container/machine. For now, just shuffle the IPs.
|
||||
rand.Shuffle(len(ips), func(i, j int) {
|
||||
ips[i], ips[j] = ips[j], ips[i]
|
||||
})
|
||||
}
|
||||
|
||||
// Create A records for each IP.
|
||||
records := make([]dns.RR, 0, len(ips))
|
||||
for _, ip := range ips {
|
||||
records = append(records, &dns.A{
|
||||
Hdr: dns.RR_Header{
|
||||
Name: name,
|
||||
Rrtype: dns.TypeA,
|
||||
Class: dns.ClassINET,
|
||||
// TODO: should we increate the TTL to some reasonably small value like 5-30 seconds to allow
|
||||
// at least some caching?
|
||||
Ttl: 0,
|
||||
},
|
||||
A: net.ParseIP(ip.String()),
|
||||
})
|
||||
}
|
||||
return records
|
||||
}
|
||||
|
||||
// parseNameserversFromResolvConf parses the nameservers from /etc/resolv.conf.
|
||||
func parseNameserversFromResolvConf() ([]netip.Addr, error) {
|
||||
config, err := dns.ClientConfigFromFile("/etc/resolv.conf")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var servers []netip.Addr
|
||||
for _, server := range config.Servers {
|
||||
if addr, err := netip.ParseAddr(server); err == nil {
|
||||
servers = append(servers, addr)
|
||||
}
|
||||
}
|
||||
|
||||
return servers, nil
|
||||
}
|
||||
|
||||
func trimInternalDomain(name string) string {
|
||||
name = dns.CanonicalName(name)
|
||||
if !dns.IsSubDomain(InternalDomain, name) {
|
||||
return name
|
||||
}
|
||||
|
||||
return strings.TrimSuffix(name, "."+InternalDomain)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user