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 notifications and saved state, never on the // records these zones publish. testHost's 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 watcher uses 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 // livedns.ErrNoAnswer when live DNS did not answer for a configured // name. The watcher saves a name's hostname state only when all of the // name's lookups succeed, so live DNS answered for a name when this // check saved its hostname state and that state holds an address. 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("%w: %s", livedns.ErrNoAnswer, name) } } return nil } // runFirstCheck 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 live DNS does not answer, the watcher is thrown away and // built again, so a failed attempt leaves nothing behind in the state // or the notifications. func runFirstCheck( t *testing.T, cfg *config.Config, prepare func(deps *testDeps), ) (*watcher.Watcher, *testDeps) { t.Helper() var ( w *watcher.Watcher deps *testDeps ) livedns.Retry(t, "first check", func(ctx context.Context) error { w, deps = newTestWatcher(t, cfg) if prepare != nil { prepare(deps) } return checkOnce(ctx, w, deps) }) return w, deps } // runCheck runs the watcher's checks once more against live DNS, // repeating them while live DNS does not answer. A failed lookup keeps // the name's saved records, so a repeat compares against the same // saved state. func runCheck(t *testing.T, w *watcher.Watcher, deps *testDeps) { t.Helper() livedns.Retry(t, "check", func(ctx context.Context) error { return checkOnce(ctx, w, 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 := runFirstCheck(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{testDomain} _, deps := runFirstCheck(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{testDomain} // The saved state lists nameservers that live DNS does not. _, deps := runFirstCheck(t, cfg, func(deps *testDeps) { deps.state.SetDomainState(testDomain, &state.DomainState{ Nameservers: []string{oldNS1, oldNS2}, }) }) 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} w, deps := runFirstCheck(t, cfg, nil) // Save, for every nameserver, an address live DNS never returns. hs, _ := deps.state.GetHostnameState(testHost) for _, nsState := range hs.RecordsByNameserver { nsState.Records = map[string][]string{"A": {oldIP}} } deps.state.SetHostnameState(testHost, hs) runCheck(t, w, deps) assertNotified(t, deps, "Record Change: "+testHost, "warning") } func TestPortStateChange(t *testing.T) { t.Parallel() cfg := defaultTestConfig(t) cfg.Hostnames = []string{testHost} w, deps := runFirstCheck(t, cfg, nil) deps.portChecker.mu.Lock() deps.portChecker.closed = true deps.portChecker.mu.Unlock() runCheck(t, w, deps) 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 := runFirstCheck(t, cfg, expiresInThreeDays) 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 w, deps := runFirstCheck(t, cfg, expiresInThreeDays) title := "TLS Expiry Warning: " + testHost warnings := countNotifications(deps, title) if warnings == 0 { t.Fatal("expected expiry warnings from the first check") } // The second check comes within the TLS interval of the first, // so it must not warn again. runCheck(t, w, deps) 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 := runFirstCheck(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} w, deps := runFirstCheck(t, cfg, nil) // Save every nameserver the first check found as failed, and add // one that live DNS does not list as having answered. 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) runCheck(t, w, deps) assertNotified(t, deps, "NS Failure: "+testHost, "error") assertNotified(t, deps, "NS Recovery: "+testHost, "success") }