2 Commits
Author SHA1 Message Date
sneak 417a87228a watcher: save state when it stops and wait for that save (closes #114)
check / check (push) Successful in 1m40s
The final save at shutdown came from the state's own stop hook, while
the watcher's stop hook only cancelled its run loop, so a check under
way could change state after that save or be cut off at exit. Run now
saves state as it returns, and the watcher's stop hook waits for Run,
bounded by the shutdown deadline. The state's own save stays; Save
holds the state lock for the whole write, so the two cannot overlap.
The start hook derives the watcher's context with WithoutCancel, so
the linter needs no exception. A new test stops a watcher built by New
and reads the change back from the state file, with no DNS.

Model: opus-5-5
2026-10-01 20:45:31 +00:00
clawbot a8f9a64600 middleware: take the client address from the right of X-Forwarded-For (closes #181)
check / check (push) Successful in 1m22s
realIP took the first X-Forwarded-For entry, which the client itself
can write, so behind a proxy that appends to the header a client chose
the address dnswatcher logs and the /metrics rate limit counts. It now
walks the entries from the right past trusted proxies, using the
existing trusted-proxy check, and takes the first that is not one; the
leftmost when all are. All X-Forwarded-For header lines are read as one
list, since a proxy may add its own line instead of appending to the
client's. An empty entry where the client address belongs falls back
to the peer address, as an empty first entry did before. X-Real-IP is
unchanged.

Model: opus-5-5
2026-10-01 22:35:17 +02:00
7 changed files with 260 additions and 43 deletions
+7 -6
View File
@@ -347,9 +347,9 @@ minute, failed logins included; beyond that it answers `429 Too Many Requests`
without checking the password. A Prometheus server scraping every 15 seconds 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 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, 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 client address is taken from the `X-Real-IP` header the proxy sets, or else
the proxy sets; a proxy that sets neither makes all its clients share one from `X-Forwarded-For`, as the last address in it that is not private or
allowance. loopback. A proxy that sets neither makes all its clients share one allowance.
**`DNSWATCHER_DNS_INTERVAL` and `DNSWATCHER_TLS_INTERVAL`** take a positive **`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 duration: a number followed by a unit such as `s`, `m` or `h`, for example
@@ -611,9 +611,10 @@ repository's `Dockerfile` and runs it. The app needs:
from a previous cycle. from a previous cycle.
4. **On change detection**: Send notifications to all configured 4. **On change detection**: Send notifications to all configured
endpoints, update in-memory state, persist to disk. endpoints, update in-memory state, persist to disk.
5. **Shutdown**: Persist final state to disk, wait for in-flight 5. **Shutdown**: The watcher stops checking and saves the final state
notification deliveries to complete, stop gracefully. The wait is to disk, and shutdown waits for that save before it goes on. Then it
bounded by the fx shutdown timeout (15s by default): deliveries still waits for in-flight notification deliveries to complete. Both waits
share the fx shutdown timeout (15s by default): deliveries still
retrying against an unreachable endpoint when that expires are retrying against an unreachable endpoint when that expires are
abandoned, and the number abandoned is logged at warn level rather abandoned, and the number abandoned is logged at warn level rather
than dropped silently. Notifications generated after shutdown has than dropped silently. Notifications generated after shutdown has
+4 -1
View File
@@ -19,6 +19,10 @@ nameserver IP address changes: https://git.eeqj.de/sneak/dnswatcher/issues/105
# Completed Steps # Completed Steps
- 2026-10-01: the watcher saves state when it stops, and shutdown waits for that
save, so it no longer relies on the state's own stop hook (closes #114).
- 2026-10-01: the client address from `X-Forwarded-For` is the last entry that
is not a trusted proxy, not the first, which the client sets (closes #181).
- 2026-10-01: a nameserver that does not answer is saved as `error` with the - 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). reason, and NS failure and NS recovery are notified (closes #104).
- 2026-10-01: a `DNSWATCHER_DNS_INTERVAL` or `DNSWATCHER_TLS_INTERVAL` that is - 2026-10-01: a `DNSWATCHER_DNS_INTERVAL` or `DNSWATCHER_TLS_INTERVAL` that is
@@ -105,7 +109,6 @@ nameserver IP address changes: https://git.eeqj.de/sneak/dnswatcher/issues/105
https://git.eeqj.de/sneak/dnswatcher/issues/66 https://git.eeqj.de/sneak/dnswatcher/issues/66
- `goimports` in `make fmt-check`, Markdown formatting: - `goimports` in `make fmt-check`, Markdown formatting:
https://git.eeqj.de/sneak/dnswatcher/issues/119 https://git.eeqj.de/sneak/dnswatcher/issues/119
- final state save at shutdown: https://git.eeqj.de/sneak/dnswatcher/issues/114
- README accuracy sweep: https://git.eeqj.de/sneak/dnswatcher/issues/108 - README accuracy sweep: https://git.eeqj.de/sneak/dnswatcher/issues/108
- README sections required by policy: - README sections required by policy:
https://git.eeqj.de/sneak/dnswatcher/issues/173 https://git.eeqj.de/sneak/dnswatcher/issues/173
+10 -1
View File
@@ -1,6 +1,9 @@
package middleware package middleware
import "time" import (
"net/http"
"time"
)
// The /metrics rate limit, exported so the tests can count requests // The /metrics rate limit, exported so the tests can count requests
// against it. // against it.
@@ -8,3 +11,9 @@ const (
MetricsRequestLimit = metricsRequestLimit MetricsRequestLimit = metricsRequestLimit
MetricsRequestWindow time.Duration = metricsRequestWindow MetricsRequestWindow time.Duration = metricsRequestWindow
) )
// RealIP is realIP, exported so the tests can check which address it
// takes as the client's.
func RealIP(r *http.Request) string {
return realIP(r)
}
+22 -6
View File
@@ -209,6 +209,12 @@ func isTrustedProxy(ip net.IP) bool {
// realIP extracts the client's real IP address from the request. // realIP extracts the client's real IP address from the request.
// Proxy headers are only trusted from RFC1918/loopback addresses. // Proxy headers are only trusted from RFC1918/loopback addresses.
//
// Each proxy adds to the end of X-Forwarded-For the address it got the
// request from, so the client can write every entry before the one the
// first trusted proxy added. The client address is therefore the
// rightmost entry that is not a trusted proxy, or the leftmost entry
// when they all are.
func realIP(r *http.Request) string { func realIP(r *http.Request) string {
addr := ipFromHostPort(r.RemoteAddr) addr := ipFromHostPort(r.RemoteAddr)
remoteIP := net.ParseIP(addr) remoteIP := net.ParseIP(addr)
@@ -223,14 +229,24 @@ func realIP(r *http.Request) string {
return ip return ip
} }
if xff := r.Header.Get("X-Forwarded-For"); xff != "" { // A proxy may add its entry as a header line of its own instead of
if parts := strings.SplitN( // appending to the line the client sent, so all lines form one list.
xff, ",", 2, //nolint:mnd entries := strings.Split(
); len(parts) > 0 { strings.Join(r.Header.Values("X-Forwarded-For"), ","), ",",
if ip := strings.TrimSpace(parts[0]); ip != "" { )
return ip client := strings.TrimSpace(entries[0])
for i := len(entries) - 1; i > 0; i-- {
entry := strings.TrimSpace(entries[i])
if !isTrustedProxy(net.ParseIP(entry)) {
client = entry
break
} }
} }
if client != "" {
return client
} }
return addr return addr
+84 -3
View File
@@ -342,9 +342,9 @@ func TestDashboardRendersWithSecurityHeaders(t *testing.T) {
} }
} }
// Addresses for the rate limit tests: a client connecting directly, a // Addresses for the rate limit and realIP tests: a client connecting
// trusted proxy, and a client behind that proxy as its X-Real-IP // directly, a trusted proxy, and a client behind that proxy as the
// header names it. // proxy's X-Real-IP or X-Forwarded-For header names it.
const ( const (
directClient = "198.51.100.1:4000" directClient = "198.51.100.1:4000"
trustedProxy = "10.0.0.1:4000" trustedProxy = "10.0.0.1:4000"
@@ -482,3 +482,84 @@ func TestMetricsRateLimitKeysOnClientAddress(t *testing.T) {
}) })
} }
} }
// TestRealIP checks which address realIP takes as the client's. Each
// element of forwardedFor is sent as an X-Forwarded-For header line of
// its own, and 198.51.100.9 is always an entry the client wrote itself.
func TestRealIP(t *testing.T) {
t.Parallel()
tests := []struct {
name string
remoteAddr string
xRealIP string
forwardedFor []string
want string
}{
{
"untrusted peer, both headers ignored",
directClient, proxiedClient, []string{"198.51.100.9"},
"198.51.100.1",
},
{
"X-Real-IP from a trusted proxy wins",
trustedProxy, proxiedClient, []string{"203.0.113.8"},
proxiedClient,
},
{
"client's own entry, then the one the proxy added",
trustedProxy, "", []string{"198.51.100.9, 203.0.113.1"},
proxiedClient,
},
{
"several trusted proxies",
trustedProxy, "",
[]string{"198.51.100.9, 203.0.113.1, 10.0.0.3, 10.0.0.2"},
proxiedClient,
},
{
"proxy adds a header line of its own",
trustedProxy, "", []string{"198.51.100.9", proxiedClient},
proxiedClient,
},
{
"every entry a trusted proxy",
trustedProxy, "", []string{"10.0.0.3, 10.0.0.2"},
"10.0.0.3",
},
{
"empty where the client address belongs",
trustedProxy, "", []string{"203.0.113.1, , 10.0.0.2"},
"10.0.0.1",
},
{
"no headers from a trusted proxy",
trustedProxy, "", nil,
"10.0.0.1",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
req := httptest.NewRequestWithContext(
t.Context(), http.MethodGet, "/", nil,
)
req.RemoteAddr = tt.remoteAddr
if tt.xRealIP != "" {
req.Header.Set("X-Real-IP", tt.xRealIP)
}
for _, line := range tt.forwardedFor {
req.Header.Add("X-Forwarded-For", line)
}
got := middleware.RealIP(req)
if got != tt.want {
t.Errorf("realIP = %q, want %q", got, tt.want)
}
})
}
}
+29 -11
View File
@@ -56,6 +56,7 @@ type Watcher struct {
tlsCheck TLSChecker tlsCheck TLSChecker
notify Notifier notify Notifier
cancel context.CancelFunc cancel context.CancelFunc
done chan struct{} // closed when Run returns
firstRun bool firstRun bool
expiryNotifiedMu sync.Mutex expiryNotifiedMu sync.Mutex
expiryNotified map[string]time.Time expiryNotified map[string]time.Time
@@ -79,31 +80,47 @@ func New(
} }
lifecycle.Append(fx.Hook{ lifecycle.Append(fx.Hook{
OnStart: func(_ context.Context) error { OnStart: func(startCtx context.Context) error {
// Use context.Background() — the fx startup context // The fx startup context expires after startup
// expires after startup completes, so deriving from it // completes, so the watcher's context drops its
// would cancel the watcher immediately. The watcher's // cancellation. The watcher's lifetime is controlled
// lifetime is controlled by w.cancel in OnStop. // by w.cancel in OnStop.
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(
context.WithoutCancel(startCtx),
)
w.cancel = cancel w.cancel = cancel
w.done = make(chan struct{})
go w.Run(ctx) //nolint:contextcheck // intentionally not derived from startCtx go func() {
defer close(w.done)
w.Run(ctx)
}()
return nil return nil
}, },
OnStop: func(_ context.Context) error { OnStop: func(ctx context.Context) error {
if w.cancel != nil {
w.cancel() w.cancel()
}
// Run saves state as it returns. Waiting for it here
// means the save is done before shutdown goes on.
select {
case <-w.done:
return nil return nil
case <-ctx.Done():
return fmt.Errorf(
"waiting for the watcher to stop: %w",
ctx.Err(),
)
}
}, },
}) })
return w, nil return w, nil
} }
// Run starts the monitoring loop with periodic scheduling. // Run starts the monitoring loop with periodic scheduling. When ctx
// is cancelled, it saves state and returns.
func (w *Watcher) Run(ctx context.Context) { func (w *Watcher) Run(ctx context.Context) {
w.log.Info( w.log.Info(
"watcher starting", "watcher starting",
@@ -125,6 +142,7 @@ func (w *Watcher) Run(ctx context.Context) {
for { for {
select { select {
case <-ctx.Done(): case <-ctx.Done():
w.saveState()
w.log.Info("watcher stopped") w.log.Info("watcher stopped")
return return
+101 -12
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"fmt" "fmt"
"log/slog" "log/slog"
"os"
"slices" "slices"
"sync" "sync"
"testing" "testing"
@@ -135,6 +136,7 @@ type testDeps struct {
notifier *mockNotifier notifier *mockNotifier
state *state.State state *state.State
config *config.Config config *config.Config
log *logger.Logger
} }
func newTestWatcher( func newTestWatcher(
@@ -143,6 +145,23 @@ func newTestWatcher(
) (*watcher.Watcher, *testDeps) { ) (*watcher.Watcher, *testDeps) {
t.Helper() 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{ deps := &testDeps{
portChecker: &mockPortChecker{}, portChecker: &mockPortChecker{},
tlsChecker: &mockTLSChecker{ tlsChecker: &mockTLSChecker{
@@ -157,30 +176,21 @@ func newTestWatcher(
t.Fatalf("globals.New: %v", err) t.Fatalf("globals.New: %v", err)
} }
log, err := logger.New(nil, logger.Params{Globals: g}) deps.log, err = logger.New(nil, logger.Params{Globals: g})
if err != nil { if err != nil {
t.Fatalf("logger.New: %v", err) t.Fatalf("logger.New: %v", err)
} }
// The watcher saves state after every check, into cfg.DataDir. // The watcher saves state after every check, into cfg.DataDir.
deps.state, err = state.New(fxtest.NewLifecycle(t), state.Params{ deps.state, err = state.New(fxtest.NewLifecycle(t), state.Params{
Logger: log, Logger: deps.log,
Config: cfg, Config: cfg,
}) })
if err != nil { if err != nil {
t.Fatalf("state.New: %v", err) t.Fatalf("state.New: %v", err)
} }
w := watcher.NewForTest( return deps
deps.config,
deps.state,
resolver.NewFromLogger(slog.Default()),
deps.portChecker,
deps.tlsChecker,
deps.notifier,
)
return w, deps
} }
func defaultTestConfig(t *testing.T) *config.Config { func defaultTestConfig(t *testing.T) *config.Config {
@@ -539,6 +549,85 @@ func TestGracefulShutdown(t *testing.T) {
} }
} }
// 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) { func TestDNSRunsBeforePortAndTLSChecks(t *testing.T) {
t.Parallel() t.Parallel()