1
2
3
4
5
6
7 package httptest
8
9 import (
10 "context"
11 "crypto/tls"
12 "crypto/x509"
13 "flag"
14 "fmt"
15 "internal/nettest"
16 "log"
17 "net"
18 "net/http"
19 "net/http/internal/testcert"
20 "os"
21 "runtime"
22 "strings"
23 "sync"
24 "testing"
25 "time"
26 _ "unsafe"
27 )
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112 type Server struct {
113
114
115
116
117
118
119
120
121
122
123
124 URL string
125
126
127
128 Listener net.Listener
129
130
131
132 EnableHTTP2 bool
133
134
135
136
137 TLS *tls.Config
138
139
140 Config *http.Server
141
142 t testing.TB
143
144
145 certificate *x509.Certificate
146
147
148 startOnce sync.Once
149
150
151 started bool
152
153
154 fakeListener *nettest.Listener
155 fakeTLSListener *nettest.Listener
156
157
158
159 wg sync.WaitGroup
160
161 mu sync.Mutex
162 closed bool
163 conns map[net.Conn]http.ConnState
164
165
166
167 client *http.Client
168 }
169
170
171
172
173
174
175
176
177 func NewTestServer(t testing.TB, handler http.Handler) *Server {
178 s := &Server{
179 t: t,
180 Config: &http.Server{Handler: testServerHandler{t: t, h: handler}},
181 }
182 t.Cleanup(func() {
183 s.Close()
184 })
185 return s
186 }
187
188 type testServerHandler struct {
189 t testing.TB
190 h http.Handler
191 }
192
193 func (h testServerHandler) ServeHTTP(w http.ResponseWriter, req *http.Request) {
194 defer func() {
195 if err := recover(); err != nil {
196 if err != http.ErrAbortHandler {
197
198
199 const size = 64 << 10
200 buf := make([]byte, size)
201 buf = buf[:runtime.Stack(buf, false)]
202 h.t.Errorf("httptest: panic in server handler: %v\n%s", err, buf)
203 }
204
205 panic(http.ErrAbortHandler)
206 }
207 }()
208 if h.h != nil {
209 h.h.ServeHTTP(w, req)
210 } else {
211 w.WriteHeader(500)
212 }
213 }
214
215 func newLocalListener() net.Listener {
216 if serveFlag != "" {
217 l, err := net.Listen("tcp", serveFlag)
218 if err != nil {
219 panic(fmt.Sprintf("httptest: failed to listen on %v: %v", serveFlag, err))
220 }
221 return l
222 }
223 l, err := net.Listen("tcp", "127.0.0.1:0")
224 if err != nil {
225 if l, err = net.Listen("tcp6", "[::1]:0"); err != nil {
226 panic(fmt.Sprintf("httptest: failed to listen on a port: %v", err))
227 }
228 }
229 return l
230 }
231
232
233
234
235
236
237
238
239
240
241 var serveFlag string
242
243 func init() {
244 if strSliceContainsPrefix(os.Args, "-httptest.serve=") || strSliceContainsPrefix(os.Args, "--httptest.serve=") {
245 flag.StringVar(&serveFlag, "httptest.serve", "", "if non-empty, httptest.NewServer serves on this address and blocks.")
246 }
247 }
248
249 func strSliceContainsPrefix(v []string, pre string) bool {
250 for _, s := range v {
251 if strings.HasPrefix(s, pre) {
252 return true
253 }
254 }
255 return false
256 }
257
258
259
260
261
262
263
264
265
266 func NewServer(handler http.Handler) *Server {
267 ts := NewUnstartedServer(handler)
268 ts.Start()
269 return ts
270 }
271
272
273
274
275
276
277
278
279
280
281
282 func NewUnstartedServer(handler http.Handler) *Server {
283 return &Server{
284 Listener: newLocalListener(),
285 Config: &http.Server{Handler: handler},
286 }
287 }
288
289 func (s *Server) startCommon(useLoopback bool) {
290 s.mu.Lock()
291 defer s.mu.Unlock()
292 if s.started {
293 panic("Server already started")
294 }
295 if s.closed {
296 panic("Start of closed Server")
297 }
298 s.started = true
299 if s.t != nil && useLoopback {
300
301
302 s.startOnce.Do(func() {})
303
304
305
306
307
308 if s.Listener != nil {
309 panic("Server.Listener is unexpectedly set")
310 }
311 s.Listener = newLocalListener()
312 }
313 s.wrap()
314 }
315
316
317
318
319 func (s *Server) Start() {
320 s.startCommon(true)
321
322 tr := &http.Transport{}
323 s.client = &http.Client{Transport: tr}
324 if s.Listener == nil {
325 return
326 }
327 dialer := net.Dialer{}
328
329
330
331 tr.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
332 if tr.Dial != nil {
333 return tr.Dial(network, addr)
334 }
335 if addr == "example.com:80" || strings.HasSuffix(addr, ".example.com:80") {
336 addr = s.Listener.Addr().String()
337 }
338 return dialer.DialContext(ctx, network, addr)
339 }
340 s.URL = "http://" + s.Listener.Addr().String()
341 s.goServe(s.Listener)
342 if serveFlag != "" {
343 fmt.Fprintln(os.Stderr, "httptest: serving on", s.URL)
344 select {}
345 }
346 }
347
348 func (s *Server) initTLS() (tlsClientConfig *tls.Config, err error) {
349 cert, err := tls.X509KeyPair(testcert.LocalhostCert, testcert.LocalhostKey)
350 if err != nil {
351 return nil, err
352 }
353
354 existingConfig := s.TLS
355 if existingConfig != nil {
356 s.TLS = existingConfig.Clone()
357 } else {
358 s.TLS = new(tls.Config)
359 }
360 if s.TLS.NextProtos == nil {
361 nextProtos := []string{"http/1.1"}
362 if s.EnableHTTP2 {
363 nextProtos = []string{"h2"}
364 }
365 s.TLS.NextProtos = nextProtos
366 }
367 if len(s.TLS.Certificates) == 0 {
368 s.TLS.Certificates = []tls.Certificate{cert}
369 }
370 s.certificate, err = x509.ParseCertificate(s.TLS.Certificates[0].Certificate[0])
371 if err != nil {
372 return nil, err
373 }
374 certpool := x509.NewCertPool()
375 certpool.AddCert(s.certificate)
376 return &tls.Config{
377 RootCAs: certpool,
378 }, nil
379 }
380
381
382
383
384 func (s *Server) StartTLS() {
385 s.startCommon(true)
386
387 s.client = &http.Client{}
388
389 tlsClientConfig, err := s.initTLS()
390 if err != nil {
391 panic(fmt.Sprintf("httptest: NewTLSServer: %v", err))
392 }
393
394 tr := &http.Transport{
395 TLSClientConfig: tlsClientConfig,
396 ForceAttemptHTTP2: s.EnableHTTP2,
397 }
398 s.client.Transport = tr
399
400 if s.Listener == nil {
401 return
402 }
403 dialer := net.Dialer{}
404 tr.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
405 if tr.Dial != nil {
406 return tr.Dial(network, addr)
407 }
408 if addr == "example.com:443" || strings.HasSuffix(addr, ".example.com:443") {
409 addr = s.Listener.Addr().String()
410 }
411 return dialer.DialContext(ctx, network, addr)
412 }
413 s.Listener = tls.NewListener(s.Listener, s.TLS)
414 s.URL = "https://" + s.Listener.Addr().String()
415 s.goServe(s.Listener)
416 }
417
418 func (s *Server) startFakeNet() {
419 s.startCommon(false)
420
421 s.client = &http.Client{}
422
423 tlsClientConfig, err := s.initTLS()
424 if err != nil {
425 panic(fmt.Sprintf("httptest: NewTestServer: %v", err))
426 }
427
428 tr := &http.Transport{
429 TLSClientConfig: tlsClientConfig,
430 ForceAttemptHTTP2: s.EnableHTTP2,
431 }
432 s.client.Transport = tr
433
434 s.fakeListener = nettest.NewListener()
435 s.fakeTLSListener = nettest.NewListener()
436
437
438 tr.TLSClientConfig.InsecureSkipVerify = true
439 tr.DialContext = func(ctx context.Context, network, address string) (net.Conn, error) {
440 return s.fakeListener.NewConn(), nil
441 }
442 tr.DialTLSContext = func(ctx context.Context, network, address string) (net.Conn, error) {
443 return tls.Client(s.fakeTLSListener.NewConn(), tr.TLSClientConfig), nil
444 }
445 s.URL = "http://example.com"
446 s.goServe(s.fakeListener)
447 s.goServe(tls.NewListener(s.fakeTLSListener, s.TLS))
448 }
449
450
451
452
453
454
455
456
457
458 func NewTLSServer(handler http.Handler) *Server {
459 ts := NewUnstartedServer(handler)
460 ts.StartTLS()
461 return ts
462 }
463
464 type closeIdleTransport interface {
465 CloseIdleConnections()
466 }
467
468
469
470 func (s *Server) Close() {
471 s.mu.Lock()
472 if !s.closed {
473 s.closed = true
474 if s.Listener != nil {
475 s.Listener.Close()
476 }
477 if s.fakeListener != nil {
478 s.fakeListener.Close()
479 s.fakeTLSListener.Close()
480 }
481 s.Config.SetKeepAlivesEnabled(false)
482 for c, st := range s.conns {
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501 if st == http.StateIdle || st == http.StateNew {
502 s.closeConn(c)
503 }
504 }
505
506 t := time.AfterFunc(5*time.Second, s.logCloseHangDebugInfo)
507 defer t.Stop()
508 }
509 s.mu.Unlock()
510
511
512
513
514 if t, ok := http.DefaultTransport.(closeIdleTransport); ok {
515 t.CloseIdleConnections()
516 }
517
518
519 if s.client != nil {
520 if t, ok := s.client.Transport.(closeIdleTransport); ok {
521 t.CloseIdleConnections()
522 }
523 }
524 s.wg.Wait()
525 }
526
527 func (s *Server) logCloseHangDebugInfo() {
528 s.mu.Lock()
529 defer s.mu.Unlock()
530 var buf strings.Builder
531 buf.WriteString("httptest.Server blocked in Close after 5 seconds, waiting for connections:\n")
532 for c, st := range s.conns {
533 fmt.Fprintf(&buf, " %T %p %v in state %v\n", c, c, c.RemoteAddr(), st)
534 }
535 log.Print(buf.String())
536 }
537
538
539 func (s *Server) CloseClientConnections() {
540 s.mu.Lock()
541 nconn := len(s.conns)
542 ch := make(chan struct{}, nconn)
543 for c := range s.conns {
544 go s.closeConnChan(c, ch)
545 }
546 s.mu.Unlock()
547
548
549
550
551
552
553
554 timer := time.NewTimer(5 * time.Second)
555 defer timer.Stop()
556 for i := 0; i < nconn; i++ {
557 select {
558 case <-ch:
559 case <-timer.C:
560
561 return
562 }
563 }
564 }
565
566
567
568 func (s *Server) Certificate() *x509.Certificate {
569 return s.certificate
570 }
571
572
573
574
575 func (s *Server) Client() *http.Client {
576 if s.t != nil {
577 s.startOnce.Do(s.startFakeNet)
578 }
579 return s.client
580 }
581
582 func (s *Server) goServe(li net.Listener) {
583 s.wg.Add(1)
584 go func() {
585 defer s.wg.Done()
586 s.Config.Serve(li)
587 }()
588 }
589
590
591
592 func (s *Server) wrap() {
593 oldHook := s.Config.ConnState
594 s.Config.ConnState = func(c net.Conn, cs http.ConnState) {
595 s.mu.Lock()
596 defer s.mu.Unlock()
597
598 switch cs {
599 case http.StateNew:
600 if _, exists := s.conns[c]; exists {
601 panic("invalid state transition")
602 }
603 if s.conns == nil {
604 s.conns = make(map[net.Conn]http.ConnState)
605 }
606
607
608 s.wg.Add(1)
609 s.conns[c] = cs
610 if s.closed {
611
612
613
614
615 s.closeConn(c)
616 }
617 case http.StateActive:
618 if oldState, ok := s.conns[c]; ok {
619 if oldState != http.StateNew && oldState != http.StateIdle {
620 panic("invalid state transition")
621 }
622 s.conns[c] = cs
623 }
624 case http.StateIdle:
625 if oldState, ok := s.conns[c]; ok {
626 if oldState != http.StateActive {
627 panic("invalid state transition")
628 }
629 s.conns[c] = cs
630 }
631 if s.closed {
632 s.closeConn(c)
633 }
634 case http.StateHijacked, http.StateClosed:
635
636
637 if _, ok := s.conns[c]; ok {
638 delete(s.conns, c)
639
640
641 defer s.wg.Done()
642 }
643 }
644 if oldHook != nil {
645 oldHook(c, cs)
646 }
647 }
648 }
649
650
651
652 func (s *Server) closeConn(c net.Conn) { s.closeConnChan(c, nil) }
653
654
655
656 func (s *Server) closeConnChan(c net.Conn, done chan<- struct{}) {
657 c.Close()
658 if done != nil {
659 done <- struct{}{}
660 }
661 }
662
View as plain text