Files
owenandClaude Sonnet 5 4895c8fd1e Add a customizable HTTP/HTTPS block page for sinkholed queries
Serves a page explaining why a domain was blocked instead of leaving a
sinkholed client with a dead connection. Binds its own HTTP/HTTPS
listeners with self-signed, per-hostname TLS certs generated on the
fly, re-evaluates the requesting client against the policy engine per
request, and renders an HTML template editable from Settings with a
live preview.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01TTKpGMQzpfDsvedu1hvSUf
2026-08-17 00:24:00 -05:00

335 lines
9.5 KiB
Go

package app
import (
"context"
"encoding/base64"
"errors"
"fmt"
"log/slog"
"os"
"time"
"github.com/owen/vibedns/internal/auditlog"
"github.com/owen/vibedns/internal/auth"
"github.com/owen/vibedns/internal/backup"
"github.com/owen/vibedns/internal/blockpage"
"github.com/owen/vibedns/internal/cache"
"github.com/owen/vibedns/internal/config"
"github.com/owen/vibedns/internal/database"
"github.com/owen/vibedns/internal/dnsengine"
"github.com/owen/vibedns/internal/metrics"
"github.com/owen/vibedns/internal/querylog"
"github.com/owen/vibedns/internal/ratelimit"
"github.com/owen/vibedns/internal/resolver"
"github.com/owen/vibedns/internal/runtimecfg"
"github.com/owen/vibedns/internal/version"
)
// keyCSRFSecret stores the CSRF signing key so tokens survive a restart.
const keyCSRFSecret = "security.csrf_key"
// App wires every component together and exposes the operations the
// management interface performs.
type App struct {
Boot config.Bootstrap
DB *database.DB
Runtime *runtimecfg.Manager
Cache *cache.Cache
Resolver *resolver.Resolver
Limiter *ratelimit.Limiter
Metrics *metrics.Metrics
QueryLog *querylog.Logger
Audit *auditlog.Logger
Auth *auth.Authenticator
Backups *backup.Manager
DNS *dnsengine.Server
BlockPage *blockpage.Server
Log *slog.Logger
cancel context.CancelFunc
started time.Time
}
// New builds the application. The database must already be migrated.
func New(ctx context.Context, boot config.Bootstrap, db *database.DB, log *slog.Logger) (*App, error) {
rt, err := runtimecfg.New(ctx, db, log)
if err != nil {
return nil, fmt.Errorf("build the initial configuration snapshot: %w", err)
}
settings := rt.Settings()
csrfKey, err := loadOrCreateCSRFKey(ctx, db)
if err != nil {
return nil, err
}
a := &App{
Boot: boot,
DB: db,
Runtime: rt,
Log: log,
started: time.Now(),
}
a.Metrics = metrics.New(version.Version)
a.Cache = cache.New(cacheConfig(settings))
a.Resolver = resolver.New(resolverConfig(settings))
a.Limiter = ratelimit.New(rateLimitConfig(settings))
a.QueryLog = querylog.New(db, log, queryLogConfig(settings))
a.Audit = auditlog.New(db, log)
a.Auth = auth.New(db, log, csrfKey)
a.Auth.SetTrustedProxies(settings.HTTP.TrustedProxies)
a.Backups = backup.New(db, log, backupConfig(settings))
a.DNS = dnsengine.New(dnsengine.Options{
Runtime: rt,
Cache: a.Cache,
Resolver: a.Resolver,
Limiter: a.Limiter,
Metrics: a.Metrics,
QueryLog: a.QueryLog,
Log: log,
})
a.BlockPage = blockpage.New(rt, log)
// Every subsystem picks up new settings from the same reload event, so a
// change in the UI takes effect without a restart.
rt.OnReload(func(s *runtimecfg.Snapshot) {
a.Cache.SetConfig(cacheConfig(s.Settings))
a.Resolver.SetConfig(resolverConfig(s.Settings))
a.Limiter.SetConfig(rateLimitConfig(s.Settings))
a.QueryLog.SetConfig(queryLogConfig(s.Settings))
a.Backups.SetConfig(backupConfig(s.Settings))
a.Auth.SetTrustedProxies(s.Settings.HTTP.TrustedProxies)
})
a.Metrics.SetGaugeSource(a.gauges)
return a, nil
}
func cacheConfig(s config.Settings) cache.Config {
return cache.Config{
Enabled: s.Cache.Enabled,
MaxEntries: s.Cache.MaxEntries,
MinTTL: uint32(s.Cache.MinTTL),
MaxTTL: uint32(s.Cache.MaxTTL),
NegativeTTL: uint32(s.Cache.NegativeTTL),
ServeStale: s.Cache.ServeStale,
StaleTTL: uint32(s.Cache.StaleTTL),
Prefetch: s.Cache.Prefetch,
PrefetchPercent: s.Cache.PrefetchPercent,
}
}
func resolverConfig(s config.Settings) resolver.Config {
return resolver.Config{
Upstreams: s.Resolver.Upstreams,
Timeout: time.Duration(s.Resolver.TimeoutMS) * time.Millisecond,
Retries: s.Resolver.Retries,
Strategy: s.Resolver.Strategy,
DNSSEC: s.Resolver.DNSSEC,
EDNSUDPSize: uint16(s.DNS.EDNSUDPSize),
MaxConcurrent: s.Resolver.MaxConcurrent,
}
}
func rateLimitConfig(s config.Settings) ratelimit.Config {
return ratelimit.Config{
Enabled: s.RateLimit.Enabled,
QPS: s.RateLimit.QPS,
Burst: s.RateLimit.Burst,
Exempt: s.RateLimit.ExemptNetworks,
}
}
func queryLogConfig(s config.Settings) querylog.Config {
return querylog.Config{
Enabled: s.QueryLog.Enabled,
RetentionDays: s.QueryLog.RetentionDays,
MaxRows: s.QueryLog.MaxRows,
CleanupMinutes: s.QueryLog.CleanupMinutes,
IgnoreNetworks: s.QueryLog.IgnoreNetworks,
IgnoreDomains: s.QueryLog.IgnoreDomains,
}
}
func backupConfig(s config.Settings) backup.Config {
return backup.Config{
Enabled: s.Backup.Enabled,
Directory: s.Backup.Directory,
IntervalHours: s.Backup.IntervalHours,
Retention: s.Backup.Retention,
}
}
// loadOrCreateCSRFKey fetches the persisted CSRF signing key, creating one on
// first run. Persisting it means tokens in open browser tabs survive a restart.
func loadOrCreateCSRFKey(ctx context.Context, db *database.DB) ([]byte, error) {
if v, ok, err := db.Setting(ctx, keyCSRFSecret); err != nil {
return nil, fmt.Errorf("read the CSRF signing key: %w", err)
} else if ok && v != "" {
key, err := base64.RawStdEncoding.DecodeString(v)
if err == nil && len(key) >= 32 {
return key, nil
}
}
encoded, err := auth.RandomKey(32)
if err != nil {
return nil, err
}
if err := db.SetSetting(ctx, keyCSRFSecret, encoded); err != nil {
return nil, fmt.Errorf("store the CSRF signing key: %w", err)
}
key, err := base64.RawStdEncoding.DecodeString(encoded)
if err != nil {
return nil, fmt.Errorf("decode the CSRF signing key: %w", err)
}
return key, nil
}
// Start launches every background worker and binds the DNS listeners.
func (a *App) Start(ctx context.Context) error {
ctx, a.cancel = context.WithCancel(ctx)
a.Runtime.Start(ctx)
a.QueryLog.Start(ctx)
a.Backups.Start(ctx)
done := ctx.Done()
go a.Cache.Run(done, func() time.Duration {
return time.Duration(a.Runtime.Settings().Cache.CleanupSeconds) * time.Second
})
go a.Limiter.Run(done, time.Minute, 10*time.Minute)
go a.pruneAuditLoop(ctx)
if err := a.DNS.Start(ctx); err != nil {
a.cancel()
return err
}
if settings := a.Runtime.Settings(); settings.BlockPage.Enabled {
if err := a.BlockPage.Start(settings.BlockPage.HTTPListen, settings.BlockPage.HTTPSListen); err != nil {
a.cancel()
_ = a.DNS.Shutdown(context.Background())
return err
}
}
return nil
}
// Shutdown stops the DNS listeners and drains the background workers.
func (a *App) Shutdown(ctx context.Context) error {
err := a.DNS.Shutdown(ctx)
if bpErr := a.BlockPage.Shutdown(ctx); bpErr != nil && err == nil {
err = bpErr
}
if a.cancel != nil {
a.cancel()
}
a.QueryLog.Stop()
a.Runtime.Stop()
a.Backups.Stop()
return err
}
// pruneAuditLoop keeps the audit log bounded.
func (a *App) pruneAuditLoop(ctx context.Context) {
t := time.NewTicker(6 * time.Hour)
defer t.Stop()
for {
select {
case <-ctx.Done():
return
case <-t.C:
max := a.Runtime.Settings().Logging.AuditMaxRows
if n, err := a.DB.PruneAuditLogs(ctx, max); err != nil {
a.Log.Warn("could not prune the audit log", "error", err)
} else if n > 0 {
a.Log.Debug("pruned audit log", "rows", n)
}
}
}
}
// Settings returns the active runtime settings.
func (a *App) Settings() config.Settings { return a.Runtime.Settings() }
// Snapshot returns the active configuration snapshot.
func (a *App) Snapshot() *runtimecfg.Snapshot { return a.Runtime.Current() }
// StartedAt returns when the application started.
func (a *App) StartedAt() time.Time { return a.started }
// Uptime returns how long the application has been running.
func (a *App) Uptime() time.Duration { return time.Since(a.started) }
// Reload rebuilds the configuration snapshot immediately.
func (a *App) Reload(ctx context.Context) error {
if err := a.Runtime.Reload(ctx); err != nil {
return Internal(err, "The configuration could not be reloaded.")
}
return nil
}
// gauges samples live values for the metrics endpoint.
func (a *App) gauges() metrics.Gauges {
snap := a.Runtime.Current()
cs := a.Cache.Stats()
rs := a.Resolver.Stats()
g := metrics.Gauges{
CacheEntries: int64(cs.Entries),
CacheBytes: cs.Bytes,
Zones: int64(snap.ZoneCount),
Records: int64(snap.RecordCount),
BlacklistDomains: int64(snap.BlacklistDomains),
AllowlistDomains: int64(snap.AllowlistDomains),
Networks: int64(snap.NetworkCount),
UpstreamsTotal: int64(rs.Upstreams),
UpstreamsHealthy: int64(rs.Healthy),
}
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
if n, err := a.DB.QueryLogCount(ctx); err == nil {
g.QueryLogRows = n
}
return g
}
// Ready reports whether the server is able to answer queries. It backs /readyz.
func (a *App) Ready(ctx context.Context) error {
if !a.DNS.Running() {
return errors.New("DNS listeners are not running")
}
if err := a.DB.PingContext(ctx); err != nil {
return fmt.Errorf("database is unreachable: %w", err)
}
if a.Runtime.Current() == nil {
return errors.New("configuration has not been loaded")
}
return nil
}
// translate converts a storage error into a user-facing one.
func translate(err error, notFound, conflict string) error {
switch {
case err == nil:
return nil
case errors.Is(err, database.ErrNotFound):
return NotFound("%s", notFound)
case errors.Is(err, database.ErrConflict):
return Conflict("%s", conflict)
default:
return Internal(err, "The change could not be saved.")
}
}
// Hostname returns the machine name, shown on the dashboard.
func Hostname() string {
h, err := os.Hostname()
if err != nil || h == "" {
return "unknown"
}
return h
}