package latest import ( "bytes" "context" "crypto/tls" "io" "net" "net/http" "net/http/httptest" "net/netip" "testing" ) func TestCoverFetcherFetchesPublicHTTPSImage(t *testing.T) { server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.TLS == nil { t.Fatal("cover request was not made over TLS") } w.Header().Set("Content-Type", "image/jpeg") io.WriteString(w, "cover-bytes") })) defer server.Close() transport := server.Client().Transport.(*http.Transport).Clone() transport.TLSClientConfig = &tls.Config{InsecureSkipVerify: true} // test server certificate transport.DialContext = func(ctx context.Context, network, _ string) (net.Conn, error) { return (&net.Dialer{}).DialContext(ctx, network, server.Listener.Addr().String()) } client := &http.Client{Transport: transport} fetcher := newCoverFetcher(client, func(context.Context, string) ([]netip.Addr, error) { return []netip.Addr{netip.MustParseAddr("198.51.100.10")}, nil }) body, contentType, err := fetcher.Fetch(context.Background(), "https://cdn.example/cover.jpg") if err != nil { t.Fatalf("Fetch: %v", err) } if string(body) != "cover-bytes" || contentType != "image/jpeg" { t.Fatalf("Fetch = (%q, %q), want (cover-bytes, image/jpeg)", body, contentType) } } func TestNewCoverFetcherRechecksResolverBeforeConnection(t *testing.T) { var requests int server := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { requests++ })) defer server.Close() _, port, err := net.SplitHostPort(server.Listener.Addr().String()) if err != nil { t.Fatalf("server address: %v", err) } resolves := 0 fetcher := NewCoverFetcherWithResolver(func(context.Context, string) ([]netip.Addr, error) { resolves++ if resolves == 1 { return []netip.Addr{netip.MustParseAddr("198.51.100.10")}, nil } return []netip.Addr{netip.MustParseAddr("127.0.0.1")}, nil }) _, _, err = fetcher.Fetch(context.Background(), "https://cdn.example:"+port+"/cover.jpg") if err == nil { t.Fatal("Fetch accepted a destination that became private") } if resolves != 2 { t.Fatalf("resolver calls = %d, want preflight and dial checks", resolves) } if requests != 0 { t.Fatalf("requests = %d, want 0", requests) } } type roundTripFunc func(*http.Request) (*http.Response, error) func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } func coverResponse(status int, contentType, location string, body []byte) *http.Response { header := make(http.Header) if contentType != "" { header.Set("Content-Type", contentType) } if location != "" { header.Set("Location", location) } return &http.Response{ StatusCode: status, Status: http.StatusText(status), Header: header, Body: io.NopCloser(bytes.NewReader(body)), ContentLength: int64(len(body)), } } func TestCoverFetcherRefusesUnsafeDestinationsBeforeRequest(t *testing.T) { var calls int client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { calls++ return coverResponse(http.StatusOK, "image/jpeg", "", []byte("must not reach network")), nil })} resolve := func(_ context.Context, host string) ([]netip.Addr, error) { switch host { case "loopback.example": return []netip.Addr{netip.MustParseAddr("127.0.0.1")}, nil case "private.example": return []netip.Addr{netip.MustParseAddr("10.0.0.1")}, nil case "linklocal.example": return []netip.Addr{netip.MustParseAddr("169.254.1.1")}, nil case "unique-local.example": return []netip.Addr{netip.MustParseAddr("fc00::1")}, nil case "cgnat.example": return []netip.Addr{netip.MustParseAddr("100.64.0.1")}, nil default: return []netip.Addr{netip.MustParseAddr("198.51.100.10")}, nil } } fetcher := newCoverFetcher(client, resolve) tests := []string{ "http://public.example/cover.jpg", "https://127.0.0.1/cover.jpg", "https://10.0.0.1/cover.jpg", "https://169.254.1.1/cover.jpg", "https://[fc00::1]/cover.jpg", "https://100.64.0.1/cover.jpg", "https://loopback.example/cover.jpg", "https://private.example/cover.jpg", "https://linklocal.example/cover.jpg", "https://unique-local.example/cover.jpg", "https://cgnat.example/cover.jpg", } for _, sourceURL := range tests { t.Run(sourceURL, func(t *testing.T) { calls = 0 if _, _, err := fetcher.Fetch(context.Background(), sourceURL); err == nil { t.Fatal("Fetch accepted refused destination") } if calls != 0 { t.Fatalf("network calls = %d, want 0", calls) } }) } } func TestCoverFetcherStopsRedirectIntoPrivateAddress(t *testing.T) { var calls int client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { calls++ if req.URL.Hostname() != "cdn.example" { t.Fatalf("redirect reached %s", req.URL) } return coverResponse(http.StatusFound, "", "https://internal.example/cover.jpg", nil), nil })} fetcher := newCoverFetcher(client, func(_ context.Context, host string) ([]netip.Addr, error) { if host == "internal.example" { return []netip.Addr{netip.MustParseAddr("192.168.1.1")}, nil } return []netip.Addr{netip.MustParseAddr("198.51.100.10")}, nil }) if _, _, err := fetcher.Fetch(context.Background(), "https://cdn.example/cover.jpg"); err == nil { t.Fatal("Fetch followed redirect into private address") } if calls != 1 { t.Fatalf("network calls = %d, want only public first hop", calls) } } func TestCoverFetcherRejectsOversizedBody(t *testing.T) { var calls int client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { calls++ response := coverResponse(http.StatusOK, "image/webp", "", bytes.Repeat([]byte("x"), maxBodyBytes+1)) response.ContentLength = -1 return response, nil })} fetcher := newCoverFetcher(client, func(context.Context, string) ([]netip.Addr, error) { return []netip.Addr{netip.MustParseAddr("198.51.100.10")}, nil }) if _, _, err := fetcher.Fetch(context.Background(), "https://cdn.example/large.webp"); err == nil { t.Fatal("Fetch accepted oversized body") } if calls != 1 { t.Fatalf("network calls = %d, want 1", calls) } } func TestCoverFetcherRejectsNonImage(t *testing.T) { var calls int client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { calls++ return coverResponse(http.StatusOK, "text/html", "", []byte("challenge")), nil })} fetcher := newCoverFetcher(client, func(context.Context, string) ([]netip.Addr, error) { return []netip.Addr{netip.MustParseAddr("198.51.100.10")}, nil }) if _, _, err := fetcher.Fetch(context.Background(), "https://cdn.example/challenge"); err == nil { t.Fatal("Fetch accepted non-image response") } if calls != 1 { t.Fatalf("network calls = %d, want 1", calls) } }