net/http: add Server.MaxHeaderValueCount setting Implement MaxHeaderValueCount support for both the HTTP/1 and HTTP/2 server. This allows servers who want to support large headers such as SSO or OIDC cookies (and therefore has to set MaxHeaderBytes to a large value), to still limit the number of headers that they are willing to accept. Also added the usual note for linknamed symbols for net/textproto.readMIMEHeader, which seems to have been missed originally. Release note changes will be added directly in x/website. For #79936 Change-Id: I600ac6c492f64f9cfa2730836be64dc76a6a6964 Reviewed-on: https://go-review.googlesource.com/c/go/+/795460 Reviewed-by: Nicholas Husin <husin@google.com> Reviewed-by: Damien Neil <dneil@google.com> LUCI-TryBot-Result: golang-scoped@luci-project-accounts.iam.gserviceaccount.com <golang-scoped@luci-project-accounts.iam.gserviceaccount.com>
diff --git a/api/go1.27.txt b/api/go1.27.txt index 2520fba..04b797d 100644 --- a/api/go1.27.txt +++ b/api/go1.27.txt
@@ -278,6 +278,9 @@ pkg math/big, method (*Int) Divide(*Int, *Int, *Int, RoundingMode) (*Int, *Int) #76821 pkg math/rand/v2, method (*Rand) N[$0 intType]($0) $0 #77853 pkg net/http, type Server struct, DisableClientPriority bool #75500 +pkg net/http, const DefaultMaxHeaderValueCount = 500 #79936 +pkg net/http, const DefaultMaxHeaderValueCount ideal-int #79936 +pkg net/http, type Server struct, MaxHeaderValueCount int #79936 pkg net/http/httptest, func NewTestServer(testing.TB, http.Handler) *Server #76608 pkg net/url, method (*URL) Clone() *URL #73450 pkg net/url, method (Values) Clone() Values #73450
diff --git a/src/net/http/http2.go b/src/net/http/http2.go index bb752a6..301aff7 100644 --- a/src/net/http/http2.go +++ b/src/net/http/http2.go
@@ -186,7 +186,8 @@ s *Server } -func (s http2ServerConfig) MaxHeaderBytes() int { return s.s.MaxHeaderBytes } +func (s http2ServerConfig) MaxHeaderBytes() int { return s.s.MaxHeaderBytes } +func (s http2ServerConfig) MaxHeaderValueCount() int { return s.s.maxHeaderValueCount() } func (s http2ServerConfig) ConnState(c net.Conn, st http2.ConnState) { if s.s.ConnState != nil { s.s.ConnState(c, ConnState(st))
diff --git a/src/net/http/internal/http2/api.go b/src/net/http/internal/http2/api.go index f208d52..1728bc3 100644 --- a/src/net/http/internal/http2/api.go +++ b/src/net/http/internal/http2/api.go
@@ -107,6 +107,7 @@ // ServerConfig is configuration from an http.Server. type ServerConfig interface { MaxHeaderBytes() int + MaxHeaderValueCount() int ConnState(net.Conn, ConnState) DoKeepAlives() bool WriteTimeout() time.Duration
diff --git a/src/net/http/internal/http2/frame.go b/src/net/http/internal/http2/frame.go index a2de8c2..e3b371a 100644 --- a/src/net/http/internal/http2/frame.go +++ b/src/net/http/internal/http2/frame.go
@@ -333,6 +333,12 @@ // If the limit is hit, MetaHeadersFrame.Truncated is set true. MaxHeaderListSize uint32 + // MaxHeaderValueCount is the maximum permitted number of + // header values. + // It's used only if ReadMetaHeaders is set; 0 means no limit. + // If the limit is hit, MetaHeadersFrame.Truncated is set true. + MaxHeaderValueCount int + // TODO: track which type of frame & with which flags was sent // last. Then return an error (unless AllowIllegalWrites) if // we're in the middle of a header block and a @@ -356,6 +362,10 @@ return fr.MaxHeaderListSize } +func (fr *Framer) maxHeaderValueCount() int { + return fr.MaxHeaderValueCount +} + func (f *Framer) startWrite(ftype FrameType, flags Flags, streamID uint32) { // Write the FrameHeader. f.wbuf = append(f.wbuf[:0], @@ -1715,6 +1725,7 @@ } var remainSize = fr.maxHeaderListSize() var sawRegular bool + var headerCount int var invalid error // pseudo header field errors hdec := fr.ReadMetaHeaders @@ -1724,6 +1735,13 @@ if VerboseLogs && fr.logReads { fr.debugReadLoggerf("http2: decoded hpack field %+v", hf) } + headerCount++ + if limit := fr.maxHeaderValueCount(); limit > 0 && headerCount > limit { + hdec.SetEmitEnabled(false) + mh.Truncated = true + remainSize = 0 + return + } if !httpguts.ValidHeaderFieldValue(hf.Value) { // Don't include the value in the error, because it may be sensitive. invalid = headerFieldValueError(hf.Name)
diff --git a/src/net/http/internal/http2/server.go b/src/net/http/internal/http2/server.go index 49abd85..aed43b2 100644 --- a/src/net/http/internal/http2/server.go +++ b/src/net/http/internal/http2/server.go
@@ -321,6 +321,7 @@ } fr.ReadMetaHeaders = hpack.NewDecoder(uint32(conf.MaxDecoderHeaderTableSize), nil) fr.MaxHeaderListSize = sc.maxHeaderListSize() + fr.MaxHeaderValueCount = sc.hs.MaxHeaderValueCount() fr.SetMaxReadFrameSize(uint32(conf.MaxReadFrameSize)) sc.framer = fr @@ -2030,6 +2031,9 @@ if len(f.PseudoFields()) > 0 { return sc.countError("trailers_pseudo", streamError(st.id, ErrCodeProtocol)) } + if f.Truncated { + return sc.countError("trailers_too_large", streamError(st.id, ErrCodeProtocol)) + } if st.trailer != nil { for _, hf := range f.RegularFields() { key := sc.canonicalHeader(hf.Name)
diff --git a/src/net/http/request.go b/src/net/http/request.go index a95e991..e07088f 100644 --- a/src/net/http/request.go +++ b/src/net/http/request.go
@@ -16,6 +16,7 @@ "fmt" "io" "maps" + "math" "mime" "mime/multipart" "net/http/httptrace" @@ -1074,6 +1075,11 @@ return req, nil } +// readMIMEHeader is defined in package [net/textproto]. +// +//go:linkname readMIMEHeader net/textproto.readMIMEHeader +func readMIMEHeader(r *textproto.Reader, maxMemory, maxHeaders int64) (textproto.MIMEHeader, error) + // readRequest should be an internal detail, // but widely used packages access it using linkname. // Notable members of the hall of shame include: @@ -1086,6 +1092,10 @@ // //go:linkname readRequest func readRequest(b *bufio.Reader) (req *Request, err error) { + return readRequestLimit(b, math.MaxInt64) +} + +func readRequestLimit(b *bufio.Reader, maxHeaders int64) (req *Request, err error) { tp := newTextprotoReader(b) defer putTextprotoReader(tp) @@ -1139,8 +1149,12 @@ } // Subsequent lines: Key: value. - mimeHeader, err := tp.ReadMIMEHeader() + mimeHeader, err := readMIMEHeader(tp, math.MaxInt64, maxHeaders) if err != nil { + // TODO: Add a distinguishable error to net/textproto. + if err.Error() == "message too large" { + return nil, errTooLarge + } return nil, err } req.Header = Header(mimeHeader)
diff --git a/src/net/http/serve_test.go b/src/net/http/serve_test.go index d879300..d8c78a3 100644 --- a/src/net/http/serve_test.go +++ b/src/net/http/serve_test.go
@@ -3355,12 +3355,16 @@ func TestRequestLimit(t *testing.T) { run(t, testRequestLimit, http3SkippedMode) } func testRequestLimit(t *testing.T, mode testMode) { + bytesPerHeader := len("header12345: val12345\r\n") + numHeaders := ((DefaultMaxHeaderBytes + 4096) / bytesPerHeader) + 1 + cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) { t.Fatalf("didn't expect to get request in Handler") - }), optQuietLog) + }), func(s *Server) { + s.MaxHeaderValueCount = numHeaders + }, optQuietLog) req, _ := NewRequest("GET", cst.ts.URL, nil) - var bytesPerHeader = len("header12345: val12345\r\n") - for i := 0; i < ((DefaultMaxHeaderBytes+4096)/bytesPerHeader)+1; i++ { + for i := range numHeaders { req.Header.Set(fmt.Sprintf("header%05d", i), fmt.Sprintf("val%05d", i)) } res, err := cst.c.Do(req) @@ -3389,6 +3393,164 @@ } } +func TestRequestHeaderValueCountLimit(t *testing.T) { + run(t, testRequestHeaderValueCountLimit, http3SkippedMode) +} +func testRequestHeaderValueCountLimit(t *testing.T, mode testMode) { + tests := []struct { + name string + limit int + setup func(req *Request) + wantStatus int + }{ + { + name: "below limit", + limit: 15, + setup: func(req *Request) { + // Send considerably below the limit, to account for the client + // automatically adding pseudo-headers and headers that it can + // infer. + for i := range 5 { + req.Header.Add(fmt.Sprintf("X-Header-%d", i), "val") + } + }, + wantStatus: 200, + }, + { + name: "above limit", + limit: 15, + setup: func(req *Request) { + for i := range 16 { + req.Header.Add(fmt.Sprintf("X-Header-%d", i), "val") + } + }, + wantStatus: 431, + }, + { + name: "comma separated values count as one", + limit: 15, + setup: func(req *Request) { + vals := make([]string, 16) + for i := range vals { + vals[i] = "val" + } + req.Header.Add("X-Comma", strings.Join(vals, ", ")) + }, + wantStatus: 200, + }, + { + name: "multiple values count as multiple", + limit: 15, + setup: func(req *Request) { + for range 16 { + req.Header.Add("X-Repeated", "val") + } + }, + wantStatus: 431, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) { + w.WriteHeader(StatusOK) + }), func(s *Server) { + s.MaxHeaderValueCount = tt.limit + }, optQuietLog) + + req, _ := NewRequest("GET", cst.ts.URL, nil) + tt.setup(req) + + res, err := cst.c.Do(req) + if err != nil { + t.Fatal(err) + } + defer res.Body.Close() + if res.StatusCode != tt.wantStatus { + t.Errorf("got status %d, want %d", res.StatusCode, tt.wantStatus) + } + }) + } +} + +func TestRequestTrailerHeaderValueCountLimit(t *testing.T) { + // HTTP/1 has a static limit for trailer headers that are not affected by + // settings such as MaxHeaderBytes and MaxHeaderValueCount. + run(t, testRequestTrailerHeaderValueCountLimit, []testMode{http2Mode}) +} +func testRequestTrailerHeaderValueCountLimit(t *testing.T, mode testMode) { + tests := []struct { + name string + limit int + setup func(req *Request) + wantErr bool + }{ + { + name: "below limit", + limit: 15, + setup: func(req *Request) { + req.Trailer = make(Header) + for i := range 14 { + req.Trailer.Add(fmt.Sprintf("X-Trailer-%d", i), "val") + } + }, + }, + { + name: "above limit", + limit: 15, + setup: func(req *Request) { + req.Trailer = make(Header) + for i := range 16 { + req.Trailer.Add(fmt.Sprintf("X-Trailer-%d", i), "val") + } + }, + wantErr: true, + }, + { + name: "comma separated values count as one", + limit: 15, + setup: func(req *Request) { + req.Trailer = make(Header) + vals := make([]string, 16) + for i := range vals { + vals[i] = "val" + } + req.Trailer.Add("X-Comma-Trailer", strings.Join(vals, ", ")) + }, + }, + { + name: "multiple values count as multiple", + limit: 15, + setup: func(req *Request) { + req.Trailer = make(Header) + for range 16 { + req.Trailer.Add("X-Repeated-Trailer", "val") + } + }, + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) { + io.Copy(io.Discard, r.Body) + }), func(s *Server) { + s.MaxHeaderValueCount = tt.limit + }, optQuietLog) + + req, _ := NewRequest("GET", cst.ts.URL, strings.NewReader("some body")) + tt.setup(req) + + res, err := cst.c.Do(req) + if (err != nil) != tt.wantErr { + t.Fatalf("got error %v, want %v", err, tt.wantErr) + } + if err == nil { + res.Body.Close() + } + }) + } +} + type neverEnding byte func (b neverEnding) Read(p []byte) (n int, err error) {
diff --git a/src/net/http/server.go b/src/net/http/server.go index 2c28abd..ae21427 100644 --- a/src/net/http/server.go +++ b/src/net/http/server.go
@@ -938,6 +938,11 @@ // This can be overridden by setting [Server.MaxHeaderBytes]. const DefaultMaxHeaderBytes = 1 << 20 // 1 MB +// DefaultMaxHeaderValueCount is the maximum permitted number of +// header values in an HTTP request. +// This can be overridden by setting [Server.MaxHeaderValueCount]. +const DefaultMaxHeaderValueCount = 500 + func (s *Server) maxHeaderBytes() int { if s.MaxHeaderBytes > 0 { return s.MaxHeaderBytes @@ -945,6 +950,13 @@ return DefaultMaxHeaderBytes } +func (s *Server) maxHeaderValueCount() int { + if s.MaxHeaderValueCount > 0 { + return s.MaxHeaderValueCount + } + return DefaultMaxHeaderValueCount +} + func (s *Server) initialReadLimitSize() int64 { return int64(s.maxHeaderBytes()) + 4096 // bufio slop } @@ -1044,7 +1056,7 @@ peek, _ := c.bufr.Peek(4) // ReadRequest will get err below c.bufr.Discard(numLeadingCRorLF(peek)) } - req, err := readRequest(c.bufr) + req, err := readRequestLimit(c.bufr, int64(c.server.maxHeaderValueCount())) if err != nil { if c.r.hitReadLimit() { return nil, errTooLarge @@ -3129,6 +3141,14 @@ // If zero, DefaultMaxHeaderBytes is used. MaxHeaderBytes int + // MaxHeaderValueCount controls the maximum number of header + // values that the server is willing to parse from a request. + // If zero, DefaultMaxHeaderValueCount is used. + // Note that comma-separated values in a single header line are + // counted once, while values sent as multiple header lines are + // counted multiple times. + MaxHeaderValueCount int + // TLSNextProto optionally specifies a function to take over // ownership of the provided TLS connection when an ALPN // protocol upgrade has occurred. The map key is the protocol
diff --git a/src/net/textproto/reader.go b/src/net/textproto/reader.go index 997e42a..b9ec465 100644 --- a/src/net/textproto/reader.go +++ b/src/net/textproto/reader.go
@@ -18,7 +18,7 @@ ) // TODO: This should be a distinguishable error (ErrMessageTooLarge) -// to allow mime/multipart to detect it. +// to allow mime/multipart and net/http to detect it. var errMessageTooLarge = errors.New("message too large") // A Reader implements convenience methods for reading requests @@ -508,11 +508,18 @@ return readMIMEHeader(r, math.MaxInt64, math.MaxInt64) } -// readMIMEHeader is accessed from mime/multipart. +// readMIMEHeader should be an internal detail, +// but widely used packages access it using linkname. +// Notable members of the hall of shame include: +// - github.com/qtgolang/SunnyNet +// +// Do not remove or change the type signature. +// See go.dev/issue/67401. +// //go:linkname readMIMEHeader // readMIMEHeader is a version of ReadMIMEHeader which takes a limit on the header size. -// It is called by the mime/multipart package. +// It is called by the mime/multipart and net/http package. func readMIMEHeader(r *Reader, maxMemory, maxHeaders int64) (MIMEHeader, error) { // Avoid lots of small slice allocations later by allocating one // large one ahead of time which we'll cut up into smaller