Files
dnswatcher/internal/watcher/watcher_test.go
T
clawbot 8259ff6a3a
check / check (push) Successful in 1m11s
tests: remove the DNS stand-ins from the watcher and resolver tests (closes #159)
The watcher tests used a stand-in resolver and the resolver timeout test
a stand-in DNS client, against the rule that DNS is never mocked.
Watcher tests that look something up in DNS now run the real resolver
against live servers, each attempt on a new watcher. A DNS change is
tested by saving values live DNS never returns (names under .invalid,
192.0.2.1) in the state a check starts from, or by marking a real
nameserver failed. The timeout test queries 192.0.2.1, where nothing
answers. The live-DNS retry and concurrency limit moved from the
resolver tests to internal/livedns, so both packages share them.
NewFromLoggerWithClient had no other use and is gone. TESTING.md now
states the README's rule.

Model: opus-5-5
2026-09-29 05:45:28 +00:00

687 lines
15 KiB
Go

package watcher_test
import (
"context"
"fmt"
"log/slog"
"slices"
"sync"
"testing"
"time"
"sneak.berlin/go/dnswatcher/internal/config"
"sneak.berlin/go/dnswatcher/internal/livedns"
"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
}
func newTestWatcher(
t *testing.T,
cfg *config.Config,
) (*watcher.Watcher, *testDeps) {
t.Helper()
deps := &testDeps{
portChecker: &mockPortChecker{},
tlsChecker: &mockTLSChecker{
notAfter: time.Now().Add(90 * 24 * time.Hour),
},
notifier: &mockNotifier{},
config: cfg,
}
deps.state = state.NewForTest()
w := watcher.NewForTest(
deps.config,
deps.state,
resolver.NewFromLogger(slog.Default()),
deps.portChecker,
deps.tlsChecker,
deps.notifier,
)
return w, 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, livedns.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
livedns.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")
}
}
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 failed, and add, as answering, one that live DNS does not list.
deps := runChecks(t, cfg, nil, func(deps *testDeps) {
hs, _ := deps.state.GetHostnameState(testHost)
for _, nsState := range hs.RecordsByNameserver {
nsState.Status = "error"
}
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")
}