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

259 lines
7.3 KiB
Go

package cli
import (
"context"
"fmt"
"io"
"log/slog"
"os"
"os/signal"
"strings"
"syscall"
"time"
"github.com/owen/vibedns/internal/api"
"github.com/owen/vibedns/internal/app"
"github.com/owen/vibedns/internal/auth"
"github.com/owen/vibedns/internal/backup"
"github.com/owen/vibedns/internal/config"
"github.com/owen/vibedns/internal/database"
"github.com/owen/vibedns/internal/version"
"github.com/owen/vibedns/internal/web"
)
func serveCommand() *Command {
return &Command{
Name: "serve",
Summary: "run the DNS server and management interface (default)",
Usage: "Usage: vibedns serve [flags]",
Run: runServe,
}
}
func runServe(ctx context.Context, env *Env, args []string) error {
c := serveCommand()
fs := newFlagSet(env, c)
env.Boot.BindFlags(fs)
if err := fs.Parse(args); err != nil {
return err
}
env.Boot.NoteFlagsSet(fs)
if err := env.Boot.Validate(); err != nil {
return Exit(2, "%v", err)
}
log := newLogger(env.Boot.LogLevel, env.Boot.LogFormat, env.Stderr)
env.Log = log
// A staged restore is applied before anything opens the database, which is
// the only moment it can be swapped safely.
if _, err := backup.ApplyPendingRestore(env.Boot.DBPath, log); err != nil {
return Exit(1, "could not apply the staged database restore: %v", err)
}
db, err := database.Open(env.Boot.DBPath)
if err != nil {
return Exit(1, "%v", err)
}
defer db.Close()
applied, err := db.Migrate(ctx)
if err != nil {
return Exit(1, "%v", err)
}
if applied > 0 {
log.Info("database migrations applied", "count", applied)
}
application, err := app.New(ctx, env.Boot, db, log)
if err != nil {
return Exit(1, "%v", err)
}
// Startup-critical addresses on the command line win over stored settings,
// which is what makes a misconfigured listen address recoverable.
if err := applyAddressOverrides(ctx, application, env.Boot); err != nil {
return Exit(1, "%v", err)
}
generated, err := ensureAdmin(ctx, application, env.Boot)
if err != nil {
return Exit(1, "%v", err)
}
// Re-read the settings after any override, then check them before binding.
settings := application.Settings()
if err := settings.Validate(); err != nil {
log.Warn("stored settings have a problem", "error", err)
}
apiServer := api.New(application, log)
webServer, err := web.New(web.Options{App: application, Log: log, API: apiServer.Handler()})
if err != nil {
return Exit(1, "could not prepare the management interface: %v", err)
}
runCtx, stop := signal.NotifyContext(ctx, os.Interrupt, syscall.SIGTERM)
defer stop()
if err := application.Start(runCtx); err != nil {
return Exit(1, "%v", err)
}
httpAddr := settings.HTTP.Listen
if err := webServer.Start(httpAddr); err != nil {
_ = application.Shutdown(context.Background())
return Exit(1, "%v", err)
}
printBanner(env, application, settings, generated)
<-runCtx.Done()
log.Info("shutting down")
shutdownCtx, cancel := context.WithTimeout(context.Background(), 20*time.Second)
defer cancel()
if err := webServer.Shutdown(shutdownCtx); err != nil {
log.Warn("management interface did not stop cleanly", "error", err)
}
if err := application.Shutdown(shutdownCtx); err != nil {
log.Warn("DNS server did not stop cleanly", "error", err)
}
if err := db.Checkpoint(shutdownCtx); err != nil {
log.Warn("could not checkpoint the write-ahead log", "error", err)
}
log.Info("stopped")
return nil
}
// applyAddressOverrides persists listen addresses supplied on the command line
// or in the environment, so the running process and the stored configuration
// agree about where it is listening.
func applyAddressOverrides(ctx context.Context, a *app.App, boot config.Bootstrap) error {
next := a.Settings()
changed := false
if boot.DNSAddrOverridden() {
if next.DNS.UDPListen != boot.DNSUDPAddr || next.DNS.TCPListen != boot.DNSTCPAddr {
next.DNS.UDPListen = boot.DNSUDPAddr
next.DNS.TCPListen = boot.DNSTCPAddr
changed = true
}
}
if boot.HTTPAddrOverridden() && next.HTTP.Listen != boot.HTTPAddr {
next.HTTP.Listen = boot.HTTPAddr
changed = true
}
if !changed {
return nil
}
next.Normalise()
if err := a.DB.SetSettings(ctx, next.ToMap()); err != nil {
return fmt.Errorf("store the listen addresses: %w", err)
}
return a.Runtime.Reload(ctx)
}
// ensureAdmin creates the administrator on first run, returning the generated
// password when one had to be invented.
func ensureAdmin(ctx context.Context, a *app.App, boot config.Bootstrap) (string, error) {
username := boot.AdminUsername
if username == "" {
username = "admin"
}
password := boot.AdminPassword
if password != "" {
if err := auth.ValidatePassword(password); err != nil {
return "", fmt.Errorf("the administrator password supplied in %s is unusable: %w",
config.EnvAdminPassword, err)
}
}
_, generated, err := a.Auth.EnsureAdmin(ctx, username, password)
if err != nil {
return "", fmt.Errorf("create the administrator account: %w", err)
}
return generated, nil
}
// printBanner writes the startup summary an operator reads once.
func printBanner(env *Env, a *app.App, settings config.Settings, generatedPassword string) {
w := env.Stdout
admin, _ := a.Admin(context.Background())
scheme := "http"
host := settings.HTTP.Listen
if strings.HasPrefix(host, "0.0.0.0:") {
host = "127.0.0.1:" + strings.TrimPrefix(host, "0.0.0.0:")
} else if strings.HasPrefix(host, "[::]:") {
host = "127.0.0.1:" + strings.TrimPrefix(host, "[::]:")
}
fmt.Fprintf(w, "\n%s %s starting\n\n", version.Name, version.Version)
fmt.Fprintf(w, " Database: %s\n", a.DB.Path())
fmt.Fprintf(w, " DNS UDP: %s\n", settings.DNS.UDPListen)
fmt.Fprintf(w, " DNS TCP: %s\n", settings.DNS.TCPListen)
fmt.Fprintf(w, " Management: %s://%s\n", scheme, host)
snap := a.Snapshot()
fmt.Fprintf(w, " Zones: %d (%d records)\n", snap.ZoneCount, snap.RecordCount)
fmt.Fprintf(w, " Filtering: %d networks, %s blocked domains\n",
snap.NetworkCount, formatCount(snap.BlacklistDomains))
if settings.DNS.Recursion {
fmt.Fprintf(w, " Recursion: enabled for %d network(s), %d upstream(s)\n",
len(settings.Resolver.AllowNetworks), len(settings.Resolver.Upstreams))
} else {
fmt.Fprintf(w, " Recursion: disabled (authoritative only)\n")
}
if generatedPassword != "" {
fmt.Fprintf(w, "\n Initial administrator:\n")
fmt.Fprintf(w, " Username: %s\n", admin.Username)
fmt.Fprintf(w, " Password: %s\n", generatedPassword)
fmt.Fprintf(w, "\n This password will not be displayed again.\n")
fmt.Fprintf(w, " Change it at %s://%s/account\n", scheme, host)
}
fmt.Fprintln(w)
}
func formatCount(n int) string {
s := fmt.Sprintf("%d", n)
if n < 1000 {
return s
}
var out []string
for len(s) > 3 {
out = append([]string{s[len(s)-3:]}, out...)
s = s[:len(s)-3]
}
return strings.Join(append([]string{s}, out...), ",")
}
// newLogger builds the structured logger.
func newLogger(level, format string, w io.Writer) *slog.Logger {
var lv slog.Level
switch strings.ToLower(level) {
case "debug":
lv = slog.LevelDebug
case "warn":
lv = slog.LevelWarn
case "error":
lv = slog.LevelError
default:
lv = slog.LevelInfo
}
opts := &slog.HandlerOptions{Level: lv}
var h slog.Handler
if strings.ToLower(format) == "json" {
h = slog.NewJSONHandler(w, opts)
} else {
h = slog.NewTextHandler(w, opts)
}
return slog.New(h)
}