diff --git a/client/client.go b/client/client.go index ba93d5a..5cd6b5d 100644 --- a/client/client.go +++ b/client/client.go @@ -63,9 +63,9 @@ func annotateTimeout(err error, timeout time.Duration) error { return err } if timeout > 0 { - return fmt.Errorf("request timed out after %s (increase or disable with --timeout)", timeout) + return fmt.Errorf("request timed out after %s (increase or disable with --timeout): %w", timeout, err) } - return fmt.Errorf("request timed out") + return fmt.Errorf("request timed out: %w", err) } func fetchConcurrentSchemes(opts Options) (*Result, error) { @@ -132,6 +132,15 @@ func fetchSingle(opts Options, target string) (*Result, error) { } func fetchSingleWithContext(ctx context.Context, opts Options, target string) (*Result, error) { + var cancel context.CancelFunc + if opts.Timeout > 0 { + ctx, cancel = context.WithTimeout(ctx, opts.Timeout) + defer func() { + if cancel != nil { + cancel() + } + }() + } transport := tunedTransport() var rt http.RoundTripper = transport if opts.HTTP3 { @@ -178,6 +187,10 @@ func fetchSingleWithContext(ctx context.Context, opts Options, target string) (* return nil, err } for name, values := range headers { + if name == "Host" { + req.Host = headers.Get("Host") + continue + } for _, value := range values { req.Header.Add(name, value) } @@ -208,6 +221,10 @@ func fetchSingleWithContext(ctx context.Context, opts Options, target string) (* continue } + if cancel != nil { + resp.Body = &cancelOnCloseReadCloser{ReadCloser: resp.Body, cancel: cancel} + cancel = nil // The caller owns the deadline until it closes the body. + } result := &Result{Request: req, Response: resp, Redirects: redirects} if timing != nil { result.Timing = timing.result() @@ -219,7 +236,8 @@ func fetchSingleWithContext(ctx context.Context, opts Options, target string) (* } func hasExplicitScheme(rawURL string) bool { - return strings.HasPrefix(rawURL, "http://") || strings.HasPrefix(rawURL, "https://") + scheme, _, ok := strings.Cut(rawURL, "://") + return ok && (strings.EqualFold(scheme, "http") || strings.EqualFold(scheme, "https")) } func resolveHostConcurrent(ctx context.Context, host string) ([]net.IP, error) { @@ -330,13 +348,12 @@ func tunedTransport() *http.Transport { } return nil, dialErr }, - ForceAttemptHTTP2: true, - DisableKeepAlives: false, - MaxIdleConns: 100, - MaxIdleConnsPerHost: 100, - IdleConnTimeout: 90 * time.Second, - TLSHandshakeTimeout: 5 * time.Second, - ResponseHeaderTimeout: 10 * time.Second, + ForceAttemptHTTP2: true, + DisableKeepAlives: false, + MaxIdleConns: 100, + MaxIdleConnsPerHost: 100, + IdleConnTimeout: 90 * time.Second, + TLSHandshakeTimeout: 5 * time.Second, } } @@ -367,6 +384,9 @@ func resolveURL(baseURL string, location string) (string, error) { func redirectRequest(method string, statusCode int, body []byte) (string, []byte) { switch statusCode { case http.StatusSeeOther: + if method == http.MethodHead { + return http.MethodHead, nil + } return http.MethodGet, nil case http.StatusMovedPermanently, http.StatusFound: if method != http.MethodGet && method != http.MethodHead { @@ -377,5 +397,5 @@ func redirectRequest(method string, statusCode int, body []byte) (string, []byte } func isRedirect(code int) bool { - return code >= 300 && code < 400 + return code == http.StatusMovedPermanently || code == http.StatusFound || code == http.StatusSeeOther || code == http.StatusTemporaryRedirect || code == http.StatusPermanentRedirect } diff --git a/client/host_header_test.go b/client/host_header_test.go new file mode 100644 index 0000000..0b1f82e --- /dev/null +++ b/client/host_header_test.go @@ -0,0 +1,21 @@ +package client + +import ( + "net/http" + "net/http/httptest" + "testing" + "time" +) + +func TestFetchUsesCustomHostHeader(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("X-Received-Host", r.Host) })) + defer server.Close() + result, err := Fetch(Options{URL: server.URL, Headers: []string{"Host: virtual.example"}, Timeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer result.Response.Body.Close() + if got := result.Response.Header.Get("X-Received-Host"); got != "virtual.example" { + t.Fatalf("got Host %q, want virtual.example", got) + } +} diff --git a/client/redirect_deadline_test.go b/client/redirect_deadline_test.go new file mode 100644 index 0000000..6e767aa --- /dev/null +++ b/client/redirect_deadline_test.go @@ -0,0 +1,58 @@ +package client + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "testing" + "time" +) + +func TestTimeoutCoversWholeRedirectChain(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + select { + case <-r.Context().Done(): + return + case <-time.After(80 * time.Millisecond): + } + if len(r.URL.Path) < 4 { + w.Header().Set("Location", r.URL.Path+"x") + w.WriteHeader(302) + return + } + w.WriteHeader(200) + })) + defer server.Close() + result, err := Fetch(Options{URL: server.URL, Timeout: 200 * time.Millisecond}) + if result != nil { + result.Response.Body.Close() + } + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("redirect chain outlived total timeout: %v", err) + } +} + +func TestTimeoutContextRemainsAliveForResponseBody(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(200) + w.(http.Flusher).Flush() + select { + case <-r.Context().Done(): + return + case <-time.After(20 * time.Millisecond): + } + _, _ = w.Write([]byte("complete")) + })) + defer server.Close() + result, err := Fetch(Options{URL: server.URL, Timeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer result.Response.Body.Close() + data, err := io.ReadAll(result.Response.Body) + if err != nil || string(data) != "complete" { + t.Fatalf("body canceled before caller read: %q %v", data, err) + } +} diff --git a/client/redirect_semantics_test.go b/client/redirect_semantics_test.go new file mode 100644 index 0000000..d7d821b --- /dev/null +++ b/client/redirect_semantics_test.go @@ -0,0 +1,56 @@ +package client + +import ( + "io" + "net/http" + "net/http/httptest" + "testing" + "time" +) + +func TestFetchDoesNotFollowNonRedirectStatuses(t *testing.T) { + for _, code := range []int{300, 304, 305, 306, 399} { + t.Run(http.StatusText(code), func(t *testing.T) { + visited := false + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/other" { + visited = true + w.WriteHeader(200) + return + } + w.Header().Set("Location", "/other") + w.WriteHeader(code) + })) + defer server.Close() + result, err := Fetch(Options{URL: server.URL, Timeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer result.Response.Body.Close() + if visited || result.Response.StatusCode != code { + t.Errorf("followed status %d", code) + } + }) + } +} + +func TestHeadRemainsHeadAcrossSeeOther(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/" { + w.Header().Set("Location", "/other") + w.WriteHeader(303) + return + } + w.Header().Set("X-Method", r.Method) + })) + defer server.Close() + result, err := Fetch(Options{URL: server.URL, Method: "HEAD", Timeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer result.Response.Body.Close() + if result.Response.Header.Get("X-Method") != "HEAD" { + t.Fatal("303 changed HEAD to GET") + } + _, _ = io.Copy(io.Discard, result.Response.Body) +} diff --git a/client/scheme_case_test.go b/client/scheme_case_test.go new file mode 100644 index 0000000..7332d9b --- /dev/null +++ b/client/scheme_case_test.go @@ -0,0 +1,23 @@ +package client + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" +) + +func TestFetchAcceptsCaseInsensitiveScheme(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(200) })) + defer server.Close() + for _, scheme := range []string{"HTTP://", "HtTp://"} { + target := scheme + strings.TrimPrefix(server.URL, "http://") + result, err := Fetch(Options{URL: target, Timeout: time.Second}) + if err != nil { + t.Errorf("%s: %v", scheme, err) + continue + } + result.Response.Body.Close() + } +} diff --git a/client/slow_headers_test.go b/client/slow_headers_test.go new file mode 100644 index 0000000..21c4782 --- /dev/null +++ b/client/slow_headers_test.go @@ -0,0 +1,32 @@ +package client + +import ( + "net/http" + "net/http/httptest" + "testing" + "time" +) + +func TestFetchAllowsSlowHeadersWithinConfiguredTimeout(t *testing.T) { + for _, timeout := range []time.Duration{0, 30 * time.Second} { + t.Run(timeout.String(), func(t *testing.T) { + t.Parallel() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + select { + case <-time.After(11 * time.Second): + w.WriteHeader(http.StatusOK) + case <-r.Context().Done(): + } + })) + defer srv.Close() + result, err := Fetch(Options{URL: srv.URL, Timeout: timeout}) + if err != nil { + t.Fatalf("headers within configured timeout %s rejected: %v", timeout, err) + } + defer result.Response.Body.Close() + if result.Response.StatusCode != http.StatusOK { + t.Fatalf("status = %d", result.Response.StatusCode) + } + }) + } +} diff --git a/client/timeout_cause_test.go b/client/timeout_cause_test.go new file mode 100644 index 0000000..197f1c8 --- /dev/null +++ b/client/timeout_cause_test.go @@ -0,0 +1,21 @@ +package client + +import ( + "context" + "errors" + "net/url" + "testing" + "time" +) + +func TestAnnotatedTimeoutPreservesCause(t *testing.T) { + original := &url.Error{Op: "Get", URL: "http://example.test", Err: context.DeadlineExceeded} + err := annotateTimeout(original, time.Second) + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatal("deadline cause was lost") + } + var urlErr *url.Error + if !errors.As(err, &urlErr) { + t.Fatal("URL error metadata was lost") + } +}