package watcher_test import ( "context" "fmt" "log/slog" "os" "slices" "strings" "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. The nameservers of testHost and // testSmallDomain stay the same between a test looking them up and its // check. Every query a check sends is one more that can be lost, so the // tests keep them few. A check asks each of a name's nameservers about // every record type, and both names have two. A domain check also looks // up each nameserver's addresses at every nameserver of the zone that // nameserver is in: testSmallDomain's nameservers are in zones with two // nameservers, while a domain whose nameservers are in, say, // cloudflare.com, which has five, makes each domain check much longer. // A test checks a domain only when it is about domains, and checks once, // from saved state it builds, rather than twice. The tests that query // testDomain's nameservers directly do no domain check. const ( testDomain = "google.com" testSmallDomain = "desec.io" testHost = "example.org" 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, or a configured domain's nameserver has no // address saved or still has oldIP, which the tests save and live DNS // never returns. 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, ) } } for _, name := range deps.config.Domains { ds, _ := deps.state.GetDomainState(name) for _, ns := range ds.Nameservers { ips := ds.NameserverAddresses[ns] if len(ips) == 0 || slices.Contains(ips, oldIP) { return fmt.Errorf( "%s: nameserver %s: %w, or the watcher saved "+ "no fresh addresses for it", name, ns, 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. // When the 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 func(deps *testDeps), ) (*watcher.Watcher, *testDeps) { t.Helper() var ( w *watcher.Watcher deps *testDeps ) livednstest.Retry(t, "watcher checks", func(ctx context.Context) error { w, deps = newTestWatcher(t, cfg) if prepare != nil { prepare(deps) } return checkOnce(ctx, w, deps) }) return w, deps } // lookupNameservers returns the nameservers live DNS lists for name, // for a test to save in the state its check starts from. func lookupNameservers(t *testing.T, name string) []string { t.Helper() res := resolver.NewFromLogger(slog.Default()) var nameservers []string livednstest.Retry(t, "LookupNS("+name+")", func(ctx context.Context) error { var err error nameservers, err = res.LookupNS(ctx, name) return err }) return nameservers } // 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{testSmallDomain} cfg.Hostnames = []string{testHost} _, deps := runChecks(t, cfg, 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{testSmallDomain} _, deps := runChecks(t, cfg, 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{testSmallDomain} // The saved state lists nameservers that live DNS does not. _, deps := runChecks(t, cfg, func(deps *testDeps) { deps.state.SetDomainState(testSmallDomain, &state.DomainState{ Nameservers: []string{oldNS1, oldNS2}, }) }) assertNotified(t, deps, "NS Change: "+testSmallDomain, "warning") ds, _ := deps.state.GetDomainState(testSmallDomain) if slices.Contains(ds.Nameservers, oldNS1) { t.Errorf("saved nameservers not updated: %v", ds.Nameservers) } } func TestNSAddressChangeDetection(t *testing.T) { t.Parallel() cfg := defaultTestConfig(t) cfg.Domains = []string{testSmallDomain} nameservers := lookupNameservers(t, testSmallDomain) // The saved state lists the nameservers live DNS lists, each at an // address live DNS never returns. _, deps := runChecks(t, cfg, func(deps *testDeps) { nsAddresses := make(map[string][]string, len(nameservers)) for _, ns := range nameservers { nsAddresses[ns] = []string{oldIP} } deps.state.SetDomainState(testSmallDomain, &state.DomainState{ Nameservers: nameservers, NameserverAddresses: nsAddresses, }) }) title := "NS Address Change: " + testSmallDomain ds, _ := deps.state.GetDomainState(testSmallDomain) // One alert per nameserver, naming it and the address it had. for _, ns := range ds.Nameservers { prefix := "Domain: " + testSmallDomain + "\nNameserver: " + ns + "\nOld: " + oldIP + "\nNew: " sent := 0 for _, n := range deps.notifier.getNotifications() { if n.Title == title && strings.HasPrefix(n.Message, prefix) { sent++ } } if sent != 1 { t.Errorf("sent %d address changes for %s, want 1", sent, ns) } } if n := countNotifications(deps, title); n != len(ds.Nameservers) { t.Errorf( "sent %d address changes for %d nameservers", n, len(ds.Nameservers), ) } if n := countNotifications(deps, "NS Change: "+testSmallDomain); n != 0 { t.Errorf("sent %d NS changes, want 0", n) } } func TestNSAddedAndRemovedIsNoAddressChange(t *testing.T) { t.Parallel() cfg := defaultTestConfig(t) cfg.Domains = []string{testSmallDomain} nameservers := lookupNameservers(t, testSmallDomain) // The saved state lists oldNS1, which live DNS does not, in place of // the first nameserver live DNS lists, so that the check finds that // one added and oldNS1 removed. Only oldNS1 has addresses saved. _, deps := runChecks(t, cfg, func(deps *testDeps) { deps.state.SetDomainState(testSmallDomain, &state.DomainState{ Nameservers: append([]string{oldNS1}, nameservers[1:]...), NameserverAddresses: map[string][]string{oldNS1: {oldIP}}, }) }) if n := countNotifications(deps, "NS Change: "+testSmallDomain); n != 1 { t.Errorf("sent %d NS changes, want 1", n) } title := "NS Address Change: " + testSmallDomain if n := countNotifications(deps, title); n != 0 { t.Errorf("sent %d address changes, want 0", n) } } func TestRecordChangeDetection(t *testing.T) { t.Parallel() cfg := defaultTestConfig(t) cfg.Hostnames = []string{testHost} nameservers := lookupNameservers(t, testHost) // The saved state has every nameserver live DNS lists answering // with an address live DNS never returns. _, deps := runChecks(t, cfg, func(deps *testDeps) { byNameserver := make(map[string]*state.NameserverRecordState) for _, ns := range nameservers { byNameserver[ns] = answered(map[string][]string{"A": {oldIP}}) } deps.state.SetHostnameState(testHost, saved(byNameserver)) }) assertNotified(t, deps, "Record Change: "+testHost, "warning") } func TestPortStateChange(t *testing.T) { t.Parallel() cfg := defaultTestConfig(t) cfg.Hostnames = []string{testHost} w, deps := runChecks(t, cfg, nil) // Every port closes, and the port checks run again. They look // nothing up. deps.portChecker.mu.Lock() deps.portChecker.closed = true deps.portChecker.mu.Unlock() w.CheckAllPorts(t.Context()) 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) assertNotified(t, deps, "TLS Expiry Warning: "+testHost, "warning") } // TestTLSExpiryWarningEachCheck runs the TLS checks three times in a // row on hostname and port state built here, for a certificate that // expires within the warning period. Each check warns once, whether the // TLS interval is a nanosecond, shorter than the time between two // checks, or a day, longer than it. func TestTLSExpiryWarningEachCheck(t *testing.T) { t.Parallel() title := "TLS Expiry Warning: " + host for _, interval := range []time.Duration{time.Nanosecond, 24 * time.Hour} { t.Run(interval.String(), func(t *testing.T) { t.Parallel() cfg := defaultTestConfig(t) cfg.Hostnames = []string{host} cfg.TLSInterval = interval // The TLS checks read the saved hostname and port state and // look nothing up, so the watcher has no resolver. deps := newTestDeps(t, cfg) w := watcher.NewForTest( cfg, deps.state, nil, deps.portChecker, deps.tlsChecker, deps.notifier, ) expiresInThreeDays(deps) deps.state.SetHostnameState(host, saved( map[string]*state.NameserverRecordState{ nsA: answered(map[string][]string{"A": {ip1}}), }, )) deps.state.SetPortState(ip1+":443", &state.PortState{ Open: true, Hostnames: []string{host}, }) for check := 1; check <= 3; check++ { w.RunTLSChecks(t.Context()) got := countNotifications(deps, title) if got != check { t.Fatalf( "after check %d: %d expiry warnings, want %d", check, got, check, ) } } }) } } 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", }, }, }) }) 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} nameservers := lookupNameservers(t, testHost) // The saved state has every nameserver live DNS lists as one that // did not answer, and, as answering, one that live DNS does not // list, which then disappears. _, deps := runChecks(t, cfg, func(deps *testDeps) { byNameserver := map[string]*state.NameserverRecordState{ oldNS1: answered(map[string][]string{"A": {oldIP}}), } for _, ns := range nameservers { byNameserver[ns] = failed() } deps.state.SetHostnameState(testHost, saved(byNameserver)) }) 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) } }