package httpfetcher import ( "errors" "net" "strconv" "testing" "time" ) // imageURLOnPort is the fake upstream's image route on testPublicHost at // port. Each port is a different host to the per-host limit, while the test // dialer sends every port to the one test server. func imageURLOnPort(port int) string { return "http://" + net.JoinHostPort(testPublicHost, strconv.Itoa(port)) + "/image" } func TestDefaultConfigMaxConnections(t *testing.T) { t.Parallel() if got := DefaultConfig().MaxConnections; got != DefaultMaxConnections { t.Errorf("MaxConnections = %d, want %d", got, DefaultMaxConnections) } } // TestFetchLimitsConnectionsToAllHostsTogether checks that MaxConnections // counts the fetches to every host together, apart from the per-host // limit: with MaxConnections at 2 and two responses open from two hosts, a // fetch from a third host, which has nothing open, waits the whole wait // timeout and fails with ErrTooManyConnections. Closing one response lets // it through. func TestFetchLimitsConnectionsToAllHostsTogether(t *testing.T) { t.Parallel() srv := startUpstream(t) cfg := DefaultConfig() cfg.MaxConnections = 2 f, _ := newServerFetcher(t, srv, cfg) f.connectionWaitTimeout = 100 * time.Millisecond first, err := f.Fetch(testContext(t), imageURLOnPort(81)) if err != nil { t.Fatalf("first Fetch() error = %v", err) } second, err := f.Fetch(testContext(t), imageURLOnPort(82)) if err != nil { t.Fatalf("second Fetch() error = %v", err) } defer func() { _ = second.Content.Close() }() start := time.Now() _, err = f.Fetch(testContext(t), imageURLOnPort(83)) if !errors.Is(err, ErrTooManyConnections) { t.Fatalf("third Fetch() error = %v, want ErrTooManyConnections", err) } if waited := time.Since(start); waited < f.connectionWaitTimeout { t.Errorf("third Fetch() failed after %v, before waiting %v", waited, f.connectionWaitTimeout) } if held := semLen(f, testPublicHost+":83"); held != 0 { t.Errorf("the refused fetch kept its host's slot: %d held", held) } err = first.Content.Close() if err != nil { t.Fatalf("close first body: %v", err) } third, err := f.Fetch(testContext(t), imageURLOnPort(83)) if err != nil { t.Fatalf("Fetch() after a response was closed: error = %v", err) } _ = third.Content.Close() } // TestFetchReleasesConnectionOnError checks that a fetch that fails after // taking its connection gives it back: with MaxConnections at 1, the slot // must be free after the failure and the next fetch must succeed. func TestFetchReleasesConnectionOnError(t *testing.T) { t.Parallel() cases := []struct { name string url string want error }{ {"upstream answers 500", upstreamURL("/status/500"), ErrUpstreamError}, {"upstream sends HTML", upstreamURL("/html"), ErrInvalidContentType}, // 198.51.100.7 (TEST-NET-2) passes the SSRF checks, and the test // dialer refuses every host but testPublicHost. {"connecting fails", "http://198.51.100.7/image", errUnexpectedDial}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { t.Parallel() srv := startUpstream(t) cfg := DefaultConfig() cfg.MaxConnections = 1 f, _ := newServerFetcher(t, srv, cfg) f.connectionWaitTimeout = 100 * time.Millisecond _, err := f.Fetch(testContext(t), tc.url) if !errors.Is(err, tc.want) { t.Fatalf("Fetch() error = %v, want %v", err, tc.want) } if held := len(f.allHostsSemaphore); held != 0 { t.Fatalf("connection still held after the error: %d held", held) } res := fetchImage(t, f, "/image") _ = res.Content.Close() }) } }