Source file src/net/http/clientserver_test.go

     1  // Copyright 2015 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  // Tests that use both the client & server, in both HTTP/1 and HTTP/2 mode.
     6  
     7  package http_test
     8  
     9  import (
    10  	"bytes"
    11  	"compress/gzip"
    12  	"context"
    13  	"crypto/rand"
    14  	"crypto/sha1"
    15  	"crypto/tls"
    16  	"fmt"
    17  	"hash"
    18  	"io"
    19  	"log"
    20  	"maps"
    21  	"net"
    22  	"net/http"
    23  	. "net/http"
    24  	"net/http/httptest"
    25  	"net/http/httptrace"
    26  	"net/http/httputil"
    27  	"net/textproto"
    28  	"net/url"
    29  	"os"
    30  	"reflect"
    31  	"runtime"
    32  	"slices"
    33  	"strconv"
    34  	"strings"
    35  	"sync"
    36  	"sync/atomic"
    37  	"testing"
    38  	"testing/synctest"
    39  	"time"
    40  
    41  	"golang.org/x/net/quic"
    42  
    43  	_ "unsafe" // for linkname
    44  
    45  	_ "golang.org/x/net/http3"
    46  )
    47  
    48  //go:linkname registerHTTP3Transport
    49  func registerHTTP3Transport(*http.Transport) <-chan *quic.Endpoint
    50  
    51  //go:linkname registerHTTP3Server
    52  func registerHTTP3Server(*http.Server) <-chan *quic.Endpoint
    53  
    54  type testMode string
    55  
    56  const (
    57  	http1Mode            = testMode("h1")            // HTTP/1.1
    58  	https1Mode           = testMode("https1")        // HTTPS/1.1
    59  	http2Mode            = testMode("h2")            // HTTP/2
    60  	http2UnencryptedMode = testMode("h2unencrypted") // HTTP/2
    61  	http3Mode            = testMode("h3")            // HTTP/3
    62  )
    63  
    64  type (
    65  	testAddMode  []testMode // default, plus these
    66  	testSkipMode []testMode // default, minus these
    67  )
    68  
    69  // http3SkippedMode is a convenient alias for []testMode{http1Mode, http2Mode},
    70  // which was the default test mode used by run and runSynctest prior to HTTP/3
    71  // development.
    72  // As we work on getting net/http tests to pass for our x/net HTTP/3
    73  // implementation, tests that still use http3SkippedMode are essentially a list
    74  // of TODOs on what work needs to be done for our HTTP/3 implementation to
    75  // reach basic feature parity with our HTTP/1 and HTTP/2 implementations
    76  var http3SkippedMode = []testMode{http1Mode, http2Mode}
    77  
    78  func (m testMode) Scheme() string {
    79  	switch m {
    80  	case http1Mode, http2UnencryptedMode:
    81  		return "http"
    82  	case https1Mode, http2Mode, http3Mode:
    83  		return "https"
    84  	}
    85  	panic("unknown testMode")
    86  }
    87  
    88  type testNotParallelOpt struct{}
    89  
    90  var (
    91  	testNotParallel = testNotParallelOpt{}
    92  )
    93  
    94  type TBRun[T any] interface {
    95  	testing.TB
    96  	Run(string, func(T)) bool
    97  }
    98  
    99  // run runs a client/server test in a variety of test configurations.
   100  //
   101  // Tests execute in HTTP/1.1 and HTTP/2 modes by default.
   102  // To run in a different set of configurations, pass a []testMode option.
   103  //
   104  // Tests call t.Parallel() by default.
   105  // To disable parallel execution, pass the testNotParallel option.
   106  func run[T TBRun[T]](t T, f func(t T, mode testMode), opts ...any) {
   107  	t.Helper()
   108  	modes := []testMode{http1Mode, http2Mode, http3Mode}
   109  	parallel := true
   110  	for _, opt := range opts {
   111  		switch opt := opt.(type) {
   112  		case testAddMode:
   113  			for _, m := range opt {
   114  				if !slices.Contains(modes, m) {
   115  					modes = append(modes, m)
   116  				}
   117  			}
   118  		case testSkipMode:
   119  			modes = slices.DeleteFunc(modes, func(m testMode) bool {
   120  				return slices.Contains(opt, m)
   121  			})
   122  		case []testMode:
   123  			modes = opt
   124  		case testNotParallelOpt:
   125  			parallel = false
   126  		default:
   127  			t.Fatalf("unknown option type %T", opt)
   128  		}
   129  	}
   130  	if t, ok := any(t).(*testing.T); ok && parallel {
   131  		setParallel(t)
   132  	}
   133  	for _, mode := range modes {
   134  		// TODO(nsh): re-enable the tests once tree re-opens.
   135  		if mode == http3Mode {
   136  			continue
   137  		}
   138  		t.Run(string(mode), func(t T) {
   139  			t.Helper()
   140  			if t, ok := any(t).(*testing.T); ok && parallel {
   141  				setParallel(t)
   142  			}
   143  			t.Cleanup(func() {
   144  				afterTest(t)
   145  			})
   146  			f(t, mode)
   147  		})
   148  	}
   149  }
   150  
   151  // runSynctest is run combined with synctest.Run.
   152  //
   153  // The TB passed to f arranges for cleanup functions to be run in the synctest bubble.
   154  func runSynctest(t *testing.T, f func(t *testing.T, mode testMode), opts ...any) {
   155  	run(t, func(t *testing.T, mode testMode) {
   156  		synctest.Test(t, func(t *testing.T) {
   157  			f(t, mode)
   158  		})
   159  	}, opts...)
   160  }
   161  
   162  type clientServerTest struct {
   163  	t  testing.TB
   164  	h2 bool
   165  	h  Handler
   166  	ts *httptest.Server
   167  	tr *Transport
   168  	c  *Client
   169  	li *fakeNetListener
   170  }
   171  
   172  func (t *clientServerTest) close() {
   173  	t.tr.CloseIdleConnections()
   174  	t.ts.Close()
   175  }
   176  
   177  func (t *clientServerTest) getURL(u string) string {
   178  	res, err := t.c.Get(u)
   179  	if err != nil {
   180  		t.t.Fatal(err)
   181  	}
   182  	defer res.Body.Close()
   183  	slurp, err := io.ReadAll(res.Body)
   184  	if err != nil {
   185  		t.t.Fatal(err)
   186  	}
   187  	return string(slurp)
   188  }
   189  
   190  func (t *clientServerTest) scheme() string {
   191  	if t.h2 {
   192  		return "https"
   193  	}
   194  	return "http"
   195  }
   196  
   197  var optQuietLog = func(ts *httptest.Server) {
   198  	ts.Config.ErrorLog = quietLog
   199  }
   200  
   201  func optWithServerLog(lg *log.Logger) func(*httptest.Server) {
   202  	return func(ts *httptest.Server) {
   203  		ts.Config.ErrorLog = lg
   204  	}
   205  }
   206  
   207  var optFakeNet = new(struct{})
   208  
   209  // newClientServerTest creates and starts an httptest.Server.
   210  //
   211  // The mode parameter selects the implementation to test:
   212  // HTTP/1, HTTP/2, etc. Tests using newClientServerTest should use
   213  // the 'run' function, which will start a subtests for each tested mode.
   214  //
   215  // The vararg opts parameter can include functions to configure the
   216  // test server or transport.
   217  //
   218  //	func(*httptest.Server) // run before starting the server
   219  //	func(*http.Transport)
   220  //
   221  // The optFakeNet option configures the server and client to use a fake network implementation,
   222  // suitable for use in testing/synctest tests.
   223  func newClientServerTest(t testing.TB, mode testMode, h Handler, opts ...any) *clientServerTest {
   224  	if mode == http2Mode || mode == http2UnencryptedMode {
   225  		CondSkipHTTP2(t)
   226  	}
   227  	cst := &clientServerTest{
   228  		t:  t,
   229  		h2: mode == http2Mode,
   230  		h:  h,
   231  	}
   232  
   233  	var transportFuncs []func(*Transport)
   234  
   235  	switch idx := slices.Index(opts, any(optFakeNet)); {
   236  	case idx >= 0:
   237  		opts = slices.Delete(opts, idx, idx+1)
   238  		cst.li = fakeNetListen()
   239  		cst.ts = &httptest.Server{
   240  			Config:   &Server{Handler: h},
   241  			Listener: cst.li,
   242  		}
   243  		transportFuncs = append(transportFuncs, func(tr *Transport) {
   244  			tr.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
   245  				return cst.li.connect(), nil
   246  			}
   247  		})
   248  	case mode == http3Mode:
   249  		// TODO: support testing HTTP/3 using fakenet.
   250  		cst.ts = &httptest.Server{
   251  			Config: &Server{Handler: h},
   252  		}
   253  	default:
   254  		cst.ts = httptest.NewUnstartedServer(h)
   255  	}
   256  
   257  	for _, opt := range opts {
   258  		switch opt := opt.(type) {
   259  		case func(*Transport):
   260  			transportFuncs = append(transportFuncs, opt)
   261  		case func(*httptest.Server):
   262  			opt(cst.ts)
   263  		case func(*Server):
   264  			opt(cst.ts.Config)
   265  		default:
   266  			t.Fatalf("unhandled option type %T", opt)
   267  		}
   268  	}
   269  
   270  	if cst.ts.Config.ErrorLog == nil {
   271  		cst.ts.Config.ErrorLog = log.New(testLogWriter{t}, "", 0)
   272  	}
   273  
   274  	p := &Protocols{}
   275  	if cst.ts.Config.Protocols == nil {
   276  		cst.ts.Config.Protocols = p
   277  	}
   278  	switch mode {
   279  	case http1Mode:
   280  		p.SetHTTP1(true)
   281  		cst.ts.Start()
   282  	case https1Mode:
   283  		p.SetHTTP1(true)
   284  		cst.ts.StartTLS()
   285  	case http2UnencryptedMode:
   286  		p.SetUnencryptedHTTP2(true)
   287  		cst.ts.Start()
   288  	case http2Mode:
   289  		p.SetHTTP2(true)
   290  		cst.ts.EnableHTTP2 = true
   291  		cst.ts.TLS = cst.ts.Config.TLSConfig
   292  		cst.ts.StartTLS()
   293  	case http3Mode:
   294  		http.ProtocolSetHTTP3(p)
   295  		cst.ts.TLS = cst.ts.Config.TLSConfig
   296  		cst.ts.StartTLS()
   297  		endpointCh := registerHTTP3Server(cst.ts.Config)
   298  
   299  		cst.ts.Config.TLSConfig = cst.ts.TLS
   300  		cst.ts.Config.Addr = ":0"
   301  		go cst.ts.Config.ListenAndServeTLS("", "")
   302  
   303  		endpoint := <-endpointCh
   304  		port := strconv.Itoa(int(endpoint.LocalAddr().Port()))
   305  		switch addr := endpoint.LocalAddr().Addr(); {
   306  		case !addr.IsUnspecified():
   307  			cst.ts.URL = "https://" + endpoint.LocalAddr().String()
   308  		case addr.Is4():
   309  			cst.ts.URL = "https://" + net.JoinHostPort("127.0.0.1", port)
   310  		case addr.Is6():
   311  			cst.ts.URL = "https://" + net.JoinHostPort("::1", port)
   312  		default:
   313  			t.Fatalf("unknown address family for %v", endpoint.LocalAddr())
   314  		}
   315  		t.Cleanup(func() {
   316  			// Give a relatively generous timeout. If the timeout is too short,
   317  			// the test might return before QUIC connections can finish closing
   318  			// asynchronously in some builders. The open connections will cause
   319  			// TestMain to detect a goroutine leak and fail.
   320  			ctx, cancel := context.WithTimeout(context.Background(), time.Second)
   321  			defer cancel()
   322  			cst.ts.Config.Shutdown(ctx)
   323  		})
   324  	default:
   325  		t.Fatalf("unknown test mode %v", mode)
   326  	}
   327  	cst.c = cst.ts.Client()
   328  	cst.tr = cst.c.Transport.(*Transport)
   329  	for _, f := range transportFuncs {
   330  		f(cst.tr)
   331  	}
   332  	if cst.tr.Protocols == nil {
   333  		cst.tr.Protocols = p
   334  	}
   335  	if mode == http3Mode {
   336  		endpointCh := registerHTTP3Transport(cst.tr)
   337  		testDoneCh := make(chan any)
   338  		var wg sync.WaitGroup
   339  		t.Cleanup(func() {
   340  			close(testDoneCh)
   341  			wg.Wait()
   342  		})
   343  		wg.Go(func() {
   344  			for {
   345  				select {
   346  				case e := <-endpointCh:
   347  					t.Cleanup(func() {
   348  						ctx, cancel := context.WithTimeout(context.Background(), time.Second)
   349  						defer cancel()
   350  						if e != nil {
   351  							e.Close(ctx)
   352  						}
   353  					})
   354  				case <-testDoneCh:
   355  					return
   356  				}
   357  			}
   358  		})
   359  	}
   360  
   361  	t.Cleanup(func() {
   362  		cst.close()
   363  	})
   364  	return cst
   365  }
   366  
   367  type testLogWriter struct {
   368  	t testing.TB
   369  }
   370  
   371  func (w testLogWriter) Write(b []byte) (int, error) {
   372  	w.t.Logf("server log: %v", strings.TrimSpace(string(b)))
   373  	return len(b), nil
   374  }
   375  
   376  // Testing the newClientServerTest helper itself.
   377  func TestNewClientServerTest(t *testing.T) {
   378  	modes := []testMode{http1Mode, https1Mode, http2Mode}
   379  	t.Run("realnet", func(t *testing.T) {
   380  		run(t, func(t *testing.T, mode testMode) {
   381  			testNewClientServerTest(t, mode)
   382  		}, modes)
   383  	})
   384  	t.Run("synctest", func(t *testing.T) {
   385  		runSynctest(t, func(t *testing.T, mode testMode) {
   386  			testNewClientServerTest(t, mode, optFakeNet)
   387  		}, modes)
   388  	})
   389  }
   390  func testNewClientServerTest(t *testing.T, mode testMode, opts ...any) {
   391  	var got struct {
   392  		sync.Mutex
   393  		proto  string
   394  		hasTLS bool
   395  	}
   396  	h := HandlerFunc(func(w ResponseWriter, r *Request) {
   397  		got.Lock()
   398  		defer got.Unlock()
   399  		got.proto = r.Proto
   400  		got.hasTLS = r.TLS != nil
   401  	})
   402  	cst := newClientServerTest(t, mode, h, opts...)
   403  	if _, err := cst.c.Head(cst.ts.URL); err != nil {
   404  		t.Fatal(err)
   405  	}
   406  	var wantProto string
   407  	var wantTLS bool
   408  	switch mode {
   409  	case http1Mode:
   410  		wantProto = "HTTP/1.1"
   411  		wantTLS = false
   412  	case https1Mode:
   413  		wantProto = "HTTP/1.1"
   414  		wantTLS = true
   415  	case http2Mode:
   416  		wantProto = "HTTP/2.0"
   417  		wantTLS = true
   418  	}
   419  	if got.proto != wantProto {
   420  		t.Errorf("req.Proto = %q, want %q", got.proto, wantProto)
   421  	}
   422  	if got.hasTLS != wantTLS {
   423  		t.Errorf("req.TLS set: %v, want %v", got.hasTLS, wantTLS)
   424  	}
   425  }
   426  
   427  func TestChunkedResponseHeaders(t *testing.T) {
   428  	run(t, testChunkedResponseHeaders, http3SkippedMode)
   429  }
   430  func testChunkedResponseHeaders(t *testing.T, mode testMode) {
   431  	log.SetOutput(io.Discard) // is noisy otherwise
   432  	defer log.SetOutput(os.Stderr)
   433  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
   434  		w.Header().Set("Content-Length", "intentional gibberish") // we check that this is deleted
   435  		w.(Flusher).Flush()
   436  		fmt.Fprintf(w, "I am a chunked response.")
   437  	}))
   438  
   439  	res, err := cst.c.Get(cst.ts.URL)
   440  	if err != nil {
   441  		t.Fatalf("Get error: %v", err)
   442  	}
   443  	defer res.Body.Close()
   444  	if g, e := res.ContentLength, int64(-1); g != e {
   445  		t.Errorf("expected ContentLength of %d; got %d", e, g)
   446  	}
   447  	wantTE := []string{"chunked"}
   448  	if mode == http2Mode {
   449  		wantTE = nil
   450  	}
   451  	if !slices.Equal(res.TransferEncoding, wantTE) {
   452  		t.Errorf("TransferEncoding = %v; want %v", res.TransferEncoding, wantTE)
   453  	}
   454  	if got, haveCL := res.Header["Content-Length"]; haveCL {
   455  		t.Errorf("Unexpected Content-Length: %q", got)
   456  	}
   457  }
   458  
   459  type reqFunc func(c *Client, url string) (*Response, error)
   460  
   461  // h12Compare is a test that compares HTTP/1 and HTTP/2 behavior
   462  // against each other.
   463  type h12Compare struct {
   464  	Handler            func(ResponseWriter, *Request)    // required
   465  	ReqFunc            reqFunc                           // optional
   466  	CheckResponse      func(proto string, res *Response) // optional
   467  	EarlyCheckResponse func(proto string, res *Response) // optional; pre-normalize
   468  	Opts               []any
   469  }
   470  
   471  func (tt h12Compare) reqFunc() reqFunc {
   472  	if tt.ReqFunc == nil {
   473  		return (*Client).Get
   474  	}
   475  	return tt.ReqFunc
   476  }
   477  
   478  func (tt h12Compare) run(t *testing.T) {
   479  	setParallel(t)
   480  	cst1 := newClientServerTest(t, http1Mode, HandlerFunc(tt.Handler), tt.Opts...)
   481  	defer cst1.close()
   482  	cst2 := newClientServerTest(t, http2Mode, HandlerFunc(tt.Handler), tt.Opts...)
   483  	defer cst2.close()
   484  
   485  	res1, err := tt.reqFunc()(cst1.c, cst1.ts.URL)
   486  	if err != nil {
   487  		t.Errorf("HTTP/1 request: %v", err)
   488  		return
   489  	}
   490  	res2, err := tt.reqFunc()(cst2.c, cst2.ts.URL)
   491  	if err != nil {
   492  		t.Errorf("HTTP/2 request: %v", err)
   493  		return
   494  	}
   495  
   496  	if fn := tt.EarlyCheckResponse; fn != nil {
   497  		fn("HTTP/1.1", res1)
   498  		fn("HTTP/2.0", res2)
   499  	}
   500  
   501  	tt.normalizeRes(t, res1, "HTTP/1.1")
   502  	tt.normalizeRes(t, res2, "HTTP/2.0")
   503  	res1body, res2body := res1.Body, res2.Body
   504  
   505  	eres1 := mostlyCopy(res1)
   506  	eres2 := mostlyCopy(res2)
   507  	if !reflect.DeepEqual(eres1, eres2) {
   508  		t.Errorf("Response headers to handler differed:\nhttp/1 (%v):\n\t%#v\nhttp/2 (%v):\n\t%#v",
   509  			cst1.ts.URL, eres1, cst2.ts.URL, eres2)
   510  	}
   511  	if !reflect.DeepEqual(res1body, res2body) {
   512  		t.Errorf("Response bodies to handler differed.\nhttp1: %v\nhttp2: %v\n", res1body, res2body)
   513  	}
   514  	if fn := tt.CheckResponse; fn != nil {
   515  		res1.Body, res2.Body = res1body, res2body
   516  		fn("HTTP/1.1", res1)
   517  		fn("HTTP/2.0", res2)
   518  	}
   519  }
   520  
   521  func mostlyCopy(r *Response) *Response {
   522  	c := *r
   523  	c.Body = nil
   524  	c.TransferEncoding = nil
   525  	c.TLS = nil
   526  	c.Request = nil
   527  	return &c
   528  }
   529  
   530  type slurpResult struct {
   531  	io.ReadCloser
   532  	body []byte
   533  	err  error
   534  }
   535  
   536  func (sr slurpResult) String() string { return fmt.Sprintf("body %q; err %v", sr.body, sr.err) }
   537  
   538  func (tt h12Compare) normalizeRes(t *testing.T, res *Response, wantProto string) {
   539  	if res.Proto == wantProto || res.Proto == "HTTP/IGNORE" {
   540  		res.Proto, res.ProtoMajor, res.ProtoMinor = "", 0, 0
   541  	} else {
   542  		t.Errorf("got %q response; want %q", res.Proto, wantProto)
   543  	}
   544  	slurp, err := io.ReadAll(res.Body)
   545  
   546  	res.Body.Close()
   547  	res.Body = slurpResult{
   548  		ReadCloser: io.NopCloser(bytes.NewReader(slurp)),
   549  		body:       slurp,
   550  		err:        err,
   551  	}
   552  	for i, v := range res.Header["Date"] {
   553  		res.Header["Date"][i] = strings.Repeat("x", len(v))
   554  	}
   555  	if res.Request == nil {
   556  		t.Errorf("for %s, no request", wantProto)
   557  	}
   558  	if (res.TLS != nil) != (wantProto == "HTTP/2.0") {
   559  		t.Errorf("TLS set = %v; want %v", res.TLS != nil, res.TLS == nil)
   560  	}
   561  }
   562  
   563  // Issue 13532
   564  func TestH12_HeadContentLengthNoBody(t *testing.T) {
   565  	h12Compare{
   566  		ReqFunc: (*Client).Head,
   567  		Handler: func(w ResponseWriter, r *Request) {
   568  		},
   569  	}.run(t)
   570  }
   571  
   572  func TestH12_HeadContentLengthSmallBody(t *testing.T) {
   573  	h12Compare{
   574  		ReqFunc: (*Client).Head,
   575  		Handler: func(w ResponseWriter, r *Request) {
   576  			io.WriteString(w, "small")
   577  		},
   578  	}.run(t)
   579  }
   580  
   581  func TestH12_HeadContentLengthLargeBody(t *testing.T) {
   582  	h12Compare{
   583  		ReqFunc: (*Client).Head,
   584  		Handler: func(w ResponseWriter, r *Request) {
   585  			chunk := strings.Repeat("x", 512<<10)
   586  			for i := 0; i < 10; i++ {
   587  				io.WriteString(w, chunk)
   588  			}
   589  		},
   590  	}.run(t)
   591  }
   592  
   593  func TestH12_200NoBody(t *testing.T) {
   594  	h12Compare{Handler: func(w ResponseWriter, r *Request) {}}.run(t)
   595  }
   596  
   597  func TestH2_204NoBody(t *testing.T) { testH12_noBody(t, 204) }
   598  func TestH2_304NoBody(t *testing.T) { testH12_noBody(t, 304) }
   599  func TestH2_404NoBody(t *testing.T) { testH12_noBody(t, 404) }
   600  
   601  func testH12_noBody(t *testing.T, status int) {
   602  	h12Compare{Handler: func(w ResponseWriter, r *Request) {
   603  		w.WriteHeader(status)
   604  	}}.run(t)
   605  }
   606  
   607  func TestH12_SmallBody(t *testing.T) {
   608  	h12Compare{Handler: func(w ResponseWriter, r *Request) {
   609  		io.WriteString(w, "small body")
   610  	}}.run(t)
   611  }
   612  
   613  func TestH12_ExplicitContentLength(t *testing.T) {
   614  	h12Compare{Handler: func(w ResponseWriter, r *Request) {
   615  		w.Header().Set("Content-Length", "3")
   616  		io.WriteString(w, "foo")
   617  	}}.run(t)
   618  }
   619  
   620  func TestH12_FlushBeforeBody(t *testing.T) {
   621  	h12Compare{Handler: func(w ResponseWriter, r *Request) {
   622  		w.(Flusher).Flush()
   623  		io.WriteString(w, "foo")
   624  	}}.run(t)
   625  }
   626  
   627  func TestH12_FlushMidBody(t *testing.T) {
   628  	h12Compare{Handler: func(w ResponseWriter, r *Request) {
   629  		io.WriteString(w, "foo")
   630  		w.(Flusher).Flush()
   631  		io.WriteString(w, "bar")
   632  	}}.run(t)
   633  }
   634  
   635  func TestH12_Head_ExplicitLen(t *testing.T) {
   636  	h12Compare{
   637  		ReqFunc: (*Client).Head,
   638  		Handler: func(w ResponseWriter, r *Request) {
   639  			if r.Method != "HEAD" {
   640  				t.Errorf("unexpected method %q", r.Method)
   641  			}
   642  			w.Header().Set("Content-Length", "1235")
   643  		},
   644  	}.run(t)
   645  }
   646  
   647  func TestH12_Head_ImplicitLen(t *testing.T) {
   648  	h12Compare{
   649  		ReqFunc: (*Client).Head,
   650  		Handler: func(w ResponseWriter, r *Request) {
   651  			if r.Method != "HEAD" {
   652  				t.Errorf("unexpected method %q", r.Method)
   653  			}
   654  			io.WriteString(w, "foo")
   655  		},
   656  	}.run(t)
   657  }
   658  
   659  func TestH12_HandlerWritesTooLittle(t *testing.T) {
   660  	h12Compare{
   661  		Handler: func(w ResponseWriter, r *Request) {
   662  			w.Header().Set("Content-Length", "3")
   663  			io.WriteString(w, "12") // one byte short
   664  		},
   665  		CheckResponse: func(proto string, res *Response) {
   666  			sr, ok := res.Body.(slurpResult)
   667  			if !ok {
   668  				t.Errorf("%s body is %T; want slurpResult", proto, res.Body)
   669  				return
   670  			}
   671  			if sr.err != io.ErrUnexpectedEOF {
   672  				t.Errorf("%s read error = %v; want io.ErrUnexpectedEOF", proto, sr.err)
   673  			}
   674  			if string(sr.body) != "12" {
   675  				t.Errorf("%s body = %q; want %q", proto, sr.body, "12")
   676  			}
   677  		},
   678  	}.run(t)
   679  }
   680  
   681  // Tests that the HTTP/1 and HTTP/2 servers prevent handlers from
   682  // writing more than they declared. This test does not test whether
   683  // the transport deals with too much data, though, since the server
   684  // doesn't make it possible to send bogus data. For those tests, see
   685  // transport_test.go (for HTTP/1) or x/net/http2/transport_test.go
   686  // (for HTTP/2).
   687  func TestHandlerWritesTooMuch(t *testing.T) { run(t, testHandlerWritesTooMuch) }
   688  func testHandlerWritesTooMuch(t *testing.T, mode testMode) {
   689  	wantBody := []byte("123")
   690  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
   691  		rc := NewResponseController(w)
   692  		w.Header().Set("Content-Length", fmt.Sprintf("%v", len(wantBody)))
   693  		rc.Flush()
   694  		w.Write(wantBody)
   695  		rc.Flush()
   696  		n, err := io.WriteString(w, "x") // too many
   697  		if err == nil {
   698  			err = rc.Flush()
   699  		}
   700  		// TODO: Check that this is ErrContentLength, not just any error.
   701  		if err == nil {
   702  			t.Errorf("for proto %q, final write = %v, %v; want _, some error", r.Proto, n, err)
   703  		}
   704  	}))
   705  
   706  	res, err := cst.c.Get(cst.ts.URL)
   707  	if err != nil {
   708  		t.Fatal(err)
   709  	}
   710  	defer res.Body.Close()
   711  
   712  	gotBody, _ := io.ReadAll(res.Body)
   713  	if !bytes.Equal(gotBody, wantBody) {
   714  		t.Fatalf("got response body: %q; want %q", gotBody, wantBody)
   715  	}
   716  }
   717  
   718  // Verify that both our HTTP/1 and HTTP/2 request and auto-decompress gzip.
   719  // Some hosts send gzip even if you don't ask for it; see golang.org/issue/13298
   720  func TestH12_AutoGzip(t *testing.T) {
   721  	h12Compare{
   722  		Handler: func(w ResponseWriter, r *Request) {
   723  			if ae := r.Header.Get("Accept-Encoding"); ae != "gzip" {
   724  				t.Errorf("%s Accept-Encoding = %q; want gzip", r.Proto, ae)
   725  			}
   726  			w.Header().Set("Content-Encoding", "gzip")
   727  			gz := gzip.NewWriter(w)
   728  			io.WriteString(gz, "I am some gzipped content. Go go go go go go go go go go go go should compress well.")
   729  			gz.Close()
   730  		},
   731  	}.run(t)
   732  }
   733  
   734  func TestH12_AutoGzip_Disabled(t *testing.T) {
   735  	h12Compare{
   736  		Opts: []any{
   737  			func(tr *Transport) { tr.DisableCompression = true },
   738  		},
   739  		Handler: func(w ResponseWriter, r *Request) {
   740  			fmt.Fprintf(w, "%q", r.Header["Accept-Encoding"])
   741  			if ae := r.Header.Get("Accept-Encoding"); ae != "" {
   742  				t.Errorf("%s Accept-Encoding = %q; want empty", r.Proto, ae)
   743  			}
   744  		},
   745  	}.run(t)
   746  }
   747  
   748  // Test304Responses verifies that 304s don't declare that they're
   749  // chunking in their response headers and aren't allowed to produce
   750  // output.
   751  func Test304Responses(t *testing.T) { run(t, test304Responses) }
   752  func test304Responses(t *testing.T, mode testMode) {
   753  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
   754  		w.WriteHeader(StatusNotModified)
   755  		_, err := w.Write([]byte("illegal body"))
   756  		if err != ErrBodyNotAllowed {
   757  			t.Errorf("on Write, expected ErrBodyNotAllowed, got %v", err)
   758  		}
   759  	}))
   760  	defer cst.close()
   761  	res, err := cst.c.Get(cst.ts.URL)
   762  	if err != nil {
   763  		t.Fatal(err)
   764  	}
   765  	if len(res.TransferEncoding) > 0 {
   766  		t.Errorf("expected no TransferEncoding; got %v", res.TransferEncoding)
   767  	}
   768  	body, err := io.ReadAll(res.Body)
   769  	if err != nil {
   770  		t.Error(err)
   771  	}
   772  	if len(body) > 0 {
   773  		t.Errorf("got unexpected body %q", string(body))
   774  	}
   775  }
   776  
   777  func TestH12_ServerEmptyContentLength(t *testing.T) {
   778  	h12Compare{
   779  		Handler: func(w ResponseWriter, r *Request) {
   780  			w.Header()["Content-Type"] = []string{""}
   781  			io.WriteString(w, "<html><body>hi</body></html>")
   782  		},
   783  	}.run(t)
   784  }
   785  
   786  func TestH12_RequestContentLength_Known_NonZero(t *testing.T) {
   787  	h12requestContentLength(t, func() io.Reader { return strings.NewReader("FOUR") }, 4)
   788  }
   789  
   790  func TestH12_RequestContentLength_Known_Zero(t *testing.T) {
   791  	h12requestContentLength(t, func() io.Reader { return nil }, 0)
   792  }
   793  
   794  func TestH12_RequestContentLength_Unknown(t *testing.T) {
   795  	h12requestContentLength(t, func() io.Reader { return struct{ io.Reader }{strings.NewReader("Stuff")} }, -1)
   796  }
   797  
   798  func h12requestContentLength(t *testing.T, bodyfn func() io.Reader, wantLen int64) {
   799  	h12Compare{
   800  		Handler: func(w ResponseWriter, r *Request) {
   801  			w.Header().Set("Got-Length", fmt.Sprint(r.ContentLength))
   802  			fmt.Fprintf(w, "Req.ContentLength=%v", r.ContentLength)
   803  		},
   804  		ReqFunc: func(c *Client, url string) (*Response, error) {
   805  			return c.Post(url, "text/plain", bodyfn())
   806  		},
   807  		CheckResponse: func(proto string, res *Response) {
   808  			if got, want := res.Header.Get("Got-Length"), fmt.Sprint(wantLen); got != want {
   809  				t.Errorf("Proto %q got length %q; want %q", proto, got, want)
   810  			}
   811  		},
   812  	}.run(t)
   813  }
   814  
   815  // Tests that closing the Request.Cancel channel also while still
   816  // reading the response body. Issue 13159.
   817  func TestCancelRequestMidBody(t *testing.T) { run(t, testCancelRequestMidBody, http3SkippedMode) }
   818  func testCancelRequestMidBody(t *testing.T, mode testMode) {
   819  	unblock := make(chan bool)
   820  	didFlush := make(chan bool, 1)
   821  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
   822  		io.WriteString(w, "Hello")
   823  		w.(Flusher).Flush()
   824  		didFlush <- true
   825  		<-unblock
   826  		io.WriteString(w, ", world.")
   827  	}))
   828  	defer close(unblock)
   829  
   830  	req, _ := NewRequest("GET", cst.ts.URL, nil)
   831  	cancel := make(chan struct{})
   832  	req.Cancel = cancel
   833  
   834  	res, err := cst.c.Do(req)
   835  	if err != nil {
   836  		t.Fatal(err)
   837  	}
   838  	defer res.Body.Close()
   839  	<-didFlush
   840  
   841  	// Read a bit before we cancel. (Issue 13626)
   842  	// We should have "Hello" at least sitting there.
   843  	firstRead := make([]byte, 10)
   844  	n, err := res.Body.Read(firstRead)
   845  	if err != nil {
   846  		t.Fatal(err)
   847  	}
   848  	firstRead = firstRead[:n]
   849  
   850  	close(cancel)
   851  
   852  	rest, err := io.ReadAll(res.Body)
   853  	all := string(firstRead) + string(rest)
   854  	if all != "Hello" {
   855  		t.Errorf("Read %q (%q + %q); want Hello", all, firstRead, rest)
   856  	}
   857  	if err != ExportErrRequestCanceled {
   858  		t.Errorf("ReadAll error = %v; want %v", err, ExportErrRequestCanceled)
   859  	}
   860  }
   861  
   862  // Tests that clients can send trailers to a server and that the server can read them.
   863  func TestTrailersClientToServer(t *testing.T) { run(t, testTrailersClientToServer) }
   864  func testTrailersClientToServer(t *testing.T, mode testMode) {
   865  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
   866  		slurp, err := io.ReadAll(r.Body)
   867  		if err != nil {
   868  			t.Errorf("Server reading request body: %v", err)
   869  		}
   870  		if string(slurp) != "foo" {
   871  			t.Errorf("Server read request body %q; want foo", slurp)
   872  		}
   873  		if r.Trailer == nil {
   874  			io.WriteString(w, "nil Trailer")
   875  		} else {
   876  			decl := slices.Sorted(maps.Keys(r.Trailer))
   877  			fmt.Fprintf(w, "decl: %v, vals: %s, %s",
   878  				decl,
   879  				r.Trailer.Get("Client-Trailer-A"),
   880  				r.Trailer.Get("Client-Trailer-B"))
   881  		}
   882  	}))
   883  
   884  	var req *Request
   885  	req, _ = NewRequest("POST", cst.ts.URL, io.MultiReader(
   886  		eofReaderFunc(func() {
   887  			req.Trailer["Client-Trailer-A"] = []string{"valuea"}
   888  		}),
   889  		strings.NewReader("foo"),
   890  		eofReaderFunc(func() {
   891  			req.Trailer["Client-Trailer-B"] = []string{"valueb"}
   892  		}),
   893  	))
   894  	req.Trailer = Header{
   895  		"Client-Trailer-A": nil, //  to be set later
   896  		"Client-Trailer-B": nil, //  to be set later
   897  	}
   898  	req.ContentLength = -1
   899  	res, err := cst.c.Do(req)
   900  	if err != nil {
   901  		t.Fatal(err)
   902  	}
   903  	if err := wantBody(res, err, "decl: [Client-Trailer-A Client-Trailer-B], vals: valuea, valueb"); err != nil {
   904  		t.Error(err)
   905  	}
   906  }
   907  
   908  // Tests that servers send trailers to a client and that the client can read them.
   909  func TestTrailersServerToClient(t *testing.T) {
   910  	run(t, func(t *testing.T, mode testMode) {
   911  		testTrailersServerToClient(t, mode, false)
   912  	}, http3SkippedMode)
   913  }
   914  func TestTrailersServerToClientFlush(t *testing.T) {
   915  	run(t, func(t *testing.T, mode testMode) {
   916  		testTrailersServerToClient(t, mode, true)
   917  	})
   918  }
   919  
   920  func testTrailersServerToClient(t *testing.T, mode testMode, flush bool) {
   921  	const body = "Some body"
   922  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
   923  		w.Header().Set("Trailer", "Server-Trailer-A, Server-Trailer-B")
   924  		w.Header().Add("Trailer", "Server-Trailer-C")
   925  
   926  		io.WriteString(w, body)
   927  		if flush {
   928  			w.(Flusher).Flush()
   929  		}
   930  
   931  		// How handlers set Trailers: declare it ahead of time
   932  		// with the Trailer header, and then mutate the
   933  		// Header() of those values later, after the response
   934  		// has been written (we wrote to w above).
   935  		w.Header().Set("Server-Trailer-A", "valuea")
   936  		w.Header().Set("Server-Trailer-C", "valuec") // skipping B
   937  		w.Header().Set("Server-Trailer-NotDeclared", "should be omitted")
   938  	}))
   939  
   940  	res, err := cst.c.Get(cst.ts.URL)
   941  	if err != nil {
   942  		t.Fatal(err)
   943  	}
   944  
   945  	wantHeader := Header{
   946  		"Content-Type": {"text/plain; charset=utf-8"},
   947  	}
   948  	wantLen := -1
   949  	if mode == http2Mode && !flush {
   950  		// In HTTP/1.1, any use of trailers forces HTTP/1.1
   951  		// chunking and a flush at the first write. That's
   952  		// unnecessary with HTTP/2's framing, so the server
   953  		// is able to calculate the length while still sending
   954  		// trailers afterwards.
   955  		wantLen = len(body)
   956  		wantHeader["Content-Length"] = []string{fmt.Sprint(wantLen)}
   957  	}
   958  	if res.ContentLength != int64(wantLen) {
   959  		t.Errorf("ContentLength = %v; want %v", res.ContentLength, wantLen)
   960  	}
   961  
   962  	delete(res.Header, "Date") // irrelevant for test
   963  	if !reflect.DeepEqual(res.Header, wantHeader) {
   964  		t.Errorf("Header = %v; want %v", res.Header, wantHeader)
   965  	}
   966  
   967  	if got, want := res.Trailer, (Header{
   968  		"Server-Trailer-A": nil,
   969  		"Server-Trailer-B": nil,
   970  		"Server-Trailer-C": nil,
   971  	}); !reflect.DeepEqual(got, want) {
   972  		t.Errorf("Trailer before body read = %v; want %v", got, want)
   973  	}
   974  
   975  	if err := wantBody(res, nil, body); err != nil {
   976  		t.Fatal(err)
   977  	}
   978  
   979  	if got, want := res.Trailer, (Header{
   980  		"Server-Trailer-A": {"valuea"},
   981  		"Server-Trailer-B": nil,
   982  		"Server-Trailer-C": {"valuec"},
   983  	}); !reflect.DeepEqual(got, want) {
   984  		t.Errorf("Trailer after body read = %v; want %v", got, want)
   985  	}
   986  }
   987  
   988  // Don't allow a Body.Read after Body.Close. Issue 13648.
   989  func TestResponseBodyReadAfterClose(t *testing.T) { run(t, testResponseBodyReadAfterClose) }
   990  func testResponseBodyReadAfterClose(t *testing.T, mode testMode) {
   991  	const body = "Some body"
   992  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
   993  		io.WriteString(w, body)
   994  	}))
   995  	res, err := cst.c.Get(cst.ts.URL)
   996  	if err != nil {
   997  		t.Fatal(err)
   998  	}
   999  	res.Body.Close()
  1000  	data, err := io.ReadAll(res.Body)
  1001  	if len(data) != 0 || err == nil {
  1002  		t.Fatalf("ReadAll returned %q, %v; want error", data, err)
  1003  	}
  1004  }
  1005  
  1006  func TestConcurrentReadWriteReqBody(t *testing.T) { run(t, testConcurrentReadWriteReqBody) }
  1007  func testConcurrentReadWriteReqBody(t *testing.T, mode testMode) {
  1008  	const reqBody = "some request body"
  1009  	const resBody = "some response body"
  1010  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1011  		var wg sync.WaitGroup
  1012  		wg.Add(2)
  1013  		didRead := make(chan bool, 1)
  1014  		// Read in one goroutine.
  1015  		go func() {
  1016  			defer wg.Done()
  1017  			data, err := io.ReadAll(r.Body)
  1018  			if string(data) != reqBody {
  1019  				t.Errorf("Handler read %q; want %q", data, reqBody)
  1020  			}
  1021  			if err != nil {
  1022  				t.Errorf("Handler Read: %v", err)
  1023  			}
  1024  			didRead <- true
  1025  		}()
  1026  		// Write in another goroutine.
  1027  		go func() {
  1028  			defer wg.Done()
  1029  			if mode != http2Mode {
  1030  				// our HTTP/1 implementation intentionally
  1031  				// doesn't permit writes during read (mostly
  1032  				// due to it being undefined); if that is ever
  1033  				// relaxed, change this.
  1034  				<-didRead
  1035  			}
  1036  			io.WriteString(w, resBody)
  1037  		}()
  1038  		wg.Wait()
  1039  	}))
  1040  	req, _ := NewRequest("POST", cst.ts.URL, strings.NewReader(reqBody))
  1041  	req.Header.Add("Expect", "100-continue") // just to complicate things
  1042  	res, err := cst.c.Do(req)
  1043  	if err != nil {
  1044  		t.Fatal(err)
  1045  	}
  1046  	data, err := io.ReadAll(res.Body)
  1047  	defer res.Body.Close()
  1048  	if err != nil {
  1049  		t.Fatal(err)
  1050  	}
  1051  	if string(data) != resBody {
  1052  		t.Errorf("read %q; want %q", data, resBody)
  1053  	}
  1054  }
  1055  
  1056  func TestConnectRequest(t *testing.T) { run(t, testConnectRequest) }
  1057  func testConnectRequest(t *testing.T, mode testMode) {
  1058  	gotc := make(chan *Request, 1)
  1059  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1060  		gotc <- r
  1061  	}))
  1062  
  1063  	u, err := url.Parse(cst.ts.URL)
  1064  	if err != nil {
  1065  		t.Fatal(err)
  1066  	}
  1067  
  1068  	tests := []struct {
  1069  		req  *Request
  1070  		want string
  1071  	}{
  1072  		{
  1073  			req: &Request{
  1074  				Method: "CONNECT",
  1075  				Header: Header{},
  1076  				URL:    u,
  1077  			},
  1078  			want: u.Host,
  1079  		},
  1080  		{
  1081  			req: &Request{
  1082  				Method: "CONNECT",
  1083  				Header: Header{},
  1084  				URL:    u,
  1085  				Host:   "example.com:123",
  1086  			},
  1087  			want: "example.com:123",
  1088  		},
  1089  	}
  1090  
  1091  	for i, tt := range tests {
  1092  		res, err := cst.c.Do(tt.req)
  1093  		if err != nil {
  1094  			t.Errorf("%d. RoundTrip = %v", i, err)
  1095  			continue
  1096  		}
  1097  		res.Body.Close()
  1098  		req := <-gotc
  1099  		if req.Method != "CONNECT" {
  1100  			t.Errorf("method = %q; want CONNECT", req.Method)
  1101  		}
  1102  		if req.Host != tt.want {
  1103  			t.Errorf("Host = %q; want %q", req.Host, tt.want)
  1104  		}
  1105  		if req.URL.Host != tt.want {
  1106  			t.Errorf("URL.Host = %q; want %q", req.URL.Host, tt.want)
  1107  		}
  1108  	}
  1109  }
  1110  
  1111  func TestTransportUserAgent(t *testing.T) { run(t, testTransportUserAgent, http3SkippedMode) }
  1112  func testTransportUserAgent(t *testing.T, mode testMode) {
  1113  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1114  		fmt.Fprintf(w, "%q", r.Header["User-Agent"])
  1115  	}))
  1116  
  1117  	either := func(a, b string) string {
  1118  		if mode == http2Mode {
  1119  			return b
  1120  		}
  1121  		return a
  1122  	}
  1123  
  1124  	tests := []struct {
  1125  		setup func(*Request)
  1126  		want  string
  1127  	}{
  1128  		{
  1129  			func(r *Request) {},
  1130  			either(`["Go-http-client/1.1"]`, `["Go-http-client/2.0"]`),
  1131  		},
  1132  		{
  1133  			func(r *Request) { r.Header.Set("User-Agent", "foo/1.2.3") },
  1134  			`["foo/1.2.3"]`,
  1135  		},
  1136  		{
  1137  			func(r *Request) { r.Header["User-Agent"] = []string{"single", "or", "multiple"} },
  1138  			`["single"]`,
  1139  		},
  1140  		{
  1141  			func(r *Request) { r.Header.Set("User-Agent", "") },
  1142  			`[]`,
  1143  		},
  1144  		{
  1145  			func(r *Request) { r.Header["User-Agent"] = nil },
  1146  			`[]`,
  1147  		},
  1148  	}
  1149  	for i, tt := range tests {
  1150  		req, _ := NewRequest("GET", cst.ts.URL, nil)
  1151  		tt.setup(req)
  1152  		res, err := cst.c.Do(req)
  1153  		if err != nil {
  1154  			t.Errorf("%d. RoundTrip = %v", i, err)
  1155  			continue
  1156  		}
  1157  		slurp, err := io.ReadAll(res.Body)
  1158  		res.Body.Close()
  1159  		if err != nil {
  1160  			t.Errorf("%d. read body = %v", i, err)
  1161  			continue
  1162  		}
  1163  		if string(slurp) != tt.want {
  1164  			t.Errorf("%d. body mismatch.\n got: %s\nwant: %s\n", i, slurp, tt.want)
  1165  		}
  1166  	}
  1167  }
  1168  
  1169  func TestStarRequestMethod(t *testing.T) {
  1170  	for _, method := range []string{"FOO", "OPTIONS"} {
  1171  		t.Run(method, func(t *testing.T) {
  1172  			run(t, func(t *testing.T, mode testMode) {
  1173  				testStarRequest(t, method, mode)
  1174  			})
  1175  		})
  1176  	}
  1177  }
  1178  func testStarRequest(t *testing.T, method string, mode testMode) {
  1179  	gotc := make(chan *Request, 1)
  1180  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1181  		w.Header().Set("foo", "bar")
  1182  		gotc <- r
  1183  		w.(Flusher).Flush()
  1184  	}))
  1185  
  1186  	u, err := url.Parse(cst.ts.URL)
  1187  	if err != nil {
  1188  		t.Fatal(err)
  1189  	}
  1190  	u.Path = "*"
  1191  
  1192  	req := &Request{
  1193  		Method: method,
  1194  		Header: Header{},
  1195  		URL:    u,
  1196  	}
  1197  
  1198  	res, err := cst.c.Do(req)
  1199  	if err != nil {
  1200  		t.Fatalf("RoundTrip = %v", err)
  1201  	}
  1202  	res.Body.Close()
  1203  
  1204  	wantFoo := "bar"
  1205  	wantLen := int64(-1)
  1206  	if method == "OPTIONS" {
  1207  		wantFoo = ""
  1208  		wantLen = 0
  1209  	}
  1210  	if res.StatusCode != 200 {
  1211  		t.Errorf("status code = %v; want %d", res.Status, 200)
  1212  	}
  1213  	if res.ContentLength != wantLen {
  1214  		t.Errorf("content length = %v; want %d", res.ContentLength, wantLen)
  1215  	}
  1216  	if got := res.Header.Get("foo"); got != wantFoo {
  1217  		t.Errorf("response \"foo\" header = %q; want %q", got, wantFoo)
  1218  	}
  1219  	select {
  1220  	case req = <-gotc:
  1221  	default:
  1222  		req = nil
  1223  	}
  1224  	if req == nil {
  1225  		if method != "OPTIONS" {
  1226  			t.Fatalf("handler never got request")
  1227  		}
  1228  		return
  1229  	}
  1230  	if req.Method != method {
  1231  		t.Errorf("method = %q; want %q", req.Method, method)
  1232  	}
  1233  	if req.URL.Path != "*" {
  1234  		t.Errorf("URL.Path = %q; want *", req.URL.Path)
  1235  	}
  1236  	if req.RequestURI != "*" {
  1237  		t.Errorf("RequestURI = %q; want *", req.RequestURI)
  1238  	}
  1239  }
  1240  
  1241  // Issue 13957
  1242  func TestTransportDiscardsUnneededConns(t *testing.T) {
  1243  	run(t, testTransportDiscardsUnneededConns, []testMode{http2Mode})
  1244  }
  1245  func testTransportDiscardsUnneededConns(t *testing.T, mode testMode) {
  1246  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1247  		fmt.Fprintf(w, "Hello, %v", r.RemoteAddr)
  1248  	}))
  1249  	defer cst.close()
  1250  
  1251  	var numOpen, numClose int32 // atomic
  1252  
  1253  	tlsConfig := &tls.Config{InsecureSkipVerify: true}
  1254  	tr := &Transport{
  1255  		TLSClientConfig: tlsConfig,
  1256  		DialTLS: func(_, addr string) (net.Conn, error) {
  1257  			time.Sleep(10 * time.Millisecond)
  1258  			rc, err := net.Dial("tcp", addr)
  1259  			if err != nil {
  1260  				return nil, err
  1261  			}
  1262  			atomic.AddInt32(&numOpen, 1)
  1263  			c := noteCloseConn{rc, func() { atomic.AddInt32(&numClose, 1) }}
  1264  			return tls.Client(c, tlsConfig), nil
  1265  		},
  1266  		Protocols: &Protocols{},
  1267  	}
  1268  	tr.Protocols.SetHTTP2(true)
  1269  	defer tr.CloseIdleConnections()
  1270  
  1271  	c := &Client{Transport: tr}
  1272  
  1273  	const N = 10
  1274  	gotBody := make(chan string, N)
  1275  	var wg sync.WaitGroup
  1276  	for i := 0; i < N; i++ {
  1277  		wg.Add(1)
  1278  		go func() {
  1279  			defer wg.Done()
  1280  			resp, err := c.Get(cst.ts.URL)
  1281  			if err != nil {
  1282  				// Try to work around spurious connection reset on loaded system.
  1283  				// See golang.org/issue/33585 and golang.org/issue/36797.
  1284  				time.Sleep(10 * time.Millisecond)
  1285  				resp, err = c.Get(cst.ts.URL)
  1286  				if err != nil {
  1287  					t.Errorf("Get: %v", err)
  1288  					return
  1289  				}
  1290  			}
  1291  			defer resp.Body.Close()
  1292  			slurp, err := io.ReadAll(resp.Body)
  1293  			if err != nil {
  1294  				t.Error(err)
  1295  			}
  1296  			gotBody <- string(slurp)
  1297  		}()
  1298  	}
  1299  	wg.Wait()
  1300  	close(gotBody)
  1301  
  1302  	var last string
  1303  	for got := range gotBody {
  1304  		if last == "" {
  1305  			last = got
  1306  			continue
  1307  		}
  1308  		if got != last {
  1309  			t.Errorf("Response body changed: %q -> %q", last, got)
  1310  		}
  1311  	}
  1312  
  1313  	var open, close int32
  1314  	for i := 0; i < 150; i++ {
  1315  		open, close = atomic.LoadInt32(&numOpen), atomic.LoadInt32(&numClose)
  1316  		if open < 1 {
  1317  			t.Fatalf("open = %d; want at least", open)
  1318  		}
  1319  		if close == open-1 {
  1320  			// Success
  1321  			return
  1322  		}
  1323  		time.Sleep(10 * time.Millisecond)
  1324  	}
  1325  	t.Errorf("%d connections opened, %d closed; want %d to close", open, close, open-1)
  1326  }
  1327  
  1328  // tests that Transport doesn't retain a pointer to the provided request.
  1329  func TestTransportGCRequest(t *testing.T) {
  1330  	run(t, func(t *testing.T, mode testMode) {
  1331  		t.Run("Body", func(t *testing.T) { testTransportGCRequest(t, mode, true) })
  1332  		t.Run("NoBody", func(t *testing.T) { testTransportGCRequest(t, mode, false) })
  1333  	})
  1334  }
  1335  func testTransportGCRequest(t *testing.T, mode testMode, body bool) {
  1336  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1337  		io.ReadAll(r.Body)
  1338  		if body {
  1339  			io.WriteString(w, "Hello.")
  1340  		}
  1341  	}))
  1342  
  1343  	didGC := make(chan struct{})
  1344  	(func() {
  1345  		body := strings.NewReader("some body")
  1346  		req, _ := NewRequest("POST", cst.ts.URL, body)
  1347  		runtime.AddCleanup(req, func(ch chan struct{}) { close(ch) }, didGC)
  1348  		res, err := cst.c.Do(req)
  1349  		if err != nil {
  1350  			t.Fatal(err)
  1351  		}
  1352  		if _, err := io.ReadAll(res.Body); err != nil {
  1353  			t.Fatal(err)
  1354  		}
  1355  		if err := res.Body.Close(); err != nil {
  1356  			t.Fatal(err)
  1357  		}
  1358  	})()
  1359  	for {
  1360  		select {
  1361  		case <-didGC:
  1362  			return
  1363  		case <-time.After(1 * time.Millisecond):
  1364  			runtime.GC()
  1365  		}
  1366  	}
  1367  }
  1368  
  1369  func TestTransportRejectsInvalidHeaders(t *testing.T) {
  1370  	run(t, testTransportRejectsInvalidHeaders, http3SkippedMode)
  1371  }
  1372  func testTransportRejectsInvalidHeaders(t *testing.T, mode testMode) {
  1373  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1374  		fmt.Fprintf(w, "Handler saw headers: %q", r.Header)
  1375  	}), optQuietLog)
  1376  	cst.tr.DisableKeepAlives = true
  1377  
  1378  	tests := []struct {
  1379  		key, val string
  1380  		ok       bool
  1381  	}{
  1382  		{"Foo", "capital-key", true}, // verify h2 allows capital keys
  1383  		{"Foo", "foo\x00bar", false}, // \x00 byte in value not allowed
  1384  		{"Foo", "two\nlines", false}, // \n byte in value not allowed
  1385  		{"bogus\nkey", "v", false},   // \n byte also not allowed in key
  1386  		{"A space", "v", false},      // spaces in keys not allowed
  1387  		{"имя", "v", false},          // key must be ascii
  1388  		{"name", "валю", true},       // value may be non-ascii
  1389  		{"", "v", false},             // key must be non-empty
  1390  		{"k", "", true},              // value may be empty
  1391  	}
  1392  	for _, tt := range tests {
  1393  		dialedc := make(chan bool, 1)
  1394  		cst.tr.Dial = func(netw, addr string) (net.Conn, error) {
  1395  			dialedc <- true
  1396  			return net.Dial(netw, addr)
  1397  		}
  1398  		req, _ := NewRequest("GET", cst.ts.URL, nil)
  1399  		req.Header[tt.key] = []string{tt.val}
  1400  		res, err := cst.c.Do(req)
  1401  		var body []byte
  1402  		if err == nil {
  1403  			body, _ = io.ReadAll(res.Body)
  1404  			res.Body.Close()
  1405  		}
  1406  		var dialed bool
  1407  		select {
  1408  		case <-dialedc:
  1409  			dialed = true
  1410  		default:
  1411  		}
  1412  
  1413  		if !tt.ok && dialed {
  1414  			t.Errorf("For key %q, value %q, transport dialed. Expected local failure. Response was: (%v, %v)\nServer replied with: %s", tt.key, tt.val, res, err, body)
  1415  		} else if (err == nil) != tt.ok {
  1416  			t.Errorf("For key %q, value %q; got err = %v; want ok=%v", tt.key, tt.val, err, tt.ok)
  1417  		}
  1418  	}
  1419  }
  1420  
  1421  func TestInterruptWithPanic(t *testing.T) {
  1422  	run(t, func(t *testing.T, mode testMode) {
  1423  		t.Run("boom", func(t *testing.T) { testInterruptWithPanic(t, mode, "boom") })
  1424  		t.Run("nil", func(t *testing.T) { t.Setenv("GODEBUG", "panicnil=1"); testInterruptWithPanic(t, mode, nil) })
  1425  		t.Run("ErrAbortHandler", func(t *testing.T) { testInterruptWithPanic(t, mode, ErrAbortHandler) })
  1426  	}, testNotParallel, http3SkippedMode)
  1427  }
  1428  func testInterruptWithPanic(t *testing.T, mode testMode, panicValue any) {
  1429  	const msg = "hello"
  1430  
  1431  	testDone := make(chan struct{})
  1432  	defer close(testDone)
  1433  
  1434  	var errorLog lockedBytesBuffer
  1435  	gotHeaders := make(chan bool, 1)
  1436  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1437  		io.WriteString(w, msg)
  1438  		w.(Flusher).Flush()
  1439  
  1440  		select {
  1441  		case <-gotHeaders:
  1442  		case <-testDone:
  1443  		}
  1444  		panic(panicValue)
  1445  	}), func(ts *httptest.Server) {
  1446  		ts.Config.ErrorLog = log.New(&errorLog, "", 0)
  1447  	})
  1448  	res, err := cst.c.Get(cst.ts.URL)
  1449  	if err != nil {
  1450  		t.Fatal(err)
  1451  	}
  1452  	gotHeaders <- true
  1453  	defer res.Body.Close()
  1454  	slurp, err := io.ReadAll(res.Body)
  1455  	if string(slurp) != msg {
  1456  		t.Errorf("client read %q; want %q", slurp, msg)
  1457  	}
  1458  	if err == nil {
  1459  		t.Errorf("client read all successfully; want some error")
  1460  	}
  1461  	logOutput := func() string {
  1462  		errorLog.Lock()
  1463  		defer errorLog.Unlock()
  1464  		return errorLog.String()
  1465  	}
  1466  	wantStackLogged := panicValue != nil && panicValue != ErrAbortHandler
  1467  
  1468  	waitCondition(t, 10*time.Millisecond, func(d time.Duration) bool {
  1469  		gotLog := logOutput()
  1470  		if !wantStackLogged {
  1471  			if gotLog == "" {
  1472  				return true
  1473  			}
  1474  			t.Fatalf("want no log output; got: %s", gotLog)
  1475  		}
  1476  		if gotLog == "" {
  1477  			if d > 0 {
  1478  				t.Logf("wanted a stack trace logged; got nothing after %v", d)
  1479  			}
  1480  			return false
  1481  		}
  1482  		if !strings.Contains(gotLog, "created by ") && strings.Count(gotLog, "\n") < 6 {
  1483  			if d > 0 {
  1484  				t.Logf("output doesn't look like a panic stack trace after %v. Got: %s", d, gotLog)
  1485  			}
  1486  			return false
  1487  		}
  1488  		return true
  1489  	})
  1490  }
  1491  
  1492  type lockedBytesBuffer struct {
  1493  	sync.Mutex
  1494  	bytes.Buffer
  1495  }
  1496  
  1497  func (b *lockedBytesBuffer) Write(p []byte) (int, error) {
  1498  	b.Lock()
  1499  	defer b.Unlock()
  1500  	return b.Buffer.Write(p)
  1501  }
  1502  
  1503  // Issue 15366
  1504  func TestH12_AutoGzipWithDumpResponse(t *testing.T) {
  1505  	h12Compare{
  1506  		Handler: func(w ResponseWriter, r *Request) {
  1507  			h := w.Header()
  1508  			h.Set("Content-Encoding", "gzip")
  1509  			h.Set("Content-Length", "23")
  1510  			io.WriteString(w, "\x1f\x8b\b\x00\x00\x00\x00\x00\x00\x00s\xf3\xf7\a\x00\xab'\xd4\x1a\x03\x00\x00\x00")
  1511  		},
  1512  		EarlyCheckResponse: func(proto string, res *Response) {
  1513  			if !res.Uncompressed {
  1514  				t.Errorf("%s: expected Uncompressed to be set", proto)
  1515  			}
  1516  			dump, err := httputil.DumpResponse(res, true)
  1517  			if err != nil {
  1518  				t.Errorf("%s: DumpResponse: %v", proto, err)
  1519  				return
  1520  			}
  1521  			if strings.Contains(string(dump), "Connection: close") {
  1522  				t.Errorf("%s: should not see \"Connection: close\" in dump; got:\n%s", proto, dump)
  1523  			}
  1524  			if !strings.Contains(string(dump), "FOO") {
  1525  				t.Errorf("%s: should see \"FOO\" in response; got:\n%s", proto, dump)
  1526  			}
  1527  		},
  1528  	}.run(t)
  1529  }
  1530  
  1531  // Issue 14607
  1532  func TestCloseIdleConnections(t *testing.T) { run(t, testCloseIdleConnections) }
  1533  func testCloseIdleConnections(t *testing.T, mode testMode) {
  1534  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1535  		w.Header().Set("X-Addr", r.RemoteAddr)
  1536  	}))
  1537  	get := func() string {
  1538  		res, err := cst.c.Get(cst.ts.URL)
  1539  		if err != nil {
  1540  			t.Fatal(err)
  1541  		}
  1542  		res.Body.Close()
  1543  		v := res.Header.Get("X-Addr")
  1544  		if v == "" {
  1545  			t.Fatal("didn't get X-Addr")
  1546  		}
  1547  		return v
  1548  	}
  1549  	a1 := get()
  1550  	cst.tr.CloseIdleConnections()
  1551  	a2 := get()
  1552  	if a1 == a2 {
  1553  		t.Errorf("didn't close connection")
  1554  	}
  1555  }
  1556  
  1557  type noteCloseConn struct {
  1558  	net.Conn
  1559  	closeFunc func()
  1560  }
  1561  
  1562  func (x noteCloseConn) Close() error {
  1563  	x.closeFunc()
  1564  	return x.Conn.Close()
  1565  }
  1566  
  1567  type testErrorReader struct{ t *testing.T }
  1568  
  1569  func (r testErrorReader) Read(p []byte) (n int, err error) {
  1570  	r.t.Error("unexpected Read call")
  1571  	return 0, io.EOF
  1572  }
  1573  
  1574  func TestNoSniffExpectRequestBody(t *testing.T) { run(t, testNoSniffExpectRequestBody) }
  1575  func testNoSniffExpectRequestBody(t *testing.T, mode testMode) {
  1576  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1577  		w.WriteHeader(StatusUnauthorized)
  1578  	}))
  1579  
  1580  	// Set ExpectContinueTimeout non-zero so RoundTrip won't try to write it.
  1581  	cst.tr.ExpectContinueTimeout = 10 * time.Second
  1582  
  1583  	req, err := NewRequest("POST", cst.ts.URL, testErrorReader{t})
  1584  	if err != nil {
  1585  		t.Fatal(err)
  1586  	}
  1587  	req.ContentLength = 0 // so transport is tempted to sniff it
  1588  	req.Header.Set("Expect", "100-continue")
  1589  	res, err := cst.tr.RoundTrip(req)
  1590  	if err != nil {
  1591  		t.Fatal(err)
  1592  	}
  1593  	defer res.Body.Close()
  1594  	if res.StatusCode != StatusUnauthorized {
  1595  		t.Errorf("status code = %v; want %v", res.StatusCode, StatusUnauthorized)
  1596  	}
  1597  }
  1598  
  1599  func TestServerUndeclaredTrailers(t *testing.T) {
  1600  	run(t, testServerUndeclaredTrailers, http3SkippedMode)
  1601  }
  1602  func testServerUndeclaredTrailers(t *testing.T, mode testMode) {
  1603  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1604  		w.Header().Set("Foo", "Bar")
  1605  		w.Header().Set("Trailer:Foo", "Baz")
  1606  		w.(Flusher).Flush()
  1607  		w.Header().Add("Trailer:Foo", "Baz2")
  1608  		w.Header().Set("Trailer:Bar", "Quux")
  1609  	}))
  1610  	res, err := cst.c.Get(cst.ts.URL)
  1611  	if err != nil {
  1612  		t.Fatal(err)
  1613  	}
  1614  	if _, err := io.Copy(io.Discard, res.Body); err != nil {
  1615  		t.Fatal(err)
  1616  	}
  1617  	res.Body.Close()
  1618  	delete(res.Header, "Date")
  1619  	delete(res.Header, "Content-Type")
  1620  
  1621  	if want := (Header{"Foo": {"Bar"}}); !reflect.DeepEqual(res.Header, want) {
  1622  		t.Errorf("Header = %#v; want %#v", res.Header, want)
  1623  	}
  1624  	if want := (Header{"Foo": {"Baz", "Baz2"}, "Bar": {"Quux"}}); !reflect.DeepEqual(res.Trailer, want) {
  1625  		t.Errorf("Trailer = %#v; want %#v", res.Trailer, want)
  1626  	}
  1627  }
  1628  
  1629  func TestBadResponseAfterReadingBody(t *testing.T) {
  1630  	run(t, testBadResponseAfterReadingBody, []testMode{http1Mode})
  1631  }
  1632  func testBadResponseAfterReadingBody(t *testing.T, mode testMode) {
  1633  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1634  		_, err := io.Copy(io.Discard, r.Body)
  1635  		if err != nil {
  1636  			t.Fatal(err)
  1637  		}
  1638  		c, _, err := w.(Hijacker).Hijack()
  1639  		if err != nil {
  1640  			t.Fatal(err)
  1641  		}
  1642  		defer c.Close()
  1643  		fmt.Fprintln(c, "some bogus crap")
  1644  	}))
  1645  
  1646  	closes := 0
  1647  	res, err := cst.c.Post(cst.ts.URL, "text/plain", countCloseReader{&closes, strings.NewReader("hello")})
  1648  	if err == nil {
  1649  		res.Body.Close()
  1650  		t.Fatal("expected an error to be returned from Post")
  1651  	}
  1652  	if closes != 1 {
  1653  		t.Errorf("closes = %d; want 1", closes)
  1654  	}
  1655  }
  1656  
  1657  func TestWriteHeader0(t *testing.T) { run(t, testWriteHeader0) }
  1658  func testWriteHeader0(t *testing.T, mode testMode) {
  1659  	gotpanic := make(chan bool, 1)
  1660  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1661  		defer close(gotpanic)
  1662  		defer func() {
  1663  			if e := recover(); e != nil {
  1664  				got := fmt.Sprintf("%T, %v", e, e)
  1665  				want := "string, invalid WriteHeader code 0"
  1666  				if got != want {
  1667  					t.Errorf("unexpected panic value:\n got: %v\nwant: %v\n", got, want)
  1668  				}
  1669  				gotpanic <- true
  1670  
  1671  				// Set an explicit 503. This also tests that the WriteHeader call panics
  1672  				// before it recorded that an explicit value was set and that bogus
  1673  				// value wasn't stuck.
  1674  				w.WriteHeader(503)
  1675  			}
  1676  		}()
  1677  		w.WriteHeader(0)
  1678  	}))
  1679  	res, err := cst.c.Get(cst.ts.URL)
  1680  	if err != nil {
  1681  		t.Fatal(err)
  1682  	}
  1683  	if res.StatusCode != 503 {
  1684  		t.Errorf("Response: %v %q; want 503", res.StatusCode, res.Status)
  1685  	}
  1686  	if !<-gotpanic {
  1687  		t.Error("expected panic in handler")
  1688  	}
  1689  }
  1690  
  1691  // Issue 23010: don't be super strict checking WriteHeader's code if
  1692  // it's not even valid to call WriteHeader then anyway.
  1693  func TestWriteHeaderNoCodeCheck(t *testing.T) {
  1694  	run(t, func(t *testing.T, mode testMode) {
  1695  		testWriteHeaderAfterWrite(t, mode, false)
  1696  	}, http3SkippedMode)
  1697  }
  1698  func TestWriteHeaderNoCodeCheck_h1hijack(t *testing.T) {
  1699  	testWriteHeaderAfterWrite(t, http1Mode, true)
  1700  }
  1701  func testWriteHeaderAfterWrite(t *testing.T, mode testMode, hijack bool) {
  1702  	var errorLog lockedBytesBuffer
  1703  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1704  		if hijack {
  1705  			conn, _, _ := w.(Hijacker).Hijack()
  1706  			defer conn.Close()
  1707  			conn.Write([]byte("HTTP/1.1 200 OK\r\nContent-Length: 6\r\n\r\nfoo"))
  1708  			w.WriteHeader(0) // verify this doesn't panic if there's already output; Issue 23010
  1709  			conn.Write([]byte("bar"))
  1710  			return
  1711  		}
  1712  		io.WriteString(w, "foo")
  1713  		w.(Flusher).Flush()
  1714  		w.WriteHeader(0) // verify this doesn't panic if there's already output; Issue 23010
  1715  		io.WriteString(w, "bar")
  1716  	}), func(ts *httptest.Server) {
  1717  		ts.Config.ErrorLog = log.New(&errorLog, "", 0)
  1718  	})
  1719  	res, err := cst.c.Get(cst.ts.URL)
  1720  	if err != nil {
  1721  		t.Fatal(err)
  1722  	}
  1723  	defer res.Body.Close()
  1724  	body, err := io.ReadAll(res.Body)
  1725  	if err != nil {
  1726  		t.Fatal(err)
  1727  	}
  1728  	if got, want := string(body), "foobar"; got != want {
  1729  		t.Errorf("got = %q; want %q", got, want)
  1730  	}
  1731  
  1732  	// Also check the stderr output:
  1733  	if mode == http2Mode {
  1734  		// TODO: also emit this log message for HTTP/2?
  1735  		// We historically haven't, so don't check.
  1736  		return
  1737  	}
  1738  	gotLog := strings.TrimSpace(errorLog.String())
  1739  	wantLog := "http: superfluous response.WriteHeader call from net/http_test.testWriteHeaderAfterWrite.func1 (clientserver_test.go:"
  1740  	if hijack {
  1741  		wantLog = "http: response.WriteHeader on hijacked connection from net/http_test.testWriteHeaderAfterWrite.func1 (clientserver_test.go:"
  1742  	}
  1743  	if !strings.HasPrefix(gotLog, wantLog) {
  1744  		t.Errorf("stderr output = %q; want %q", gotLog, wantLog)
  1745  	}
  1746  }
  1747  
  1748  func TestBidiStreamReverseProxy(t *testing.T) {
  1749  	run(t, testBidiStreamReverseProxy, []testMode{http2Mode})
  1750  }
  1751  func testBidiStreamReverseProxy(t *testing.T, mode testMode) {
  1752  	backend := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1753  		if _, err := io.Copy(w, r.Body); err != nil {
  1754  			log.Printf("bidi backend copy: %v", err)
  1755  		}
  1756  	}))
  1757  
  1758  	backURL, err := url.Parse(backend.ts.URL)
  1759  	if err != nil {
  1760  		t.Fatal(err)
  1761  	}
  1762  	rp := httputil.NewSingleHostReverseProxy(backURL)
  1763  	rp.Transport = backend.tr
  1764  	proxy := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1765  		rp.ServeHTTP(w, r)
  1766  	}))
  1767  
  1768  	bodyRes := make(chan any, 1) // error or hash.Hash
  1769  	pr, pw := io.Pipe()
  1770  	req, _ := NewRequest("PUT", proxy.ts.URL, pr)
  1771  	const size = 4 << 20
  1772  	go func() {
  1773  		h := sha1.New()
  1774  		_, err := io.CopyN(io.MultiWriter(h, pw), rand.Reader, size)
  1775  		go pw.Close()
  1776  		if err != nil {
  1777  			t.Errorf("body copy: %v", err)
  1778  			bodyRes <- err
  1779  		} else {
  1780  			bodyRes <- h
  1781  		}
  1782  	}()
  1783  	res, err := backend.c.Do(req)
  1784  	if err != nil {
  1785  		t.Fatal(err)
  1786  	}
  1787  	defer res.Body.Close()
  1788  	hgot := sha1.New()
  1789  	n, err := io.Copy(hgot, res.Body)
  1790  	if err != nil {
  1791  		t.Fatal(err)
  1792  	}
  1793  	if n != size {
  1794  		t.Fatalf("got %d bytes; want %d", n, size)
  1795  	}
  1796  	select {
  1797  	case v := <-bodyRes:
  1798  		switch v := v.(type) {
  1799  		default:
  1800  			t.Fatalf("body copy: %v", err)
  1801  		case hash.Hash:
  1802  			if !bytes.Equal(v.Sum(nil), hgot.Sum(nil)) {
  1803  				t.Errorf("written bytes didn't match received bytes")
  1804  			}
  1805  		}
  1806  	case <-time.After(10 * time.Second):
  1807  		t.Fatal("timeout")
  1808  	}
  1809  
  1810  }
  1811  
  1812  // Always use HTTP/1.1 for WebSocket upgrades.
  1813  func TestH12_WebSocketUpgrade(t *testing.T) {
  1814  	h12Compare{
  1815  		Handler: func(w ResponseWriter, r *Request) {
  1816  			h := w.Header()
  1817  			h.Set("Foo", "bar")
  1818  		},
  1819  		ReqFunc: func(c *Client, url string) (*Response, error) {
  1820  			req, _ := NewRequest("GET", url, nil)
  1821  			req.Header.Set("Connection", "Upgrade")
  1822  			req.Header.Set("Upgrade", "WebSocket")
  1823  			return c.Do(req)
  1824  		},
  1825  		EarlyCheckResponse: func(proto string, res *Response) {
  1826  			if res.Proto != "HTTP/1.1" {
  1827  				t.Errorf("%s: expected HTTP/1.1, got %q", proto, res.Proto)
  1828  			}
  1829  			res.Proto = "HTTP/IGNORE" // skip later checks that Proto must be 1.1 vs 2.0
  1830  		},
  1831  		Opts: []any{
  1832  			func(s *Server) {
  1833  				// Configure servers to support HTTP/1 and HTTP/2,
  1834  				// so we can verify that we use HTTP/1
  1835  				// even when HTTP/2 is an option.
  1836  				s.Protocols = &Protocols{}
  1837  				s.Protocols.SetHTTP1(true)
  1838  				s.Protocols.SetHTTP2(true)
  1839  			},
  1840  		},
  1841  	}.run(t)
  1842  }
  1843  
  1844  func TestIdentityTransferEncoding(t *testing.T) { run(t, testIdentityTransferEncoding) }
  1845  func testIdentityTransferEncoding(t *testing.T, mode testMode) {
  1846  	const body = "body"
  1847  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1848  		gotBody, _ := io.ReadAll(r.Body)
  1849  		if got, want := string(gotBody), body; got != want {
  1850  			t.Errorf("got request body = %q; want %q", got, want)
  1851  		}
  1852  		w.Header().Set("Transfer-Encoding", "identity")
  1853  		w.WriteHeader(StatusOK)
  1854  		w.(Flusher).Flush()
  1855  		io.WriteString(w, body)
  1856  	}))
  1857  	req, _ := NewRequest("GET", cst.ts.URL, strings.NewReader(body))
  1858  	res, err := cst.c.Do(req)
  1859  	if err != nil {
  1860  		t.Fatal(err)
  1861  	}
  1862  	defer res.Body.Close()
  1863  	gotBody, err := io.ReadAll(res.Body)
  1864  	if err != nil {
  1865  		t.Fatal(err)
  1866  	}
  1867  	if got, want := string(gotBody), body; got != want {
  1868  		t.Errorf("got response body = %q; want %q", got, want)
  1869  	}
  1870  }
  1871  
  1872  func TestEarlyHintsRequest(t *testing.T) { run(t, testEarlyHintsRequest) }
  1873  func testEarlyHintsRequest(t *testing.T, mode testMode) {
  1874  	var wg sync.WaitGroup
  1875  	wg.Add(1)
  1876  	cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
  1877  		h := w.Header()
  1878  
  1879  		h.Add("Content-Length", "123") // must be ignored
  1880  		h.Add("Link", "</style.css>; rel=preload; as=style")
  1881  		h.Add("Link", "</script.js>; rel=preload; as=script")
  1882  		w.WriteHeader(StatusEarlyHints)
  1883  
  1884  		wg.Wait()
  1885  
  1886  		h.Add("Link", "</foo.js>; rel=preload; as=script")
  1887  		w.WriteHeader(StatusEarlyHints)
  1888  
  1889  		w.Write([]byte("Hello"))
  1890  	}))
  1891  
  1892  	checkLinkHeaders := func(t *testing.T, expected, got []string) {
  1893  		t.Helper()
  1894  
  1895  		if len(expected) != len(got) {
  1896  			t.Errorf("got %d expected %d", len(got), len(expected))
  1897  		}
  1898  
  1899  		for i := range expected {
  1900  			if expected[i] != got[i] {
  1901  				t.Errorf("got %q expected %q", got[i], expected[i])
  1902  			}
  1903  		}
  1904  	}
  1905  
  1906  	checkExcludedHeaders := func(t *testing.T, header textproto.MIMEHeader) {
  1907  		t.Helper()
  1908  
  1909  		for _, h := range []string{"Content-Length", "Transfer-Encoding"} {
  1910  			if v, ok := header[h]; ok {
  1911  				t.Errorf("%s is %q; must not be sent", h, v)
  1912  			}
  1913  		}
  1914  	}
  1915  
  1916  	var respCounter uint8
  1917  	trace := &httptrace.ClientTrace{
  1918  		Got1xxResponse: func(code int, header textproto.MIMEHeader) error {
  1919  			switch respCounter {
  1920  			case 0:
  1921  				checkLinkHeaders(t, []string{"</style.css>; rel=preload; as=style", "</script.js>; rel=preload; as=script"}, header["Link"])
  1922  				checkExcludedHeaders(t, header)
  1923  
  1924  				wg.Done()
  1925  			case 1:
  1926  				checkLinkHeaders(t, []string{"</style.css>; rel=preload; as=style", "</script.js>; rel=preload; as=script", "</foo.js>; rel=preload; as=script"}, header["Link"])
  1927  				checkExcludedHeaders(t, header)
  1928  
  1929  			default:
  1930  				t.Error("Unexpected 1xx response")
  1931  			}
  1932  
  1933  			respCounter++
  1934  
  1935  			return nil
  1936  		},
  1937  	}
  1938  	req, _ := NewRequestWithContext(httptrace.WithClientTrace(context.Background(), trace), "GET", cst.ts.URL, nil)
  1939  
  1940  	res, err := cst.c.Do(req)
  1941  	if err != nil {
  1942  		t.Fatal(err)
  1943  	}
  1944  	defer res.Body.Close()
  1945  
  1946  	checkLinkHeaders(t, []string{"</style.css>; rel=preload; as=style", "</script.js>; rel=preload; as=script", "</foo.js>; rel=preload; as=script"}, res.Header["Link"])
  1947  	if cl := res.Header.Get("Content-Length"); cl != "123" {
  1948  		t.Errorf("Content-Length is %q; want 123", cl)
  1949  	}
  1950  
  1951  	body, _ := io.ReadAll(res.Body)
  1952  	if string(body) != "Hello" {
  1953  		t.Errorf("Read body %q; want Hello", body)
  1954  	}
  1955  }
  1956  
  1957  // TestClientServerTLSConnWrapper verifies that the Transport and Server can
  1958  // negotiate an HTTP/2 connection using a net.Conn that has a
  1959  // "ConnectionState() tls.ConnectionState" method but is not a *tls.Conn.
  1960  func TestClientServerTLSConnWrapper(t *testing.T) {
  1961  	synctest.Test(t, func(t *testing.T) {
  1962  		protocols := &Protocols{}
  1963  		protocols.SetHTTP1(true)
  1964  		protocols.SetHTTP2(true)
  1965  
  1966  		li := fakeNetListen()
  1967  		server := &Server{
  1968  			Handler: HandlerFunc(func(w ResponseWriter, r *Request) {
  1969  				if r.TLS == nil {
  1970  					t.Fatal("server request has no TLS ConnectionState")
  1971  				}
  1972  			}),
  1973  			Protocols: protocols,
  1974  		}
  1975  		defer server.Close()
  1976  		go server.Serve(&testListener{
  1977  			accept: func() (net.Conn, error) {
  1978  				conn, err := li.Accept()
  1979  				if err != nil {
  1980  					return nil, err
  1981  				}
  1982  				return &testTLSConn{
  1983  					Conn: conn,
  1984  					state: tls.ConnectionState{
  1985  						Version:            tls.VersionTLS13,
  1986  						CipherSuite:        tls.TLS_AES_128_GCM_SHA256,
  1987  						NegotiatedProtocol: "h2",
  1988  					},
  1989  				}, nil
  1990  			},
  1991  			close: li.Close,
  1992  			addr:  li.Addr(),
  1993  		})
  1994  
  1995  		tr := &Transport{
  1996  			DialTLS: func(network, address string) (net.Conn, error) {
  1997  				return &testTLSConn{
  1998  					Conn: li.connect(),
  1999  					state: tls.ConnectionState{
  2000  						Version:            tls.VersionTLS13,
  2001  						CipherSuite:        tls.TLS_AES_128_GCM_SHA256,
  2002  						NegotiatedProtocol: "h2",
  2003  					},
  2004  				}, nil
  2005  			},
  2006  			Protocols: protocols,
  2007  		}
  2008  
  2009  		req, _ := NewRequest("GET", "https://example.tld", nil)
  2010  		resp, err := tr.RoundTrip(req)
  2011  		if err != nil {
  2012  			t.Fatal(err)
  2013  		}
  2014  		resp.Body.Close()
  2015  		if resp.StatusCode != 200 {
  2016  			t.Errorf("response status %v, want 200", resp.StatusCode)
  2017  		}
  2018  		if resp.TLS == nil {
  2019  			t.Fatal("server request has no TLS ConnectionState")
  2020  		}
  2021  	})
  2022  }
  2023  
  2024  type testListener struct {
  2025  	accept func() (net.Conn, error)
  2026  	close  func() error
  2027  	addr   net.Addr
  2028  }
  2029  
  2030  func (li *testListener) Accept() (net.Conn, error) { return li.accept() }
  2031  func (li *testListener) Close() error              { return li.close() }
  2032  func (li *testListener) Addr() net.Addr            { return li.addr }
  2033  
  2034  type testTLSConn struct {
  2035  	net.Conn
  2036  	state tls.ConnectionState
  2037  }
  2038  
  2039  func (c *testTLSConn) Handshake() error                     { return nil }
  2040  func (c *testTLSConn) ConnectionState() tls.ConnectionState { return c.state }
  2041  

View as plain text