Source file
src/net/http/client_test.go
1
2
3
4
5
6
7 package http_test
8
9 import (
10 "bytes"
11 "context"
12 "crypto/tls"
13 "encoding/base64"
14 "errors"
15 "fmt"
16 "internal/testenv"
17 "io"
18 "log"
19 "net"
20 . "net/http"
21 "net/http/cookiejar"
22 "net/http/httptest"
23 "net/url"
24 "reflect"
25 "runtime"
26 "strconv"
27 "strings"
28 "sync"
29 "sync/atomic"
30 "testing"
31 "testing/synctest"
32 "time"
33 )
34
35 var robotsTxtHandler = HandlerFunc(func(w ResponseWriter, r *Request) {
36 w.Header().Set("Last-Modified", "sometime")
37 fmt.Fprintf(w, "User-agent: go\nDisallow: /something/")
38 })
39
40
41
42 func pedanticReadAll(r io.Reader) (b []byte, err error) {
43 var bufa [64]byte
44 buf := bufa[:]
45 for {
46 n, err := r.Read(buf)
47 if n == 0 && err == nil {
48 return nil, fmt.Errorf("Read: n=0 with err=nil")
49 }
50 b = append(b, buf[:n]...)
51 if err == io.EOF {
52 n, err := r.Read(buf)
53 if n != 0 || err != io.EOF {
54 return nil, fmt.Errorf("Read: n=%d err=%#v after EOF", n, err)
55 }
56 return b, nil
57 }
58 if err != nil {
59 return b, err
60 }
61 }
62 }
63
64 func TestClient(t *testing.T) {
65 runSynctest(t, testClient, []testMode{http1Mode, https1Mode, http2UnencryptedMode, http2Mode})
66 }
67 func testClient(t *testing.T, mode testMode) {
68 ts := newClientServerTest(t, mode, robotsTxtHandler).ts
69
70 c := ts.Client()
71 r, err := c.Get(ts.URL)
72 var b []byte
73 if err == nil {
74 b, err = pedanticReadAll(r.Body)
75 r.Body.Close()
76 }
77 if err != nil {
78 t.Error(err)
79 } else if s := string(b); !strings.HasPrefix(s, "User-agent:") {
80 t.Errorf("Incorrect page body (did not begin with User-agent): %q", s)
81 }
82 }
83
84 func TestClientHead(t *testing.T) { runSynctest(t, testClientHead) }
85 func testClientHead(t *testing.T, mode testMode) {
86 cst := newClientServerTest(t, mode, robotsTxtHandler)
87 r, err := cst.c.Head(cst.ts.URL)
88 if err != nil {
89 t.Fatal(err)
90 }
91 if _, ok := r.Header["Last-Modified"]; !ok {
92 t.Error("Last-Modified header not found.")
93 }
94 }
95
96 type recordingTransport struct {
97 req *Request
98 }
99
100 func (t *recordingTransport) RoundTrip(req *Request) (resp *Response, err error) {
101 t.req = req
102 return nil, errors.New("dummy impl")
103 }
104
105 func TestGetRequestFormat(t *testing.T) {
106 setParallel(t)
107 defer afterTest(t)
108 tr := &recordingTransport{}
109 client := &Client{Transport: tr}
110 url := "http://dummy.faketld/"
111 client.Get(url)
112 if tr.req.Method != "GET" {
113 t.Errorf("expected method %q; got %q", "GET", tr.req.Method)
114 }
115 if tr.req.URL.String() != url {
116 t.Errorf("expected URL %q; got %q", url, tr.req.URL.String())
117 }
118 if tr.req.Header == nil {
119 t.Errorf("expected non-nil request Header")
120 }
121 }
122
123 func TestPostRequestFormat(t *testing.T) {
124 defer afterTest(t)
125 tr := &recordingTransport{}
126 client := &Client{Transport: tr}
127
128 url := "http://dummy.faketld/"
129 json := `{"key":"value"}`
130 b := strings.NewReader(json)
131 client.Post(url, "application/json", b)
132
133 if tr.req.Method != "POST" {
134 t.Errorf("got method %q, want %q", tr.req.Method, "POST")
135 }
136 if tr.req.URL.String() != url {
137 t.Errorf("got URL %q, want %q", tr.req.URL.String(), url)
138 }
139 if tr.req.Header == nil {
140 t.Fatalf("expected non-nil request Header")
141 }
142 if tr.req.Close {
143 t.Error("got Close true, want false")
144 }
145 if g, e := tr.req.ContentLength, int64(len(json)); g != e {
146 t.Errorf("got ContentLength %d, want %d", g, e)
147 }
148 }
149
150 func TestPostFormRequestFormat(t *testing.T) {
151 defer afterTest(t)
152 tr := &recordingTransport{}
153 client := &Client{Transport: tr}
154
155 urlStr := "http://dummy.faketld/"
156 form := make(url.Values)
157 form.Set("foo", "bar")
158 form.Add("foo", "bar2")
159 form.Set("bar", "baz")
160 client.PostForm(urlStr, form)
161
162 if tr.req.Method != "POST" {
163 t.Errorf("got method %q, want %q", tr.req.Method, "POST")
164 }
165 if tr.req.URL.String() != urlStr {
166 t.Errorf("got URL %q, want %q", tr.req.URL.String(), urlStr)
167 }
168 if tr.req.Header == nil {
169 t.Fatalf("expected non-nil request Header")
170 }
171 if g, e := tr.req.Header.Get("Content-Type"), "application/x-www-form-urlencoded"; g != e {
172 t.Errorf("got Content-Type %q, want %q", g, e)
173 }
174 if tr.req.Close {
175 t.Error("got Close true, want false")
176 }
177
178 expectedBody := "foo=bar&foo=bar2&bar=baz"
179 expectedBody1 := "bar=baz&foo=bar&foo=bar2"
180 if g, e := tr.req.ContentLength, int64(len(expectedBody)); g != e {
181 t.Errorf("got ContentLength %d, want %d", g, e)
182 }
183 bodyb, err := io.ReadAll(tr.req.Body)
184 if err != nil {
185 t.Fatalf("ReadAll on req.Body: %v", err)
186 }
187 if g := string(bodyb); g != expectedBody && g != expectedBody1 {
188 t.Errorf("got body %q, want %q or %q", g, expectedBody, expectedBody1)
189 }
190 }
191
192 func TestClientRedirects(t *testing.T) { runSynctest(t, testClientRedirects) }
193 func testClientRedirects(t *testing.T, mode testMode) {
194 var ts *httptest.Server
195 ts = newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
196 n, _ := strconv.Atoi(r.FormValue("n"))
197
198 if n == 7 {
199 if g, e := r.Referer(), ts.URL+"/?n=6"; e != g {
200 t.Errorf("on request ?n=7, expected referer of %q; got %q", e, g)
201 }
202 }
203 if n < 15 {
204 Redirect(w, r, fmt.Sprintf("/?n=%d", n+1), StatusTemporaryRedirect)
205 return
206 }
207 fmt.Fprintf(w, "n=%d", n)
208 })).ts
209
210 c := ts.Client()
211 _, err := c.Get(ts.URL)
212 if e, g := `Get "/?n=10": stopped after 10 redirects`, fmt.Sprintf("%v", err); e != g {
213 t.Errorf("with default client Get, expected error %q, got %q", e, g)
214 }
215
216
217 _, err = c.Head(ts.URL)
218 if e, g := `Head "/?n=10": stopped after 10 redirects`, fmt.Sprintf("%v", err); e != g {
219 t.Errorf("with default client Head, expected error %q, got %q", e, g)
220 }
221
222
223 greq, _ := NewRequest("GET", ts.URL, nil)
224 _, err = c.Do(greq)
225 if e, g := `Get "/?n=10": stopped after 10 redirects`, fmt.Sprintf("%v", err); e != g {
226 t.Errorf("with default client Do, expected error %q, got %q", e, g)
227 }
228
229
230 greq.Method = ""
231 _, err = c.Do(greq)
232 if e, g := `Get "/?n=10": stopped after 10 redirects`, fmt.Sprintf("%v", err); e != g {
233 t.Errorf("with default client Do and empty Method, expected error %q, got %q", e, g)
234 }
235
236 var checkErr error
237 var lastVia []*Request
238 var lastReq *Request
239 c.CheckRedirect = func(req *Request, via []*Request) error {
240 lastReq = req
241 lastVia = via
242 return checkErr
243 }
244 res, err := c.Get(ts.URL)
245 if err != nil {
246 t.Fatalf("Get error: %v", err)
247 }
248 res.Body.Close()
249 finalURL := res.Request.URL.String()
250 if e, g := "<nil>", fmt.Sprintf("%v", err); e != g {
251 t.Errorf("with custom client, expected error %q, got %q", e, g)
252 }
253 if !strings.HasSuffix(finalURL, "/?n=15") {
254 t.Errorf("expected final url to end in /?n=15; got url %q", finalURL)
255 }
256 if e, g := 15, len(lastVia); e != g {
257 t.Errorf("expected lastVia to have contained %d elements; got %d", e, g)
258 }
259
260
261 creq, _ := NewRequest("HEAD", ts.URL, nil)
262 cancel := make(chan struct{})
263 creq.Cancel = cancel
264 if _, err := c.Do(creq); err != nil {
265 t.Fatal(err)
266 }
267 if lastReq == nil {
268 t.Fatal("didn't see redirect")
269 }
270 if lastReq.Cancel != cancel {
271 t.Errorf("expected lastReq to have the cancel channel set on the initial req")
272 }
273
274 checkErr = errors.New("no redirects allowed")
275 res, err = c.Get(ts.URL)
276 if urlError, ok := err.(*url.Error); !ok || urlError.Err != checkErr {
277 t.Errorf("with redirects forbidden, expected a *url.Error with our 'no redirects allowed' error inside; got %#v (%q)", err, err)
278 }
279 if res == nil {
280 t.Fatalf("Expected a non-nil Response on CheckRedirect failure (https://golang.org/issue/3795)")
281 }
282 res.Body.Close()
283 if res.Header.Get("Location") == "" {
284 t.Errorf("no Location header in Response")
285 }
286 }
287
288
289 func TestClientRedirectsContext(t *testing.T) { runSynctest(t, testClientRedirectsContext) }
290 func testClientRedirectsContext(t *testing.T, mode testMode) {
291 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
292 Redirect(w, r, "/", StatusTemporaryRedirect)
293 })).ts
294
295 ctx, cancel := context.WithCancel(context.Background())
296 c := ts.Client()
297 c.CheckRedirect = func(req *Request, via []*Request) error {
298 cancel()
299 select {
300 case <-req.Context().Done():
301 return nil
302 case <-time.After(5 * time.Second):
303 return errors.New("redirected request's context never expired after root request canceled")
304 }
305 }
306 req, _ := NewRequestWithContext(ctx, "GET", ts.URL, nil)
307 _, err := c.Do(req)
308 ue, ok := err.(*url.Error)
309 if !ok {
310 t.Fatalf("got error %T; want *url.Error", err)
311 }
312 if ue.Err != context.Canceled {
313 t.Errorf("url.Error.Err = %v; want %v", ue.Err, context.Canceled)
314 }
315 }
316
317 type redirectTest struct {
318 suffix string
319 want int
320 redirectBody string
321 }
322
323 func TestPostRedirects(t *testing.T) {
324 postRedirectTests := []redirectTest{
325 {"/", 200, "first"},
326 {"/?code=301&next=302", 200, "c301"},
327 {"/?code=302&next=302", 200, "c302"},
328 {"/?code=303&next=301", 200, "c303wc301"},
329 {"/?code=304", 304, "c304"},
330 {"/?code=305", 305, "c305"},
331 {"/?code=307&next=303,308,302", 200, "c307"},
332 {"/?code=308&next=302,301", 200, "c308"},
333 {"/?code=404", 404, "c404"},
334 }
335
336 wantSegments := []string{
337 `POST / "first"`,
338 `POST /?code=301&next=302 "c301"`,
339 `GET /?code=302 ""`,
340 `GET / ""`,
341 `POST /?code=302&next=302 "c302"`,
342 `GET /?code=302 ""`,
343 `GET / ""`,
344 `POST /?code=303&next=301 "c303wc301"`,
345 `GET /?code=301 ""`,
346 `GET / ""`,
347 `POST /?code=304 "c304"`,
348 `POST /?code=305 "c305"`,
349 `POST /?code=307&next=303,308,302 "c307"`,
350 `POST /?code=303&next=308,302 "c307"`,
351 `GET /?code=308&next=302 ""`,
352 `GET /?code=302 ""`,
353 `GET / ""`,
354 `POST /?code=308&next=302,301 "c308"`,
355 `POST /?code=302&next=301 "c308"`,
356 `GET /?code=301 ""`,
357 `GET / ""`,
358 `POST /?code=404 "c404"`,
359 }
360 want := strings.Join(wantSegments, "\n")
361 runSynctest(t, func(t *testing.T, mode testMode) {
362 testRedirectsByMethod(t, mode, "POST", postRedirectTests, want)
363 }, http3SkippedMode)
364 }
365
366 func TestDeleteRedirects(t *testing.T) {
367 deleteRedirectTests := []redirectTest{
368 {"/", 200, "first"},
369 {"/?code=301&next=302,308", 200, "c301"},
370 {"/?code=302&next=302", 200, "c302"},
371 {"/?code=303", 200, "c303"},
372 {"/?code=307&next=301,308,303,302,304", 304, "c307"},
373 {"/?code=308&next=307", 200, "c308"},
374 {"/?code=404", 404, "c404"},
375 }
376
377 wantSegments := []string{
378 `DELETE / "first"`,
379 `DELETE /?code=301&next=302,308 "c301"`,
380 `GET /?code=302&next=308 ""`,
381 `GET /?code=308 ""`,
382 `GET / ""`,
383 `DELETE /?code=302&next=302 "c302"`,
384 `GET /?code=302 ""`,
385 `GET / ""`,
386 `DELETE /?code=303 "c303"`,
387 `GET / ""`,
388 `DELETE /?code=307&next=301,308,303,302,304 "c307"`,
389 `DELETE /?code=301&next=308,303,302,304 "c307"`,
390 `GET /?code=308&next=303,302,304 ""`,
391 `GET /?code=303&next=302,304 ""`,
392 `GET /?code=302&next=304 ""`,
393 `GET /?code=304 ""`,
394 `DELETE /?code=308&next=307 "c308"`,
395 `DELETE /?code=307 "c308"`,
396 `DELETE / "c308"`,
397 `DELETE /?code=404 "c404"`,
398 }
399 want := strings.Join(wantSegments, "\n")
400 runSynctest(t, func(t *testing.T, mode testMode) {
401 testRedirectsByMethod(t, mode, "DELETE", deleteRedirectTests, want)
402 }, http3SkippedMode)
403 }
404
405 func TestQueryRedirects(t *testing.T) {
406
407
408
409 queryRedirectTests := []redirectTest{
410 {"/", 200, "first"},
411 {"/?code=302&next=302", 200, "c302"},
412 {"/?code=301&next=302,308", 200, "c301"},
413 {"/?code=303&next=301", 200, "c303"},
414 {"/?code=307&next=303,302", 200, "c307"},
415 {"/?code=404", 404, "c404"},
416 }
417
418 wantSegments := []string{
419 `QUERY / "first"`,
420 `QUERY /?code=302&next=302 "c302"`,
421 `QUERY /?code=302 "c302"`,
422 `QUERY / "c302"`,
423 `QUERY /?code=301&next=302,308 "c301"`,
424 `QUERY /?code=302&next=308 "c301"`,
425 `QUERY /?code=308 "c301"`,
426 `QUERY / "c301"`,
427 `QUERY /?code=303&next=301 "c303"`,
428 `GET /?code=301 ""`,
429 `GET / ""`,
430 `QUERY /?code=307&next=303,302 "c307"`,
431 `QUERY /?code=303&next=302 "c307"`,
432 `GET /?code=302 ""`,
433 `GET / ""`,
434 `QUERY /?code=404 "c404"`,
435 }
436 want := strings.Join(wantSegments, "\n")
437 runSynctest(t, func(t *testing.T, mode testMode) {
438 testRedirectsByMethod(t, mode, "QUERY", queryRedirectTests, want)
439 }, http3SkippedMode)
440 }
441
442 func testRedirectsByMethod(t *testing.T, mode testMode, method string, table []redirectTest, want string) {
443 var log struct {
444 sync.Mutex
445 bytes.Buffer
446 }
447 var ts *httptest.Server
448 ts = newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
449 log.Lock()
450 slurp, _ := io.ReadAll(r.Body)
451 fmt.Fprintf(&log.Buffer, "%s %s %q", r.Method, r.RequestURI, slurp)
452 if cl := r.Header.Get("Content-Length"); r.Method == "GET" && len(slurp) == 0 && (r.ContentLength != 0 || cl != "") {
453 fmt.Fprintf(&log.Buffer, " (but with body=%T, content-length = %v, %q)", r.Body, r.ContentLength, cl)
454 }
455 log.WriteByte('\n')
456 log.Unlock()
457 urlQuery := r.URL.Query()
458 if v := urlQuery.Get("code"); v != "" {
459 location := ts.URL
460 if final := urlQuery.Get("next"); final != "" {
461 first, rest, _ := strings.Cut(final, ",")
462 location = fmt.Sprintf("%s?code=%s", location, first)
463 if rest != "" {
464 location = fmt.Sprintf("%s&next=%s", location, rest)
465 }
466 }
467 code, _ := strconv.Atoi(v)
468 if code/100 == 3 {
469 w.Header().Set("Location", location)
470 }
471 w.WriteHeader(code)
472 }
473 })).ts
474
475 c := ts.Client()
476 for _, tt := range table {
477 content := tt.redirectBody
478 req, _ := NewRequest(method, ts.URL+tt.suffix, strings.NewReader(content))
479 req.GetBody = func() (io.ReadCloser, error) { return io.NopCloser(strings.NewReader(content)), nil }
480 res, err := c.Do(req)
481 if err != nil {
482 t.Fatal(err)
483 }
484 if res.StatusCode != tt.want {
485 t.Errorf("POST %s: status code = %d; want %d", tt.suffix, res.StatusCode, tt.want)
486 }
487 }
488 log.Lock()
489 got := log.String()
490 log.Unlock()
491
492 got = strings.TrimSpace(got)
493 want = strings.TrimSpace(want)
494
495 if got != want {
496 got, want, lines := removeCommonLines(got, want)
497 t.Errorf("Log differs after %d common lines.\n\nGot:\n%s\n\nWant:\n%s\n", lines, got, want)
498 }
499 }
500
501 func removeCommonLines(a, b string) (asuffix, bsuffix string, commonLines int) {
502 for {
503 nl := strings.IndexByte(a, '\n')
504 if nl < 0 {
505 return a, b, commonLines
506 }
507 line := a[:nl+1]
508 if !strings.HasPrefix(b, line) {
509 return a, b, commonLines
510 }
511 commonLines++
512 a = a[len(line):]
513 b = b[len(line):]
514 }
515 }
516
517 func TestClientRedirectUseResponse(t *testing.T) { runSynctest(t, testClientRedirectUseResponse) }
518 func testClientRedirectUseResponse(t *testing.T, mode testMode) {
519 const body = "Hello, world."
520 var ts *httptest.Server
521 ts = newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
522 if strings.Contains(r.URL.Path, "/other") {
523 io.WriteString(w, "wrong body")
524 } else {
525 scheme := "http"
526 if r.TLS != nil {
527 scheme = "https"
528 }
529 w.Header().Set("Location", fmt.Sprintf("%s://%s/other", scheme, r.Host))
530 w.WriteHeader(StatusFound)
531 io.WriteString(w, body)
532 }
533 })).ts
534
535 c := ts.Client()
536 c.CheckRedirect = func(req *Request, via []*Request) error {
537 if req.Response == nil {
538 t.Error("expected non-nil Request.Response")
539 }
540 return ErrUseLastResponse
541 }
542 res, err := c.Get(ts.URL)
543 if err != nil {
544 t.Fatal(err)
545 }
546 if res.StatusCode != StatusFound {
547 t.Errorf("status = %d; want %d", res.StatusCode, StatusFound)
548 }
549 defer res.Body.Close()
550 slurp, err := io.ReadAll(res.Body)
551 if err != nil {
552 t.Fatal(err)
553 }
554 if string(slurp) != body {
555 t.Errorf("body = %q; want %q", slurp, body)
556 }
557 }
558
559
560
561 func TestClientRedirectNoLocation(t *testing.T) { runNoSynctest(t, testClientRedirectNoLocation) }
562 func testClientRedirectNoLocation(t *testing.T, mode testMode) {
563 for _, code := range []int{301, 308} {
564 synctest.Subtest(t, fmt.Sprint(code), func(t *testing.T) {
565 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
566 w.Header().Set("Foo", "Bar")
567 w.WriteHeader(code)
568 }))
569 res, err := cst.c.Get(cst.ts.URL)
570 if err != nil {
571 t.Fatal(err)
572 }
573 res.Body.Close()
574 if res.StatusCode != code {
575 t.Errorf("status = %d; want %d", res.StatusCode, code)
576 }
577 if got := res.Header.Get("Foo"); got != "Bar" {
578 t.Errorf("Foo header = %q; want Bar", got)
579 }
580 })
581 }
582 }
583
584
585 func TestClientRedirect308NoGetBody(t *testing.T) { runSynctest(t, testClientRedirect308NoGetBody) }
586 func testClientRedirect308NoGetBody(t *testing.T, mode testMode) {
587 const fakeURL = "https://localhost:1234/"
588 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
589 w.Header().Set("Location", fakeURL)
590 w.WriteHeader(308)
591 })).ts
592 req, err := NewRequest("POST", ts.URL, strings.NewReader("some body"))
593 if err != nil {
594 t.Fatal(err)
595 }
596 c := ts.Client()
597 req.GetBody = nil
598 res, err := c.Do(req)
599 if err != nil {
600 t.Fatal(err)
601 }
602 res.Body.Close()
603 if res.StatusCode != 308 {
604 t.Errorf("status = %d; want %d", res.StatusCode, 308)
605 }
606 if got := res.Header.Get("Location"); got != fakeURL {
607 t.Errorf("Location header = %q; want %q", got, fakeURL)
608 }
609 }
610
611 var expectedCookies = []*Cookie{
612 {Name: "ChocolateChip", Value: "tasty"},
613 {Name: "First", Value: "Hit"},
614 {Name: "Second", Value: "Hit"},
615 }
616
617 var echoCookiesRedirectHandler = HandlerFunc(func(w ResponseWriter, r *Request) {
618 for _, cookie := range r.Cookies() {
619 SetCookie(w, cookie)
620 }
621 if r.URL.Path == "/" {
622 SetCookie(w, expectedCookies[1])
623 Redirect(w, r, "/second", StatusMovedPermanently)
624 } else {
625 SetCookie(w, expectedCookies[2])
626 w.Write([]byte("hello"))
627 }
628 })
629
630 func TestHostMismatchCookies(t *testing.T) { runSynctest(t, testHostMismatchCookies) }
631 func testHostMismatchCookies(t *testing.T, mode testMode) {
632 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
633 for _, c := range r.Cookies() {
634 c.Value = "SetOnServer"
635 SetCookie(w, c)
636 }
637 })).ts
638
639 reqURL, _ := url.Parse(ts.URL)
640 hostURL := *reqURL
641 hostURL.Host = "cookies.example.com"
642
643 c := ts.Client()
644 c.Jar = new(TestJar)
645 c.Jar.SetCookies(reqURL, []*Cookie{{Name: "First", Value: "SetOnClient"}})
646 c.Jar.SetCookies(&hostURL, []*Cookie{{Name: "Second", Value: "SetOnClient"}})
647
648 req, _ := NewRequest("GET", ts.URL, NoBody)
649 req.Host = hostURL.Host
650 resp, err := c.Do(req)
651 if err != nil {
652 t.Fatalf("Get: %v", err)
653 }
654 resp.Body.Close()
655
656 matchReturnedCookies(t, []*Cookie{{Name: "First", Value: "SetOnClient"}}, c.Jar.Cookies(reqURL))
657 matchReturnedCookies(t, []*Cookie{{Name: "Second", Value: "SetOnServer"}}, c.Jar.Cookies(&hostURL))
658 }
659
660 func TestClientSendsCookieFromJar(t *testing.T) {
661 defer afterTest(t)
662 tr := &recordingTransport{}
663 client := &Client{Transport: tr}
664 client.Jar = &TestJar{perURL: make(map[string][]*Cookie)}
665 us := "http://dummy.faketld/"
666 u, _ := url.Parse(us)
667 client.Jar.SetCookies(u, expectedCookies)
668
669 client.Get(us)
670 matchReturnedCookies(t, expectedCookies, tr.req.Cookies())
671
672 client.Head(us)
673 matchReturnedCookies(t, expectedCookies, tr.req.Cookies())
674
675 client.Post(us, "text/plain", strings.NewReader("body"))
676 matchReturnedCookies(t, expectedCookies, tr.req.Cookies())
677
678 client.PostForm(us, url.Values{})
679 matchReturnedCookies(t, expectedCookies, tr.req.Cookies())
680
681 req, _ := NewRequest("GET", us, nil)
682 client.Do(req)
683 matchReturnedCookies(t, expectedCookies, tr.req.Cookies())
684
685 req, _ = NewRequest("POST", us, nil)
686 client.Do(req)
687 matchReturnedCookies(t, expectedCookies, tr.req.Cookies())
688 }
689
690
691
692 type TestJar struct {
693 m sync.Mutex
694 perURL map[string][]*Cookie
695 }
696
697 func (j *TestJar) SetCookies(u *url.URL, cookies []*Cookie) {
698 j.m.Lock()
699 defer j.m.Unlock()
700 if j.perURL == nil {
701 j.perURL = make(map[string][]*Cookie)
702 }
703 j.perURL[u.Host] = cookies
704 }
705
706 func (j *TestJar) Cookies(u *url.URL) []*Cookie {
707 j.m.Lock()
708 defer j.m.Unlock()
709 return j.perURL[u.Host]
710 }
711
712 func TestRedirectCookiesJar(t *testing.T) { runSynctest(t, testRedirectCookiesJar) }
713 func testRedirectCookiesJar(t *testing.T, mode testMode) {
714 var ts *httptest.Server
715 ts = newClientServerTest(t, mode, echoCookiesRedirectHandler).ts
716 c := ts.Client()
717 c.Jar = new(TestJar)
718 u, _ := url.Parse(ts.URL)
719 c.Jar.SetCookies(u, []*Cookie{expectedCookies[0]})
720 resp, err := c.Get(ts.URL)
721 if err != nil {
722 t.Fatalf("Get: %v", err)
723 }
724 resp.Body.Close()
725 matchReturnedCookies(t, expectedCookies, resp.Cookies())
726 }
727
728 func matchReturnedCookies(t *testing.T, expected, given []*Cookie) {
729 if len(given) != len(expected) {
730 t.Logf("Received cookies: %v", given)
731 t.Errorf("Expected %d cookies, got %d", len(expected), len(given))
732 }
733 for _, ec := range expected {
734 foundC := false
735 for _, c := range given {
736 if ec.Name == c.Name && ec.Value == c.Value {
737 foundC = true
738 break
739 }
740 }
741 if !foundC {
742 t.Errorf("Missing cookie %v", ec)
743 }
744 }
745 }
746
747 func TestJarCalls(t *testing.T) { runSynctest(t, testJarCalls, []testMode{http1Mode}) }
748 func testJarCalls(t *testing.T, mode testMode) {
749 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
750 pathSuffix := r.RequestURI[1:]
751 if r.RequestURI == "/nosetcookie" {
752 return
753 }
754 SetCookie(w, &Cookie{Name: "name" + pathSuffix, Value: "val" + pathSuffix})
755 if r.RequestURI == "/" {
756 Redirect(w, r, "http://secondhost.fake/secondpath", 302)
757 }
758 })).ts
759 jar := new(RecordingJar)
760 c := ts.Client()
761 c.Jar = jar
762 c.Transport.(*Transport).Dial = func(_ string, _ string) (net.Conn, error) {
763 return net.Dial("tcp", ts.Listener.Addr().String())
764 }
765 _, err := c.Get("http://firsthost.fake/")
766 if err != nil {
767 t.Fatal(err)
768 }
769 _, err = c.Get("http://firsthost.fake/nosetcookie")
770 if err != nil {
771 t.Fatal(err)
772 }
773 got := jar.log.String()
774 want := `Cookies("http://firsthost.fake/")
775 SetCookie("http://firsthost.fake/", [name=val])
776 Cookies("http://secondhost.fake/secondpath")
777 SetCookie("http://secondhost.fake/secondpath", [namesecondpath=valsecondpath])
778 Cookies("http://firsthost.fake/nosetcookie")
779 `
780 if got != want {
781 t.Errorf("Got Jar calls:\n%s\nWant:\n%s", got, want)
782 }
783 }
784
785
786
787 type RecordingJar struct {
788 mu sync.Mutex
789 log bytes.Buffer
790 }
791
792 func (j *RecordingJar) SetCookies(u *url.URL, cookies []*Cookie) {
793 j.logf("SetCookie(%q, %v)\n", u, cookies)
794 }
795
796 func (j *RecordingJar) Cookies(u *url.URL) []*Cookie {
797 j.logf("Cookies(%q)\n", u)
798 return nil
799 }
800
801 func (j *RecordingJar) logf(format string, args ...any) {
802 j.mu.Lock()
803 defer j.mu.Unlock()
804 fmt.Fprintf(&j.log, format, args...)
805 }
806
807 func TestStreamingGet(t *testing.T) { runSynctest(t, testStreamingGet) }
808 func testStreamingGet(t *testing.T, mode testMode) {
809 say := make(chan string)
810 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
811 w.(Flusher).Flush()
812 for str := range say {
813 w.Write([]byte(str))
814 w.(Flusher).Flush()
815 }
816 }))
817
818 c := cst.c
819 res, err := c.Get(cst.ts.URL)
820 if err != nil {
821 t.Fatal(err)
822 }
823 var buf [10]byte
824 for _, str := range []string{"i", "am", "also", "known", "as", "comet"} {
825 say <- str
826 n, err := io.ReadFull(res.Body, buf[:len(str)])
827 if err != nil {
828 t.Fatalf("ReadFull on %q: %v", str, err)
829 }
830 if n != len(str) {
831 t.Fatalf("Receiving %q, only read %d bytes", str, n)
832 }
833 got := string(buf[0:n])
834 if got != str {
835 t.Fatalf("Expected %q, got %q", str, got)
836 }
837 }
838 close(say)
839 _, err = io.ReadFull(res.Body, buf[0:1])
840 if err != io.EOF {
841 t.Fatalf("at end expected EOF, got %v", err)
842 }
843 }
844
845 type writeCountingConn struct {
846 net.Conn
847 count *int
848 }
849
850 func (c *writeCountingConn) Write(p []byte) (int, error) {
851 *c.count++
852 return c.Conn.Write(p)
853 }
854
855
856
857 func TestClientWrites(t *testing.T) { runNoSynctest(t, testClientWrites, []testMode{http1Mode}) }
858 func testClientWrites(t *testing.T, mode testMode) {
859 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
860 }), optRealNet).ts
861
862 writes := 0
863 dialer := func(netz string, addr string) (net.Conn, error) {
864 c, err := net.Dial(netz, addr)
865 if err == nil {
866 c = &writeCountingConn{c, &writes}
867 }
868 return c, err
869 }
870 c := ts.Client()
871 c.Transport.(*Transport).Dial = dialer
872
873 _, err := c.Get(ts.URL)
874 if err != nil {
875 t.Fatal(err)
876 }
877 if writes != 1 {
878 t.Errorf("Get request did %d Write calls, want 1", writes)
879 }
880
881 writes = 0
882 _, err = c.PostForm(ts.URL, url.Values{"foo": {"bar"}})
883 if err != nil {
884 t.Fatal(err)
885 }
886 if writes != 1 {
887 t.Errorf("Post request did %d Write calls, want 1", writes)
888 }
889 }
890
891 func TestClientInsecureTransport(t *testing.T) {
892 runNoSynctest(t, testClientInsecureTransport, []testMode{https1Mode, http2Mode})
893 }
894 func testClientInsecureTransport(t *testing.T, mode testMode) {
895 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
896 w.Write([]byte("Hello"))
897 }), optRealNet)
898 ts := cst.ts
899 errLog := new(strings.Builder)
900 ts.Config.ErrorLog = log.New(errLog, "", 0)
901
902
903
904
905 c := ts.Client()
906 for _, insecure := range []bool{true, false} {
907 c.Transport.(*Transport).TLSClientConfig = &tls.Config{
908 InsecureSkipVerify: insecure,
909 NextProtos: cst.tr.TLSClientConfig.NextProtos,
910 }
911 req, _ := NewRequest("GET", ts.URL, nil)
912 req.Header.Set("Connection", "close")
913 res, err := c.Do(req)
914 if (err == nil) != insecure {
915 t.Errorf("insecure=%v: got unexpected err=%v", insecure, err)
916 }
917 if res != nil {
918 res.Body.Close()
919 }
920 }
921
922 cst.close()
923 if !strings.Contains(errLog.String(), "TLS handshake error") {
924 t.Errorf("expected an error log message containing 'TLS handshake error'; got %q", errLog)
925 }
926 }
927
928 func TestClientErrorWithRequestURI(t *testing.T) {
929 defer afterTest(t)
930 req, _ := NewRequest("GET", "http://localhost:1234/", nil)
931 req.RequestURI = "/this/field/is/illegal/and/should/error/"
932 _, err := DefaultClient.Do(req)
933 if err == nil {
934 t.Fatalf("expected an error")
935 }
936 if !strings.Contains(err.Error(), "RequestURI") {
937 t.Errorf("wanted error mentioning RequestURI; got error: %v", err)
938 }
939 }
940
941 func TestClientWithCorrectTLSServerName(t *testing.T) {
942 runSynctest(t, testClientWithCorrectTLSServerName, []testMode{https1Mode, http2Mode})
943 }
944 func testClientWithCorrectTLSServerName(t *testing.T, mode testMode) {
945 const serverName = "example.com"
946 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
947 if r.TLS.ServerName != serverName {
948 t.Errorf("expected client to set ServerName %q, got: %q", serverName, r.TLS.ServerName)
949 }
950 })).ts
951
952 c := ts.Client()
953 c.Transport.(*Transport).TLSClientConfig.ServerName = serverName
954 if _, err := c.Get("https://" + serverName); err != nil {
955 t.Fatalf("expected successful TLS connection, got error: %v", err)
956 }
957 }
958
959 func TestClientWithIncorrectTLSServerName(t *testing.T) {
960 runNoSynctest(t, testClientWithIncorrectTLSServerName, []testMode{https1Mode, http2Mode})
961 }
962 func testClientWithIncorrectTLSServerName(t *testing.T, mode testMode) {
963 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {}), optRealNet)
964 ts := cst.ts
965 errLog := new(strings.Builder)
966 ts.Config.ErrorLog = log.New(errLog, "", 0)
967
968 c := ts.Client()
969 c.Transport.(*Transport).TLSClientConfig.ServerName = "badserver"
970 _, err := c.Get(ts.URL)
971 if err == nil {
972 t.Fatalf("expected an error")
973 }
974 if !strings.Contains(err.Error(), "127.0.0.1") || !strings.Contains(err.Error(), "badserver") {
975 t.Errorf("wanted error mentioning 127.0.0.1 and badserver; got error: %v", err)
976 }
977
978 cst.close()
979 if !strings.Contains(errLog.String(), "TLS handshake error") {
980 t.Errorf("expected an error log message containing 'TLS handshake error'; got %q", errLog)
981 }
982 }
983
984
985
986
987
988
989
990
991
992
993 func TestTransportUsesTLSConfigServerName(t *testing.T) {
994 runSynctest(t, testTransportUsesTLSConfigServerName, []testMode{https1Mode, http2Mode})
995 }
996 func testTransportUsesTLSConfigServerName(t *testing.T, mode testMode) {
997 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
998 w.Write([]byte("Hello"))
999 })).ts
1000
1001 c := ts.Client()
1002 tr := c.Transport.(*Transport)
1003 tr.TLSClientConfig.ServerName = "example.com"
1004 tr.Dial = func(netw, addr string) (net.Conn, error) {
1005 return net.Dial(netw, ts.Listener.Addr().String())
1006 }
1007 res, err := c.Get("https://some-other-host.tld/")
1008 if err != nil {
1009 t.Fatal(err)
1010 }
1011 res.Body.Close()
1012 }
1013
1014 func TestResponseSetsTLSConnectionState(t *testing.T) {
1015 runSynctest(t, testResponseSetsTLSConnectionState, []testMode{https1Mode})
1016 }
1017 func testResponseSetsTLSConnectionState(t *testing.T, mode testMode) {
1018 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1019 w.Write([]byte("Hello"))
1020 })).ts
1021
1022 c := ts.Client()
1023 tr := c.Transport.(*Transport)
1024 tr.TLSClientConfig.CipherSuites = []uint16{tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256}
1025 tr.TLSClientConfig.MaxVersion = tls.VersionTLS12
1026 tr.Dial = func(netw, addr string) (net.Conn, error) {
1027 return net.Dial(netw, ts.Listener.Addr().String())
1028 }
1029 res, err := c.Get("https://example.com/")
1030 if err != nil {
1031 t.Fatal(err)
1032 }
1033 defer res.Body.Close()
1034 if res.TLS == nil {
1035 t.Fatal("Response didn't set TLS Connection State.")
1036 }
1037 if got, want := res.TLS.CipherSuite, tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256; got != want {
1038 t.Errorf("TLS Cipher Suite = %d; want %d", got, want)
1039 }
1040 }
1041
1042
1043
1044
1045 func TestHTTPSClientDetectsHTTPServer(t *testing.T) {
1046 runNoSynctest(t, testHTTPSClientDetectsHTTPServer, []testMode{http1Mode})
1047 }
1048 func testHTTPSClientDetectsHTTPServer(t *testing.T, mode testMode) {
1049 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {}), optRealNet).ts
1050 ts.Config.ErrorLog = quietLog
1051
1052 _, err := Get(strings.Replace(ts.URL, "http", "https", 1))
1053 if got := err.Error(); !strings.Contains(got, "HTTP response to HTTPS client") {
1054 t.Fatalf("error = %q; want error indicating HTTP response to HTTPS request", got)
1055 }
1056 }
1057
1058
1059 func TestClientHeadContentLength(t *testing.T) { runSynctest(t, testClientHeadContentLength) }
1060 func testClientHeadContentLength(t *testing.T, mode testMode) {
1061 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1062 if v := r.FormValue("cl"); v != "" {
1063 w.Header().Set("Content-Length", v)
1064 }
1065 }))
1066 tests := []struct {
1067 suffix string
1068 want int64
1069 }{
1070 {"/?cl=1234", 1234},
1071 {"/?cl=0", 0},
1072 {"", -1},
1073 }
1074 for _, tt := range tests {
1075 req, _ := NewRequest("HEAD", cst.ts.URL+tt.suffix, nil)
1076 res, err := cst.c.Do(req)
1077 if err != nil {
1078 t.Fatal(err)
1079 }
1080 if res.ContentLength != tt.want {
1081 t.Errorf("Content-Length = %d; want %d", res.ContentLength, tt.want)
1082 }
1083 bs, err := io.ReadAll(res.Body)
1084 if err != nil {
1085 t.Fatal(err)
1086 }
1087 if len(bs) != 0 {
1088 t.Errorf("Unexpected content: %q", bs)
1089 }
1090 }
1091 }
1092
1093 func TestEmptyPasswordAuth(t *testing.T) { runSynctest(t, testEmptyPasswordAuth) }
1094 func testEmptyPasswordAuth(t *testing.T, mode testMode) {
1095 gopher := "gopher"
1096 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1097 auth := r.Header.Get("Authorization")
1098 if strings.HasPrefix(auth, "Basic ") {
1099 encoded := auth[6:]
1100 decoded, err := base64.StdEncoding.DecodeString(encoded)
1101 if err != nil {
1102 t.Fatal(err)
1103 }
1104 expected := gopher + ":"
1105 s := string(decoded)
1106 if expected != s {
1107 t.Errorf("Invalid Authorization header. Got %q, wanted %q", s, expected)
1108 }
1109 } else {
1110 t.Errorf("Invalid auth %q", auth)
1111 }
1112 })).ts
1113 defer ts.Close()
1114 req, err := NewRequest("GET", ts.URL, nil)
1115 if err != nil {
1116 t.Fatal(err)
1117 }
1118 req.URL.User = url.User(gopher)
1119 c := ts.Client()
1120 resp, err := c.Do(req)
1121 if err != nil {
1122 t.Fatal(err)
1123 }
1124 defer resp.Body.Close()
1125 }
1126
1127 func TestBasicAuth(t *testing.T) {
1128 defer afterTest(t)
1129 tr := &recordingTransport{}
1130 client := &Client{Transport: tr}
1131
1132 url := "http://My%20User:My%20Pass@dummy.faketld/"
1133 expected := "My User:My Pass"
1134 client.Get(url)
1135
1136 if tr.req.Method != "GET" {
1137 t.Errorf("got method %q, want %q", tr.req.Method, "GET")
1138 }
1139 if tr.req.URL.String() != url {
1140 t.Errorf("got URL %q, want %q", tr.req.URL.String(), url)
1141 }
1142 if tr.req.Header == nil {
1143 t.Fatalf("expected non-nil request Header")
1144 }
1145 auth := tr.req.Header.Get("Authorization")
1146 if strings.HasPrefix(auth, "Basic ") {
1147 encoded := auth[6:]
1148 decoded, err := base64.StdEncoding.DecodeString(encoded)
1149 if err != nil {
1150 t.Fatal(err)
1151 }
1152 s := string(decoded)
1153 if expected != s {
1154 t.Errorf("Invalid Authorization header. Got %q, wanted %q", s, expected)
1155 }
1156 } else {
1157 t.Errorf("Invalid auth %q", auth)
1158 }
1159 }
1160
1161 func TestBasicAuthHeadersPreserved(t *testing.T) {
1162 defer afterTest(t)
1163 tr := &recordingTransport{}
1164 client := &Client{Transport: tr}
1165
1166
1167 url := "http://My%20User@dummy.faketld/"
1168 req, err := NewRequest("GET", url, nil)
1169 if err != nil {
1170 t.Fatal(err)
1171 }
1172 req.SetBasicAuth("My User", "My Pass")
1173 expected := "My User:My Pass"
1174 client.Do(req)
1175
1176 if tr.req.Method != "GET" {
1177 t.Errorf("got method %q, want %q", tr.req.Method, "GET")
1178 }
1179 if tr.req.URL.String() != url {
1180 t.Errorf("got URL %q, want %q", tr.req.URL.String(), url)
1181 }
1182 if tr.req.Header == nil {
1183 t.Fatalf("expected non-nil request Header")
1184 }
1185 auth := tr.req.Header.Get("Authorization")
1186 if strings.HasPrefix(auth, "Basic ") {
1187 encoded := auth[6:]
1188 decoded, err := base64.StdEncoding.DecodeString(encoded)
1189 if err != nil {
1190 t.Fatal(err)
1191 }
1192 s := string(decoded)
1193 if expected != s {
1194 t.Errorf("Invalid Authorization header. Got %q, wanted %q", s, expected)
1195 }
1196 } else {
1197 t.Errorf("Invalid auth %q", auth)
1198 }
1199
1200 }
1201
1202 func TestStripPasswordFromError(t *testing.T) {
1203 client := &Client{Transport: &recordingTransport{}}
1204 testCases := []struct {
1205 desc string
1206 in string
1207 out string
1208 }{
1209 {
1210 desc: "Strip password from error message",
1211 in: "http://user:password@dummy.faketld/",
1212 out: `Get "http://user:***@dummy.faketld/": dummy impl`,
1213 },
1214 {
1215 desc: "Don't Strip password from domain name",
1216 in: "http://user:password@password.faketld/",
1217 out: `Get "http://user:***@password.faketld/": dummy impl`,
1218 },
1219 {
1220 desc: "Don't Strip password from path",
1221 in: "http://user:password@dummy.faketld/password",
1222 out: `Get "http://user:***@dummy.faketld/password": dummy impl`,
1223 },
1224 {
1225 desc: "Strip escaped password",
1226 in: "http://user:pa%2Fssword@dummy.faketld/",
1227 out: `Get "http://user:***@dummy.faketld/": dummy impl`,
1228 },
1229 }
1230 for _, tC := range testCases {
1231 t.Run(tC.desc, func(t *testing.T) {
1232 _, err := client.Get(tC.in)
1233 if err.Error() != tC.out {
1234 t.Errorf("Unexpected output for %q: expected %q, actual %q",
1235 tC.in, tC.out, err.Error())
1236 }
1237 })
1238 }
1239 }
1240
1241 func TestClientTimeout(t *testing.T) { runSynctest(t, testClientTimeout, http3SkippedMode) }
1242 func testClientTimeout(t *testing.T, mode testMode) {
1243 sawSlow := false
1244 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1245 switch r.URL.Path {
1246 case "/":
1247 Redirect(w, r, "/slow", StatusFound)
1248 case "/slow":
1249 w.WriteHeader(200)
1250 w.Write([]byte("hello"))
1251 NewResponseController(w).Flush()
1252 sawSlow = true
1253 <-r.Context().Done()
1254 }
1255 }))
1256
1257
1258 timeout := 10 * time.Second
1259 cst.c.Timeout = timeout
1260
1261 res, err := cst.c.Get(cst.ts.URL + "/")
1262 if err != nil {
1263 t.Fatal(err)
1264 }
1265
1266 synctest.Wait()
1267 if !sawSlow {
1268 t.Fatal("handler never got /slow request, but client returned response")
1269 }
1270
1271 _, err = io.ReadAll(res.Body)
1272 res.Body.Close()
1273
1274 if err == nil {
1275 t.Fatal("expected error from ReadAll")
1276 }
1277 ne, ok := err.(net.Error)
1278 if !ok {
1279 t.Errorf("error value from ReadAll was %T; expected some net.Error", err)
1280 } else if !ne.Timeout() {
1281 t.Errorf("net.Error.Timeout = false; want true")
1282 }
1283 if !errors.Is(err, context.DeadlineExceeded) {
1284 t.Errorf("ReadAll error = %q; expected some context.DeadlineExceeded", err)
1285 }
1286 if got := ne.Error(); !strings.Contains(got, "(Client.Timeout") {
1287 t.Errorf("error string = %q; missing timeout substring", got)
1288 }
1289 }
1290
1291
1292 func TestClientTimeout_Headers(t *testing.T) {
1293
1294
1295
1296 runSynctest(t, testClientTimeout_Headers, http3SkippedMode)
1297 }
1298 func testClientTimeout_Headers(t *testing.T, mode testMode) {
1299 donec := make(chan bool, 1)
1300 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1301 <-donec
1302 }), optQuietLog)
1303
1304
1305
1306
1307
1308
1309
1310 defer func() { donec <- true }()
1311
1312 cst.c.Timeout = 5 * time.Millisecond
1313 res, err := cst.c.Get(cst.ts.URL)
1314 if err == nil {
1315 res.Body.Close()
1316 t.Fatal("got response from Get; expected error")
1317 }
1318 if _, ok := err.(*url.Error); !ok {
1319 t.Fatalf("Got error of type %T; want *url.Error", err)
1320 }
1321 ne, ok := err.(net.Error)
1322 if !ok {
1323 t.Fatalf("Got error of type %T; want some net.Error", err)
1324 }
1325 if !ne.Timeout() {
1326 t.Error("net.Error.Timeout = false; want true")
1327 }
1328 if !errors.Is(err, context.DeadlineExceeded) {
1329 t.Errorf("ReadAll error = %q; expected some context.DeadlineExceeded", err)
1330 }
1331 if got := ne.Error(); !strings.Contains(got, "Client.Timeout exceeded") {
1332 if runtime.GOOS == "windows" && runtime.GOARCH == "arm64" {
1333 testenv.SkipFlaky(t, 43120)
1334 }
1335 t.Errorf("error string = %q; missing timeout substring", got)
1336 }
1337 }
1338
1339
1340
1341 func TestClientTimeoutCancel(t *testing.T) { runSynctest(t, testClientTimeoutCancel, http3SkippedMode) }
1342 func testClientTimeoutCancel(t *testing.T, mode testMode) {
1343 testDone := make(chan struct{})
1344 ctx, cancel := context.WithCancel(context.Background())
1345
1346 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1347 w.(Flusher).Flush()
1348 <-testDone
1349 }))
1350 defer close(testDone)
1351
1352 cst.c.Timeout = 1 * time.Hour
1353 req, _ := NewRequest("GET", cst.ts.URL, nil)
1354 req.Cancel = ctx.Done()
1355 res, err := cst.c.Do(req)
1356 if err != nil {
1357 t.Fatal(err)
1358 }
1359 cancel()
1360 _, err = io.Copy(io.Discard, res.Body)
1361 if err != ExportErrRequestCanceled {
1362 t.Fatalf("error = %v; want errRequestCanceled", err)
1363 }
1364 }
1365
1366
1367 func TestClientTimeoutDoesNotExpire(t *testing.T) { runSynctest(t, testClientTimeoutDoesNotExpire) }
1368 func testClientTimeoutDoesNotExpire(t *testing.T, mode testMode) {
1369 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1370 w.Write([]byte("body"))
1371 }))
1372
1373 cst.c.Timeout = 1 * time.Hour
1374 req, _ := NewRequest("GET", cst.ts.URL, nil)
1375 res, err := cst.c.Do(req)
1376 if err != nil {
1377 t.Fatal(err)
1378 }
1379 if _, err = io.Copy(io.Discard, res.Body); err != nil {
1380 t.Fatalf("io.Copy(io.Discard, res.Body) = %v, want nil", err)
1381 }
1382 if err = res.Body.Close(); err != nil {
1383 t.Fatalf("res.Body.Close() = %v, want nil", err)
1384 }
1385 }
1386
1387 func TestClientRedirectEatsBody_h1(t *testing.T) { runSynctest(t, testClientRedirectEatsBody) }
1388 func testClientRedirectEatsBody(t *testing.T, mode testMode) {
1389 saw := make(chan string, 2)
1390 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1391 saw <- r.RemoteAddr
1392 if r.URL.Path == "/" {
1393 Redirect(w, r, "/foo", StatusFound)
1394 }
1395 }))
1396
1397 res, err := cst.c.Get(cst.ts.URL)
1398 if err != nil {
1399 t.Fatal(err)
1400 }
1401 _, err = io.ReadAll(res.Body)
1402 res.Body.Close()
1403 if err != nil {
1404 t.Fatal(err)
1405 }
1406
1407 var first string
1408 select {
1409 case first = <-saw:
1410 default:
1411 t.Fatal("server didn't see a request")
1412 }
1413
1414 var second string
1415 select {
1416 case second = <-saw:
1417 default:
1418 t.Fatal("server didn't see a second request")
1419 }
1420
1421 if first != second {
1422 t.Fatal("server saw different client ports before & after the redirect")
1423 }
1424 }
1425
1426
1427 type eofReaderFunc func()
1428
1429 func (f eofReaderFunc) Read(p []byte) (n int, err error) {
1430 f()
1431 return 0, io.EOF
1432 }
1433
1434 func TestReferer(t *testing.T) {
1435 tests := []struct {
1436 lastReq, newReq, explicitRef string
1437 want string
1438 }{
1439
1440 {lastReq: "http://gopher@test.com", newReq: "http://link.com", want: "http://test.com"},
1441 {lastReq: "https://gopher@test.com", newReq: "https://link.com", want: "https://test.com"},
1442
1443
1444 {lastReq: "http://gopher:go@test.com", newReq: "http://link.com", want: "http://test.com"},
1445 {lastReq: "https://gopher:go@test.com", newReq: "https://link.com", want: "https://test.com"},
1446
1447
1448 {lastReq: "http://test.com", newReq: "http://link.com", want: "http://test.com"},
1449 {lastReq: "https://test.com", newReq: "https://link.com", want: "https://test.com"},
1450
1451
1452 {lastReq: "https://test.com", newReq: "http://link.com", want: ""},
1453 {lastReq: "https://gopher:go@test.com", newReq: "http://link.com", want: ""},
1454
1455
1456 {lastReq: "https://test.com", newReq: "http://link.com", explicitRef: "https://foo.com", want: ""},
1457 {lastReq: "https://gopher:go@test.com", newReq: "http://link.com", explicitRef: "https://foo.com", want: ""},
1458
1459
1460 {lastReq: "https://test.com", newReq: "https://link.com", explicitRef: "https://foo.com", want: "https://foo.com"},
1461 {lastReq: "https://gopher:go@test.com", newReq: "https://link.com", explicitRef: "https://foo.com", want: "https://foo.com"},
1462 }
1463 for _, tt := range tests {
1464 l, err := url.Parse(tt.lastReq)
1465 if err != nil {
1466 t.Fatal(err)
1467 }
1468 n, err := url.Parse(tt.newReq)
1469 if err != nil {
1470 t.Fatal(err)
1471 }
1472 r := ExportRefererForURL(l, n, tt.explicitRef)
1473 if r != tt.want {
1474 t.Errorf("refererForURL(%q, %q) = %q; want %q", tt.lastReq, tt.newReq, r, tt.want)
1475 }
1476 }
1477 }
1478
1479
1480
1481 type issue15577Tripper struct{}
1482
1483 func (issue15577Tripper) RoundTrip(*Request) (*Response, error) {
1484 resp := &Response{
1485 StatusCode: 303,
1486 Header: map[string][]string{"Location": {"http://www.example.com/"}},
1487 Body: io.NopCloser(strings.NewReader("")),
1488 }
1489 return resp, nil
1490 }
1491
1492
1493 func TestClientRedirectResponseWithoutRequest(t *testing.T) {
1494 c := &Client{
1495 CheckRedirect: func(*Request, []*Request) error { return fmt.Errorf("no redirects!") },
1496 Transport: issue15577Tripper{},
1497 }
1498
1499 c.Get("http://dummy.tld")
1500 }
1501
1502
1503
1504
1505
1506 func TestClientCopyHeadersOnRedirect(t *testing.T) {
1507 runSynctest(t, testClientCopyHeadersOnRedirect, http3SkippedMode)
1508 }
1509 func testClientCopyHeadersOnRedirect(t *testing.T, mode testMode) {
1510 const (
1511 ua = "some-agent/1.2"
1512 xfoo = "foo-val"
1513 )
1514 var (
1515 targetHost = "target.example.com"
1516 redirectHost = "example.com"
1517 targetURL = mode.Scheme() + "://" + targetHost + "/"
1518 redirectURL = mode.Scheme() + "://" + redirectHost + "/"
1519 )
1520 mux := NewServeMux()
1521 mux.Handle(targetHost+"/", HandlerFunc(func(w ResponseWriter, r *Request) {
1522 want := Header{
1523 "User-Agent": []string{ua},
1524 "X-Foo": []string{xfoo},
1525 "Referer": []string{redirectURL},
1526 "Accept-Encoding": []string{"gzip"},
1527 "Cookie": []string{"foo=bar"},
1528 "Authorization": []string{"secretpassword"},
1529 }
1530 if !reflect.DeepEqual(r.Header, want) {
1531 t.Errorf("Request.Header = %#v; want %#v", r.Header, want)
1532 }
1533 if t.Failed() {
1534 w.Header().Set("Result", "got errors")
1535 } else {
1536 w.Header().Set("Result", "ok")
1537 }
1538 }))
1539 mux.Handle(redirectHost+"/", HandlerFunc(func(w ResponseWriter, r *Request) {
1540 Redirect(w, r, targetURL, StatusFound)
1541 }))
1542 ts := newClientServerTest(t, mode, mux).ts
1543
1544 c := ts.Client()
1545 c.CheckRedirect = func(r *Request, via []*Request) error {
1546 want := Header{
1547 "User-Agent": []string{ua},
1548 "X-Foo": []string{xfoo},
1549 "Referer": []string{redirectURL},
1550 "Cookie": []string{"foo=bar"},
1551 "Authorization": []string{"secretpassword"},
1552 }
1553 if !reflect.DeepEqual(r.Header, want) {
1554 t.Errorf("CheckRedirect Request.Header = %#v; want %#v", r.Header, want)
1555 }
1556 return nil
1557 }
1558
1559 req, _ := NewRequest("GET", redirectURL, nil)
1560 req.Header.Add("User-Agent", ua)
1561 req.Header.Add("X-Foo", xfoo)
1562 req.Header.Add("Cookie", "foo=bar")
1563 req.Header.Add("Authorization", "secretpassword")
1564 res, err := c.Do(req)
1565 if err != nil {
1566 t.Fatal(err)
1567 }
1568 defer res.Body.Close()
1569 if res.StatusCode != 200 {
1570 t.Fatal(res.Status)
1571 }
1572 if got := res.Header.Get("Result"); got != "ok" {
1573 t.Errorf("result = %q; want ok", got)
1574 }
1575 }
1576
1577
1578
1579 func TestClientStripHeadersOnRepeatedRedirect(t *testing.T) {
1580 runSynctest(t, testClientStripHeadersOnRepeatedRedirect, http3SkippedMode)
1581 }
1582 func testClientStripHeadersOnRepeatedRedirect(t *testing.T, mode testMode) {
1583 var proto string
1584 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1585 if r.Host+r.URL.Path != "a.example.com/" {
1586 if h := r.Header.Get("Authorization"); h != "" {
1587 t.Errorf("on request to %v%v, Authorization=%q, want no header", r.Host, r.URL.Path, h)
1588 } else if h := r.Header.Get("Proxy-Authorization"); h != "" {
1589 t.Errorf("on request to %v%v, Proxy-Authorization=%q, want no header", r.Host, r.URL.Path, h)
1590 }
1591 }
1592
1593
1594
1595 switch r.Host + r.URL.Path {
1596 case "a.example.com/":
1597 Redirect(w, r, proto+"://b.example.com/", StatusFound)
1598 case "b.example.com/":
1599 Redirect(w, r, proto+"://b.example.com/redirect", StatusFound)
1600 case "b.example.com/redirect":
1601 Redirect(w, r, proto+"://a.example.com/redirect", StatusFound)
1602 case "a.example.com/redirect":
1603 w.Header().Set("X-Done", "true")
1604 default:
1605 t.Errorf("unexpected request to %v", r.URL)
1606 }
1607 })).ts
1608 proto, _, _ = strings.Cut(ts.URL, ":")
1609
1610 c := ts.Client()
1611 c.Transport.(*Transport).Dial = func(_ string, _ string) (net.Conn, error) {
1612 return net.Dial("tcp", ts.Listener.Addr().String())
1613 }
1614
1615 req, _ := NewRequest("GET", proto+"://a.example.com/", nil)
1616 req.Header.Add("Cookie", "foo=bar")
1617 req.Header.Add("Authorization", "secretpassword")
1618 req.Header.Add("Proxy-Authorization", "secretpassword")
1619 res, err := c.Do(req)
1620 if err != nil {
1621 t.Fatal(err)
1622 }
1623 defer res.Body.Close()
1624 if res.Header.Get("X-Done") != "true" {
1625 t.Fatalf("response missing expected header: X-Done=true")
1626 }
1627 }
1628
1629 func TestClientStripHeadersOnPostToGetRedirect(t *testing.T) {
1630 runSynctest(t, testClientStripHeadersOnPostToGetRedirect)
1631 }
1632 func testClientStripHeadersOnPostToGetRedirect(t *testing.T, mode testMode) {
1633 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1634 if r.Method == "POST" {
1635 Redirect(w, r, "/redirected", StatusFound)
1636 return
1637 } else if r.Method != "GET" {
1638 t.Errorf("unexpected request method: %v", r.Method)
1639 return
1640 }
1641 for key, val := range r.Header {
1642 if strings.HasPrefix(key, "Content-") {
1643 t.Errorf("unexpected request body header after redirect: %v: %v", key, val)
1644 }
1645 }
1646 })).ts
1647
1648 c := ts.Client()
1649
1650 req, _ := NewRequest("POST", ts.URL, strings.NewReader("hello world"))
1651 req.Header.Set("Content-Encoding", "a")
1652 req.Header.Set("Content-Language", "b")
1653 req.Header.Set("Content-Length", "c")
1654 req.Header.Set("Content-Type", "d")
1655 res, err := c.Do(req)
1656 if err != nil {
1657 t.Fatal(err)
1658 }
1659 defer res.Body.Close()
1660 }
1661
1662
1663 func TestClientCopyHostOnRedirect(t *testing.T) { runSynctest(t, testClientCopyHostOnRedirect) }
1664 func testClientCopyHostOnRedirect(t *testing.T, mode testMode) {
1665
1666 virtual := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1667 t.Errorf("Virtual host received request %v", r.URL)
1668 w.WriteHeader(403)
1669 io.WriteString(w, "should not see this response")
1670 })).ts
1671 defer virtual.Close()
1672 virtualHost := strings.TrimPrefix(virtual.URL, "http://")
1673 virtualHost = strings.TrimPrefix(virtualHost, "https://")
1674 t.Logf("Virtual host is %v", virtualHost)
1675
1676
1677 const wantBody = "response body"
1678 var tsURL string
1679 var tsHost string
1680 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1681 switch r.URL.Path {
1682 case "/":
1683
1684 if r.Host != virtualHost {
1685 t.Errorf("Serving /: Request.Host = %#v; want %#v", r.Host, virtualHost)
1686 w.WriteHeader(404)
1687 return
1688 }
1689 w.Header().Set("Location", "/hop")
1690 w.WriteHeader(302)
1691 case "/hop":
1692
1693 if r.Host != virtualHost {
1694 t.Errorf("Serving /hop: Request.Host = %#v; want %#v", r.Host, virtualHost)
1695 w.WriteHeader(404)
1696 return
1697 }
1698 w.Header().Set("Location", tsURL+"/final")
1699 w.WriteHeader(302)
1700 case "/final":
1701 if r.Host != tsHost {
1702 t.Errorf("Serving /final: Request.Host = %#v; want %#v", r.Host, tsHost)
1703 w.WriteHeader(404)
1704 return
1705 }
1706 w.WriteHeader(200)
1707 io.WriteString(w, wantBody)
1708 default:
1709 t.Errorf("Serving unexpected path %q", r.URL.Path)
1710 w.WriteHeader(404)
1711 }
1712 })).ts
1713 tsURL = ts.URL
1714 tsHost = strings.TrimPrefix(ts.URL, "http://")
1715 tsHost = strings.TrimPrefix(tsHost, "https://")
1716 t.Logf("Server host is %v", tsHost)
1717
1718 c := ts.Client()
1719 req, _ := NewRequest("GET", ts.URL, nil)
1720 req.Host = virtualHost
1721 resp, err := c.Do(req)
1722 if err != nil {
1723 t.Fatal(err)
1724 }
1725 defer resp.Body.Close()
1726 if resp.StatusCode != 200 {
1727 t.Fatal(resp.Status)
1728 }
1729 if got, err := io.ReadAll(resp.Body); err != nil || string(got) != wantBody {
1730 t.Errorf("body = %q; want %q", got, wantBody)
1731 }
1732 }
1733
1734
1735 func TestClientAltersCookiesOnRedirect(t *testing.T) {
1736 runSynctest(t, testClientAltersCookiesOnRedirect)
1737 }
1738 func testClientAltersCookiesOnRedirect(t *testing.T, mode testMode) {
1739 cookieMap := func(cs []*Cookie) map[string][]string {
1740 m := make(map[string][]string)
1741 for _, c := range cs {
1742 m[c.Name] = append(m[c.Name], c.Value)
1743 }
1744 return m
1745 }
1746
1747 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1748 var want map[string][]string
1749 got := cookieMap(r.Cookies())
1750
1751 c, _ := r.Cookie("Cycle")
1752 switch c.Value {
1753 case "0":
1754 want = map[string][]string{
1755 "Cookie1": {"OldValue1a", "OldValue1b"},
1756 "Cookie2": {"OldValue2"},
1757 "Cookie3": {"OldValue3a", "OldValue3b"},
1758 "Cookie4": {"OldValue4"},
1759 "Cycle": {"0"},
1760 }
1761 SetCookie(w, &Cookie{Name: "Cycle", Value: "1", Path: "/"})
1762 SetCookie(w, &Cookie{Name: "Cookie2", Path: "/", MaxAge: -1})
1763 Redirect(w, r, "/", StatusFound)
1764 case "1":
1765 want = map[string][]string{
1766 "Cookie1": {"OldValue1a", "OldValue1b"},
1767 "Cookie3": {"OldValue3a", "OldValue3b"},
1768 "Cookie4": {"OldValue4"},
1769 "Cycle": {"1"},
1770 }
1771 SetCookie(w, &Cookie{Name: "Cycle", Value: "2", Path: "/"})
1772 SetCookie(w, &Cookie{Name: "Cookie3", Value: "NewValue3", Path: "/"})
1773 SetCookie(w, &Cookie{Name: "Cookie4", Value: "NewValue4", Path: "/"})
1774 Redirect(w, r, "/", StatusFound)
1775 case "2":
1776 want = map[string][]string{
1777 "Cookie1": {"OldValue1a", "OldValue1b"},
1778 "Cookie3": {"NewValue3"},
1779 "Cookie4": {"NewValue4"},
1780 "Cycle": {"2"},
1781 }
1782 SetCookie(w, &Cookie{Name: "Cycle", Value: "3", Path: "/"})
1783 SetCookie(w, &Cookie{Name: "Cookie5", Value: "NewValue5", Path: "/"})
1784 Redirect(w, r, "/", StatusFound)
1785 case "3":
1786 want = map[string][]string{
1787 "Cookie1": {"OldValue1a", "OldValue1b"},
1788 "Cookie3": {"NewValue3"},
1789 "Cookie4": {"NewValue4"},
1790 "Cookie5": {"NewValue5"},
1791 "Cycle": {"3"},
1792 }
1793
1794 default:
1795 t.Errorf("unexpected redirect cycle")
1796 return
1797 }
1798
1799 if !reflect.DeepEqual(got, want) {
1800 t.Errorf("redirect %s, Cookie = %v, want %v", c.Value, got, want)
1801 }
1802 })).ts
1803
1804 jar, _ := cookiejar.New(nil)
1805 c := ts.Client()
1806 c.Jar = jar
1807
1808 u, _ := url.Parse(ts.URL)
1809 req, _ := NewRequest("GET", ts.URL, nil)
1810 req.AddCookie(&Cookie{Name: "Cookie1", Value: "OldValue1a"})
1811 req.AddCookie(&Cookie{Name: "Cookie1", Value: "OldValue1b"})
1812 req.AddCookie(&Cookie{Name: "Cookie2", Value: "OldValue2"})
1813 req.AddCookie(&Cookie{Name: "Cookie3", Value: "OldValue3a"})
1814 req.AddCookie(&Cookie{Name: "Cookie3", Value: "OldValue3b"})
1815 jar.SetCookies(u, []*Cookie{{Name: "Cookie4", Value: "OldValue4", Path: "/"}})
1816 jar.SetCookies(u, []*Cookie{{Name: "Cycle", Value: "0", Path: "/"}})
1817 res, err := c.Do(req)
1818 if err != nil {
1819 t.Fatal(err)
1820 }
1821 defer res.Body.Close()
1822 if res.StatusCode != 200 {
1823 t.Fatal(res.Status)
1824 }
1825 }
1826
1827
1828 func TestShouldCopyHeaderOnRedirect(t *testing.T) {
1829 tests := []struct {
1830 initialURL string
1831 destURL string
1832 want bool
1833 }{
1834
1835 {"http://foo.com/", "http://bar.com/", false},
1836 {"http://foo.com/", "http://bar.com/", false},
1837 {"http://foo.com/", "http://bar.com/", false},
1838 {"http://foo.com/", "https://foo.com/", true},
1839 {"http://foo.com:1234/", "http://foo.com:4321/", true},
1840 {"http://foo.com/", "http://bar.com/", false},
1841 {"http://foo.com/", "http://[::1%25.foo.com]/", false},
1842
1843
1844 {"http://foo.com/", "http://foo.com/", true},
1845 {"http://foo.com/", "http://sub.foo.com/", true},
1846 {"http://foo.com/", "http://notfoo.com/", false},
1847 {"http://foo.com/", "https://foo.com/", true},
1848 {"http://foo.com:80/", "http://foo.com/", true},
1849 {"http://foo.com:80/", "http://sub.foo.com/", true},
1850 {"http://foo.com:443/", "https://foo.com/", true},
1851 {"http://foo.com:443/", "https://sub.foo.com/", true},
1852 {"http://foo.com:1234/", "http://foo.com/", true},
1853
1854 {"http://foo.com/", "http://foo.com/", true},
1855 {"http://foo.com/", "http://sub.foo.com/", true},
1856 {"http://foo.com/", "http://notfoo.com/", false},
1857 {"http://foo.com/", "https://foo.com/", true},
1858 {"http://foo.com:80/", "http://foo.com/", true},
1859 {"http://foo.com:80/", "http://sub.foo.com/", true},
1860 {"http://foo.com:443/", "https://foo.com/", true},
1861 {"http://foo.com:443/", "https://sub.foo.com/", true},
1862 {"http://foo.com:1234/", "http://foo.com/", true},
1863
1864 {"http://foobar.com/", "http://fooBAR.com/", true},
1865
1866 {"http://example.com/", "http://evil。example.com/", false},
1867 {"http://example.com/", "http://example.com/", false},
1868 {"http://süb.example.com/", "http://sÜb.example.com/", false},
1869 }
1870 for i, tt := range tests {
1871 u0, err := url.Parse(tt.initialURL)
1872 if err != nil {
1873 t.Errorf("%d. initial URL %q parse error: %v", i, tt.initialURL, err)
1874 continue
1875 }
1876 u1, err := url.Parse(tt.destURL)
1877 if err != nil {
1878 t.Errorf("%d. dest URL %q parse error: %v", i, tt.destURL, err)
1879 continue
1880 }
1881 got := Export_shouldCopyHeaderOnRedirect(u0, u1)
1882 if got != tt.want {
1883 t.Errorf("%d. shouldCopyHeaderOnRedirect(%q => %q) = %v; want %v",
1884 i, tt.initialURL, tt.destURL, got, tt.want)
1885 }
1886 }
1887 }
1888
1889 func TestClientRedirectTypes(t *testing.T) { runSynctest(t, testClientRedirectTypes) }
1890 func testClientRedirectTypes(t *testing.T, mode testMode) {
1891 tests := [...]struct {
1892 method string
1893 serverStatus int
1894 wantMethod string
1895 }{
1896 0: {method: "POST", serverStatus: 301, wantMethod: "GET"},
1897 1: {method: "POST", serverStatus: 302, wantMethod: "GET"},
1898 2: {method: "POST", serverStatus: 303, wantMethod: "GET"},
1899 3: {method: "POST", serverStatus: 307, wantMethod: "POST"},
1900 4: {method: "POST", serverStatus: 308, wantMethod: "POST"},
1901
1902 5: {method: "HEAD", serverStatus: 301, wantMethod: "HEAD"},
1903 6: {method: "HEAD", serverStatus: 302, wantMethod: "HEAD"},
1904 7: {method: "HEAD", serverStatus: 303, wantMethod: "HEAD"},
1905 8: {method: "HEAD", serverStatus: 307, wantMethod: "HEAD"},
1906 9: {method: "HEAD", serverStatus: 308, wantMethod: "HEAD"},
1907
1908 10: {method: "GET", serverStatus: 301, wantMethod: "GET"},
1909 11: {method: "GET", serverStatus: 302, wantMethod: "GET"},
1910 12: {method: "GET", serverStatus: 303, wantMethod: "GET"},
1911 13: {method: "GET", serverStatus: 307, wantMethod: "GET"},
1912 14: {method: "GET", serverStatus: 308, wantMethod: "GET"},
1913
1914 15: {method: "DELETE", serverStatus: 301, wantMethod: "GET"},
1915 16: {method: "DELETE", serverStatus: 302, wantMethod: "GET"},
1916 17: {method: "DELETE", serverStatus: 303, wantMethod: "GET"},
1917 18: {method: "DELETE", serverStatus: 307, wantMethod: "DELETE"},
1918 19: {method: "DELETE", serverStatus: 308, wantMethod: "DELETE"},
1919
1920 20: {method: "PUT", serverStatus: 301, wantMethod: "GET"},
1921 21: {method: "PUT", serverStatus: 302, wantMethod: "GET"},
1922 22: {method: "PUT", serverStatus: 303, wantMethod: "GET"},
1923 23: {method: "PUT", serverStatus: 307, wantMethod: "PUT"},
1924 24: {method: "PUT", serverStatus: 308, wantMethod: "PUT"},
1925
1926 25: {method: "MADEUPMETHOD", serverStatus: 301, wantMethod: "GET"},
1927 26: {method: "MADEUPMETHOD", serverStatus: 302, wantMethod: "GET"},
1928 27: {method: "MADEUPMETHOD", serverStatus: 303, wantMethod: "GET"},
1929 28: {method: "MADEUPMETHOD", serverStatus: 307, wantMethod: "MADEUPMETHOD"},
1930 29: {method: "MADEUPMETHOD", serverStatus: 308, wantMethod: "MADEUPMETHOD"},
1931 }
1932
1933 handlerc := make(chan HandlerFunc, 1)
1934
1935 ts := newClientServerTest(t, mode, HandlerFunc(func(rw ResponseWriter, req *Request) {
1936 h := <-handlerc
1937 h(rw, req)
1938 })).ts
1939
1940 c := ts.Client()
1941 for i, tt := range tests {
1942 handlerc <- func(w ResponseWriter, r *Request) {
1943 w.Header().Set("Location", ts.URL)
1944 w.WriteHeader(tt.serverStatus)
1945 }
1946
1947 req, err := NewRequest(tt.method, ts.URL, nil)
1948 if err != nil {
1949 t.Errorf("#%d: NewRequest: %v", i, err)
1950 continue
1951 }
1952
1953 c.CheckRedirect = func(req *Request, via []*Request) error {
1954 if got, want := req.Method, tt.wantMethod; got != want {
1955 return fmt.Errorf("#%d: got next method %q; want %q", i, got, want)
1956 }
1957 handlerc <- func(rw ResponseWriter, req *Request) {
1958
1959 }
1960 return nil
1961 }
1962
1963 res, err := c.Do(req)
1964 if err != nil {
1965 t.Errorf("#%d: Response: %v", i, err)
1966 continue
1967 }
1968
1969 res.Body.Close()
1970 }
1971 }
1972
1973
1974
1975
1976 type issue18239Body struct {
1977 readCalls *int32
1978 closeCalls *int32
1979 readErr error
1980 }
1981
1982 func (b issue18239Body) Read([]byte) (int, error) {
1983 atomic.AddInt32(b.readCalls, 1)
1984 return 0, b.readErr
1985 }
1986
1987 func (b issue18239Body) Close() error {
1988 atomic.AddInt32(b.closeCalls, 1)
1989 return nil
1990 }
1991
1992
1993
1994 func TestTransportBodyReadError(t *testing.T) { runSynctest(t, testTransportBodyReadError) }
1995 func testTransportBodyReadError(t *testing.T, mode testMode) {
1996 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
1997 if r.URL.Path == "/ping" {
1998 return
1999 }
2000 buf := make([]byte, 1)
2001 n, err := r.Body.Read(buf)
2002 w.Header().Set("X-Body-Read", fmt.Sprintf("%v, %v", n, err))
2003 })).ts
2004 c := ts.Client()
2005 tr := c.Transport.(*Transport)
2006
2007
2008
2009
2010 res, err := c.Get(ts.URL + "/ping")
2011 if err != nil {
2012 t.Fatal(err)
2013 }
2014 res.Body.Close()
2015
2016 var readCallsAtomic int32
2017 var closeCallsAtomic int32
2018 someErr := errors.New("some body read error")
2019 body := issue18239Body{&readCallsAtomic, &closeCallsAtomic, someErr}
2020
2021 req, err := NewRequest("POST", ts.URL, body)
2022 if err != nil {
2023 t.Fatal(err)
2024 }
2025 req = req.WithT(t)
2026 _, err = tr.RoundTrip(req)
2027 if err != someErr {
2028 t.Errorf("Got error: %v; want Request.Body read error: %v", err, someErr)
2029 }
2030
2031
2032
2033
2034 readCalls := atomic.LoadInt32(&readCallsAtomic)
2035 closeCalls := atomic.LoadInt32(&closeCallsAtomic)
2036 if readCalls != 1 {
2037 t.Errorf("read calls = %d; want 1", readCalls)
2038 }
2039 if closeCalls != 1 {
2040 t.Errorf("close calls = %d; want 1", closeCalls)
2041 }
2042 }
2043
2044
2045 func TestRedirectGetBody(t *testing.T) { runSynctest(t, testRedirectGetBody) }
2046
2047 func testRedirectGetBody(t *testing.T, mode testMode) {
2048 ts := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
2049 b, err := io.ReadAll(r.Body)
2050 if err != nil {
2051 t.Error(err)
2052 }
2053 if err = r.Body.Close(); err != nil {
2054 t.Error(err)
2055 }
2056 if s := string(b); s != "hello" {
2057 t.Errorf("expected hello, got %s", s)
2058 }
2059 if r.URL.Path == "/first" {
2060 Redirect(w, r, "/second", StatusTemporaryRedirect)
2061 return
2062 }
2063 w.Write([]byte("world"))
2064 })).ts
2065 c := ts.Client()
2066 c.Transport = &roundTripperGetBody{c.Transport, t}
2067 req, err := NewRequest("POST", ts.URL+"/first", strings.NewReader("hello"))
2068 if err != nil {
2069 t.Fatal(err)
2070 }
2071 res, err := c.Do(req.WithT(t))
2072 if err != nil {
2073 t.Fatal(err)
2074 }
2075 b, err := io.ReadAll(res.Body)
2076 if err != nil {
2077 t.Fatal(err)
2078 }
2079 if err = res.Body.Close(); err != nil {
2080 t.Fatal(err)
2081 }
2082 if s := string(b); s != "world" {
2083 t.Fatalf("expected world, got %s", s)
2084 }
2085 }
2086
2087 type roundTripperGetBody struct {
2088 Transport RoundTripper
2089 t *testing.T
2090 }
2091
2092 func (r *roundTripperGetBody) RoundTrip(req *Request) (*Response, error) {
2093 if req.GetBody == nil {
2094 r.t.Error("missing Request.GetBody")
2095 }
2096 return r.Transport.RoundTrip(req)
2097 }
2098
2099 type roundTripperWithoutCloseIdle struct{}
2100
2101 func (roundTripperWithoutCloseIdle) RoundTrip(*Request) (*Response, error) { panic("unused") }
2102
2103 type roundTripperWithCloseIdle func()
2104
2105 func (roundTripperWithCloseIdle) RoundTrip(*Request) (*Response, error) { panic("unused") }
2106 func (f roundTripperWithCloseIdle) CloseIdleConnections() { f() }
2107
2108 func TestClientCloseIdleConnections(t *testing.T) {
2109 c := &Client{Transport: roundTripperWithoutCloseIdle{}}
2110 c.CloseIdleConnections()
2111
2112 closed := false
2113 var tr RoundTripper = roundTripperWithCloseIdle(func() {
2114 closed = true
2115 })
2116 c = &Client{Transport: tr}
2117 c.CloseIdleConnections()
2118 if !closed {
2119 t.Error("not closed")
2120 }
2121 }
2122
2123 type testRoundTripper func(*Request) (*Response, error)
2124
2125 func (t testRoundTripper) RoundTrip(req *Request) (*Response, error) {
2126 return t(req)
2127 }
2128
2129 func TestClientPropagatesTimeoutToContext(t *testing.T) {
2130 c := &Client{
2131 Timeout: 5 * time.Second,
2132 Transport: testRoundTripper(func(req *Request) (*Response, error) {
2133 ctx := req.Context()
2134 deadline, ok := ctx.Deadline()
2135 if !ok {
2136 t.Error("no deadline")
2137 } else {
2138 t.Logf("deadline in %v", deadline.Sub(time.Now()).Round(time.Second/10))
2139 }
2140 return nil, errors.New("not actually making a request")
2141 }),
2142 }
2143 c.Get("https://example.tld/")
2144 }
2145
2146
2147
2148 func TestClientDoCanceledVsTimeout(t *testing.T) { runNoSynctest(t, testClientDoCanceledVsTimeout) }
2149 func testClientDoCanceledVsTimeout(t *testing.T, mode testMode) {
2150 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
2151 w.Write([]byte("Hello, World!"))
2152 }))
2153
2154 cases := []string{"timeout", "canceled"}
2155
2156 for _, name := range cases {
2157 synctest.Subtest(t, name, func(t *testing.T) {
2158 var ctx context.Context
2159 var cancel func()
2160 if name == "timeout" {
2161 ctx, cancel = context.WithTimeout(context.Background(), -time.Nanosecond)
2162 } else {
2163 ctx, cancel = context.WithCancel(context.Background())
2164 cancel()
2165 }
2166 defer cancel()
2167
2168 req, _ := NewRequestWithContext(ctx, "GET", cst.ts.URL, nil)
2169 _, err := cst.c.Do(req)
2170 if err == nil {
2171 t.Fatal("Unexpectedly got a nil error")
2172 }
2173
2174 ue := err.(*url.Error)
2175
2176 var wantIsTimeout bool
2177 var wantErr error = context.Canceled
2178 if name == "timeout" {
2179 wantErr = context.DeadlineExceeded
2180 wantIsTimeout = true
2181 }
2182 if g, w := ue.Timeout(), wantIsTimeout; g != w {
2183 t.Fatalf("url.Timeout() = %t, want %t", g, w)
2184 }
2185 if g, w := ue.Err, wantErr; g != w {
2186 t.Errorf("url.Error.Err = %v; want %v", g, w)
2187 }
2188 if got := errors.Is(err, context.DeadlineExceeded); got != wantIsTimeout {
2189 t.Errorf("errors.Is(err, context.DeadlineExceeded) = %v, want %v", got, wantIsTimeout)
2190 }
2191 })
2192 }
2193 }
2194
2195 type nilBodyRoundTripper struct{}
2196
2197 func (nilBodyRoundTripper) RoundTrip(req *Request) (*Response, error) {
2198 return &Response{
2199 StatusCode: StatusOK,
2200 Status: StatusText(StatusOK),
2201 Body: nil,
2202 Request: req,
2203 }, nil
2204 }
2205
2206 func TestClientPopulatesNilResponseBody(t *testing.T) {
2207 c := &Client{Transport: nilBodyRoundTripper{}}
2208
2209 resp, err := c.Get("http://localhost/anything")
2210 if err != nil {
2211 t.Fatalf("Client.Get rejected Response with nil Body: %v", err)
2212 }
2213
2214 if resp.Body == nil {
2215 t.Fatalf("Client failed to provide a non-nil Body as documented")
2216 }
2217 defer func() {
2218 if err := resp.Body.Close(); err != nil {
2219 t.Fatalf("error from Close on substitute Response.Body: %v", err)
2220 }
2221 }()
2222
2223 if b, err := io.ReadAll(resp.Body); err != nil {
2224 t.Errorf("read error from substitute Response.Body: %v", err)
2225 } else if len(b) != 0 {
2226 t.Errorf("substitute Response.Body was unexpectedly non-empty: %q", b)
2227 }
2228 }
2229
2230
2231 func TestClientCallsCloseOnlyOnce(t *testing.T) { runSynctest(t, testClientCallsCloseOnlyOnce) }
2232 func testClientCallsCloseOnlyOnce(t *testing.T, mode testMode) {
2233 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
2234 w.WriteHeader(StatusNoContent)
2235 }))
2236
2237
2238
2239 for i := 0; i < 50 && !t.Failed(); i++ {
2240 body := &issue40382Body{t: t, n: 256 << 10}
2241 req, err := NewRequest(MethodPost, cst.ts.URL, body)
2242 if err != nil {
2243 t.Fatal(err)
2244 }
2245 resp, err := cst.tr.RoundTrip(req)
2246 if err != nil {
2247 t.Fatal(err)
2248 }
2249 resp.Body.Close()
2250 }
2251 }
2252
2253
2254
2255
2256 type issue40382Body struct {
2257 t *testing.T
2258 n int
2259 closeCallsAtomic int32
2260 }
2261
2262 func (b *issue40382Body) Read(p []byte) (int, error) {
2263 switch {
2264 case b.n == 0:
2265 return 0, io.EOF
2266 case b.n < len(p):
2267 p = p[:b.n]
2268 fallthrough
2269 default:
2270 for i := range p {
2271 p[i] = 'x'
2272 }
2273 b.n -= len(p)
2274 return len(p), nil
2275 }
2276 }
2277
2278 func (b *issue40382Body) Close() error {
2279 if atomic.AddInt32(&b.closeCallsAtomic, 1) == 2 {
2280 b.t.Error("Body closed more than once")
2281 }
2282 return nil
2283 }
2284
2285 func TestProbeZeroLengthBody(t *testing.T) { runSynctest(t, testProbeZeroLengthBody) }
2286 func testProbeZeroLengthBody(t *testing.T, mode testMode) {
2287 reqc := make(chan struct{})
2288 cst := newClientServerTest(t, mode, HandlerFunc(func(w ResponseWriter, r *Request) {
2289 close(reqc)
2290 if _, err := io.Copy(w, r.Body); err != nil {
2291 t.Errorf("error copying request body: %v", err)
2292 }
2293 }))
2294
2295 bodyr, bodyw := io.Pipe()
2296 var gotBody string
2297 var wg sync.WaitGroup
2298 wg.Add(1)
2299 go func() {
2300 defer wg.Done()
2301 req, _ := NewRequest("GET", cst.ts.URL, bodyr)
2302 res, err := cst.c.Do(req)
2303 if err != nil {
2304 t.Error(err)
2305 return
2306 }
2307 defer res.Body.Close()
2308 b, err := io.ReadAll(res.Body)
2309 if err != nil {
2310 t.Error(err)
2311 }
2312 gotBody = string(b)
2313 }()
2314
2315 select {
2316 case <-reqc:
2317
2318 case <-time.After(60 * time.Second):
2319 t.Errorf("request not sent after 60s")
2320 }
2321
2322
2323 const content = "body"
2324 bodyw.Write([]byte(content))
2325 bodyw.Close()
2326 wg.Wait()
2327 if gotBody != content {
2328 t.Fatalf("server got body %q, want %q", gotBody, content)
2329 }
2330 }
2331
View as plain text