Files
dnswatcher/internal/watcher/watcher_test.go
T
sneak 032e07d823
check / check (push) Successful in 1m16s
watcher: save state when it stops and wait for that save (closes #114)
The final save at shutdown came from the state's own stop hook, while
the watcher's stop hook only cancelled its run loop, so a check under
way could change state after that save or be cut off at exit. Run now
saves state as it returns, and the watcher's stop hook waits for Run,
bounded by the shutdown deadline. The state's own save stays; Save
holds the state lock for the whole write, so the two cannot overlap.
The start hook derives the watcher's context with WithoutCancel, so
the linter needs no exception. A new test stops a watcher built by New
and reads the change back from the state file, with no DNS.

Model: opus-5-5
2026-10-01 21:12:29 +00:00

804 lines
18 KiB
Go

package watcher_test
import (
"context"
"fmt"
"log/slog"
"os"
"slices"
"sync"
"testing"
"time"
"go.uber.org/fx/fxtest"
"sneak.berlin/go/dnswatcher/internal/config"
"sneak.berlin/go/dnswatcher/internal/globals"
"sneak.berlin/go/dnswatcher/internal/livednstest"
"sneak.berlin/go/dnswatcher/internal/logger"
"sneak.berlin/go/dnswatcher/internal/portcheck"
"sneak.berlin/go/dnswatcher/internal/resolver"
"sneak.berlin/go/dnswatcher/internal/state"
"sneak.berlin/go/dnswatcher/internal/tlscheck"
"sneak.berlin/go/dnswatcher/internal/watcher"
)
// The watcher looks these names up in live DNS with the real resolver,
// so tests assert on what the watcher does with the answers, never on
// the records these zones publish. testHost's nameservers and addresses
// stay the same from one check to the next, which the tests that check
// it twice rely on.
const (
testDomain = "google.com"
testHost = "cloudflare.com"
testIssuer = "DigiCert"
)
// Saved-state values that live DNS never returns: nameserver names
// under .invalid and a documentation address.
const (
oldNS1 = "ns1.example.invalid."
oldNS2 = "ns2.example.invalid."
oldIP = "192.0.2.1"
)
// --- Stand-ins for the port checker, TLS checker and notifier ---
//
// DNS has none: the watchers built here use the real resolver (see
// TESTING.md).
// mockPortChecker reports every port open until closed is set.
type mockPortChecker struct {
mu sync.Mutex
closed bool
calls int
}
func (m *mockPortChecker) CheckPort(
_ context.Context,
_ string,
_ int,
) (*portcheck.PortResult, error) {
m.mu.Lock()
defer m.mu.Unlock()
m.calls++
return &portcheck.PortResult{Open: !m.closed}, nil
}
// mockTLSChecker returns a certificate for the requested hostname that
// expires at notAfter.
type mockTLSChecker struct {
mu sync.Mutex
notAfter time.Time
calls int
}
func (m *mockTLSChecker) CheckCertificate(
_ context.Context,
_ string,
hostname string,
) (*tlscheck.CertificateInfo, error) {
m.mu.Lock()
defer m.mu.Unlock()
m.calls++
return &tlscheck.CertificateInfo{
CommonName: hostname,
Issuer: testIssuer,
NotAfter: m.notAfter,
SubjectAlternativeNames: []string{hostname},
}, nil
}
type notification struct {
Title string
Message string
Priority string
}
type mockNotifier struct {
mu sync.Mutex
notifications []notification
}
func (m *mockNotifier) SendNotification(
_ context.Context,
title, message, priority string,
) {
m.mu.Lock()
defer m.mu.Unlock()
m.notifications = append(m.notifications, notification{
Title: title,
Message: message,
Priority: priority,
})
}
func (m *mockNotifier) getNotifications() []notification {
m.mu.Lock()
defer m.mu.Unlock()
result := make([]notification, len(m.notifications))
copy(result, m.notifications)
return result
}
// --- Helpers to build a Watcher and run its checks against live DNS ---
type testDeps struct {
portChecker *mockPortChecker
tlsChecker *mockTLSChecker
notifier *mockNotifier
state *state.State
config *config.Config
log *logger.Logger
}
func newTestWatcher(
t *testing.T,
cfg *config.Config,
) (*watcher.Watcher, *testDeps) {
t.Helper()
deps := newTestDeps(t, cfg)
w := watcher.NewForTest(
deps.config,
deps.state,
resolver.NewFromLogger(slog.Default()),
deps.portChecker,
deps.tlsChecker,
deps.notifier,
)
return w, deps
}
func newTestDeps(t *testing.T, cfg *config.Config) *testDeps {
t.Helper()
deps := &testDeps{
portChecker: &mockPortChecker{},
tlsChecker: &mockTLSChecker{
notAfter: time.Now().Add(90 * 24 * time.Hour),
},
notifier: &mockNotifier{},
config: cfg,
}
g, err := globals.New(nil)
if err != nil {
t.Fatalf("globals.New: %v", err)
}
deps.log, err = logger.New(nil, logger.Params{Globals: g})
if err != nil {
t.Fatalf("logger.New: %v", err)
}
// The watcher saves state after every check, into cfg.DataDir.
deps.state, err = state.New(fxtest.NewLifecycle(t), state.Params{
Logger: deps.log,
Config: cfg,
})
if err != nil {
t.Fatalf("state.New: %v", err)
}
return deps
}
func defaultTestConfig(t *testing.T) *config.Config {
t.Helper()
return &config.Config{
DNSInterval: time.Hour,
TLSInterval: 12 * time.Hour,
TLSExpiryWarning: 7,
DataDir: t.TempDir(),
}
}
// checkOnce runs the watcher's checks once and returns an error when a
// configured name has no hostname state saved by this check, or that
// state holds no address. Either live DNS gave no answer for the name,
// or the watcher saved no fresh result for it.
func checkOnce(
ctx context.Context,
w *watcher.Watcher,
deps *testDeps,
) error {
started := time.Now()
w.RunOnce(ctx)
names := slices.Concat(deps.config.Domains, deps.config.Hostnames)
for _, name := range names {
hs, ok := deps.state.GetHostnameState(name)
if !ok || hs.LastChecked.Before(started) ||
len(addresses(hs)) == 0 {
return fmt.Errorf(
"%s: %w, or the watcher saved no fresh "+
"result for it",
name, livednstest.ErrNoAnswer,
)
}
}
return nil
}
// runChecks builds a watcher, lets prepare set up the saved state and
// stand-ins it starts from, and runs its checks once against live DNS.
// If change is not nil, change then alters the saved state or stand-ins
// and the checks run a second time. When either check finds no fresh
// address for a name (see checkOnce), the watcher is thrown away and
// all of this runs again on a new one, so a failed attempt leaves
// nothing behind in the saved state, the stand-ins or the notifications.
func runChecks(
t *testing.T,
cfg *config.Config,
prepare, change func(deps *testDeps),
) *testDeps {
t.Helper()
var deps *testDeps
livednstest.Retry(t, "watcher checks", func(ctx context.Context) error {
var w *watcher.Watcher
w, deps = newTestWatcher(t, cfg)
if prepare != nil {
prepare(deps)
}
err := checkOnce(ctx, w, deps)
if err != nil || change == nil {
return err
}
change(deps)
return checkOnce(ctx, w, deps)
})
return deps
}
// addresses returns the A and AAAA values saved for a hostname.
func addresses(hs *state.HostnameState) []string {
var ips []string
for _, nsState := range hs.RecordsByNameserver {
ips = append(ips, nsState.Records["A"]...)
ips = append(ips, nsState.Records["AAAA"]...)
}
return ips
}
// assertNotified checks that a notification with this title and
// priority was sent.
func assertNotified(
t *testing.T,
deps *testDeps,
title, priority string,
) {
t.Helper()
notifications := deps.notifier.getNotifications()
for _, n := range notifications {
if n.Title == title && n.Priority == priority {
return
}
}
t.Errorf(
"expected %s notification %q, got: %v",
priority, title, notifications,
)
}
// countNotifications counts the notifications sent with this title.
func countNotifications(deps *testDeps, title string) int {
count := 0
for _, n := range deps.notifier.getNotifications() {
if n.Title == title {
count++
}
}
return count
}
func TestFirstRunBaseline(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Domains = []string{testDomain}
cfg.Hostnames = []string{testHost}
deps := runChecks(t, cfg, nil, nil)
assertNoNotifications(t, deps)
assertStatePopulated(t, deps)
}
func assertNoNotifications(
t *testing.T,
deps *testDeps,
) {
t.Helper()
notifications := deps.notifier.getNotifications()
if len(notifications) != 0 {
t.Errorf(
"expected 0 notifications on first run, got %d",
len(notifications),
)
}
}
func assertStatePopulated(
t *testing.T,
deps *testDeps,
) {
t.Helper()
snap := deps.state.GetSnapshot()
if len(snap.Domains) != 1 {
t.Errorf(
"expected 1 domain in state, got %d",
len(snap.Domains),
)
}
// Hostnames includes both explicit hostnames and domains
// (domains now also get hostname state for port/TLS checks).
if len(snap.Hostnames) < 1 {
t.Errorf(
"expected at least 1 hostname in state, got %d",
len(snap.Hostnames),
)
}
}
func TestDomainPortAndTLSChecks(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Domains = []string{testDomain}
deps := runChecks(t, cfg, nil, nil)
snap := deps.state.GetSnapshot()
// Domain should have port state populated
if len(snap.Ports) == 0 {
t.Error("expected port state for domain, got none")
}
// Domain should have certificate state populated
if len(snap.Certificates) == 0 {
t.Error("expected certificate state for domain, got none")
}
// Verify port checker was actually called
deps.portChecker.mu.Lock()
calls := deps.portChecker.calls
deps.portChecker.mu.Unlock()
if calls == 0 {
t.Error("expected port checker to be called for domain")
}
// Verify TLS checker was actually called
deps.tlsChecker.mu.Lock()
tlsCalls := deps.tlsChecker.calls
deps.tlsChecker.mu.Unlock()
if tlsCalls == 0 {
t.Error("expected TLS checker to be called for domain")
}
}
func TestNSChangeDetection(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Domains = []string{testDomain}
// The saved state lists nameservers that live DNS does not.
deps := runChecks(t, cfg, func(deps *testDeps) {
deps.state.SetDomainState(testDomain, &state.DomainState{
Nameservers: []string{oldNS1, oldNS2},
})
}, nil)
assertNotified(t, deps, "NS Change: "+testDomain, "warning")
ds, _ := deps.state.GetDomainState(testDomain)
if slices.Contains(ds.Nameservers, oldNS1) {
t.Errorf("saved nameservers not updated: %v", ds.Nameservers)
}
}
func TestRecordChangeDetection(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Hostnames = []string{testHost}
// Between the checks, save for every nameserver an address live DNS
// never returns.
deps := runChecks(t, cfg, nil, func(deps *testDeps) {
hs, _ := deps.state.GetHostnameState(testHost)
for _, nsState := range hs.RecordsByNameserver {
nsState.Records = map[string][]string{"A": {oldIP}}
}
deps.state.SetHostnameState(testHost, hs)
})
assertNotified(t, deps, "Record Change: "+testHost, "warning")
}
func TestPortStateChange(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Hostnames = []string{testHost}
// Between the checks, every port closes.
deps := runChecks(t, cfg, nil, func(deps *testDeps) {
deps.portChecker.mu.Lock()
deps.portChecker.closed = true
deps.portChecker.mu.Unlock()
})
hs, _ := deps.state.GetHostnameState(testHost)
assertNotified(
t, deps, "Port Change: "+addresses(hs)[0]+":443", "warning",
)
}
// expiresInThreeDays makes the TLS checker return certificates that
// expire within the seven-day warning period.
func expiresInThreeDays(deps *testDeps) {
deps.tlsChecker.notAfter = time.Now().Add(3 * 24 * time.Hour)
}
func TestTLSExpiryWarning(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Hostnames = []string{testHost}
deps := runChecks(t, cfg, expiresInThreeDays, nil)
assertNotified(t, deps, "TLS Expiry Warning: "+testHost, "warning")
}
func TestTLSExpiryWarningDedup(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Hostnames = []string{testHost}
cfg.TLSInterval = 24 * time.Hour
title := "TLS Expiry Warning: " + testHost
// The second check comes within the TLS interval of the first,
// so it must not warn again.
var warnings int
deps := runChecks(t, cfg, expiresInThreeDays, func(deps *testDeps) {
warnings = countNotifications(deps, title)
})
if warnings == 0 {
t.Fatal("expected expiry warnings from the first check")
}
got := countNotifications(deps, title)
if got != warnings {
t.Errorf(
"expected %d expiry warnings (dedup), got %d",
warnings, got,
)
}
}
func TestGracefulShutdown(t *testing.T) {
t.Parallel()
// No domains or hostnames: stopping does not involve DNS.
cfg := defaultTestConfig(t)
cfg.DNSInterval = 100 * time.Millisecond
cfg.TLSInterval = 100 * time.Millisecond
w, _ := newTestWatcher(t, cfg)
ctx, cancel := context.WithCancel(t.Context())
done := make(chan struct{})
go func() {
w.Run(ctx)
close(done)
}()
time.Sleep(250 * time.Millisecond)
cancel()
select {
case <-done:
// Shut down cleanly
case <-time.After(5 * time.Second):
t.Error("watcher did not shut down within timeout")
}
}
// TestStopSavesState stops a watcher built by New the way fx stops it,
// and checks that a change made to the state after the last check is in
// the state file afterwards. The state's own stop hook never runs here,
// so only the watcher can have saved it. Nothing is configured to
// check, so no DNS is involved.
func TestStopSavesState(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
deps := newTestDeps(t, cfg)
lc := fxtest.NewLifecycle(t)
_, err := watcher.New(lc, watcher.Params{
Logger: deps.log,
Config: cfg,
State: deps.state,
Resolver: resolver.NewFromLogger(slog.Default()),
PortCheck: deps.portChecker,
TLSCheck: deps.tlsChecker,
Notify: deps.notifier,
})
if err != nil {
t.Fatalf("watcher.New: %v", err)
}
lc.RequireStart()
// The first check saves state once. Wait for that save before
// changing the state, so the change can reach the file only
// through the save made at stop.
deadline := time.Now().Add(5 * time.Second)
for {
_, err = os.Stat(cfg.StatePath())
if err == nil {
break
}
if time.Now().After(deadline) {
t.Fatalf("the first check saved no state: %v", err)
}
time.Sleep(10 * time.Millisecond)
}
deps.state.SetDomainState(testDomain, &state.DomainState{
Nameservers: []string{oldNS1},
})
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
defer cancel()
err = lc.Stop(ctx)
if err != nil {
t.Fatalf("stopping the watcher: %v", err)
}
saved, err := state.New(fxtest.NewLifecycle(t), state.Params{
Logger: deps.log,
Config: cfg,
})
if err != nil {
t.Fatalf("state.New: %v", err)
}
err = saved.Load()
if err != nil {
t.Fatalf("loading the state file: %v", err)
}
ds, ok := saved.GetDomainState(testDomain)
if !ok || !slices.Equal(ds.Nameservers, []string{oldNS1}) {
t.Errorf(
"state file after stop has %+v for %s, want nameservers %v",
ds, testDomain, []string{oldNS1},
)
}
}
func TestDNSRunsBeforePortAndTLSChecks(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Hostnames = []string{testHost}
// The saved state says the last check found testHost at oldIP.
deps := runChecks(t, cfg, func(deps *testDeps) {
deps.state.SetHostnameState(testHost, &state.HostnameState{
RecordsByNameserver: map[string]*state.NameserverRecordState{
oldNS1: {
Records: map[string][]string{"A": {oldIP}},
Status: "ok",
},
},
})
}, nil)
snap := deps.state.GetSnapshot()
if _, ok := snap.Ports[oldIP+":80"]; ok {
t.Error("port check used stale DNS: found " + oldIP + ":80")
}
// Port and TLS checks must use the addresses this check found.
for _, ip := range addresses(snap.Hostnames[testHost]) {
if _, ok := snap.Ports[ip+":80"]; !ok {
t.Error("port check used stale DNS: missing " + ip + ":80")
}
certKey := ip + ":443:" + testHost
if _, ok := snap.Certificates[certKey]; !ok {
t.Error("TLS check used stale DNS: missing " + certKey)
}
}
}
func TestSendTestNotification_Enabled(t *testing.T) {
t.Parallel()
// No domains or hostnames: the startup notification does not
// involve DNS.
cfg := defaultTestConfig(t)
cfg.SendTestNotification = true
w, deps := newTestWatcher(t, cfg)
w.RunOnce(t.Context())
// RunOnce does not send the test notification — it is
// sent by Run after RunOnce completes. Call the exported
// RunOnce then check that no test notification was sent
// (only Run triggers it). We test the full path via Run.
notifications := deps.notifier.getNotifications()
if len(notifications) != 0 {
t.Errorf(
"RunOnce should not send test notification, got %d",
len(notifications),
)
}
}
func TestSendTestNotification_ViaRun(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.SendTestNotification = true
cfg.DNSInterval = 24 * time.Hour
cfg.TLSInterval = 24 * time.Hour
w, deps := newTestWatcher(t, cfg)
ctx, cancel := context.WithCancel(t.Context())
done := make(chan struct{})
go func() {
w.Run(ctx)
close(done)
}()
// Wait for the initial scan and test notification.
time.Sleep(500 * time.Millisecond)
cancel()
<-done
notifications := deps.notifier.getNotifications()
found := false
for _, n := range notifications {
if n.Priority == "success" &&
n.Title == "✅ dnswatcher startup complete" {
found = true
}
}
if !found {
t.Errorf(
"expected startup test notification, got: %v",
notifications,
)
}
}
func TestSendTestNotification_Disabled(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.SendTestNotification = false
cfg.DNSInterval = 24 * time.Hour
cfg.TLSInterval = 24 * time.Hour
w, deps := newTestWatcher(t, cfg)
ctx, cancel := context.WithCancel(t.Context())
done := make(chan struct{})
go func() {
w.Run(ctx)
close(done)
}()
time.Sleep(500 * time.Millisecond)
cancel()
<-done
notifications := deps.notifier.getNotifications()
for _, n := range notifications {
if n.Title == "✅ dnswatcher startup complete" {
t.Error(
"test notification should not be sent when disabled",
)
}
}
}
func TestNSFailureAndRecovery(t *testing.T) {
t.Parallel()
cfg := defaultTestConfig(t)
cfg.Hostnames = []string{testHost}
// Between the checks, save every nameserver the first check found
// as one that did not answer, and add, as answering, one that live
// DNS does not list, which then disappears.
deps := runChecks(t, cfg, nil, func(deps *testDeps) {
hs, _ := deps.state.GetHostnameState(testHost)
for ns := range hs.RecordsByNameserver {
hs.RecordsByNameserver[ns] = failed()
}
hs.RecordsByNameserver[oldNS1] = &state.NameserverRecordState{
Records: map[string][]string{"A": {oldIP}},
Status: "ok",
}
deps.state.SetHostnameState(testHost, hs)
})
assertNotified(t, deps, "NS Failure: "+testHost, "error")
assertNotified(t, deps, "NS Recovery: "+testHost, "success")
// A nameserver that did not answer has no records to compare, so
// its recovery is not also a record change.
if n := countNotifications(deps, "Record Change: "+testHost); n != 0 {
t.Errorf("sent %d record changes on recovery, want 0", n)
}
}