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