1
2
3
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
32
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
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
139 cnew := dial()
140 defer cnew.Close()
141
142
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()
152 }
153
154
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
169
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
191
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
203
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
221
222 func TestServerZeroValueClose(t *testing.T) {
223 ts := &Server{
224 Listener: onlyCloseListener{},
225 Config: &http.Server{},
226 }
227
228 ts.Close()
229 }
230
231
232
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
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
271
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