diff --git a/bootstrap/bootstrap.go b/bootstrap/bootstrap.go index 199c2397..a3312a70 100644 --- a/bootstrap/bootstrap.go +++ b/bootstrap/bootstrap.go @@ -22,6 +22,7 @@ import ( "log/slog" "net/http" "os" + "time" "github.com/alecthomas/kingpin/v2" "github.com/prometheus/common/promslog" @@ -43,6 +44,9 @@ var ( errNegativeMaxRequests = errors.New("web max requests must be greater than or equal to zero") ) +// defaultReadHeaderTimeout applies when Config.ReadHeaderTimeout is left unset. +const defaultReadHeaderTimeout = time.Minute + // MetricsHandlerFactory builds an exporter-specific metrics handler after the // common toolkit flags have been parsed. type MetricsHandlerFactory func(*Bootstrap) (http.Handler, error) @@ -82,6 +86,9 @@ type Config struct { MetricsHandler http.Handler // MetricsHandlerFactory builds the metrics handler after flags are parsed. MetricsHandlerFactory MetricsHandlerFactory + // ReadHeaderTimeout bounds request-header reads, mitigating Slowloris + // (gosec G112). Defaults to one minute when zero. + ReadHeaderTimeout time.Duration } // Runner manages generic exporter startup around flag parsing, landing page @@ -250,7 +257,15 @@ func (t *Runner) newServer(metricsHandler http.Handler) (*http.Server, error) { mux.Handle("/", landingPage) } - return &http.Server{Handler: mux}, nil + readHeaderTimeout := defaultReadHeaderTimeout + if t.provided.ReadHeaderTimeout > 0 { + readHeaderTimeout = t.provided.ReadHeaderTimeout + } + + return &http.Server{ + Handler: mux, + ReadHeaderTimeout: readHeaderTimeout, + }, nil } func (t *Runner) defaultLandingConfig() web.LandingConfig { diff --git a/bootstrap/bootstrap_test.go b/bootstrap/bootstrap_test.go index dfa43ec1..4525eb56 100644 --- a/bootstrap/bootstrap_test.go +++ b/bootstrap/bootstrap_test.go @@ -14,10 +14,12 @@ package bootstrap import ( + "net" "net/http" "net/http/httptest" "strings" "testing" + "time" "github.com/alecthomas/kingpin/v2" "github.com/prometheus/common/promslog" @@ -129,3 +131,117 @@ func TestNewServerRegistersMetricsAndLandingPage(t *testing.T) { t.Fatalf("unexpected landing body: %q", body) } } + +// TestNewServerReadHeaderTimeout checks newServer maps Config.ReadHeaderTimeout +// onto the server, defaulting to one minute when unset. +func TestNewServerReadHeaderTimeout(t *testing.T) { + for _, tc := range []struct { + name string + configured time.Duration + want time.Duration + }{ + {name: "unset defaults", configured: 0, want: defaultReadHeaderTimeout}, + {name: "explicit default", configured: time.Minute, want: time.Minute}, + {name: "sub-second value", configured: 250 * time.Millisecond, want: 250 * time.Millisecond}, + {name: "smallest positive value", configured: time.Nanosecond, want: time.Nanosecond}, + {name: "large value", configured: time.Hour, want: time.Hour}, + {name: "negative defaults", configured: -1, want: defaultReadHeaderTimeout}, + {name: "large negative defaults", configured: -time.Hour, want: defaultReadHeaderTimeout}, + } { + t.Run(tc.name, func(t *testing.T) { + tk := New(Config{ + App: kingpin.New("test", ""), + DefaultAddress: ":9100", + Logger: promslog.NewNopLogger(), + ReadHeaderTimeout: tc.configured, + MetricsHandler: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }), + }) + if err := tk.parse([]string{"--web.listen-address=:0"}); err != nil { + t.Fatalf("unexpected parse error: %v", err) + } + handler, err := tk.resolveMetricsHandler() + if err != nil { + t.Fatalf("unexpected handler resolution error: %v", err) + } + server, err := tk.newServer(handler) + if err != nil { + t.Fatalf("unexpected server creation error: %v", err) + } + if server.ReadHeaderTimeout != tc.want { + t.Fatalf("unexpected ReadHeaderTimeout: got %v, want %v", server.ReadHeaderTimeout, tc.want) + } + }) + } +} + +// TestNewServerReadHeaderTimeoutClosesStalledConnection checks the timeout is +// effective end-to-end: a connection with incomplete headers is closed while a +// well-formed request succeeds. +func TestNewServerReadHeaderTimeoutClosesStalledConnection(t *testing.T) { + const readHeaderTimeout = 250 * time.Millisecond + + tk := New(Config{ + App: kingpin.New("test", ""), + Name: "test_exporter", + Description: "test description", + DefaultAddress: ":9100", + Logger: promslog.NewNopLogger(), + ReadHeaderTimeout: readHeaderTimeout, + MetricsHandler: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }), + }) + if err := tk.parse([]string{"--web.listen-address=:0"}); err != nil { + t.Fatalf("unexpected parse error: %v", err) + } + handler, err := tk.resolveMetricsHandler() + if err != nil { + t.Fatalf("unexpected handler resolution error: %v", err) + } + server, err := tk.newServer(handler) + if err != nil { + t.Fatalf("unexpected server creation error: %v", err) + } + if server.ReadHeaderTimeout != readHeaderTimeout { + t.Fatalf("unexpected ReadHeaderTimeout: got %v, want %v", server.ReadHeaderTimeout, readHeaderTimeout) + } + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { _ = ln.Close() }) + go func() { _ = server.Serve(ln) }() + t.Cleanup(func() { _ = server.Close() }) + + // A well-formed request still succeeds. + resp, err := http.Get("http://" + ln.Addr().String() + "/metrics") + if err != nil { + t.Fatalf("well-formed request failed: %v", err) + } + _ = resp.Body.Close() + + // Start request headers but never terminate them (no final CRLF). + conn, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + t.Fatalf("dial: %v", err) + } + t.Cleanup(func() { _ = conn.Close() }) + if _, err := conn.Write([]byte("GET /metrics HTTP/1.1\r\nHost: localhost\r\n")); err != nil { + t.Fatalf("write partial request: %v", err) + } + + // Read in a goroutine so the test never hangs if the server keeps it open. + done := make(chan struct{}) + go func() { + _, _ = conn.Read(make([]byte, 1)) // unblocks on server-side close + close(done) + }() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("server did not close the stalled connection within 5s; ReadHeaderTimeout not effective") + } +}