Add an integration test of the image proxy flow (closes #80) #195
@@ -0,0 +1,66 @@
|
|||||||
|
package httpfetcher
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"net"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestNewUsesCheckedDialerWithoutDialContext checks that a fetcher built
|
||||||
|
// without DialContext, as pixa builds it, refuses to connect to a local
|
||||||
|
// server.
|
||||||
|
func TestNewUsesCheckedDialerWithoutDialContext(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
srv := startUpstream(t)
|
||||||
|
transport := transportOf(t, New(DefaultConfig()))
|
||||||
|
|
||||||
|
addr := srv.Listener.Addr().String()
|
||||||
|
|
||||||
|
_, err := transport.DialContext(testContext(t), "tcp", addr)
|
||||||
|
if !errors.Is(err, ErrSSRFBlocked) {
|
||||||
|
t.Fatalf("DialContext(%s) error = %v, want ErrSSRFBlocked", addr, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDialContextReplacesOnlyTheDialer checks that a fetcher built with
|
||||||
|
// DialContext connects through it, while the URL check still refuses a
|
||||||
|
// loopback URL and the redirect check a redirect to a link-local address.
|
||||||
|
func TestDialContextReplacesOnlyTheDialer(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
|
||||||
|
srv := startUpstream(t)
|
||||||
|
dialer := &recordingDialer{target: srv.Listener.Addr().String()}
|
||||||
|
|
||||||
|
cfg := DefaultConfig()
|
||||||
|
cfg.AllowHTTP = true
|
||||||
|
cfg.DialContext = dialer.dialContext
|
||||||
|
f := New(cfg)
|
||||||
|
|
||||||
|
if body := fetchBody(t, f, "/image"); body != imagePayload {
|
||||||
|
t.Errorf("body = %q, want %q", body, imagePayload)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := f.Fetch(testContext(t), "http://127.0.0.1/image")
|
||||||
|
if !errors.Is(err, ErrSSRFBlocked) {
|
||||||
|
t.Errorf("Fetch(loopback URL) error = %v, want ErrSSRFBlocked", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = f.Fetch(testContext(t), upstreamURL("/redirect/private"))
|
||||||
|
if !errors.Is(err, ErrSSRFBlocked) {
|
||||||
|
t.Errorf("Fetch(/redirect/private) error = %v, want ErrSSRFBlocked", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The upstream server is reached through DialContext, and nothing else
|
||||||
|
// is asked of it.
|
||||||
|
dialed := dialer.dialedAddrs()
|
||||||
|
if len(dialed) == 0 {
|
||||||
|
t.Error("DialContext was never called")
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, addr := range dialed {
|
||||||
|
if addr != net.JoinHostPort(testPublicHost, "80") {
|
||||||
|
t.Errorf("DialContext was asked to connect to %s", addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -137,6 +137,11 @@ type Config struct {
|
|||||||
// BlockedNetworks are operator-supplied CIDR ranges refused by the
|
// BlockedNetworks are operator-supplied CIDR ranges refused by the
|
||||||
// dialer, in addition to the always-enforced built-in ranges.
|
// dialer, in addition to the always-enforced built-in ranges.
|
||||||
BlockedNetworks []netip.Prefix
|
BlockedNetworks []netip.Prefix
|
||||||
|
// DialContext, when set, makes the fetcher's connections in place of
|
||||||
|
// the dialer that refuses internal addresses; the URL and redirect
|
||||||
|
// checks still run. Only tests set it, to reach a local server; the
|
||||||
|
// config file and the environment cannot.
|
||||||
|
DialContext func(ctx context.Context, network, addr string) (net.Conn, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// DefaultConfig returns a Config with sensible defaults.
|
// DefaultConfig returns a Config with sensible defaults.
|
||||||
@@ -190,13 +195,19 @@ func New(config *Config) *HTTPFetcher {
|
|||||||
config = DefaultConfig()
|
config = DefaultConfig()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Create transport with SSRF-safe dialer. The dialer re-resolves and
|
// Unless config.DialContext replaces it, the transport connects with
|
||||||
// re-checks at connect time (closing the DNS-rebinding window) against
|
// the SSRF-safe dialer, which re-resolves and re-checks at connect time
|
||||||
// both the built-in ranges and the operator-supplied blocklist.
|
// (closing the DNS-rebinding window) against both the built-in ranges
|
||||||
transport := &http.Transport{
|
// and the operator-supplied blocklist.
|
||||||
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
|
dialContext := config.DialContext
|
||||||
|
if dialContext == nil {
|
||||||
|
dialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||||
return dialSSRFSafe(ctx, network, addr, config.BlockedNetworks)
|
return dialSSRFSafe(ctx, network, addr, config.BlockedNetworks)
|
||||||
},
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
transport := &http.Transport{
|
||||||
|
DialContext: dialContext,
|
||||||
TLSHandshakeTimeout: DefaultTLSTimeout,
|
TLSHandshakeTimeout: DefaultTLSTimeout,
|
||||||
MaxIdleConns: DefaultMaxIdleConns,
|
MaxIdleConns: DefaultMaxIdleConns,
|
||||||
IdleConnTimeout: DefaultIdleConnTimeout,
|
IdleConnTimeout: DefaultIdleConnTimeout,
|
||||||
|
|||||||
Reference in New Issue
Block a user