Source file
src/net/http/clientserver_test.go
1
2
3
4
5
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"
44
45 _ "golang.org/x/net/http3"
46 )
47
48
49 func registerHTTP3Transport(*http.Transport) <-chan *quic.Endpoint
50
51
52 func registerHTTP3Server(*http.Server) <-chan *quic.Endpoint
53
54 type testMode string
55
56 const (
57 http1Mode = testMode("h1")
58 https1Mode = testMode("https1")
59 http2Mode = testMode("h2")
60 http2UnencryptedMode = testMode("h2unencrypted")
61 http3Mode = testMode("h3")
62 )
63
64 type (
65 testAddMode []testMode
66 testSkipMode []testMode
67 )
68
69
70
71
72
73
74
75
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
100
101
102
103
104
105
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
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
152
153
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
210
211
212
213
214
215
216
217
218
219
220
221
222
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
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
317
318
319
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
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)
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")
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
462
463 type h12Compare struct {
464 Handler func(ResponseWriter, *Request)
465 ReqFunc reqFunc
466 CheckResponse func(proto string, res *Response)
467 EarlyCheckResponse func(proto string, res *Response)
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
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")
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
682
683
684
685
686
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")
697 if err == nil {
698 err = rc.Flush()
699 }
700
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
719
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
749
750
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
816
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
842
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
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,
896 "Client-Trailer-B": nil,
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
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
932
933
934
935 w.Header().Set("Server-Trailer-A", "valuea")
936 w.Header().Set("Server-Trailer-C", "valuec")
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
951
952
953
954
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")
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
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
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
1027 go func() {
1028 defer wg.Done()
1029 if mode != http2Mode {
1030
1031
1032
1033
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")
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
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
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
1283
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
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
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},
1383 {"Foo", "foo\x00bar", false},
1384 {"Foo", "two\nlines", false},
1385 {"bogus\nkey", "v", false},
1386 {"A space", "v", false},
1387 {"имя", "v", false},
1388 {"name", "валю", true},
1389 {"", "v", false},
1390 {"k", "", true},
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
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
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
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
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
1672
1673
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
1692
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)
1709 conn.Write([]byte("bar"))
1710 return
1711 }
1712 io.WriteString(w, "foo")
1713 w.(Flusher).Flush()
1714 w.WriteHeader(0)
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
1733 if mode == http2Mode {
1734
1735
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)
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
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"
1830 },
1831 Opts: []any{
1832 func(s *Server) {
1833
1834
1835
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")
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
1958
1959
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