package httpfetcher import ( "context" "errors" "net" "strconv" "sync" "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() } // TestFetchFreesHostSlotWhenContextEndsWaitingForConnection checks that a // fetch whose request context ends while it waits for a connection shared // by all hosts gives its host's slot back. With MaxConnections at 1 and one // response open, a fetch from another host takes that host's slot and waits; // its context ends long before the 10 second wait timeout. func TestFetchFreesHostSlotWhenContextEndsWaitingForConnection(t *testing.T) { t.Parallel() srv := startUpstream(t) cfg := DefaultConfig() cfg.MaxConnections = 1 f, _ := newServerFetcher(t, srv, cfg) first, err := f.Fetch(testContext(t), imageURLOnPort(81)) if err != nil { t.Fatalf("first Fetch() error = %v", err) } defer func() { _ = first.Content.Close() }() ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond) defer cancel() _, err = f.Fetch(ctx, imageURLOnPort(82)) if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("second Fetch() error = %v, want context.DeadlineExceeded", err) } if held := semLen(f, testPublicHost+":82"); held != 0 { t.Errorf("the fetch kept its host's slot after its context ended: "+ "%d held", held) } } // TestFetchRemovesIdleHostSemaphores checks that a host's semaphore is // removed once no fetch holds or waits for one of its slots: after 100 // concurrent fetches from 50 hosts have all finished, no semaphore is left. func TestFetchRemovesIdleHostSemaphores(t *testing.T) { t.Parallel() srv := startUpstream(t) f, _ := newServerFetcher(t, srv, nil) ctx := testContext(t) var wg sync.WaitGroup for i := range 100 { wg.Go(func() { res, err := f.Fetch(ctx, imageURLOnPort(1+i%50)) if err != nil { t.Errorf("Fetch() error = %v", err) return } _ = res.Content.Close() }) } wg.Wait() if n := hostSemCount(f); n != 0 { t.Errorf("%d host semaphores left after every fetch finished, want 0", n) } } // TestFetchRemovesHostSemaphoreWhenNoConnection checks that a fetch that // ends without a connection leaves no semaphore behind: when it is refused // after waiting for a connection shared by all hosts, when its context ends // while it waits for its host's slot, and when its context ends while it // waits for a connection shared by all hosts, long before the 10 second // wait timeout. func TestFetchRemovesHostSemaphoreWhenNoConnection(t *testing.T) { t.Parallel() srv := startUpstream(t) cfg := DefaultConfig() cfg.MaxConnections = 1 cfg.MaxConnectionsPerHost = 1 f, _ := newServerFetcher(t, srv, cfg) f.connectionWaitTimeout = 100 * time.Millisecond open, err := f.Fetch(testContext(t), imageURLOnPort(81)) if err != nil { t.Fatalf("first Fetch() error = %v", err) } _, err = f.Fetch(testContext(t), imageURLOnPort(82)) if !errors.Is(err, ErrTooManyConnections) { t.Fatalf("Fetch() from another host: error = %v, "+ "want ErrTooManyConnections", err) } ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond) defer cancel() _, err = f.Fetch(ctx, imageURLOnPort(81)) if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("Fetch() from the busy host: error = %v, "+ "want context.DeadlineExceeded", err) } // Back to the 10 second wait, so the next fetch's context ends first. f.connectionWaitTimeout = ConnectionWaitTimeout ctx, cancel = context.WithTimeout(t.Context(), 100*time.Millisecond) defer cancel() _, err = f.Fetch(ctx, imageURLOnPort(83)) if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("Fetch() from a host with nothing open: error = %v, "+ "want context.DeadlineExceeded", err) } err = open.Content.Close() if err != nil { t.Fatalf("close first body: %v", err) } if n := hostSemCount(f); n != 0 { t.Errorf("%d host semaphores left after every fetch finished, want 0", n) } } // 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() }) } }