Files
2026-08-16 21:18:45 -05:00

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