316 lines
7.9 KiB
Go
316 lines
7.9 KiB
Go
package dnsengine
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"log/slog"
|
|
"net/netip"
|
|
"os"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/miekg/dns"
|
|
|
|
"github.com/owen/vibedns/internal/cache"
|
|
"github.com/owen/vibedns/internal/metrics"
|
|
"github.com/owen/vibedns/internal/models"
|
|
"github.com/owen/vibedns/internal/querylog"
|
|
"github.com/owen/vibedns/internal/ratelimit"
|
|
"github.com/owen/vibedns/internal/resolver"
|
|
"github.com/owen/vibedns/internal/runtimecfg"
|
|
)
|
|
|
|
// Server runs the UDP and TCP DNS listeners.
|
|
type Server struct {
|
|
runtime *runtimecfg.Manager
|
|
cache *cache.Cache
|
|
resolver *resolver.Resolver
|
|
limiter *ratelimit.Limiter
|
|
metrics *metrics.Metrics
|
|
qlog *querylog.Logger
|
|
log *slog.Logger
|
|
hostname string
|
|
|
|
ctx context.Context
|
|
cancel context.CancelFunc
|
|
|
|
mu sync.Mutex
|
|
udp *dns.Server
|
|
tcp *dns.Server
|
|
running bool
|
|
udpAddr string
|
|
tcpAddr string
|
|
startErr chan error
|
|
wg sync.WaitGroup
|
|
}
|
|
|
|
// Options bundles the dependencies the DNS server needs.
|
|
type Options struct {
|
|
Runtime *runtimecfg.Manager
|
|
Cache *cache.Cache
|
|
Resolver *resolver.Resolver
|
|
Limiter *ratelimit.Limiter
|
|
Metrics *metrics.Metrics
|
|
QueryLog *querylog.Logger
|
|
Log *slog.Logger
|
|
}
|
|
|
|
// New creates a DNS server. Call Start to bind the listeners.
|
|
func New(opts Options) *Server {
|
|
host, err := os.Hostname()
|
|
if err != nil || host == "" {
|
|
host = "vibedns"
|
|
}
|
|
s := &Server{
|
|
runtime: opts.Runtime,
|
|
cache: opts.Cache,
|
|
resolver: opts.Resolver,
|
|
limiter: opts.Limiter,
|
|
metrics: opts.Metrics,
|
|
qlog: opts.QueryLog,
|
|
log: opts.Log,
|
|
hostname: host,
|
|
startErr: make(chan error, 2),
|
|
}
|
|
// Refreshing an ageing cache entry keeps popular names warm without the
|
|
// client ever waiting on the upstream.
|
|
s.cache.SetPrefetcher(s.prefetch)
|
|
return s
|
|
}
|
|
|
|
// Start binds both listeners. It returns once they are accepting queries, or
|
|
// with an error explaining which address could not be bound.
|
|
func (s *Server) Start(ctx context.Context) error {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
if s.running {
|
|
return errors.New("DNS server is already running")
|
|
}
|
|
|
|
s.ctx, s.cancel = context.WithCancel(ctx)
|
|
snap := s.runtime.Current()
|
|
s.udpAddr = snap.Settings.DNS.UDPListen
|
|
s.tcpAddr = snap.Settings.DNS.TCPListen
|
|
|
|
udpSize := snap.Settings.DNS.EDNSUDPSize
|
|
if udpSize < dns.MinMsgSize {
|
|
udpSize = dns.MinMsgSize
|
|
}
|
|
idle := time.Duration(snap.Settings.DNS.TCPIdleSeconds) * time.Second
|
|
|
|
ready := make(chan struct{}, 2)
|
|
s.udp = &dns.Server{
|
|
Addr: s.udpAddr,
|
|
Net: "udp",
|
|
Handler: s,
|
|
UDPSize: udpSize,
|
|
NotifyStartedFunc: func() { ready <- struct{}{} },
|
|
}
|
|
s.tcp = &dns.Server{
|
|
Addr: s.tcpAddr,
|
|
Net: "tcp",
|
|
Handler: s,
|
|
IdleTimeout: func() time.Duration { return idle },
|
|
ReadTimeout: idle + 2*time.Second,
|
|
NotifyStartedFunc: func() { ready <- struct{}{} },
|
|
}
|
|
|
|
errCh := make(chan error, 2)
|
|
s.wg.Add(2)
|
|
go func() {
|
|
defer s.wg.Done()
|
|
if err := s.udp.ListenAndServe(); err != nil {
|
|
errCh <- fmt.Errorf("DNS UDP listener on %s: %w", s.udpAddr, describeBindError(err, s.udpAddr))
|
|
}
|
|
}()
|
|
go func() {
|
|
defer s.wg.Done()
|
|
if err := s.tcp.ListenAndServe(); err != nil {
|
|
errCh <- fmt.Errorf("DNS TCP listener on %s: %w", s.tcpAddr, describeBindError(err, s.tcpAddr))
|
|
}
|
|
}()
|
|
|
|
// Wait for both listeners to report ready, or for one to fail.
|
|
started := 0
|
|
deadline := time.After(10 * time.Second)
|
|
for started < 2 {
|
|
select {
|
|
case <-ready:
|
|
started++
|
|
case err := <-errCh:
|
|
s.cancel()
|
|
return err
|
|
case <-deadline:
|
|
s.cancel()
|
|
return fmt.Errorf("DNS listeners did not become ready within 10 seconds")
|
|
}
|
|
}
|
|
|
|
s.running = true
|
|
s.log.Info("DNS listeners started", "udp", s.udpAddr, "tcp", s.tcpAddr)
|
|
|
|
// Surface a listener that dies later.
|
|
go func() {
|
|
select {
|
|
case err := <-errCh:
|
|
s.log.Error("DNS listener stopped unexpectedly", "error", err)
|
|
case <-s.ctx.Done():
|
|
}
|
|
}()
|
|
return nil
|
|
}
|
|
|
|
// describeBindError turns a raw bind failure into something actionable.
|
|
func describeBindError(err error, addr string) error {
|
|
msg := err.Error()
|
|
switch {
|
|
case strings.Contains(msg, "permission denied"):
|
|
return fmt.Errorf("%w (binding a port below 1024 needs root, or grant the "+
|
|
"binary CAP_NET_BIND_SERVICE with: setcap 'cap_net_bind_service=+ep' ./vibedns)", err)
|
|
case strings.Contains(msg, "address already in use"):
|
|
return fmt.Errorf("%w (another DNS server is already listening on %s; on many "+
|
|
"systems that is systemd-resolved, which can be disabled with: "+
|
|
"systemctl disable --now systemd-resolved)", err, addr)
|
|
default:
|
|
return err
|
|
}
|
|
}
|
|
|
|
// Shutdown stops both listeners.
|
|
func (s *Server) Shutdown(ctx context.Context) error {
|
|
s.mu.Lock()
|
|
udp, tcp, running := s.udp, s.tcp, s.running
|
|
s.running = false
|
|
s.mu.Unlock()
|
|
|
|
if !running {
|
|
return nil
|
|
}
|
|
if s.cancel != nil {
|
|
s.cancel()
|
|
}
|
|
|
|
var firstErr error
|
|
if udp != nil {
|
|
if err := udp.ShutdownContext(ctx); err != nil && firstErr == nil {
|
|
firstErr = err
|
|
}
|
|
}
|
|
if tcp != nil {
|
|
if err := tcp.ShutdownContext(ctx); err != nil && firstErr == nil {
|
|
firstErr = err
|
|
}
|
|
}
|
|
s.wg.Wait()
|
|
s.log.Info("DNS listeners stopped")
|
|
return firstErr
|
|
}
|
|
|
|
// Running reports whether the listeners are up.
|
|
func (s *Server) Running() bool {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
return s.running
|
|
}
|
|
|
|
// ListenAddrs returns the bound addresses for the status page.
|
|
func (s *Server) ListenAddrs() (udp, tcp string) {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
return s.udpAddr, s.tcpAddr
|
|
}
|
|
|
|
// prefetch refreshes a cache entry in the background.
|
|
func (s *Server) prefetch(k cache.Key) {
|
|
snap := s.runtime.Current()
|
|
if !snap.Settings.DNS.Recursion {
|
|
return
|
|
}
|
|
ctx, cancel := context.WithTimeout(s.ctx, s.forwardBudget(snap))
|
|
defer cancel()
|
|
|
|
req := new(dns.Msg)
|
|
req.SetQuestion(k.Name, k.Type)
|
|
req.Question[0].Qclass = k.Class
|
|
req.RecursionDesired = true
|
|
if k.DO {
|
|
req.SetEdns0(uint16(snap.Settings.DNS.EDNSUDPSize), true)
|
|
}
|
|
|
|
res, err := s.resolver.Resolve(ctx, req)
|
|
if err != nil {
|
|
s.log.Debug("cache prefetch failed", "name", k.Name, "error", err)
|
|
return
|
|
}
|
|
s.cache.Put(k, res.Msg)
|
|
}
|
|
|
|
// Resolve performs a query through the full pipeline on behalf of the UI's
|
|
// "test a lookup" tool, without going over the network.
|
|
func (s *Server) Resolve(ctx context.Context, name string, qtype uint16, client netip.Addr, do bool) (*dns.Msg, string, error) {
|
|
snap := s.runtime.Current()
|
|
req := new(dns.Msg)
|
|
req.SetQuestion(dns.Fqdn(name), qtype)
|
|
req.RecursionDesired = true
|
|
if do {
|
|
req.SetEdns0(uint16(snap.Settings.DNS.EDNSUDPSize), true)
|
|
}
|
|
resp, out := s.respond(req, client, snap)
|
|
if resp == nil {
|
|
return nil, "", errors.New("no response was produced")
|
|
}
|
|
return resp, out.source, nil
|
|
}
|
|
|
|
// logQueryOutcome writes one query log entry.
|
|
func (s *Server) logQueryOutcome(req *dns.Msg, client netip.Addr, protocol string,
|
|
out outcome, snap *runtimecfg.Snapshot, elapsed time.Duration) {
|
|
|
|
if s.qlog == nil || !s.qlog.Enabled() {
|
|
return
|
|
}
|
|
q := req.Question[0]
|
|
rcode := "DROPPED"
|
|
if out.rcode >= 0 {
|
|
rcode = dns.RcodeToString[out.rcode]
|
|
}
|
|
|
|
e := models.QueryLogEntry{
|
|
Timestamp: time.Now(),
|
|
ClientIP: client.String(),
|
|
QName: strings.ToLower(dns.Fqdn(q.Name)),
|
|
QType: qtypeName(req),
|
|
Rcode: rcode,
|
|
Source: out.source,
|
|
CacheHit: out.cacheHit,
|
|
Blocked: out.blocked,
|
|
Protocol: protocol,
|
|
DurationUS: elapsed.Microseconds(),
|
|
AnswerCount: out.answerCount,
|
|
}
|
|
d := out.decision
|
|
e.NetworkID = d.NetworkID()
|
|
e.NetworkName = d.NetworkName()
|
|
e.PolicyID = d.PolicyID()
|
|
e.PolicyName = d.PolicyName()
|
|
if out.blocked {
|
|
e.BlacklistID = d.ListRef()
|
|
e.BlacklistName = d.ListName
|
|
e.MatchedRule = d.MatchedDomain
|
|
}
|
|
s.qlog.Log(e)
|
|
}
|
|
|
|
// logQuery records a query that never produced a response, such as one dropped
|
|
// by the rate limiter.
|
|
func (s *Server) logQuery(req *dns.Msg, client netip.Addr, protocol string,
|
|
out outcome, snap *runtimecfg.Snapshot, start time.Time) {
|
|
|
|
if len(req.Question) == 0 {
|
|
return
|
|
}
|
|
s.logQueryOutcome(req, client, protocol, out, snap, time.Since(start))
|
|
}
|