Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 31 additions & 11 deletions client/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -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()
Expand All @@ -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) {
Expand Down Expand Up @@ -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,
}
}

Expand Down Expand Up @@ -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 {
Expand All @@ -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
}
21 changes: 21 additions & 0 deletions client/host_header_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
58 changes: 58 additions & 0 deletions client/redirect_deadline_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
56 changes: 56 additions & 0 deletions client/redirect_semantics_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
23 changes: 23 additions & 0 deletions client/scheme_case_test.go
Original file line number Diff line number Diff line change
@@ -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()
}
}
32 changes: 32 additions & 0 deletions client/slow_headers_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
}
21 changes: 21 additions & 0 deletions client/timeout_cause_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
}
Loading