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