package watcher_test import ( "bytes" "context" "log/slog" "reflect" "strings" "testing" "time" "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" ) // TestCancelledCheckSavesNothing runs a check with its context already // cancelled, which is how the rest of a check runs once shutdown cuts it // short. The real resolver drops the DNS lookup without sending a query, // and the real port and TLS checkers fail without connecting. The port // and certificate state the last check saved must stay as it was, and // nothing may be notified. func TestCancelledCheckSavesNothing(t *testing.T) { t.Parallel() cfg := defaultTestConfig(t) cfg.Hostnames = []string{host} // newTestWatcher's watcher has stand-in checkers. This one, on the // same state and notifier, has the real ones. _, deps := newTestWatcher(t, cfg) w := watcher.NewForTest( cfg, deps.state, resolver.NewFromLogger(slog.Default()), portcheck.NewStandalone(), tlscheck.NewStandalone(), deps.notifier, ) // The last check found host at a local address, with both ports // open and a good certificate. const localIP = "127.0.0.1" deps.state.SetHostnameState(host, hostnameState( map[string]map[string][]string{nsA: {"A": {localIP}}}, )) ports := map[string]*state.PortState{ localIP + ":80": {Open: true, Hostnames: []string{host}}, localIP + ":443": {Open: true, Hostnames: []string{host}}, } for key, ps := range ports { deps.state.SetPortState(key, ps) } certKey := localIP + ":443:" + host cert := &state.CertificateState{CommonName: host, Status: "ok"} deps.state.SetCertificateState(certKey, cert) ctx, cancel := context.WithCancel(t.Context()) cancel() w.RunOnce(ctx) for key, want := range ports { got, _ := deps.state.GetPortState(key) if !reflect.DeepEqual(got, want) { t.Errorf("port %s saved as %+v, want %+v", key, got, want) } } got, _ := deps.state.GetCertificateState(certKey) if !reflect.DeepEqual(got, cert) { t.Errorf("certificate saved as %+v, want %+v", got, cert) } notifications := deps.notifier.getNotifications() if len(notifications) != 0 { t.Errorf("sent %v, want no notifications", notifications) } } // newLoggingWatcher returns a watcher for a domain and a hostname, with // the real resolver, that writes what it logs at warning level or above // into the returned buffer. func newLoggingWatcher(t *testing.T) (*watcher.Watcher, *bytes.Buffer) { t.Helper() cfg := defaultTestConfig(t) cfg.Domains = []string{testSmallDomain} cfg.Hostnames = []string{host} w, _ := newTestWatcher(t, cfg) logs := &bytes.Buffer{} w.SetLogger(slog.New(slog.NewJSONHandler( logs, &slog.HandlerOptions{Level: slog.LevelWarn}, ))) return w, logs } // TestLookupCutShortIsNotLogged checks a domain and a hostname, looks // up a nameserver's addresses and follows a CNAME, with the context // cancelled, as shutdown leaves it. The real resolver fails each lookup // without sending a query. Shutdown cutting a lookup short is not a // failure, so nothing may be logged at warning level or above. func TestLookupCutShortIsNotLogged(t *testing.T) { t.Parallel() w, logs := newLoggingWatcher(t) ctx, cancel := context.WithCancel(t.Context()) cancel() w.RunOnce(ctx) w.ResolveNameserverAddresses(ctx, []string{nsA}, nil) w.ResolveCNAMEAddresses(ctx, host, cnameState(), nil) if logs.Len() > 0 { t.Errorf("logged at warning level or above:\n%s", logs) } } // TestLookupOutOfTimeIsLoggedAsError does what // TestLookupCutShortIsNotLogged does, with the context's deadline passed // instead. A lookup that ran out of time did fail, so the domain's NS // lookup, the hostname's lookup, the nameserver's address lookup and the // CNAME's are each logged as an error. func TestLookupOutOfTimeIsLoggedAsError(t *testing.T) { t.Parallel() w, logs := newLoggingWatcher(t) ctx, cancel := context.WithDeadline(t.Context(), time.Now()) t.Cleanup(cancel) w.RunOnce(ctx) w.ResolveNameserverAddresses(ctx, []string{nsA}, nil) w.ResolveCNAMEAddresses(ctx, host, cnameState(), nil) const want = 4 lines := strings.Count(logs.String(), "\n") errorLines := strings.Count(logs.String(), `"level":"ERROR"`) if lines != want || errorLines != want { t.Errorf("logged:\n%s\nwant %d lines, each at error level", logs, want) } }