1 Commits

Author SHA1 Message Date
clawbot
e97a4e523f feat: add security response headers middleware (closes #98)
All checks were successful
check / check (push) Successful in 35s
Add SecurityHeaders() to internal/middleware and register it in the
global middleware stack so every response - dashboard, embedded static
assets, healthchecks, JSON API, and metrics - carries the six response
headers required by REPO_POLICIES.md before tagging 1.0:

  Strict-Transport-Security: max-age=31536000; includeSubDomains
  Content-Security-Policy:   default-src 'self'; script-src 'none';
                             style-src 'self'; img-src 'self';
                             font-src 'none'; connect-src 'none';
                             object-src 'none'; base-uri 'none';
                             form-action 'none'; frame-ancestors 'none'
  X-Frame-Options:           DENY
  X-Content-Type-Options:    nosniff
  Referrer-Policy:           no-referrer
  Permissions-Policy:        unused browser features denied

The dashboard template ships no JavaScript, no inline styles, no inline
event handlers and no images, and its only subresource is the embedded
stylesheet at /s/css/tailwind.min.css, so the policy needs neither
unsafe-inline nor unsafe-eval. frame-ancestors 'none' is the primary
anti-framing control with X-Frame-Options as the legacy fallback.

HSTS is emitted unconditionally rather than gated on r.TLS, because the
service runs behind a TLS-terminating proxy and the browser must still
enforce HTTPS end to end.

The headers are set before the request reaches the next handler, so
they are present on error responses too, including recovered panics and
request timeouts.

Tests cover each header's exact value, the CSP's required and forbidden
directives, presence on a 500 response, and a render of the real
dashboard through the middleware confirming the page still references
its stylesheet.
2026-08-09 01:47:47 +00:00
10 changed files with 543 additions and 781 deletions

View File

@@ -182,6 +182,46 @@ dnswatcher exposes a lightweight HTTP API for operational visibility:
| `GET /api/v1/status` | Current monitoring state |
| `GET /metrics` | Prometheus metrics (optional) |
### Security Headers
Every response — the dashboard, the static assets under `/s/...`, the
healthchecks, the JSON API, and `/metrics` — carries the following
headers, set by a global middleware:
| Header | Value |
|-----------------------------|---------------------------------------|
| `Strict-Transport-Security` | `max-age=31536000; includeSubDomains` |
| `Content-Security-Policy` | see below |
| `X-Frame-Options` | `DENY` |
| `X-Content-Type-Options` | `nosniff` |
| `Referrer-Policy` | `no-referrer` |
| `Permissions-Policy` | all unused browser features denied |
The content security policy is:
```
default-src 'self'; script-src 'none'; style-src 'self'; img-src 'self';
font-src 'none'; connect-src 'none'; object-src 'none'; base-uri 'none';
form-action 'none'; frame-ancestors 'none'
```
The dashboard ships no JavaScript (the 30-second refresh is a
`<meta http-equiv="refresh">`), no inline styles, no inline event
handlers, and no images; its only subresource is the embedded stylesheet
at `/s/css/tailwind.min.css`, which `style-src 'self'` permits. The
policy therefore needs neither `unsafe-inline` nor `unsafe-eval`.
`frame-ancestors 'none'` is the primary anti-framing control, with
`X-Frame-Options: DENY` retained as the legacy fallback.
HSTS is emitted unconditionally, including over plain HTTP. dnswatcher is
expected to run behind a TLS-terminating reverse proxy, and the browser
must still be told to enforce HTTPS end to end, so the header is never
gated on whether the request itself arrived over TLS.
`Referrer-Policy: no-referrer` is stricter than the
`strict-origin-when-cross-origin` baseline: the dashboard has no
cross-origin navigation needs, and its URL may name internal hosts.
---
## Architecture
@@ -194,7 +234,8 @@ internal/
globals/globals.go Build-time variables (version)
logger/logger.go slog structured logging (TTY detection)
healthcheck/healthcheck.go Health check service
middleware/middleware.go HTTP middleware (logging, CORS, metrics auth)
middleware/middleware.go HTTP middleware (logging, CORS, security
headers, metrics auth)
handlers/handlers.go HTTP request handlers
server/
server.go HTTP server lifecycle
@@ -218,8 +259,7 @@ internal/
- **Structured logging**: All logs use `log/slog` with JSON output in
production (TTY detection for development).
- **Graceful shutdown**: All background goroutines respect context
cancellation and the fx lifecycle. In-flight notification deliveries
are drained on shutdown, bounded by the shutdown timeout.
cancellation and the fx lifecycle.
---
@@ -453,14 +493,8 @@ docker run -d \
from a previous cycle.
4. **On change detection**: Send notifications to all configured
endpoints, update in-memory state, persist to disk.
5. **Shutdown**: Persist final state to disk, wait for in-flight
notification deliveries to complete, stop gracefully. The wait is
bounded by the fx shutdown timeout (15s by default): deliveries still
retrying against an unreachable endpoint when that expires are
abandoned, and the number abandoned is logged at warn level rather
than dropped silently. Notifications generated after shutdown has
begun are refused and logged, so a late burst cannot extend the
shutdown.
5. **Shutdown**: Persist final state to disk, complete in-flight
notifications, stop gracefully.
---

19
TODO.md
View File

@@ -25,15 +25,16 @@ confirm make check still passes.
# Completed Steps
- 2026-08-09: in-flight notification deliveries are now drained at
shutdown (#106): `notify.New` registers an fx `OnStop` hook that waits
on a `sync.WaitGroup` of tracked delivery goroutines, bounded by the
`OnStop` context; on expiry the outstanding count is logged at warn
level and parked retry backoffs are released instead of being dropped
silently, and deliveries submitted after the drain begins are refused
so shutdown cannot be extended indefinitely; an `OnStop` context that
is already expired on entry with nothing outstanding drains quietly
rather than warning about deliveries that were never abandoned
- 2026-08-09: security response headers middleware
(`SecurityHeaders()` in `internal/middleware/middleware.go`)
registered globally in `internal/server/routes.go`, so HSTS, CSP,
`X-Frame-Options`, `X-Content-Type-Options`, `Referrer-Policy`, and
`Permissions-Policy` are set on every response including `/s/...` and
`/metrics`; the CSP needs no `unsafe-inline`/`unsafe-eval` because the
dashboard ships no JavaScript and no inline styles; HSTS is emitted
unconditionally per policy (TLS-terminating proxy in front). Remaining
1.0 hardening items — `http.Server` timeouts, request body limits,
rate limiting, CORS scoping — are tracked separately
- 2026-08-07: golangci-lint bumped to v2.12.2 (commit-pinned installs
in `Dockerfile` and `script/bootstrap`); `.golangci.yml` set to the
org-standard v2-schema config used across the org's repos

View File

@@ -21,6 +21,60 @@ import (
// corsMaxAge is the maximum age for CORS preflight responses.
const corsMaxAge = 300
// Security response header values applied to every response.
//
// The CSP is as strict as the dashboard allows: the template ships no
// JavaScript, no inline styles, no inline event handlers and no images,
// and its only subresource is the embedded stylesheet at
// /s/css/tailwind.min.css, which style-src 'self' permits. Neither
// unsafe-inline nor unsafe-eval is used. frame-ancestors 'none' is the
// primary anti-framing control; X-Frame-Options is the legacy fallback.
const (
// hstsValue is emitted unconditionally, including over plain HTTP,
// because the service runs behind a TLS-terminating proxy and the
// browser must still enforce HTTPS end to end.
hstsValue = "max-age=31536000; includeSubDomains"
cspValue = "default-src 'self'; " +
"script-src 'none'; " +
"style-src 'self'; " +
"img-src 'self'; " +
"font-src 'none'; " +
"connect-src 'none'; " +
"object-src 'none'; " +
"base-uri 'none'; " +
"form-action 'none'; " +
"frame-ancestors 'none'"
frameOptionsValue = "DENY"
contentTypeOptionsValue = "nosniff"
// referrerPolicyValue is stricter than the policy minimum of
// strict-origin-when-cross-origin: the dashboard has no
// cross-origin navigation needs and its URL may name internal
// hosts.
referrerPolicyValue = "no-referrer"
permissionsPolicyValue = "accelerometer=(), " +
"autoplay=(), " +
"camera=(), " +
"display-capture=(), " +
"encrypted-media=(), " +
"fullscreen=(), " +
"geolocation=(), " +
"gyroscope=(), " +
"magnetometer=(), " +
"microphone=(), " +
"midi=(), " +
"payment=(), " +
"picture-in-picture=(), " +
"publickey-credentials-get=(), " +
"screen-wake-lock=(), " +
"usb=(), " +
"xr-spatial-tracking=()"
)
// Params contains dependencies for Middleware.
type Params struct {
fx.In
@@ -186,6 +240,37 @@ func (m *Middleware) CORS() func(http.Handler) http.Handler {
})
}
// SecurityHeaders returns middleware that sets the security response
// headers required for production internet exposure on every response.
//
// The headers are set before the request reaches the next handler so
// that they are present on every response, including panics recovered
// by chi's Recoverer and timeouts produced by chi's Timeout.
func (m *Middleware) SecurityHeaders() func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(
writer http.ResponseWriter,
request *http.Request,
) {
header := writer.Header()
header.Set("Strict-Transport-Security", hstsValue)
header.Set("Content-Security-Policy", cspValue)
header.Set("X-Frame-Options", frameOptionsValue)
header.Set(
"X-Content-Type-Options",
contentTypeOptionsValue,
)
header.Set("Referrer-Policy", referrerPolicyValue)
header.Set(
"Permissions-Policy",
permissionsPolicyValue,
)
next.ServeHTTP(writer, request)
})
}
}
// MetricsAuth returns basic auth middleware for /metrics.
func (m *Middleware) MetricsAuth() func(http.Handler) http.Handler {
if m.params.Config.MetricsUsername == "" {

View File

@@ -0,0 +1,333 @@
package middleware_test
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/go-chi/chi/v5"
"sneak.berlin/go/dnswatcher/internal/config"
"sneak.berlin/go/dnswatcher/internal/globals"
"sneak.berlin/go/dnswatcher/internal/handlers"
"sneak.berlin/go/dnswatcher/internal/logger"
"sneak.berlin/go/dnswatcher/internal/middleware"
"sneak.berlin/go/dnswatcher/internal/notify"
"sneak.berlin/go/dnswatcher/internal/state"
)
// Expected security header values, spelled out literally so that any
// change to the middleware has to be made deliberately here as well.
const (
wantHSTS = "max-age=31536000; includeSubDomains"
wantCSP = "default-src 'self'; " +
"script-src 'none'; " +
"style-src 'self'; " +
"img-src 'self'; " +
"font-src 'none'; " +
"connect-src 'none'; " +
"object-src 'none'; " +
"base-uri 'none'; " +
"form-action 'none'; " +
"frame-ancestors 'none'"
wantFrameOptions = "DENY"
wantContentTypeOptions = "nosniff"
wantReferrerPolicy = "no-referrer"
wantPermissionsPolicy = "accelerometer=(), " +
"autoplay=(), " +
"camera=(), " +
"display-capture=(), " +
"encrypted-media=(), " +
"fullscreen=(), " +
"geolocation=(), " +
"gyroscope=(), " +
"magnetometer=(), " +
"microphone=(), " +
"midi=(), " +
"payment=(), " +
"picture-in-picture=(), " +
"publickey-credentials-get=(), " +
"screen-wake-lock=(), " +
"usb=(), " +
"xr-spatial-tracking=()"
)
// stylesheetPath is the only subresource the dashboard loads.
const stylesheetPath = "/s/css/tailwind.min.css"
// newTestLogger builds a logger for direct component construction.
func newTestLogger(t *testing.T) *logger.Logger {
t.Helper()
glob, err := globals.New(nil)
if err != nil {
t.Fatalf("globals.New: %v", err)
}
log, err := logger.New(nil, logger.Params{Globals: glob})
if err != nil {
t.Fatalf("logger.New: %v", err)
}
return log
}
// newTestMiddleware builds a Middleware without an fx application.
func newTestMiddleware(t *testing.T) *middleware.Middleware {
t.Helper()
glob, err := globals.New(nil)
if err != nil {
t.Fatalf("globals.New: %v", err)
}
mw, err := middleware.New(nil, middleware.Params{
Logger: newTestLogger(t),
Globals: glob,
Config: &config.Config{},
})
if err != nil {
t.Fatalf("middleware.New: %v", err)
}
return mw
}
// serveWithSecurityHeaders runs a GET through SecurityHeaders and
// returns the recorded response.
func serveWithSecurityHeaders(
t *testing.T,
target string,
handler http.Handler,
) *httptest.ResponseRecorder {
t.Helper()
mw := newTestMiddleware(t)
rec := httptest.NewRecorder()
req := httptest.NewRequestWithContext(
t.Context(), http.MethodGet, target, nil,
)
mw.SecurityHeaders()(handler).ServeHTTP(rec, req)
return rec
}
// okHandler writes a trivial 200 response.
func okHandler() http.Handler {
return http.HandlerFunc(func(
writer http.ResponseWriter,
_ *http.Request,
) {
writer.WriteHeader(http.StatusOK)
})
}
func TestSecurityHeaders(t *testing.T) {
t.Parallel()
tests := []struct {
name string
header string
want string
}{
{
"hsts",
"Strict-Transport-Security",
wantHSTS,
},
{
"csp",
"Content-Security-Policy",
wantCSP,
},
{
"frame options",
"X-Frame-Options",
wantFrameOptions,
},
{
"content type options",
"X-Content-Type-Options",
wantContentTypeOptions,
},
{
"referrer policy",
"Referrer-Policy",
wantReferrerPolicy,
},
{
"permissions policy",
"Permissions-Policy",
wantPermissionsPolicy,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
rec := serveWithSecurityHeaders(t, "/", okHandler())
got := rec.Header().Get(tt.header)
if got != tt.want {
t.Errorf(
"%s = %q, want %q",
tt.header, got, tt.want,
)
}
})
}
}
// TestSecurityHeadersCSPDirectives guards the properties the repo
// policy requires of the content security policy itself.
func TestSecurityHeadersCSPDirectives(t *testing.T) {
t.Parallel()
rec := serveWithSecurityHeaders(t, "/", okHandler())
csp := rec.Header().Get("Content-Security-Policy")
forbidden := []string{"unsafe-inline", "unsafe-eval"}
for _, directive := range forbidden {
if strings.Contains(csp, directive) {
t.Errorf("CSP must not contain %q: %q", directive, csp)
}
}
required := []string{
"default-src 'self'",
"script-src 'none'",
"style-src 'self'",
"frame-ancestors 'none'",
}
for _, directive := range required {
if !strings.Contains(csp, directive) {
t.Errorf("CSP must contain %q: %q", directive, csp)
}
}
}
// TestSecurityHeadersOnErrorResponse verifies the headers are emitted
// even when the wrapped handler fails, since they are set before the
// handler runs.
func TestSecurityHeadersOnErrorResponse(t *testing.T) {
t.Parallel()
failing := http.HandlerFunc(func(
writer http.ResponseWriter,
_ *http.Request,
) {
http.Error(
writer,
"boom",
http.StatusInternalServerError,
)
})
rec := serveWithSecurityHeaders(t, "/api/v1/status", failing)
if rec.Code != http.StatusInternalServerError {
t.Fatalf("status = %d, want 500", rec.Code)
}
if got := rec.Header().Get(
"X-Content-Type-Options",
); got != wantContentTypeOptions {
t.Errorf(
"X-Content-Type-Options = %q, want %q",
got, wantContentTypeOptions,
)
}
if got := rec.Header().Get(
"Strict-Transport-Security",
); got != wantHSTS {
t.Errorf(
"Strict-Transport-Security = %q, want %q",
got, wantHSTS,
)
}
}
// newTestHandlers builds real Handlers with empty monitoring state.
func newTestHandlers(t *testing.T) *handlers.Handlers {
t.Helper()
glob, err := globals.New(nil)
if err != nil {
t.Fatalf("globals.New: %v", err)
}
log := newTestLogger(t)
notifier, err := notify.New(nil, notify.Params{
Logger: log,
Config: &config.Config{},
})
if err != nil {
t.Fatalf("notify.New: %v", err)
}
hnd, err := handlers.New(nil, handlers.Params{
Logger: log,
Globals: glob,
State: state.NewForTest(),
Notify: notifier,
})
if err != nil {
t.Fatalf("handlers.New: %v", err)
}
return hnd
}
// TestDashboardRendersWithSecurityHeaders renders the real dashboard
// through the middleware and checks that the policy still permits the
// one stylesheet the page loads.
func TestDashboardRendersWithSecurityHeaders(t *testing.T) {
t.Parallel()
mw := newTestMiddleware(t)
hnd := newTestHandlers(t)
router := chi.NewRouter()
router.Use(mw.SecurityHeaders())
router.Get("/", hnd.HandleDashboard())
rec := httptest.NewRecorder()
req := httptest.NewRequestWithContext(
t.Context(), http.MethodGet, "/", nil,
)
router.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rec.Code)
}
body := rec.Body.String()
if !strings.Contains(body, stylesheetPath) {
t.Errorf("dashboard does not reference %q", stylesheetPath)
}
if !strings.Contains(body, "dnswatcher") {
t.Errorf("dashboard body looks empty: %d bytes", len(body))
}
csp := rec.Header().Get("Content-Security-Policy")
if csp != wantCSP {
t.Errorf("CSP = %q, want %q", csp, wantCSP)
}
// The stylesheet is same-origin, so style-src 'self' allows it.
if !strings.Contains(csp, "style-src 'self'") {
t.Errorf("CSP would block %q: %q", stylesheetPath, csp)
}
}

View File

@@ -32,27 +32,11 @@ func NewRequestForTest(
// NewTestService creates a Service suitable for unit testing.
// It discards log output and uses the given transport.
func NewTestService(transport http.RoundTripper) *Service {
return newService(slog.New(slog.DiscardHandler), transport)
}
// NewTestServiceWithLogger creates a Service that writes to the
// given handler, so tests can assert on emitted log records.
func NewTestServiceWithLogger(
transport http.RoundTripper,
handler slog.Handler,
) *Service {
return newService(slog.New(handler), transport)
}
// Drain exports drain for testing.
func (svc *Service) Drain(ctx context.Context) {
svc.drain(ctx)
}
// OutstandingDeliveries reports how many delivery goroutines
// are currently tracked as in flight.
func (svc *Service) OutstandingDeliveries() int64 {
return svc.outstanding.Load()
return &Service{
log: slog.New(slog.DiscardHandler),
transport: transport,
history: NewAlertHistory(),
}
}
// SetNtfyURL sets the ntfy URL on a Service for testing.

View File

@@ -12,8 +12,6 @@ import (
"log/slog"
"net/http"
"net/url"
"sync"
"sync/atomic"
"time"
"go.uber.org/fx"
@@ -117,41 +115,19 @@ type Service struct {
history *AlertHistory
retryConfig RetryConfig
sleepFn func(time.Duration) <-chan time.Time
// Shutdown draining state. drainMu guards draining and
// serialises it against the counter increment in
// startDelivery; inFlight tracks the delivery goroutines
// themselves and outstanding mirrors its count so a timed
// out drain can report how many were abandoned.
drainMu sync.Mutex
draining bool
inFlight sync.WaitGroup
outstanding atomic.Int64
abandon chan struct{}
abandonOnce sync.Once
}
// newService builds a Service with the fields every Service
// needs regardless of how it was constructed.
func newService(
log *slog.Logger,
transport http.RoundTripper,
) *Service {
return &Service{
log: log,
transport: transport,
history: NewAlertHistory(),
abandon: make(chan struct{}),
}
}
// New creates a new notify Service.
func New(
lifecycle fx.Lifecycle,
_ fx.Lifecycle,
params Params,
) (*Service, error) {
svc := newService(params.Logger.Get(), http.DefaultTransport)
svc.config = params.Config
svc := &Service{
log: params.Logger.Get(),
transport: http.DefaultTransport,
config: params.Config,
history: NewAlertHistory(),
}
if params.Config.NtfyTopic != "" {
u, err := ValidateWebhookURL(
@@ -192,14 +168,6 @@ func New(
svc.mattermostWebhookURL = u
}
lifecycle.Append(fx.Hook{
OnStop: func(ctx context.Context) error {
svc.drain(ctx)
return nil
},
})
return svc, nil
}
@@ -226,32 +194,6 @@ func (svc *Service) SendNotification(
svc.dispatchMattermost(ctx, title, message, priority)
}
// dispatch delivers a notification to one endpoint on a
// tracked background goroutine.
//
// The delivery context is detached from ctx with
// context.WithoutCancel so that a cancelled caller does not
// kill a delivery already under way; the shutdown drain, not
// the caller, decides how long deliveries may keep running.
func (svc *Service) dispatch(
ctx context.Context,
endpoint string,
send func(context.Context) error,
) {
notifyCtx := context.WithoutCancel(ctx)
svc.startDelivery(endpoint, func() {
err := svc.deliverWithRetry(notifyCtx, endpoint, send)
if err != nil {
svc.log.Error(
"failed to send notification after retries",
"endpoint", endpoint,
"error", err,
)
}
})
}
func (svc *Service) dispatchNtfy(
ctx context.Context,
title, message, priority string,
@@ -260,11 +202,26 @@ func (svc *Service) dispatchNtfy(
return
}
svc.dispatch(ctx, "ntfy", func(c context.Context) error {
return svc.sendNtfy(
c, svc.ntfyURL, title, message, priority,
go func() {
notifyCtx := context.WithoutCancel(ctx)
err := svc.deliverWithRetry(
notifyCtx, "ntfy",
func(c context.Context) error {
return svc.sendNtfy(
c, svc.ntfyURL,
title, message, priority,
)
},
)
})
if err != nil {
svc.log.Error(
"failed to send ntfy notification "+
"after retries",
"error", err,
)
}
}()
}
func (svc *Service) dispatchSlack(
@@ -275,11 +232,26 @@ func (svc *Service) dispatchSlack(
return
}
svc.dispatch(ctx, "slack", func(c context.Context) error {
return svc.sendSlack(
c, svc.slackWebhookURL, title, message, priority,
go func() {
notifyCtx := context.WithoutCancel(ctx)
err := svc.deliverWithRetry(
notifyCtx, "slack",
func(c context.Context) error {
return svc.sendSlack(
c, svc.slackWebhookURL,
title, message, priority,
)
},
)
})
if err != nil {
svc.log.Error(
"failed to send slack notification "+
"after retries",
"error", err,
)
}
}()
}
func (svc *Service) dispatchMattermost(
@@ -290,15 +262,26 @@ func (svc *Service) dispatchMattermost(
return
}
svc.dispatch(
ctx, "mattermost",
func(c context.Context) error {
return svc.sendSlack(
c, svc.mattermostWebhookURL,
title, message, priority,
go func() {
notifyCtx := context.WithoutCancel(ctx)
err := svc.deliverWithRetry(
notifyCtx, "mattermost",
func(c context.Context) error {
return svc.sendSlack(
c, svc.mattermostWebhookURL,
title, message, priority,
)
},
)
if err != nil {
svc.log.Error(
"failed to send mattermost notification "+
"after retries",
"error", err,
)
},
)
}
}()
}
func (svc *Service) sendNtfy(

View File

@@ -2,7 +2,6 @@ package notify
import (
"context"
"fmt"
"math"
"math/rand/v2"
"time"
@@ -122,14 +121,6 @@ func (svc *Service) deliverWithRetry(
select {
case <-ctx.Done():
return ctx.Err()
case <-svc.abandon:
// Shutdown drained past its deadline; stop
// sleeping rather than outlive the process.
// A nil channel (Service built without a
// constructor) simply never fires.
return fmt.Errorf(
"%w: %s", ErrDeliveryAbandoned, endpoint,
)
case <-svc.sleepFunc(delay):
}
}

View File

@@ -1,119 +0,0 @@
package notify
import (
"context"
"errors"
)
// ErrDeliveryAbandoned is returned by a retry loop that was
// cut short because shutdown drained past its deadline.
var ErrDeliveryAbandoned = errors.New(
"notification delivery abandoned at shutdown",
)
// startDelivery runs fn on its own goroutine while tracking it,
// so that drain can wait for it during shutdown.
//
// The WaitGroup counter is incremented here, on the caller's
// goroutine, before the worker exists: incrementing it inside
// the worker would race with drain's Wait and could let
// shutdown sail past a delivery that had not started yet.
//
// Once draining has begun the delivery is refused outright
// rather than queued, so a steady stream of newly submitted
// notifications cannot keep extending the drain.
func (svc *Service) startDelivery(endpoint string, fn func()) {
svc.drainMu.Lock()
if svc.draining {
svc.drainMu.Unlock()
svc.log.Warn(
"notification not dispatched: shutdown in progress",
"endpoint", endpoint,
)
return
}
svc.outstanding.Add(1)
// WaitGroup.Go increments the counter synchronously, here,
// and only then starts the goroutine.
svc.inFlight.Go(func() {
// Runs before the WaitGroup counter is decremented, so
// a drain that times out reports an accurate count.
defer svc.outstanding.Add(-1)
fn()
})
svc.drainMu.Unlock()
}
// drain waits for in-flight notification deliveries to finish.
//
// It first stops accepting new deliveries, then waits until
// either every outstanding delivery has completed or ctx
// expires — whichever comes first. ctx is the context fx
// passes to the OnStop hook, so a permanently dead webhook
// cannot hang shutdown indefinitely.
//
// When the deadline arrives with deliveries still outstanding,
// the count is logged at warn level and the abandon channel is
// closed, which releases any retry loop sleeping in backoff.
// Deliveries already inside an HTTP round trip are bounded by
// the existing httpClientTimeout instead.
//
// A ctx that is already expired on entry is not by itself cause
// for alarm: if nothing is outstanding there is nothing to
// abandon, and the drain says so at debug level rather than
// warning about deliveries that do not exist.
func (svc *Service) drain(ctx context.Context) {
svc.drainMu.Lock()
svc.draining = true
svc.drainMu.Unlock()
done := make(chan struct{})
go func() {
svc.inFlight.Wait()
close(done)
}()
select {
case <-done:
svc.log.Debug(
"all in-flight notifications completed",
)
case <-ctx.Done():
// outstanding is decremented before the WaitGroup
// counter, and startDelivery can no longer add to it
// now that draining is set, so a zero here means every
// delivery really did finish. ctx expiring in that
// state (an OnStop context that was already cancelled
// on entry is the usual way) abandons nothing, so it
// must not close abandon or warn about it.
abandoned := svc.outstanding.Load()
if abandoned == 0 {
svc.log.Debug(
"all in-flight notifications completed",
)
return
}
svc.abandonOnce.Do(func() {
if svc.abandon != nil {
close(svc.abandon)
}
})
svc.log.Warn(
"shutdown deadline reached with notifications "+
"still in flight; abandoning them",
"abandoned", abandoned,
"error", ctx.Err(),
)
}
}

View File

@@ -1,531 +0,0 @@
package notify_test
import (
"bytes"
"context"
"log/slog"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"go.uber.org/fx"
"sneak.berlin/go/dnswatcher/internal/config"
"sneak.berlin/go/dnswatcher/internal/globals"
"sneak.berlin/go/dnswatcher/internal/logger"
"sneak.berlin/go/dnswatcher/internal/notify"
)
// Timings used by the drain tests. They stay in the same
// 10-100ms band as the retry tests so the suite never waits on
// a real backoff delay.
const (
// inFlightHold is how long a delivery is kept mid-request
// before the handler is released.
inFlightHold = 30 * time.Millisecond
// drainDeadline bounds a drain that is expected to time
// out.
drainDeadline = 50 * time.Millisecond
// drainSlack is the upper bound on how long a bounded
// drain may take; generous enough for a loaded CI box,
// still far below the 20s test ceiling.
drainSlack = 2 * time.Second
// settleDelay is how long to wait before asserting that
// something did *not* happen.
settleDelay = 50 * time.Millisecond
// idleDrainBound is the upper bound on a drain that has
// nothing in flight. It is deliberately far above the cost
// of the goroutine hop through inFlight.Wait() — which
// reached 57ms on a loaded box under -race with the package's
// parallel tests — and far below drainSlack, the deadline
// such a drain is given. A drain that blocked until its
// deadline instead of returning on the WaitGroup therefore
// still fails this bound, but scheduling delay alone cannot.
idleDrainBound = 500 * time.Millisecond
)
// syncBuffer is an io.Writer safe for concurrent use, so log
// output written from delivery goroutines can be inspected.
type syncBuffer struct {
mu sync.Mutex
buf bytes.Buffer
}
func (sb *syncBuffer) Write(p []byte) (int, error) {
sb.mu.Lock()
defer sb.mu.Unlock()
return sb.buf.Write(p) //nolint:wrapcheck // test helper
}
func (sb *syncBuffer) String() string {
sb.mu.Lock()
defer sb.mu.Unlock()
return sb.buf.String()
}
// newLoggingService returns a Service writing JSON logs into
// the returned buffer.
func newLoggingService(
transport http.RoundTripper,
) (*notify.Service, *syncBuffer) {
logs := &syncBuffer{}
handler := slog.NewJSONHandler(logs, nil)
return notify.NewTestServiceWithLogger(transport, handler),
logs
}
// blockingNtfyServer returns a server whose handler signals on
// entered, waits for release, and then responds 200.
func blockingNtfyServer(
entered chan<- struct{},
release <-chan struct{},
served *atomic.Bool,
) *httptest.Server {
var once sync.Once
return httptest.NewServer(
http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
once.Do(func() { close(entered) })
<-release
served.Store(true)
w.WriteHeader(http.StatusOK)
}),
)
}
// TestDrainWaitsForInFlightDelivery verifies that a delivery
// already under way when shutdown starts is allowed to finish.
func TestDrainWaitsForInFlightDelivery(t *testing.T) {
t.Parallel()
var served atomic.Bool
entered := make(chan struct{})
release := make(chan struct{})
srv := blockingNtfyServer(entered, release, &served)
defer srv.Close()
topicURL, _ := url.Parse(srv.URL)
svc := notify.NewTestService(http.DefaultTransport)
svc.SetNtfyURL(topicURL)
svc.SendNotification(
context.Background(), "t", "m", prioInfo,
)
// Make sure the delivery really is mid-request before the
// drain begins.
select {
case <-entered:
case <-time.After(drainSlack):
t.Fatal("delivery never reached the endpoint")
}
// As in TestDrainBoundedByContextDeadline: start is captured
// before the clock it is compared against, here the timer
// holding the delivery open, so elapsed covers the whole hold
// and the lower bound cannot come out short from scheduling
// delay alone.
start := time.Now()
timer := time.AfterFunc(inFlightHold, func() {
close(release)
})
defer timer.Stop()
ctx, cancel := context.WithTimeout(
context.Background(), drainSlack,
)
defer cancel()
svc.Drain(ctx)
elapsed := time.Since(start)
if !served.Load() {
t.Error(
"drain returned before the in-flight delivery " +
"completed",
)
}
if elapsed < inFlightHold {
t.Errorf(
"drain took %v, want at least %v",
elapsed, inFlightHold,
)
}
if got := svc.OutstandingDeliveries(); got != 0 {
t.Errorf("outstanding deliveries = %d, want 0", got)
}
}
// neverFires returns a channel that never delivers, standing in
// for a long backoff sleep without actually sleeping.
func neverFires(_ time.Duration) <-chan time.Time {
return make(chan time.Time)
}
// TestDrainBoundedByContextDeadline verifies that a delivery
// stuck retrying against a dead endpoint does not hold shutdown
// past the OnStop context deadline, and that the abandoned
// deliveries are logged at warn level rather than dropped
// silently.
func TestDrainBoundedByContextDeadline(t *testing.T) {
t.Parallel()
var requests atomic.Int64
srv := httptest.NewServer(
http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
requests.Add(1)
w.WriteHeader(http.StatusInternalServerError)
}),
)
defer srv.Close()
topicURL, _ := url.Parse(srv.URL)
svc, logs := newLoggingService(http.DefaultTransport)
svc.SetNtfyURL(topicURL)
// Never let the backoff sleep complete: the delivery is
// parked in its retry wait until shutdown releases it.
svc.SetSleepFunc(neverFires)
svc.SetRetryConfig(notify.RetryConfig{
MaxRetries: 5,
BaseDelay: time.Hour,
MaxDelay: time.Hour,
})
svc.SendNotification(
context.Background(), "t", "m", prioError,
)
waitForCondition(t, func() bool {
return requests.Load() >= 1 &&
svc.OutstandingDeliveries() == 1
})
// start must be captured *before* the deadline clock starts,
// so that the measured interval is a superset of the deadline
// interval. Capturing it after context.WithTimeout would
// make elapsed structurally smaller than drainDeadline and
// the lower bound below unfalsifiable-by-luck: it would fail
// whenever the two statements were separated by any
// scheduling delay, and pass otherwise, regardless of what
// the drain did.
start := time.Now()
ctx, cancel := context.WithTimeout(
context.Background(), drainDeadline,
)
defer cancel()
// The upper bound is enforced by a watchdog rather than by
// measuring after the fact: a drain that is not bounded at
// all never returns here (the delivery is parked in a backoff
// that never fires), so an unbounded drain must fail this
// test promptly instead of hanging the package until the test
// binary's 30s timeout.
returned := make(chan struct{})
go func() {
defer close(returned)
svc.Drain(ctx)
}()
select {
case <-returned:
case <-time.After(drainSlack):
t.Fatalf(
"drain did not return within %v; its %v deadline "+
"did not bound it",
drainSlack, drainDeadline,
)
}
// The lower bound is the real assertion: the drain must have
// waited for its whole deadline rather than giving up on the
// outstanding delivery early. With start captured above, an
// early return is the only thing that can make it fail.
if elapsed := time.Since(start); elapsed < drainDeadline {
t.Errorf(
"drain returned after %v, before its %v deadline",
elapsed, drainDeadline,
)
}
assertAbandonLogged(t, logs.String())
// The abandoned delivery must stop retrying rather than
// outlive the drain.
waitForCondition(t, func() bool {
return svc.OutstandingDeliveries() == 0
})
}
// assertAbandonLogged checks that the drain logged the
// abandoned deliveries at warn level with a count.
func assertAbandonLogged(t *testing.T, output string) {
t.Helper()
if !strings.Contains(output, `"level":"WARN"`) {
t.Errorf(
"abandoned deliveries not logged at warn level; "+
"log output: %s",
output,
)
}
if !strings.Contains(output, `"abandoned":1`) {
t.Errorf(
"abandoned delivery count not logged; "+
"log output: %s",
output,
)
}
}
// TestDrainRefusesNewDeliveries verifies that notifications
// submitted after the drain has begun are refused and logged,
// so a stream of new work cannot extend shutdown indefinitely.
func TestDrainRefusesNewDeliveries(t *testing.T) {
t.Parallel()
var requests atomic.Int64
srv := httptest.NewServer(
http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
requests.Add(1)
w.WriteHeader(http.StatusOK)
}),
)
defer srv.Close()
target, _ := url.Parse(srv.URL)
svc, logs := newLoggingService(http.DefaultTransport)
svc.SetNtfyURL(target)
svc.SetSlackWebhookURL(target)
svc.SetMattermostWebhookURL(target)
ctx, cancel := context.WithTimeout(
context.Background(), drainSlack,
)
defer cancel()
// Nothing is in flight, so this returns immediately and
// leaves the service refusing further deliveries.
svc.Drain(ctx)
for range 3 {
svc.SendNotification(
context.Background(), "t", "m", prioInfo,
)
}
time.Sleep(settleDelay)
if got := requests.Load(); got != 0 {
t.Errorf(
"%d requests reached the endpoint after drain, "+
"want 0",
got,
)
}
if got := svc.OutstandingDeliveries(); got != 0 {
t.Errorf("outstanding deliveries = %d, want 0", got)
}
output := logs.String()
if !strings.Contains(output, "shutdown in progress") {
t.Errorf(
"refused deliveries not logged; log output: %s",
output,
)
}
}
// recordingLifecycle is a minimal fx.Lifecycle that records the
// hooks appended to it, so the wiring done by notify.New can be
// inspected without standing up a whole fx application.
type recordingLifecycle struct {
hooks []fx.Hook
}
func (l *recordingLifecycle) Append(hook fx.Hook) {
l.hooks = append(l.hooks, hook)
}
// newNotifyService builds a Service through the real
// constructor, wired to the given lifecycle.
func newNotifyService(
t *testing.T,
lifecycle fx.Lifecycle,
ntfyTopic string,
) *notify.Service {
t.Helper()
g, err := globals.New(nil)
if err != nil {
t.Fatalf("globals.New: %v", err)
}
log, err := logger.New(nil, logger.Params{Globals: g})
if err != nil {
t.Fatalf("logger.New: %v", err)
}
svc, err := notify.New(lifecycle, notify.Params{
Logger: log,
Config: &config.Config{NtfyTopic: ntfyTopic},
})
if err != nil {
t.Fatalf("notify.New: %v", err)
}
return svc
}
// TestNewRegistersDrainingStopHook verifies that notify.New
// wires an OnStop hook into the fx lifecycle and that the hook
// waits for in-flight deliveries.
func TestNewRegistersDrainingStopHook(t *testing.T) {
t.Parallel()
var served atomic.Bool
entered := make(chan struct{})
release := make(chan struct{})
srv := blockingNtfyServer(entered, release, &served)
defer srv.Close()
lifecycle := &recordingLifecycle{}
svc := newNotifyService(t, lifecycle, srv.URL)
if len(lifecycle.hooks) != 1 {
t.Fatalf(
"appended %d lifecycle hooks, want 1",
len(lifecycle.hooks),
)
}
stop := lifecycle.hooks[0].OnStop
if stop == nil {
t.Fatal("lifecycle hook has no OnStop function")
}
svc.SendNotification(
context.Background(), "t", "m", prioInfo,
)
select {
case <-entered:
case <-time.After(drainSlack):
t.Fatal("delivery never reached the endpoint")
}
timer := time.AfterFunc(inFlightHold, func() {
close(release)
})
defer timer.Stop()
ctx, cancel := context.WithTimeout(
context.Background(), drainSlack,
)
defer cancel()
err := stop(ctx)
if err != nil {
t.Fatalf("OnStop returned error: %v", err)
}
if !served.Load() {
t.Error(
"OnStop returned before the in-flight delivery " +
"completed",
)
}
}
// TestDrainWithoutDeliveriesReturnsImmediately verifies the
// common case: nothing in flight, shutdown is not delayed.
func TestDrainWithoutDeliveriesReturnsImmediately(t *testing.T) {
t.Parallel()
svc := notify.NewTestService(http.DefaultTransport)
// Captured before the deadline clock, as elsewhere in this
// file; for an upper bound that is the conservative
// direction, since the measured interval can then only be
// longer than the drain itself.
start := time.Now()
ctx, cancel := context.WithTimeout(
context.Background(), drainSlack,
)
defer cancel()
svc.Drain(ctx)
if elapsed := time.Since(start); elapsed > idleDrainBound {
t.Errorf(
"drain of an idle service took %v, want well "+
"under its %v deadline",
elapsed, drainSlack,
)
}
}
// TestDrainWithCancelledContextDoesNotWarn verifies that an
// OnStop context that is already dead on entry does not produce
// an "abandoning them" warning when there was nothing in flight
// to abandon. The expired context wins the select immediately,
// so only the outstanding count can tell the difference between
// a genuine timeout and a shutdown that had simply already run
// out of time with no work left.
func TestDrainWithCancelledContextDoesNotWarn(t *testing.T) {
t.Parallel()
svc, logs := newLoggingService(http.DefaultTransport)
ctx, cancel := context.WithCancel(context.Background())
cancel()
svc.Drain(ctx)
if output := logs.String(); strings.Contains(
output, `"level":"WARN"`,
) {
t.Errorf(
"drain with nothing in flight warned about "+
"abandoned deliveries; log output: %s",
output,
)
}
}

View File

@@ -21,6 +21,7 @@ func (s *Server) SetupRoutes() {
// Global middleware
s.router.Use(chimw.Recoverer)
s.router.Use(chimw.RequestID)
s.router.Use(s.mw.SecurityHeaders())
s.router.Use(s.mw.Logging())
s.router.Use(s.mw.CORS())
s.router.Use(chimw.Timeout(requestTimeout))