...

Source file src/net/http/httptest/server_test.go

Documentation: net/http/httptest

     1  // Copyright 2012 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  package httptest
     6  
     7  import (
     8  	"bufio"
     9  	"internal/testenv"
    10  	"io"
    11  	"net"
    12  	"net/http"
    13  	"os"
    14  	"regexp"
    15  	"strings"
    16  	"sync"
    17  	"testing"
    18  	"testing/synctest"
    19  )
    20  
    21  type newServerFunc func(*testing.T, http.Handler) *Server
    22  
    23  var newServers = map[string]newServerFunc{
    24  	"NewServer": func(t *testing.T, h http.Handler) *Server {
    25  		return NewServer(h)
    26  	},
    27  	"NewTLSServer": func(t *testing.T, h http.Handler) *Server {
    28  		return NewTLSServer(h)
    29  	},
    30  
    31  	// The manual variants of newServer create a Server manually by only filling
    32  	// in the exported fields of Server.
    33  	"NewServerManual": func(t *testing.T, h http.Handler) *Server {
    34  		ts := &Server{Listener: newLocalListener(), Config: &http.Server{Handler: h}}
    35  		ts.Start()
    36  		return ts
    37  	},
    38  	"NewTLSServerManual": func(t *testing.T, h http.Handler) *Server {
    39  		ts := &Server{Listener: newLocalListener(), Config: &http.Server{Handler: h}}
    40  		ts.StartTLS()
    41  		return ts
    42  	},
    43  
    44  	"NewTestServerMemory": func(t *testing.T, h http.Handler) *Server {
    45  		return NewTestServer(t, h)
    46  	},
    47  	"NewTestServerLoopback": func(t *testing.T, h http.Handler) *Server {
    48  		ts := NewTestServer(t, h)
    49  		ts.Start()
    50  		return ts
    51  	},
    52  	"NewTestServerLoopbackTLS": func(t *testing.T, h http.Handler) *Server {
    53  		ts := NewTestServer(t, h)
    54  		ts.StartTLS()
    55  		return ts
    56  	},
    57  }
    58  
    59  func TestServer(t *testing.T) {
    60  	for _, name := range []string{"NewServer", "NewServerManual", "NewTestServerLoopback"} {
    61  		t.Run(name, func(t *testing.T) {
    62  			newServer := newServers[name]
    63  			t.Run("Server", func(t *testing.T) { testServer(t, newServer) })
    64  			t.Run("GetAfterClose", func(t *testing.T) { testGetAfterClose(t, newServer) })
    65  			t.Run("ServerCloseBlocking", func(t *testing.T) { testServerCloseBlocking(t, newServer) })
    66  			t.Run("ServerCloseClientConnections", func(t *testing.T) { testServerCloseClientConnections(t, newServer) })
    67  			t.Run("ServerClientTransportType", func(t *testing.T) { testServerClientTransportType(t, newServer) })
    68  		})
    69  	}
    70  	for _, name := range []string{"NewTLSServer", "NewTLSServerManual", "NewTestServerMemory", "NewTestServerLoopbackTLS"} {
    71  		t.Run(name, func(t *testing.T) {
    72  			newServer := newServers[name]
    73  			t.Run("ServerClient", func(t *testing.T) { testServerClient(t, newServer) })
    74  			t.Run("TLSServerClientTransportType", func(t *testing.T) { testTLSServerClientTransportType(t, newServer) })
    75  		})
    76  	}
    77  }
    78  
    79  func testServer(t *testing.T, newServer newServerFunc) {
    80  	ts := newServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
    81  		w.Write([]byte("hello"))
    82  	}))
    83  	defer ts.Close()
    84  	res, err := http.Get(ts.URL)
    85  	if err != nil {
    86  		t.Fatal(err)
    87  	}
    88  	got, err := io.ReadAll(res.Body)
    89  	res.Body.Close()
    90  	if err != nil {
    91  		t.Fatal(err)
    92  	}
    93  	if string(got) != "hello" {
    94  		t.Errorf("got %q, want hello", string(got))
    95  	}
    96  }
    97  
    98  // Issue 12781
    99  func testGetAfterClose(t *testing.T, newServer newServerFunc) {
   100  	ts := newServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   101  		w.Write([]byte("hello"))
   102  	}))
   103  
   104  	res, err := http.Get(ts.URL)
   105  	if err != nil {
   106  		t.Fatal(err)
   107  	}
   108  	got, err := io.ReadAll(res.Body)
   109  	res.Body.Close()
   110  	if err != nil {
   111  		t.Fatal(err)
   112  	}
   113  	if string(got) != "hello" {
   114  		t.Fatalf("got %q, want hello", string(got))
   115  	}
   116  
   117  	ts.Close()
   118  
   119  	res, err = http.Get(ts.URL)
   120  	if err == nil {
   121  		body, _ := io.ReadAll(res.Body)
   122  		t.Fatalf("Unexpected response after close: %v, %v, %s", res.Status, res.Header, body)
   123  	}
   124  }
   125  
   126  func testServerCloseBlocking(t *testing.T, newServer newServerFunc) {
   127  	ts := newServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   128  		w.Write([]byte("hello"))
   129  	}))
   130  	dial := func() net.Conn {
   131  		c, err := net.Dial("tcp", ts.Listener.Addr().String())
   132  		if err != nil {
   133  			t.Fatal(err)
   134  		}
   135  		return c
   136  	}
   137  
   138  	// Keep one connection in StateNew (connected, but not sending anything)
   139  	cnew := dial()
   140  	defer cnew.Close()
   141  
   142  	// Keep one connection in StateIdle (idle after a request)
   143  	cidle := dial()
   144  	defer cidle.Close()
   145  	cidle.Write([]byte("HEAD / HTTP/1.1\r\nHost: foo\r\n\r\n"))
   146  	_, err := http.ReadResponse(bufio.NewReader(cidle), nil)
   147  	if err != nil {
   148  		t.Fatal(err)
   149  	}
   150  
   151  	ts.Close() // test we don't hang here forever.
   152  }
   153  
   154  // Issue 14290
   155  func testServerCloseClientConnections(t *testing.T, newServer newServerFunc) {
   156  	var s *Server
   157  	s = newServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   158  		s.CloseClientConnections()
   159  	}))
   160  	defer s.Close()
   161  	res, err := http.Get(s.URL)
   162  	if err == nil {
   163  		res.Body.Close()
   164  		t.Fatalf("Unexpected response: %#v", res)
   165  	}
   166  }
   167  
   168  // Tests that the Server.Client method works and returns an http.Client that can hit
   169  // NewTLSServer without cert warnings.
   170  func testServerClient(t *testing.T, newTLSServer newServerFunc) {
   171  	ts := newTLSServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   172  		w.Write([]byte("hello"))
   173  	}))
   174  	defer ts.Close()
   175  	client := ts.Client()
   176  	res, err := client.Get(ts.URL)
   177  	if err != nil {
   178  		t.Fatal(err)
   179  	}
   180  	got, err := io.ReadAll(res.Body)
   181  	res.Body.Close()
   182  	if err != nil {
   183  		t.Fatal(err)
   184  	}
   185  	if string(got) != "hello" {
   186  		t.Errorf("got %q, want hello", string(got))
   187  	}
   188  }
   189  
   190  // Tests that the Server.Client.Transport interface is implemented
   191  // by a *http.Transport.
   192  func testServerClientTransportType(t *testing.T, newServer newServerFunc) {
   193  	ts := newServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   194  	}))
   195  	defer ts.Close()
   196  	client := ts.Client()
   197  	if _, ok := client.Transport.(*http.Transport); !ok {
   198  		t.Errorf("got %T, want *http.Transport", client.Transport)
   199  	}
   200  }
   201  
   202  // Tests that the TLS Server.Client.Transport interface is implemented
   203  // by a *http.Transport.
   204  func testTLSServerClientTransportType(t *testing.T, newTLSServer newServerFunc) {
   205  	ts := newTLSServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   206  	}))
   207  	defer ts.Close()
   208  	client := ts.Client()
   209  	if _, ok := client.Transport.(*http.Transport); !ok {
   210  		t.Errorf("got %T, want *http.Transport", client.Transport)
   211  	}
   212  }
   213  
   214  type onlyCloseListener struct {
   215  	net.Listener
   216  }
   217  
   218  func (onlyCloseListener) Close() error { return nil }
   219  
   220  // Issue 19729: panic in Server.Close for values created directly
   221  // without a constructor (so the unexported client field is nil).
   222  func TestServerZeroValueClose(t *testing.T) {
   223  	ts := &Server{
   224  		Listener: onlyCloseListener{},
   225  		Config:   &http.Server{},
   226  	}
   227  
   228  	ts.Close() // tests that it doesn't panic
   229  }
   230  
   231  // Issue 51799: test hijacking a connection and then closing it
   232  // concurrently with closing the server.
   233  func TestCloseHijackedConnection(t *testing.T) {
   234  	hijacked := make(chan net.Conn)
   235  	ts := NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   236  		defer close(hijacked)
   237  		hj, ok := w.(http.Hijacker)
   238  		if !ok {
   239  			t.Fatal("failed to hijack")
   240  		}
   241  		c, _, err := hj.Hijack()
   242  		if err != nil {
   243  			t.Fatal(err)
   244  		}
   245  		hijacked <- c
   246  	}))
   247  
   248  	var wg sync.WaitGroup
   249  	wg.Add(1)
   250  	go func() {
   251  		defer wg.Done()
   252  		req, err := http.NewRequest("GET", ts.URL, nil)
   253  		if err != nil {
   254  			t.Log(err)
   255  		}
   256  		// Use a client not associated with the Server.
   257  		var c http.Client
   258  		resp, err := c.Do(req)
   259  		if err != nil {
   260  			t.Log(err)
   261  			return
   262  		}
   263  		resp.Body.Close()
   264  	}()
   265  
   266  	wg.Add(1)
   267  	conn := <-hijacked
   268  	go func(conn net.Conn) {
   269  		defer wg.Done()
   270  		// Close the connection and then inform the Server that
   271  		// we closed it.
   272  		conn.Close()
   273  		ts.Config.ConnState(conn, http.StateClosed)
   274  	}(conn)
   275  
   276  	wg.Add(1)
   277  	go func() {
   278  		defer wg.Done()
   279  		ts.Close()
   280  	}()
   281  	wg.Wait()
   282  }
   283  
   284  func TestTLSServerWithHTTP2(t *testing.T) {
   285  	modes := []struct {
   286  		name      string
   287  		wantProto string
   288  	}{
   289  		{"http1", "HTTP/1.1"},
   290  		{"http2", "HTTP/2.0"},
   291  	}
   292  
   293  	for _, tt := range modes {
   294  		t.Run(tt.name, func(t *testing.T) {
   295  			cst := NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   296  				w.Header().Set("X-Proto", r.Proto)
   297  			}))
   298  
   299  			switch tt.name {
   300  			case "http2":
   301  				cst.EnableHTTP2 = true
   302  				cst.StartTLS()
   303  			default:
   304  				cst.Start()
   305  			}
   306  
   307  			defer cst.Close()
   308  
   309  			res, err := cst.Client().Get(cst.URL)
   310  			if err != nil {
   311  				t.Fatalf("Failed to make request: %v", err)
   312  			}
   313  			if g, w := res.Header.Get("X-Proto"), tt.wantProto; g != w {
   314  				t.Fatalf("X-Proto header mismatch:\n\tgot:  %q\n\twant: %q", g, w)
   315  			}
   316  		})
   317  	}
   318  }
   319  
   320  func TestClientExampleCom(t *testing.T) {
   321  	modes := []struct {
   322  		proto string
   323  		host  string
   324  	}{
   325  		{"http", "example.com"},
   326  		{"http", "foo.example.com"},
   327  		{"https", "example.com"},
   328  		{"https", "foo.example.com"},
   329  	}
   330  
   331  	for _, tt := range modes {
   332  		t.Run(tt.proto+" "+tt.host, func(t *testing.T) {
   333  			cst := NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   334  				w.Header().Set("requested-hostname", r.Host)
   335  			}))
   336  			switch tt.proto {
   337  			case "https":
   338  				cst.EnableHTTP2 = true
   339  				cst.StartTLS()
   340  			default:
   341  				cst.Start()
   342  			}
   343  
   344  			defer cst.Close()
   345  
   346  			res, err := cst.Client().Get(tt.proto + "://" + tt.host)
   347  			if err != nil {
   348  				t.Fatalf("Failed to make request: %v", err)
   349  			}
   350  			if got, want := res.Header.Get("requested-hostname"), tt.host; got != want {
   351  				t.Fatalf("Requested hostname mismatch\ngot: %q\nwant: %q", got, want)
   352  			}
   353  		})
   354  	}
   355  }
   356  
   357  func TestServerInMemoryNetwork(t *testing.T) {
   358  	synctest.Test(t, func(t *testing.T) {
   359  		ts := NewTestServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   360  		}))
   361  
   362  		for _, u := range []string{
   363  			"http://example.tld/",
   364  			"https://example.tld/",
   365  			"https://go.dev/",
   366  			"http://127.0.0.1/",
   367  			"http://[::1]/",
   368  			"https://127.0.0.1/",
   369  		} {
   370  			resp, err := ts.Client().Get(u)
   371  			if err != nil {
   372  				t.Errorf("Get(%q): %v", u, err)
   373  				continue
   374  			}
   375  			resp.Body.Close()
   376  			if resp.StatusCode != 200 {
   377  				t.Errorf("Get(%q): Response.StatusCode = %v, want 200", u, resp.StatusCode)
   378  			}
   379  			if gotTLS, wantTLS := resp.TLS != nil, strings.HasPrefix(u, "https://"); gotTLS != wantTLS {
   380  				t.Errorf("Get(%q): TLS: %v; want %v", u, gotTLS, wantTLS)
   381  			}
   382  		}
   383  	})
   384  }
   385  
   386  func TestServerNilHandler(t *testing.T) {
   387  	ts := NewTestServer(t, nil)
   388  	resp, err := ts.Client().Get("http://example.tld/")
   389  	if err != nil {
   390  		t.Fatalf("Get: %v", err)
   391  	}
   392  	resp.Body.Close()
   393  	if got, want := resp.StatusCode, 500; got != want {
   394  		t.Errorf("Response.StatusCode = %v, want %v", got, want)
   395  	}
   396  
   397  }
   398  
   399  func TestServerPanicErrAbortHandler(t *testing.T) {
   400  	synctest.Test(t, func(t *testing.T) {
   401  		ts := NewTestServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   402  			panic(http.ErrAbortHandler)
   403  		}))
   404  		resp, err := ts.Client().Get("http://example.com/")
   405  		if err == nil {
   406  			resp.Body.Close()
   407  			t.Errorf("request succeeded; want failure")
   408  		}
   409  	})
   410  }
   411  
   412  func TestServerPanicFailsTest(t *testing.T) {
   413  	runTest(t, func() {
   414  		synctest.Test(t, func(t *testing.T) {
   415  			ts := NewTestServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
   416  				panic("PANIC MESSAGE")
   417  			}))
   418  			ts.Client().Get("http://example.com/")
   419  		})
   420  	}, `--- FAIL: TestServerPanicFailsTest.*
   421      .*: httptest: panic in server handler: PANIC MESSAGE
   422  `)
   423  }
   424  
   425  func runTest(t *testing.T, f func(), pattern string) {
   426  	if os.Getenv("GO_WANT_HELPER_PROCESS") == "1" {
   427  		f()
   428  		return
   429  	}
   430  	t.Helper()
   431  	re := regexp.MustCompile(pattern)
   432  	testenv.MustHaveExec(t)
   433  	cmd := testenv.Command(t, testenv.Executable(t), "-test.run=^"+regexp.QuoteMeta(t.Name())+"$", "-test.count=1")
   434  	cmd = testenv.CleanCmdEnv(cmd)
   435  	cmd.Env = append(cmd.Env, "GO_WANT_HELPER_PROCESS=1")
   436  	out, _ := cmd.CombinedOutput()
   437  	if !re.Match(out) {
   438  		t.Errorf("got output:\n%s\nwant matching:\n%s", out, pattern)
   439  	}
   440  }
   441  

View as plain text