3 Commits
Author SHA1 Message Date
clawbot c998c5977b watcher: notify NS query failure and recovery (closes #104)
check / check (push) Successful in 1m37s
LookupAllRecords now returns each nameserver's response, so the
watcher saves its status: ok when it answered, NXDOMAIN and no records
included, and error with the reason when it timed out, answered
SERVFAIL or REFUSED, or could not be reached. A nameserver that starts
failing sends NS Failure and one that answers again sends NS Recovery.
A failing nameserver is left out of the record change and
inconsistency comparisons. The resolver used to report REFUSED and
network errors as an answer with no records; they are now errors. A
lookup cut short by its context now returns an error instead of a
failure of the nameserver it was querying.

Model: opus-5-5
2026-10-01 20:22:06 +00:00
clawbot fcd4f7e2c2 config: stop startup on an invalid DNS or TLS interval (closes #177)
check / check (push) Successful in 1m3s
DNSWATCHER_DNS_INTERVAL and DNSWATCHER_TLS_INTERVAL were parsed with
time.ParseDuration and silently replaced by the default when that failed,
so a value like 5 or 1d gave hourly checks with no hint why, and zero or
negative values were accepted. Both now go through parseInterval, which
returns an error naming the variable and the value, and startup stops the
same way it does for invalid targets. An unset or empty variable still gets
its default from setupViper. The three tests that pinned the old fallback
are replaced. The README says what a valid value looks like.

Model: opus-5-5
2026-10-01 22:19:04 +02:00
clawbot ed0f56f144 metrics: rate limit /metrics per client address before Basic Auth (closes #101)
check / check (push) Successful in 1m18s
/metrics is behind a password, and REPO_POLICIES.md requires rate
limiting on password logins. Each client address may now send it 30
requests a minute, counted by httprate before Basic Auth, so failed
logins use up the allowance and a request over it gets 429 without
the password being checked. The address is the one the existing
trusted-proxy logic in internal/middleware works out, with IPv6
addresses grouped by /64; an IPv4 address a proxy reports in
IPv6-mapped form counts as the plain IPv4 address. A Prometheus
server scraping every 15 seconds sends 4 requests a minute.

Model: opus-5-5
2026-10-01 22:09:14 +02:00
20 changed files with 939 additions and 129 deletions
+38 -10
View File
@@ -71,18 +71,25 @@ rejected.
did on the previous check (additions, removals, value changes). did on the previous check (additions, removals, value changes).
- **NS query failure**: A nameserver that previously responded - **NS query failure**: A nameserver that previously responded
becomes unreachable (timeout, SERVFAIL, REFUSED, network error). becomes unreachable (timeout, SERVFAIL, REFUSED, network error).
This is distinct from "responded with no records." This is distinct from "responded with no records": a nameserver
that answers NXDOMAIN or with no records has responded. The alert
is sent once, on the check where it starts failing. A failing
nameserver gives no records, so it is not reported as a record
change or compared for inconsistency. A nameserver that is already
failing on the first check that sees it is recorded silently.
- **NS recovery**: A previously-unreachable nameserver starts - **NS recovery**: A previously-unreachable nameserver starts
responding again. responding again. Its records are not compared with those from
before it failed, so a change made while it was failing is not
reported as a record change.
- **Inconsistency detected**: Two nameservers return different record - **Inconsistency detected**: Two nameservers return different record
sets for the same hostname and did not already differ on the previous sets for the same hostname and did not already differ on the previous
check. Every pair of nameservers is compared. The alert is sent once check. Every pair of nameservers is compared. The alert is sent once
for each such pair, on the check where they start to disagree, and not for each such pair, on the check where they start to disagree, and not
again while they keep disagreeing, including after a restart. A again while they keep disagreeing, including after a restart. A
nameserver that was not in the previous check (newly added, or back nameserver that was not in the previous check (newly added, or back
after dropping out) and answers differently is reported on the check after dropping out), or failed on it, and answers differently is
where it appears. If a pair agrees again and later disagrees, the reported on the check where it answers. If a pair agrees again and
alert is sent again. later disagrees, the alert is sent again.
### TCP Port Monitoring ### TCP Port Monitoring
@@ -267,7 +274,7 @@ internal/
logger/logger.go slog structured logging (TTY detection) logger/logger.go slog structured logging (TTY detection)
healthcheck/healthcheck.go Health check service healthcheck/healthcheck.go Health check service
middleware/middleware.go HTTP middleware (logging, CORS, security middleware/middleware.go HTTP middleware (logging, CORS, security
headers, metrics auth) headers, metrics auth and rate limit)
handlers/handlers.go HTTP request handlers handlers/handlers.go HTTP request handlers
server/ server/
server.go HTTP server lifecycle server.go HTTP server lifecycle
@@ -320,8 +327,8 @@ the following precedence (highest to lowest):
| `DNSWATCHER_SLACK_WEBHOOK` | Slack incoming webhook URL | `""` | | `DNSWATCHER_SLACK_WEBHOOK` | Slack incoming webhook URL | `""` |
| `DNSWATCHER_MATTERMOST_WEBHOOK` | Mattermost incoming webhook URL | `""` | | `DNSWATCHER_MATTERMOST_WEBHOOK` | Mattermost incoming webhook URL | `""` |
| `DNSWATCHER_NTFY_TOPIC` | ntfy topic URL | `""` | | `DNSWATCHER_NTFY_TOPIC` | ntfy topic URL | `""` |
| `DNSWATCHER_DNS_INTERVAL` | DNS check interval | `1h` | | `DNSWATCHER_DNS_INTERVAL` | DNS check interval, a positive duration such as `30m`; empty means the default, anything else stops startup | `1h` |
| `DNSWATCHER_TLS_INTERVAL` | TLS check interval | `12h` | | `DNSWATCHER_TLS_INTERVAL` | TLS check interval, a positive duration such as `6h`; empty means the default, anything else stops startup | `12h` |
| `DNSWATCHER_TLS_EXPIRY_WARNING` | Days before expiry to warn | `7` | | `DNSWATCHER_TLS_EXPIRY_WARNING` | Days before expiry to warn | `7` |
| `DNSWATCHER_SENTRY_DSN` | Sentry DSN for error reporting | `""` | | `DNSWATCHER_SENTRY_DSN` | Sentry DSN for error reporting | `""` |
| `DNSWATCHER_MAINTENANCE_MODE` | Enable maintenance mode | `false` | | `DNSWATCHER_MAINTENANCE_MODE` | Enable maintenance mode | `false` |
@@ -335,6 +342,23 @@ is a misconfiguration, so dnswatcher fails fast with a clear error message
rather than running silently. Set `DNSWATCHER_TARGETS` to a comma-separated rather than running silently. Set `DNSWATCHER_TARGETS` to a comma-separated
list of DNS names before starting. list of DNS names before starting.
**`/metrics` is rate limited.** Each client address may send it 30 requests a
minute, failed logins included; beyond that it answers `429 Too Many Requests`
without checking the password. A Prometheus server scraping every 15 seconds
sends 4 a minute. IPv6 addresses in one /64 count as one client. When the
request comes from a private or loopback address, such as a reverse proxy's,
the client address is taken from the `X-Real-IP` or `X-Forwarded-For` header
the proxy sets; a proxy that sets neither makes all its clients share one
allowance.
**`DNSWATCHER_DNS_INTERVAL` and `DNSWATCHER_TLS_INTERVAL`** take a positive
duration: a number followed by a unit such as `s`, `m` or `h`, for example
`90s`, `30m`, `1h` or `1h30m`. There is no unit for days; write `24h`. An
unset or empty variable (`DNSWATCHER_DNS_INTERVAL=`) means the default. If
either is set to anything else, including a bare number or a zero or negative
duration, dnswatcher refuses to start with an error naming the variable and
the value.
### Example `.env` ### Example `.env`
```sh ```sh
@@ -442,9 +466,13 @@ The `status` field for each per-nameserver entry and certificate entry
tracks reachability: tracks reachability:
| Status | Meaning | | Status | Meaning |
|-------------|-------------------------------------------------| |-------------|------------------------------------------------------------|
| `ok` | Query succeeded, records are current | | `ok` | Query succeeded, records are current |
| `error` | Query failed (timeout, SERVFAIL, network error) | | `error` | Query failed (timeout, SERVFAIL, REFUSED, network error) |
A nameserver that answers NXDOMAIN or with no records has status `ok` and
empty `records`. A nameserver whose query failed has status `error`, empty
`records`, and the reason in `error`.
--- ---
+7 -7
View File
@@ -15,11 +15,16 @@ on the 1.0 milestone: https://git.eeqj.de/sneak/dnswatcher/milestone/7
# Next Step # Next Step
NS failure and NS recovery notifications: nameserver IP address changes: https://git.eeqj.de/sneak/dnswatcher/issues/105
https://git.eeqj.de/sneak/dnswatcher/issues/104
# Completed Steps # Completed Steps
- 2026-10-01: a nameserver that does not answer is saved as `error` with the
reason, and NS failure and NS recovery are notified (closes #104).
- 2026-10-01: a `DNSWATCHER_DNS_INTERVAL` or `DNSWATCHER_TLS_INTERVAL` that is
not a positive duration stops startup; empty means the default (closes #177).
- 2026-10-01: `/metrics` allows each client address 30 requests a minute,
counted before Basic Auth, and answers 429 beyond that (closes #101).
- 2026-10-01: the image built by `make docker` reports the `git describe` - 2026-10-01: the image built by `make docker` reports the `git describe`
version, not `dev`, and the startup log now shows it (closes #109). version, not `dev`, and the startup log now shows it (closes #109).
- 2026-10-01: two notify shutdown tests always release the delivery they hold, - 2026-10-01: two notify shutdown tests always release the delivery they hold,
@@ -92,13 +97,8 @@ https://git.eeqj.de/sneak/dnswatcher/issues/104
# Future Steps # Future Steps
- nameserver IP address changes: https://git.eeqj.de/sneak/dnswatcher/issues/105
- `DNSWATCHER_SENTRY_DSN` does nothing: - `DNSWATCHER_SENTRY_DSN` does nothing:
https://git.eeqj.de/sneak/dnswatcher/issues/107 https://git.eeqj.de/sneak/dnswatcher/issues/107
- invalid DNS or TLS interval silently replaced by the default:
https://git.eeqj.de/sneak/dnswatcher/issues/177
- rate limit on `/metrics` Basic Auth:
https://git.eeqj.de/sneak/dnswatcher/issues/101
- trial run of the finished image: - trial run of the finished image:
https://git.eeqj.de/sneak/dnswatcher/issues/149 https://git.eeqj.de/sneak/dnswatcher/issues/149
- 1.0 readiness: run it with a real config and read the logs: - 1.0 readiness: run it with a real config and read the logs:
+3
View File
@@ -6,6 +6,7 @@ require (
github.com/99designs/basicauth-go v0.0.0-20230316000542-bf6f9cbbf0f8 github.com/99designs/basicauth-go v0.0.0-20230316000542-bf6f9cbbf0f8
github.com/go-chi/chi/v5 v5.2.5 github.com/go-chi/chi/v5 v5.2.5
github.com/go-chi/cors v1.2.2 github.com/go-chi/cors v1.2.2
github.com/go-chi/httprate v0.16.0
github.com/joho/godotenv v1.5.1 github.com/joho/godotenv v1.5.1
github.com/miekg/dns v1.1.72 github.com/miekg/dns v1.1.72
github.com/prometheus/client_golang v1.23.2 github.com/prometheus/client_golang v1.23.2
@@ -22,6 +23,7 @@ require (
github.com/davecgh/go-spew v1.1.1 // indirect github.com/davecgh/go-spew v1.1.1 // indirect
github.com/fsnotify/fsnotify v1.9.0 // indirect github.com/fsnotify/fsnotify v1.9.0 // indirect
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
github.com/klauspost/cpuid/v2 v2.2.10 // indirect
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
github.com/pelletier/go-toml/v2 v2.2.4 // indirect github.com/pelletier/go-toml/v2 v2.2.4 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect
@@ -34,6 +36,7 @@ require (
github.com/spf13/cast v1.10.0 // indirect github.com/spf13/cast v1.10.0 // indirect
github.com/spf13/pflag v1.0.10 // indirect github.com/spf13/pflag v1.0.10 // indirect
github.com/subosito/gotenv v1.6.0 // indirect github.com/subosito/gotenv v1.6.0 // indirect
github.com/zeebo/xxh3 v1.0.2 // indirect
go.uber.org/dig v1.19.0 // indirect go.uber.org/dig v1.19.0 // indirect
go.uber.org/multierr v1.10.0 // indirect go.uber.org/multierr v1.10.0 // indirect
go.uber.org/zap v1.26.0 // indirect go.uber.org/zap v1.26.0 // indirect
+8
View File
@@ -14,6 +14,8 @@ github.com/go-chi/chi/v5 v5.2.5 h1:Eg4myHZBjyvJmAFjFvWgrqDTXFyOzjj7YIm3L3mu6Ug=
github.com/go-chi/chi/v5 v5.2.5/go.mod h1:X7Gx4mteadT3eDOMTsXzmI4/rwUpOwBHLpAfupzFJP0= github.com/go-chi/chi/v5 v5.2.5/go.mod h1:X7Gx4mteadT3eDOMTsXzmI4/rwUpOwBHLpAfupzFJP0=
github.com/go-chi/cors v1.2.2 h1:Jmey33TE+b+rB7fT8MUy1u0I4L+NARQlK6LhzKPSyQE= github.com/go-chi/cors v1.2.2 h1:Jmey33TE+b+rB7fT8MUy1u0I4L+NARQlK6LhzKPSyQE=
github.com/go-chi/cors v1.2.2/go.mod h1:sSbTewc+6wYHBBCW7ytsFSn836hqM7JxpglAy2Vzc58= github.com/go-chi/cors v1.2.2/go.mod h1:sSbTewc+6wYHBBCW7ytsFSn836hqM7JxpglAy2Vzc58=
github.com/go-chi/httprate v0.16.0 h1:8V5DH9j6pSK6UQoBsTpvMyFxycqaKEIToyPKzHJjUa8=
github.com/go-chi/httprate v0.16.0/go.mod h1:A8lo+qRhk+s9LiuP5saS7XCGDXRXMcrueq0NfIuCa/I=
github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs= github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs=
github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
@@ -22,6 +24,8 @@ github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0=
github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4= github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4=
github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo=
github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ= github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ=
github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
@@ -62,6 +66,10 @@ github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8= github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8=
github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU= github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU=
github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ=
github.com/zeebo/assert v1.3.0/go.mod h1:Pq9JiuJQpG8JLJdtkwrJESF0Foym2/D9XMU5ciN/wJ0=
github.com/zeebo/xxh3 v1.0.2 h1:xZmwmqxHZA8AI603jOQ0tMqmBr9lPeFwGg6d+xy9DC0=
github.com/zeebo/xxh3 v1.0.2/go.mod h1:5NWz9Sef7zIDm2JHfFlcQvNekmcEl9ekUZQQKCYaDcA=
go.uber.org/dig v1.19.0 h1:BACLhebsYdpQ7IROQ1AGPjrXcP5dF80U3gKoFzbaq/4= go.uber.org/dig v1.19.0 h1:BACLhebsYdpQ7IROQ1AGPjrXcP5dF80U3gKoFzbaq/4=
go.uber.org/dig v1.19.0/go.mod h1:Us0rSJiThwCv2GteUN0Q7OKvU7n5J4dxZ9JKUXozFdE= go.uber.org/dig v1.19.0/go.mod h1:Us0rSJiThwCv2GteUN0Q7OKvU7n5J4dxZ9JKUXozFdE=
go.uber.org/fx v1.24.0 h1:wE8mruvpg2kiiL1Vqd0CC+tr0/24XIB10Iwp2lLWzkg= go.uber.org/fx v1.24.0 h1:wE8mruvpg2kiiL1Vqd0CC+tr0/24XIB10Iwp2lLWzkg=
+27 -8
View File
@@ -28,6 +28,13 @@ var ErrNoTargets = errors.New(
"no monitoring targets configured: set DNSWATCHER_TARGETS environment variable", "no monitoring targets configured: set DNSWATCHER_TARGETS environment variable",
) )
// ErrInvalidInterval is returned when DNSWATCHER_DNS_INTERVAL or
// DNSWATCHER_TLS_INTERVAL is set but is not a positive duration. An empty
// value counts as unset and means the default.
var ErrInvalidInterval = errors.New(
"interval must be a positive duration such as 30m or 1h",
)
// Params contains dependencies for Config. // Params contains dependencies for Config.
type Params struct { type Params struct {
fx.In fx.In
@@ -125,18 +132,14 @@ func buildConfig(
} }
} }
dnsInterval, err := time.ParseDuration( dnsInterval, err := parseInterval("DNS_INTERVAL")
viper.GetString("DNS_INTERVAL"),
)
if err != nil { if err != nil {
dnsInterval = defaultDNSInterval return nil, err
} }
tlsInterval, err := time.ParseDuration( tlsInterval, err := parseInterval("TLS_INTERVAL")
viper.GetString("TLS_INTERVAL"),
)
if err != nil { if err != nil {
tlsInterval = defaultTLSInterval return nil, err
} }
domains, hostnames, err := parseAndValidateTargets() domains, hostnames, err := parseAndValidateTargets()
@@ -168,6 +171,22 @@ func buildConfig(
return cfg, nil return cfg, nil
} }
// parseInterval reads the DNSWATCHER_-prefixed setting key as a duration. A
// value that does not parse, or is zero or negative, is an error naming the
// variable and the value; an unset variable has its default from setupViper.
func parseInterval(key string) (time.Duration, error) {
value := viper.GetString(key)
interval, err := time.ParseDuration(value)
if err != nil || interval <= 0 {
return 0, fmt.Errorf(
"invalid DNSWATCHER_%s %q: %w", key, value, ErrInvalidInterval,
)
}
return interval, nil
}
func parseAndValidateTargets() ([]string, []string, error) { func parseAndValidateTargets() ([]string, []string, error) {
domains, hostnames, err := ClassifyTargets( domains, hostnames, err := ClassifyTargets(
parseCSV(viper.GetString("TARGETS")), parseCSV(viper.GetString("TARGETS")),
+26 -18
View File
@@ -1,6 +1,7 @@
package config_test package config_test
import ( import (
"strconv"
"testing" "testing"
"time" "time"
@@ -113,33 +114,40 @@ func TestNew_OnlyEmptyCSVSegments(t *testing.T) {
assert.ErrorIs(t, err, config.ErrNoTargets) assert.ErrorIs(t, err, config.ErrNoTargets)
} }
func TestNew_InvalidDNSInterval_FallsBackToDefault(t *testing.T) { // TestNew_InvalidIntervalStopsStartup checks values that must stop startup;
viper.Reset() // TestNew_DefaultValues and TestNew_EmptyIntervalMeansDefault check that an
t.Setenv("DNSWATCHER_TARGETS", "example.com") // unset or empty interval means the default.
t.Setenv("DNSWATCHER_DNS_INTERVAL", "banana") func TestNew_InvalidIntervalStopsStartup(t *testing.T) {
variables := []string{"DNSWATCHER_DNS_INTERVAL", "DNSWATCHER_TLS_INTERVAL"}
cfg, err := config.New(nil, newTestParams(t)) values := []string{
require.NoError(t, err) "banana", // not a duration
assert.Equal(t, time.Hour, cfg.DNSInterval, "5", // no unit
"invalid DNS interval should fall back to 1h default") "1d", // days are not a unit time.ParseDuration knows
"0", // zero
"-1h", // negative
} }
func TestNew_InvalidTLSInterval_FallsBackToDefault(t *testing.T) { for _, variable := range variables {
for _, value := range values {
t.Run(variable+"="+value, func(t *testing.T) {
viper.Reset() viper.Reset()
t.Setenv("DNSWATCHER_TARGETS", "example.com") t.Setenv("DNSWATCHER_TARGETS", "example.com")
t.Setenv("DNSWATCHER_TLS_INTERVAL", "notaduration") t.Setenv(variable, value)
cfg, err := config.New(nil, newTestParams(t)) _, err := config.New(nil, newTestParams(t))
require.NoError(t, err) require.ErrorIs(t, err, config.ErrInvalidInterval)
assert.Equal(t, 12*time.Hour, cfg.TLSInterval, require.ErrorContains(t, err, variable)
"invalid TLS interval should fall back to 12h default") require.ErrorContains(t, err, strconv.Quote(value))
})
}
}
} }
func TestNew_BothIntervalsInvalid(t *testing.T) { func TestNew_EmptyIntervalMeansDefault(t *testing.T) {
viper.Reset() viper.Reset()
t.Setenv("DNSWATCHER_TARGETS", "example.com") t.Setenv("DNSWATCHER_TARGETS", "example.com")
t.Setenv("DNSWATCHER_DNS_INTERVAL", "xyz") t.Setenv("DNSWATCHER_DNS_INTERVAL", "")
t.Setenv("DNSWATCHER_TLS_INTERVAL", "abc") t.Setenv("DNSWATCHER_TLS_INTERVAL", "")
cfg, err := config.New(nil, newTestParams(t)) cfg, err := config.New(nil, newTestParams(t))
require.NoError(t, err) require.NoError(t, err)
+10
View File
@@ -0,0 +1,10 @@
package middleware
import "time"
// The /metrics rate limit, exported so the tests can count requests
// against it.
const (
MetricsRequestLimit = metricsRequestLimit
MetricsRequestWindow time.Duration = metricsRequestWindow
)
+39
View File
@@ -5,12 +5,14 @@ import (
"log/slog" "log/slog"
"net" "net"
"net/http" "net/http"
"net/netip"
"strings" "strings"
"time" "time"
"github.com/99designs/basicauth-go" "github.com/99designs/basicauth-go"
"github.com/go-chi/chi/v5/middleware" "github.com/go-chi/chi/v5/middleware"
"github.com/go-chi/cors" "github.com/go-chi/cors"
"github.com/go-chi/httprate"
"go.uber.org/fx" "go.uber.org/fx"
"sneak.berlin/go/dnswatcher/internal/config" "sneak.berlin/go/dnswatcher/internal/config"
@@ -21,6 +23,17 @@ import (
// corsMaxAge is the maximum age for CORS preflight responses. // corsMaxAge is the maximum age for CORS preflight responses.
const corsMaxAge = 300 const corsMaxAge = 300
// Rate limit for /metrics: each client address may send
// metricsRequestLimit requests per metricsRequestWindow. Every request
// counts, so password guessing gets at most 30 tries a minute per
// address. One Prometheus server scraping every 15 seconds sends 4
// requests a minute, and two scraping every 5 seconds from one address
// send 24, so normal scraping stays under the limit.
const (
metricsRequestLimit = 30
metricsRequestWindow = time.Minute
)
// Security response header values applied to every response. // Security response header values applied to every response.
// //
// The CSP is as strict as the dashboard allows: the template ships no // The CSP is as strict as the dashboard allows: the template ships no
@@ -268,6 +281,32 @@ func (m *Middleware) SecurityHeaders() func(http.Handler) http.Handler {
} }
} }
// MetricsRateLimit returns middleware for /metrics that answers 429
// Too Many Requests to a client address over the rate limit. The
// address is the one realIP works out, so a client that is not a
// trusted proxy cannot get a fresh allowance by sending its own
// X-Real-IP or X-Forwarded-For. CanonicalizeIP counts all IPv6
// addresses in one /64 as one client, since a client usually holds a
// whole /64. An IPv4 address a proxy reports in IPv6-mapped form
// (::ffff:203.0.113.1) is turned back into plain IPv4 first, as every
// such address is in the same /64.
func (m *Middleware) MetricsRateLimit() func(http.Handler) http.Handler {
return httprate.LimitBy(
metricsRequestLimit,
metricsRequestWindow,
func(request *http.Request) (string, error) {
ip := realIP(request)
addr, err := netip.ParseAddr(ip)
if err == nil {
ip = addr.Unmap().String()
}
return httprate.CanonicalizeIP(ip), nil
},
)
}
// MetricsAuth returns basic auth middleware for /metrics. // MetricsAuth returns basic auth middleware for /metrics.
func (m *Middleware) MetricsAuth() func(http.Handler) http.Handler { func (m *Middleware) MetricsAuth() func(http.Handler) http.Handler {
if m.params.Config.MetricsUsername == "" { if m.params.Config.MetricsUsername == "" {
+142
View File
@@ -5,6 +5,7 @@ import (
"net/http/httptest" "net/http/httptest"
"strings" "strings"
"testing" "testing"
"time"
"github.com/go-chi/chi/v5" "github.com/go-chi/chi/v5"
"go.uber.org/fx/fxtest" "go.uber.org/fx/fxtest"
@@ -340,3 +341,144 @@ func TestDashboardRendersWithSecurityHeaders(t *testing.T) {
t.Errorf("CSP would block %q: %q", stylesheetPath, csp) t.Errorf("CSP would block %q: %q", stylesheetPath, csp)
} }
} }
// Addresses for the rate limit tests: a client connecting directly, a
// trusted proxy, and a client behind that proxy as its X-Real-IP
// header names it.
const (
directClient = "198.51.100.1:4000"
trustedProxy = "10.0.0.1:4000"
proxiedClient = "203.0.113.1"
)
// statusFrom sends a GET through handler as if from remoteAddr, with
// an X-Real-IP header when xRealIP is not empty, and returns the
// response status.
func statusFrom(
t *testing.T,
handler http.Handler,
remoteAddr string,
xRealIP string,
) int {
t.Helper()
req := httptest.NewRequestWithContext(
t.Context(), http.MethodGet, "/metrics", nil,
)
req.RemoteAddr = remoteAddr
if xRealIP != "" {
req.Header.Set("X-Real-IP", xRealIP)
}
rec := httptest.NewRecorder()
handler.ServeHTTP(rec, req)
return rec.Code
}
// TestMetricsRateLimitAllowsScraping checks that one address can send,
// within one window, what two Prometheus servers scraping every 5
// seconds send in that time, without being turned away.
func TestMetricsRateLimitAllowsScraping(t *testing.T) {
t.Parallel()
const scrapeInterval = 5 * time.Second
scrapes := 2 * int(middleware.MetricsRequestWindow/scrapeInterval)
limited := newTestMiddleware(t).MetricsRateLimit()(okHandler())
for i := range scrapes {
got := statusFrom(t, limited, directClient, "")
if got != http.StatusOK {
t.Fatalf(
"scrape %d of %d: status = %d, want 200",
i+1, scrapes, got,
)
}
}
}
// TestMetricsRateLimitKeysOnClientAddress checks which requests share
// an allowance. Each case uses up the allowance of one client, then
// sends one more request.
func TestMetricsRateLimitKeysOnClientAddress(t *testing.T) {
t.Parallel()
tests := []struct {
name string
usedRemoteAddr string
usedXRealIP string
nextRemoteAddr string
nextXRealIP string
want int
}{
{
"same address",
directClient, "",
directClient, "",
http.StatusTooManyRequests,
},
{
"another address",
directClient, "",
"198.51.100.2:4000", "",
http.StatusOK,
},
{
"own X-Real-IP from an untrusted address",
directClient, "",
directClient, "203.0.113.9",
http.StatusTooManyRequests,
},
{
"same client behind the proxy",
trustedProxy, proxiedClient,
trustedProxy, proxiedClient,
http.StatusTooManyRequests,
},
{
"another client behind the proxy",
trustedProxy, proxiedClient,
trustedProxy, "203.0.113.2",
http.StatusOK,
},
{
"another client behind the proxy, IPv6-mapped",
trustedProxy, "::ffff:203.0.113.1",
trustedProxy, "::ffff:203.0.113.2",
http.StatusOK,
},
{
"same IPv6 /64",
"[2001:db8::1]:4000", "",
"[2001:db8::2]:4000", "",
http.StatusTooManyRequests,
},
{
"another IPv6 /64",
"[2001:db8::1]:4000", "",
"[2001:db8:0:1::1]:4000", "",
http.StatusOK,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
limited := newTestMiddleware(t).MetricsRateLimit()(okHandler())
for range middleware.MetricsRequestLimit {
statusFrom(t, limited, tt.usedRemoteAddr, tt.usedXRealIP)
}
got := statusFrom(
t, limited, tt.nextRemoteAddr, tt.nextXRealIP,
)
if got != tt.want {
t.Errorf("status = %d, want %d", got, tt.want)
}
})
}
}
+14 -1
View File
@@ -1,8 +1,21 @@
package resolver package resolver
import "github.com/miekg/dns" import (
"context"
"github.com/miekg/dns"
)
// ExtractRecordValue exports extractRecordValue for testing. // ExtractRecordValue exports extractRecordValue for testing.
func ExtractRecordValue(rr dns.RR) string { func ExtractRecordValue(rr dns.RR) string {
return extractRecordValue(rr) return extractRecordValue(rr)
} }
// QueryEachNS exports queryEachNS for testing.
func (r *Resolver) QueryEachNS(
ctx context.Context,
nameservers []string,
hostname string,
) (map[string]*NameserverResponse, error) {
return r.queryEachNS(ctx, nameservers, hostname)
}
+22 -14
View File
@@ -504,7 +504,9 @@ func (r *Resolver) queryAllTypes(
type queryState struct { type queryState struct {
gotNXDomain bool gotNXDomain bool
gotSERVFAIL bool gotSERVFAIL bool
gotRefused bool
gotTimeout bool gotTimeout bool
netErr error
hasRecords bool hasRecords bool
} }
@@ -542,8 +544,13 @@ func (r *Resolver) querySingleType(
) { ) {
msg, err := r.queryDNS(ctx, nsIP, hostname, qtype) msg, err := r.queryDNS(ctx, nsIP, hostname, qtype)
if err != nil { if err != nil {
if isTimeout(err) { switch {
case isTimeout(err):
state.gotTimeout = true state.gotTimeout = true
case errors.Is(err, ErrRefused):
state.gotRefused = true
default:
state.netErr = err
} }
return return
@@ -603,6 +610,12 @@ func classifyResponse(resp *NameserverResponse, state queryState) {
case state.gotSERVFAIL && !state.hasRecords: case state.gotSERVFAIL && !state.hasRecords:
resp.Status = StatusError resp.Status = StatusError
resp.Error = "server returned SERVFAIL" resp.Error = "server returned SERVFAIL"
case state.gotRefused && !state.hasRecords:
resp.Status = StatusError
resp.Error = "server returned REFUSED"
case state.netErr != nil && !state.hasRecords:
resp.Status = StatusError
resp.Error = "network error: " + state.netErr.Error()
case !state.hasRecords && !state.gotNXDomain: case !state.hasRecords && !state.gotNXDomain:
resp.Status = StatusNoData resp.Status = StatusNoData
} }
@@ -682,11 +695,14 @@ func (r *Resolver) queryEachNS(
results := make(map[string]*NameserverResponse) results := make(map[string]*NameserverResponse)
for _, ns := range nameservers { for _, ns := range nameservers {
resp, err := r.QueryNameserver(ctx, ns, hostname)
// A query the context cut short says nothing about the
// nameserver, so it must not be returned as its failure.
if checkCtx(ctx) != nil { if checkCtx(ctx) != nil {
return nil, ErrContextCanceled return nil, ErrContextCanceled
} }
resp, err := r.QueryNameserver(ctx, ns, hostname)
if err != nil { if err != nil {
results[ns] = &NameserverResponse{ results[ns] = &NameserverResponse{
Nameserver: ns, Nameserver: ns,
@@ -714,21 +730,13 @@ func (r *Resolver) LookupNS(
// LookupAllRecords performs iterative resolution to find all DNS // LookupAllRecords performs iterative resolution to find all DNS
// records for the given hostname, keyed by authoritative nameserver. // records for the given hostname, keyed by authoritative nameserver.
// Each nameserver's response carries its status and error with its
// records.
func (r *Resolver) LookupAllRecords( func (r *Resolver) LookupAllRecords(
ctx context.Context, ctx context.Context,
hostname string, hostname string,
) (map[string]map[string][]string, error) { ) (map[string]*NameserverResponse, error) {
results, err := r.QueryAllNameservers(ctx, hostname) return r.QueryAllNameservers(ctx, hostname)
if err != nil {
return nil, err
}
out := make(map[string]map[string][]string, len(results))
for ns, resp := range results {
out[ns] = resp.Records
}
return out, nil
} }
// ResolveIPAddresses resolves a hostname to all IPv4 and IPv6 // ResolveIPAddresses resolves a hostname to all IPv4 and IPv6
+65 -1
View File
@@ -2,6 +2,7 @@ package resolver_test
import ( import (
"context" "context"
"fmt"
"log/slog" "log/slog"
"net" "net"
"os" "os"
@@ -13,6 +14,7 @@ import (
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"sneak.berlin/go/dnswatcher/internal/livednstest"
"sneak.berlin/go/dnswatcher/internal/resolver" "sneak.berlin/go/dnswatcher/internal/resolver"
) )
@@ -231,6 +233,45 @@ func TestQueryNameserver_NXDomain(t *testing.T) {
assert.Equal(t, resolver.StatusNXDomain, resp.Status) assert.Equal(t, resolver.StatusNXDomain, resp.Status)
} }
// TestQueryNameserver_Refused asks a google.com nameserver about
// cloudflare.com, a zone it does not serve, which it refuses. Refusing
// is a failure to answer, not an answer with no records.
func TestQueryNameserver_Refused(t *testing.T) {
t.Parallel()
r := newTestResolver(t)
ns := findOneNSForDomain(t, r, "google.com")
var resp *resolver.NameserverResponse
livednstest.Retry(
t,
"QueryNameserver("+ns+", cloudflare.com)",
func(ctx context.Context) error {
var err error
resp, err = r.QueryNameserver(ctx, ns, "cloudflare.com")
if err != nil {
return err
}
// A timeout or a network error is no reply at all.
if resp.Status == resolver.StatusTimeout ||
strings.HasPrefix(resp.Error, "network error") {
return fmt.Errorf(
"%w: %s: %s",
livednstest.ErrNoAnswer, ns, resp.Error,
)
}
return nil
},
)
assert.Equal(t, resolver.StatusError, resp.Status)
assert.Equal(t, "server returned REFUSED", resp.Error)
}
func TestQueryNameserver_RecordsSorted(t *testing.T) { func TestQueryNameserver_RecordsSorted(t *testing.T) {
t.Parallel() t.Parallel()
@@ -518,6 +559,29 @@ func TestQueryAllNameservers_ContextCanceled(t *testing.T) {
assert.Error(t, err) assert.Error(t, err)
} }
// TestQueryEachNS_CanceledDuringQuery cancels the context while a
// nameserver is being queried, as shutdown does. A lookup cut short
// says nothing about the nameserver, so it must return an error, not a
// failed response for it.
func TestQueryEachNS_CanceledDuringQuery(t *testing.T) {
t.Parallel()
r := newTestResolver(t)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
// Finding the nameserver's address alone starts at the root
// servers and takes several round trips, so a cancel a few
// milliseconds in lands during the query.
time.AfterFunc(5*time.Millisecond, cancel)
results, err := r.QueryEachNS(
ctx, []string{"ns1.google.com."}, "google.com",
)
require.ErrorIs(t, err, resolver.ErrContextCanceled)
assert.Nil(t, results)
}
// ---------------------------------------------------------------- // ----------------------------------------------------------------
// Timeout tests // Timeout tests
// ---------------------------------------------------------------- // ----------------------------------------------------------------
@@ -530,7 +594,7 @@ func TestQueryNameserverIP_Timeout(t *testing.T) {
// Nothing answers at 192.0.2.1, a documentation address. The // Nothing answers at 192.0.2.1, a documentation address. The
// resolver tries each query twice, and the first try gives up // resolver tries each query twice, and the first try gives up
// after two seconds. A deadline that ends during the first try // after two seconds. A deadline that ends during the first try
// makes the status vary from run to run between nodata and // makes the status vary from run to run between error and
// timeout, so the deadline must outlast the first try. // timeout, so the deadline must outlast the first try.
ctx, cancel := context.WithTimeout( ctx, cancel := context.WithTimeout(
context.Background(), 3*time.Second, context.Background(), 3*time.Second,
+4 -1
View File
@@ -64,9 +64,12 @@ func (s *Server) SetupRoutes() {
// Prometheus scraper is not a browser. It is mounted rather than // Prometheus scraper is not a browser. It is mounted rather than
// added with Get so that every method on /metrics, OPTIONS // added with Get so that every method on /metrics, OPTIONS
// included, ends here instead of falling through to the public // included, ends here instead of falling through to the public
// router and its CORS. // router and its CORS. The rate limit comes before Basic Auth, so
// failed logins count against it and a request over the limit
// never reaches the password check.
if s.params.Config.MetricsUsername != "" { if s.params.Config.MetricsUsername != "" {
metrics := chi.NewRouter() metrics := chi.NewRouter()
metrics.Use(s.mw.MetricsRateLimit())
metrics.Use(s.mw.MetricsAuth()) metrics.Use(s.mw.MetricsAuth())
metrics.Get("/", promhttp.Handler().ServeHTTP) metrics.Get("/", promhttp.Handler().ServeHTTP)
s.router.Mount("/metrics", metrics) s.router.Mount("/metrics", metrics)
+69
View File
@@ -219,3 +219,72 @@ func TestPreflightAllowsOnlyWhatPublicRoutesServe(t *testing.T) {
} }
} }
} }
// metricsRequest builds a GET for /metrics from remoteAddr that logs
// in with the given password.
func metricsRequest(
t *testing.T,
remoteAddr string,
password string,
) *http.Request {
t.Helper()
req := httptest.NewRequestWithContext(
t.Context(), http.MethodGet, "/metrics", nil,
)
req.RemoteAddr = remoteAddr
req.SetBasicAuth(metricsUsername, password)
return req
}
// TestMetricsRateLimitComesBeforeAuth checks that failed logins to
// /metrics count against the rate limit; that once an address is over
// it, even the right password gets 429, with the same body as a wrong
// one; and that another address still gets in.
func TestMetricsRateLimitComesBeforeAuth(t *testing.T) {
viper.Reset()
t.Setenv("DNSWATCHER_TARGETS", "example.com")
t.Setenv("DNSWATCHER_METRICS_USERNAME", metricsUsername)
t.Setenv("DNSWATCHER_METRICS_PASSWORD", metricsPassword)
const (
guesser = "198.51.100.1:4000"
other = "198.51.100.2:4000"
// Far more guesses than the rate limit allows.
maxGuesses = 1000
)
srv := routedServer(t)
var guess *httptest.ResponseRecorder
for range maxGuesses {
guess = serve(srv, metricsRequest(t, guesser, "wrong"))
if guess.Code != http.StatusUnauthorized {
break
}
}
if guess.Code != http.StatusTooManyRequests {
t.Fatalf("wrong password: status = %d, want 429", guess.Code)
}
right := serve(srv, metricsRequest(t, guesser, metricsPassword))
if right.Code != http.StatusTooManyRequests {
t.Errorf("right password: status = %d, want 429", right.Code)
}
if right.Body.String() != guess.Body.String() {
t.Errorf(
"429 body with right password = %q, with wrong one = %q",
right.Body.String(), guess.Body.String(),
)
}
rec := serve(srv, metricsRequest(t, other, metricsPassword))
if rec.Code != http.StatusOK {
t.Errorf("another address: status = %d, want 200", rec.Code)
}
}
+11 -4
View File
@@ -6,6 +6,7 @@ import (
"time" "time"
"sneak.berlin/go/dnswatcher/internal/config" "sneak.berlin/go/dnswatcher/internal/config"
"sneak.berlin/go/dnswatcher/internal/resolver"
"sneak.berlin/go/dnswatcher/internal/state" "sneak.berlin/go/dnswatcher/internal/state"
) )
@@ -33,8 +34,7 @@ func NewForTest(
// NewlyDisagreeingPairs exports newlyDisagreeingPairs for testing. // NewlyDisagreeingPairs exports newlyDisagreeingPairs for testing.
func NewlyDisagreeingPairs( func NewlyDisagreeingPairs(
prev *state.HostnameState, prev, current *state.HostnameState,
current map[string]map[string][]string,
) [][2]string { ) [][2]string {
return newlyDisagreeingPairs(prev, current) return newlyDisagreeingPairs(prev, current)
} }
@@ -43,8 +43,15 @@ func NewlyDisagreeingPairs(
func (w *Watcher) DetectHostnameChanges( func (w *Watcher) DetectHostnameChanges(
ctx context.Context, ctx context.Context,
hostname string, hostname string,
prev *state.HostnameState, prev, current *state.HostnameState,
current map[string]map[string][]string,
) { ) {
w.detectHostnameChanges(ctx, hostname, prev, current) w.detectHostnameChanges(ctx, hostname, prev, current)
} }
// BuildHostnameState exports buildHostnameState for testing.
func BuildHostnameState(
results map[string]*resolver.NameserverResponse,
now time.Time,
) *state.HostnameState {
return buildHostnameState(results, now)
}
+7 -4
View File
@@ -105,7 +105,9 @@ func TestNewlyDisagreeingPairs(t *testing.T) {
prev := hostnameState(tt.loaded) prev := hostnameState(tt.loaded)
for i, current := range tt.checks { for i, records := range tt.checks {
current := hostnameState(records)
got := watcher.NewlyDisagreeingPairs(prev, current) got := watcher.NewlyDisagreeingPairs(prev, current)
if !slices.Equal(got, tt.want[i]) { if !slices.Equal(got, tt.want[i]) {
t.Errorf( t.Errorf(
@@ -114,7 +116,7 @@ func TestNewlyDisagreeingPairs(t *testing.T) {
) )
} }
prev = hostnameState(current) prev = current
} }
}) })
} }
@@ -162,8 +164,9 @@ func TestInconsistencyAlert(t *testing.T) {
prev := hostnameState(tt.loaded) prev := hostnameState(tt.loaded)
for range 3 { for range 3 {
w.DetectHostnameChanges(t.Context(), host, prev, disagree) current := hostnameState(disagree)
prev = hostnameState(disagree) w.DetectHostnameChanges(t.Context(), host, prev, current)
prev = current
} }
got := 0 got := 0
+3 -2
View File
@@ -5,6 +5,7 @@ import (
"context" "context"
"sneak.berlin/go/dnswatcher/internal/portcheck" "sneak.berlin/go/dnswatcher/internal/portcheck"
"sneak.berlin/go/dnswatcher/internal/resolver"
"sneak.berlin/go/dnswatcher/internal/tlscheck" "sneak.berlin/go/dnswatcher/internal/tlscheck"
) )
@@ -17,11 +18,11 @@ type DNSResolver interface {
) ([]string, error) ) ([]string, error)
// LookupAllRecords queries all record types for a hostname, // LookupAllRecords queries all record types for a hostname,
// returning results keyed by nameserver then record type. // returning each nameserver's response keyed by nameserver.
LookupAllRecords( LookupAllRecords(
ctx context.Context, ctx context.Context,
hostname string, hostname string,
) (map[string]map[string][]string, error) ) (map[string]*resolver.NameserverResponse, error)
// ResolveIPAddresses resolves a hostname to all IP addresses. // ResolveIPAddresses resolves a hostname to all IP addresses.
ResolveIPAddresses( ResolveIPAddresses(
+333
View File
@@ -0,0 +1,333 @@
package watcher_test
import (
"context"
"fmt"
"log/slog"
"strings"
"testing"
"time"
"sneak.berlin/go/dnswatcher/internal/livednstest"
"sneak.berlin/go/dnswatcher/internal/resolver"
"sneak.berlin/go/dnswatcher/internal/state"
"sneak.berlin/go/dnswatcher/internal/watcher"
)
// answered is what a check saves for a nameserver that answered with
// these records.
func answered(records map[string][]string) *state.NameserverRecordState {
return &state.NameserverRecordState{Records: records, Status: "ok"}
}
// failed is what a check saves for a nameserver that did not answer.
func failed() *state.NameserverRecordState {
return &state.NameserverRecordState{
Records: map[string][]string{},
Status: "error",
Error: "all queries timed out",
}
}
// saved builds the hostname state a check saves.
func saved(
byNameserver map[string]*state.NameserverRecordState,
) *state.HostnameState {
return &state.HostnameState{RecordsByNameserver: byNameserver}
}
// alertCounts counts the hostname alerts sent, by kind.
type alertCounts struct {
failures, recoveries, recordChanges, inconsistencies int
}
// countAlerts runs the hostname change detection from the state loaded
// at startup through each check in turn, and counts the alerts sent.
func countAlerts(
t *testing.T,
loaded *state.HostnameState,
checks []*state.HostnameState,
) alertCounts {
t.Helper()
// The hostname change detection uses only the notifier.
notifier := &mockNotifier{}
w := watcher.NewForTest(nil, nil, nil, nil, nil, notifier)
prev := loaded
for _, current := range checks {
w.DetectHostnameChanges(t.Context(), host, prev, current)
prev = current
}
var got alertCounts
for _, n := range notifier.getNotifications() {
kind, _, _ := strings.Cut(n.Title, ":")
switch kind {
case "NS Failure":
got.failures++
case "NS Recovery":
got.recoveries++
case "Record Change":
got.recordChanges++
case "Inconsistency":
got.inconsistencies++
}
}
return got
}
func TestNSFailureAndRecoveryAlerts(t *testing.T) {
t.Parallel()
records := map[string][]string{"A": {ip1}}
bothAnswer := saved(map[string]*state.NameserverRecordState{
nsA: answered(records), nsB: answered(records),
})
bFails := saved(map[string]*state.NameserverRecordState{
nsA: answered(records), nsB: failed(),
})
onlyA := saved(map[string]*state.NameserverRecordState{
nsA: answered(records),
})
bAnswersNoRecords := saved(map[string]*state.NameserverRecordState{
nsA: answered(records), nsB: answered(map[string][]string{}),
})
bAnswersDifferently := saved(map[string]*state.NameserverRecordState{
nsA: answered(records), nsB: answered(map[string][]string{"A": {ip2}}),
})
// Each case starts from the state loaded at startup and runs the
// checks in order.
tests := []struct {
name string
loaded *state.HostnameState
checks []*state.HostnameState
want alertCounts
}{
{
"failure lasting several checks alerts once",
bothAnswer, []*state.HostnameState{bFails, bFails, bFails},
alertCounts{failures: 1},
},
{
"recovery alerts once",
bFails, []*state.HostnameState{bothAnswer, bothAnswer},
alertCounts{recoveries: 1},
},
{
"failing again after recovering alerts again",
bothAnswer, []*state.HostnameState{bFails, bothAnswer, bFails},
alertCounts{failures: 2, recoveries: 1},
},
{
"nameserver failing when first seen does not alert",
onlyA, []*state.HostnameState{bFails, bFails},
alertCounts{},
},
{
"answer with no records is a record change, not a failure",
bothAnswer, []*state.HostnameState{bAnswersNoRecords},
alertCounts{recordChanges: 1, inconsistencies: 1},
},
{
"recovered nameserver that answers differently disagrees",
bFails, []*state.HostnameState{bAnswersDifferently},
alertCounts{recoveries: 1, inconsistencies: 1},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
got := countAlerts(t, tt.loaded, tt.checks)
if got != tt.want {
t.Errorf("sent %+v, want %+v", got, tt.want)
}
})
}
}
func TestNSFailureAlertNamesHostnameNameserverAndReason(t *testing.T) {
t.Parallel()
records := map[string][]string{"A": {ip1}}
notifier := &mockNotifier{}
w := watcher.NewForTest(nil, nil, nil, nil, nil, notifier)
w.DetectHostnameChanges(
t.Context(), host,
saved(map[string]*state.NameserverRecordState{nsA: answered(records)}),
saved(map[string]*state.NameserverRecordState{nsA: failed()}),
)
notifications := notifier.getNotifications()
if len(notifications) != 1 {
t.Fatalf("sent %v, want one NS Failure", notifications)
}
msg := notifications[0].Message
if !strings.Contains(msg, host) || !strings.Contains(msg, nsA) ||
!strings.Contains(msg, failed().Error) {
t.Errorf(
"message %q does not name %s, %s and the reason",
msg, host, nsA,
)
}
}
// TestNameserverThatNeverAnswers asks a nameserver address where
// nothing answers, 192.0.2.1, and checks what the watcher saves for it.
// The deadline outlasts the resolver's first two-second try, as in the
// resolver's timeout test.
func TestNameserverThatNeverAnswers(t *testing.T) {
t.Parallel()
ctx, cancel := context.WithTimeout(t.Context(), 3*time.Second)
t.Cleanup(cancel)
res := resolver.NewFromLogger(slog.Default())
resp, err := res.QueryNameserverIP(ctx, nsA, "192.0.2.1", host)
if err != nil {
t.Fatal(err)
}
hs := watcher.BuildHostnameState(
map[string]*resolver.NameserverResponse{nsA: resp}, time.Now(),
)
got := hs.RecordsByNameserver[nsA]
if got.Status != "error" || got.Error == "" {
t.Errorf(
"saved status %q, error %q; want status error with a reason",
got.Status, got.Error,
)
}
}
// TestNameserverThatAnswersNXDOMAIN asks a real nameserver about a name
// that does not exist and checks what the watcher saves for it: NXDOMAIN
// is an answer, so the nameserver is saved as ok with no error.
func TestNameserverThatAnswersNXDOMAIN(t *testing.T) {
t.Parallel()
res := resolver.NewFromLogger(slog.Default())
name := "this-surely-does-not-exist-xyz." + testDomain
var (
ns string
resp *resolver.NameserverResponse
)
livednstest.Retry(t, "QueryNameserver("+name+")", func(ctx context.Context) error {
nameservers, err := res.LookupNS(ctx, testDomain)
if err != nil {
return err
}
ns = nameservers[0]
resp, err = res.QueryNameserver(ctx, ns, name)
if err != nil {
return err
}
// A timeout or a failure is no answer to check.
if resp.Status == resolver.StatusTimeout ||
resp.Status == resolver.StatusError {
return fmt.Errorf(
"%w: %s: %s", livednstest.ErrNoAnswer, ns, resp.Error,
)
}
return nil
})
if resp.Status != resolver.StatusNXDomain {
t.Fatalf("%s answered %q for %s, want NXDOMAIN", ns, resp.Status, name)
}
hs := watcher.BuildHostnameState(
map[string]*resolver.NameserverResponse{ns: resp}, time.Now(),
)
got := hs.RecordsByNameserver[ns]
if got.Status != "ok" || got.Error != "" {
t.Errorf(
"saved status %q, error %q; want status ok with no error",
got.Status, got.Error,
)
}
}
// TestNameserverThatRefuses asks a google.com nameserver about
// cloudflare.com, a zone it does not serve, which it refuses, and checks
// what the watcher saves for it: REFUSED is no answer, so the nameserver
// is saved as error with the reason.
func TestNameserverThatRefuses(t *testing.T) {
t.Parallel()
const reason = "server returned REFUSED"
res := resolver.NewFromLogger(slog.Default())
var (
ns string
resp *resolver.NameserverResponse
)
livednstest.Retry(
t,
"QueryNameserver(cloudflare.com)",
func(ctx context.Context) error {
nameservers, err := res.LookupNS(ctx, testDomain)
if err != nil {
return err
}
ns = nameservers[0]
resp, err = res.QueryNameserver(ctx, ns, "cloudflare.com")
if err != nil {
return err
}
// A timeout or a network error is no reply at all.
if resp.Status == resolver.StatusTimeout ||
strings.HasPrefix(resp.Error, "network error") {
return fmt.Errorf(
"%w: %s: %s", livednstest.ErrNoAnswer, ns, resp.Error,
)
}
return nil
},
)
if resp.Error != reason {
t.Fatalf(
"%s answered %q (%s) for cloudflare.com, want REFUSED",
ns, resp.Status, resp.Error,
)
}
hs := watcher.BuildHostnameState(
map[string]*resolver.NameserverResponse{ns: resp}, time.Now(),
)
got := hs.RecordsByNameserver[ns]
if got.Status != failed().Status || got.Error != reason {
t.Errorf(
"saved status %q, error %q; want status %q, error %q",
got.Status, got.Error, failed().Status, reason,
)
}
}
+84 -39
View File
@@ -13,6 +13,7 @@ import (
"sneak.berlin/go/dnswatcher/internal/config" "sneak.berlin/go/dnswatcher/internal/config"
"sneak.berlin/go/dnswatcher/internal/logger" "sneak.berlin/go/dnswatcher/internal/logger"
"sneak.berlin/go/dnswatcher/internal/resolver"
"sneak.berlin/go/dnswatcher/internal/state" "sneak.berlin/go/dnswatcher/internal/state"
"sneak.berlin/go/dnswatcher/internal/tlscheck" "sneak.berlin/go/dnswatcher/internal/tlscheck"
) )
@@ -227,7 +228,7 @@ func (w *Watcher) checkDomain(
// Also look up A/AAAA records for the apex domain so that // Also look up A/AAAA records for the apex domain so that
// port and TLS checks (which read HostnameState) can find // port and TLS checks (which read HostnameState) can find
// the domain's IP addresses. // the domain's IP addresses.
records, err := w.resolver.LookupAllRecords(ctx, domain) results, err := w.resolver.LookupAllRecords(ctx, domain)
if err != nil { if err != nil {
w.log.Error( w.log.Error(
"failed to lookup records for domain", "failed to lookup records for domain",
@@ -238,12 +239,13 @@ func (w *Watcher) checkDomain(
return return
} }
newState := buildHostnameState(results, now)
prevHS, hasPrevHS := w.state.GetHostnameState(domain) prevHS, hasPrevHS := w.state.GetHostnameState(domain)
if hasPrevHS && !w.firstRun { if hasPrevHS && !w.firstRun {
w.detectHostnameChanges(ctx, domain, prevHS, records) w.detectHostnameChanges(ctx, domain, prevHS, newState)
} }
newState := buildHostnameState(records, now)
w.state.SetHostnameState(domain, newState) w.state.SetHostnameState(domain, newState)
} }
@@ -292,7 +294,7 @@ func (w *Watcher) checkHostname(
ctx context.Context, ctx context.Context,
hostname string, hostname string,
) { ) {
records, err := w.resolver.LookupAllRecords(ctx, hostname) results, err := w.resolver.LookupAllRecords(ctx, hostname)
if err != nil { if err != nil {
w.log.Error( w.log.Error(
"failed to lookup records", "failed to lookup records",
@@ -303,19 +305,22 @@ func (w *Watcher) checkHostname(
return return
} }
now := time.Now().UTC() newState := buildHostnameState(results, time.Now().UTC())
prev, hasPrev := w.state.GetHostnameState(hostname)
prev, hasPrev := w.state.GetHostnameState(hostname)
if hasPrev && !w.firstRun { if hasPrev && !w.firstRun {
w.detectHostnameChanges(ctx, hostname, prev, records) w.detectHostnameChanges(ctx, hostname, prev, newState)
} }
newState := buildHostnameState(records, now)
w.state.SetHostnameState(hostname, newState) w.state.SetHostnameState(hostname, newState)
} }
// buildHostnameState saves each nameserver's response. A nameserver
// that answered, even with NXDOMAIN or no records, is saved as ok; one
// that timed out or failed is saved as error with the reason, and its
// empty record set is not an answer.
func buildHostnameState( func buildHostnameState(
records map[string]map[string][]string, results map[string]*resolver.NameserverResponse,
now time.Time, now time.Time,
) *state.HostnameState { ) *state.HostnameState {
hs := &state.HostnameState{ hs := &state.HostnameState{
@@ -325,12 +330,20 @@ func buildHostnameState(
LastChecked: now, LastChecked: now,
} }
for ns, recs := range records { for ns, resp := range results {
hs.RecordsByNameserver[ns] = &state.NameserverRecordState{ nsState := &state.NameserverRecordState{
Records: recs, Records: resp.Records,
Status: statusOK, Status: statusOK,
LastChecked: now, LastChecked: now,
} }
if resp.Status == resolver.StatusTimeout ||
resp.Status == resolver.StatusError {
nsState.Status = statusError
nsState.Error = resp.Error
}
hs.RecordsByNameserver[ns] = nsState
} }
return hs return hs
@@ -339,27 +352,29 @@ func buildHostnameState(
func (w *Watcher) detectHostnameChanges( func (w *Watcher) detectHostnameChanges(
ctx context.Context, ctx context.Context,
hostname string, hostname string,
prev *state.HostnameState, prev, current *state.HostnameState,
current map[string]map[string][]string,
) { ) {
w.detectRecordChanges(ctx, hostname, prev, current) w.detectRecordChanges(ctx, hostname, prev, current)
w.detectNSDisappearances(ctx, hostname, prev, current) w.detectNSDisappearances(ctx, hostname, prev, current)
w.detectNSFailures(ctx, hostname, prev, current)
w.detectInconsistencies(ctx, hostname, prev, current) w.detectInconsistencies(ctx, hostname, prev, current)
} }
// detectRecordChanges compares each nameserver's records with those of
// the previous check. Only answers are compared: a nameserver that
// failed on either check has no records to compare.
func (w *Watcher) detectRecordChanges( func (w *Watcher) detectRecordChanges(
ctx context.Context, ctx context.Context,
hostname string, hostname string,
prev *state.HostnameState, prev, current *state.HostnameState,
current map[string]map[string][]string,
) { ) {
for ns, recs := range current { for ns, cur := range current.RecordsByNameserver {
prevNS, ok := prev.RecordsByNameserver[ns] prevNS, ok := prev.RecordsByNameserver[ns]
if !ok { if !ok || prevNS.Status != statusOK || cur.Status != statusOK {
continue continue
} }
if recordsEqual(prevNS.Records, recs) { if recordsEqual(prevNS.Records, cur.Records) {
continue continue
} }
@@ -367,7 +382,7 @@ func (w *Watcher) detectRecordChanges(
"Hostname: %s\nNameserver: %s\n"+ "Hostname: %s\nNameserver: %s\n"+
"Old: %v\nNew: %v", "Old: %v\nNew: %v",
hostname, ns, hostname, ns,
prevNS.Records, recs, prevNS.Records, cur.Records,
) )
w.notify.SendNotification( w.notify.SendNotification(
@@ -382,11 +397,10 @@ func (w *Watcher) detectRecordChanges(
func (w *Watcher) detectNSDisappearances( func (w *Watcher) detectNSDisappearances(
ctx context.Context, ctx context.Context,
hostname string, hostname string,
prev *state.HostnameState, prev, current *state.HostnameState,
current map[string]map[string][]string,
) { ) {
for ns, prevNS := range prev.RecordsByNameserver { for ns, prevNS := range prev.RecordsByNameserver {
if _, ok := current[ns]; ok || prevNS.Status != statusOK { if _, ok := current.RecordsByNameserver[ns]; ok || prevNS.Status != statusOK {
continue continue
} }
@@ -402,13 +416,36 @@ func (w *Watcher) detectNSDisappearances(
"error", "error",
) )
} }
}
for ns := range current { // detectNSFailures notifies when a nameserver that answered on the
// previous check fails, and when one that failed answers again. A
// nameserver missing from the previous check is not compared.
func (w *Watcher) detectNSFailures(
ctx context.Context,
hostname string,
prev, current *state.HostnameState,
) {
for ns, cur := range current.RecordsByNameserver {
prevNS, ok := prev.RecordsByNameserver[ns] prevNS, ok := prev.RecordsByNameserver[ns]
if !ok || prevNS.Status != statusError { if !ok {
continue continue
} }
switch {
case prevNS.Status == statusOK && cur.Status == statusError:
msg := fmt.Sprintf(
"Hostname: %s\nNameserver: %s\nError: %s",
hostname, ns, cur.Error,
)
w.notify.SendNotification(
ctx,
"NS Failure: "+hostname,
msg,
"error",
)
case prevNS.Status == statusError && cur.Status == statusOK:
msg := fmt.Sprintf( msg := fmt.Sprintf(
"Hostname: %s\nNameserver: %s recovered", "Hostname: %s\nNameserver: %s recovered",
hostname, ns, hostname, ns,
@@ -422,12 +459,12 @@ func (w *Watcher) detectNSDisappearances(
) )
} }
} }
}
func (w *Watcher) detectInconsistencies( func (w *Watcher) detectInconsistencies(
ctx context.Context, ctx context.Context,
hostname string, hostname string,
prev *state.HostnameState, prev, current *state.HostnameState,
current map[string]map[string][]string,
) { ) {
for _, pair := range newlyDisagreeingPairs(prev, current) { for _, pair := range newlyDisagreeingPairs(prev, current) {
ns1, ns2 := pair[0], pair[1] ns1, ns2 := pair[0], pair[1]
@@ -435,8 +472,8 @@ func (w *Watcher) detectInconsistencies(
msg := fmt.Sprintf( msg := fmt.Sprintf(
"Hostname: %s\n%s: %v\n%s: %v", "Hostname: %s\n%s: %v\n%s: %v",
hostname, hostname,
ns1, current[ns1], ns1, current.RecordsByNameserver[ns1].Records,
ns2, current[ns2], ns2, current.RecordsByNameserver[ns2].Records,
) )
w.notify.SendNotification( w.notify.SendNotification(
@@ -448,18 +485,21 @@ func (w *Watcher) detectInconsistencies(
} }
} }
// newlyDisagreeingPairs returns every pair of nameservers whose records // newlyDisagreeingPairs returns every pair of nameservers that answered
// differ in current, in sorted order of name, except pairs where both // in current and whose records differ there, in sorted order of name,
// nameservers were in prev and already differed there. A nameserver // except pairs where both nameservers answered in prev and already
// missing from prev is paired with every nameserver it differs from. // differed there. A nameserver missing from prev, or that failed there,
// is paired with every nameserver it differs from. A nameserver that
// failed in current has no records to compare and is in no pair.
func newlyDisagreeingPairs( func newlyDisagreeingPairs(
prev *state.HostnameState, prev, current *state.HostnameState,
current map[string]map[string][]string,
) [][2]string { ) [][2]string {
nameservers := make([]string, 0, len(current)) nameservers := make([]string, 0, len(current.RecordsByNameserver))
for ns := range current { for ns, cur := range current.RecordsByNameserver {
if cur.Status == statusOK {
nameservers = append(nameservers, ns) nameservers = append(nameservers, ns)
} }
}
sort.Strings(nameservers) sort.Strings(nameservers)
@@ -467,14 +507,19 @@ func newlyDisagreeingPairs(
for i, ns1 := range nameservers { for i, ns1 := range nameservers {
for _, ns2 := range nameservers[i+1:] { for _, ns2 := range nameservers[i+1:] {
if recordsEqual(current[ns1], current[ns2]) { if recordsEqual(
current.RecordsByNameserver[ns1].Records,
current.RecordsByNameserver[ns2].Records,
) {
continue continue
} }
prev1, ok1 := prev.RecordsByNameserver[ns1] prev1, ok1 := prev.RecordsByNameserver[ns1]
prev2, ok2 := prev.RecordsByNameserver[ns2] prev2, ok2 := prev.RecordsByNameserver[ns2]
if ok1 && ok2 && !recordsEqual(prev1.Records, prev2.Records) { if ok1 && ok2 &&
prev1.Status == statusOK && prev2.Status == statusOK &&
!recordsEqual(prev1.Records, prev2.Records) {
continue continue
} }
+10 -3
View File
@@ -687,11 +687,12 @@ func TestNSFailureAndRecovery(t *testing.T) {
cfg.Hostnames = []string{testHost} cfg.Hostnames = []string{testHost}
// Between the checks, save every nameserver the first check found // Between the checks, save every nameserver the first check found
// as failed, and add, as answering, one that live DNS does not list. // 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) { deps := runChecks(t, cfg, nil, func(deps *testDeps) {
hs, _ := deps.state.GetHostnameState(testHost) hs, _ := deps.state.GetHostnameState(testHost)
for _, nsState := range hs.RecordsByNameserver { for ns := range hs.RecordsByNameserver {
nsState.Status = "error" hs.RecordsByNameserver[ns] = failed()
} }
hs.RecordsByNameserver[oldNS1] = &state.NameserverRecordState{ hs.RecordsByNameserver[oldNS1] = &state.NameserverRecordState{
@@ -704,4 +705,10 @@ func TestNSFailureAndRecovery(t *testing.T) {
assertNotified(t, deps, "NS Failure: "+testHost, "error") assertNotified(t, deps, "NS Failure: "+testHost, "error")
assertNotified(t, deps, "NS Recovery: "+testHost, "success") 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)
}
} }