Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
497f5bf6ec |
@@ -1,421 +0,0 @@
|
|||||||
package httpfetcher
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"slices"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
// testPublicHost is a TEST-NET-1 (RFC 5737) literal. isPrivateIP treats it as
|
|
||||||
// public, so validateURL and the redirect check accept it with no DNS lookup,
|
|
||||||
// while the recording dialer routes it to the local httptest server. The
|
|
||||||
// address is reserved for documentation and is never routed on the network.
|
|
||||||
const testPublicHost = "192.0.2.10"
|
|
||||||
|
|
||||||
// imagePayload is the body served by the fake upstream's image route.
|
|
||||||
const imagePayload = "fake-jpeg-bytes"
|
|
||||||
|
|
||||||
// errUnexpectedDial reports a dial to any host other than testPublicHost, which
|
|
||||||
// would mean SSRF protection let a forbidden target reach the transport.
|
|
||||||
var errUnexpectedDial = errors.New("unexpected dial target")
|
|
||||||
|
|
||||||
// upstreamURL builds a fetch URL on the fake public host for the given path.
|
|
||||||
func upstreamURL(path string) string {
|
|
||||||
return "http://" + testPublicHost + path
|
|
||||||
}
|
|
||||||
|
|
||||||
// recordingDialer records every address the transport asks it to dial and
|
|
||||||
// routes connections for testPublicHost to a real local server, so the SSRF
|
|
||||||
// checks run against a public-looking host while bytes go to httptest.
|
|
||||||
type recordingDialer struct {
|
|
||||||
target string
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
dialed []string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *recordingDialer) dialContext(
|
|
||||||
ctx context.Context,
|
|
||||||
network, addr string,
|
|
||||||
) (net.Conn, error) {
|
|
||||||
d.mu.Lock()
|
|
||||||
d.dialed = append(d.dialed, addr)
|
|
||||||
d.mu.Unlock()
|
|
||||||
|
|
||||||
host, _, err := net.SplitHostPort(addr)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if host != testPublicHost {
|
|
||||||
return nil, fmt.Errorf("%w: %s", errUnexpectedDial, addr)
|
|
||||||
}
|
|
||||||
|
|
||||||
var dialer net.Dialer
|
|
||||||
|
|
||||||
return dialer.DialContext(ctx, network, d.target)
|
|
||||||
}
|
|
||||||
|
|
||||||
// dialedAddrs returns a copy of the addresses the dialer was asked to reach.
|
|
||||||
func (d *recordingDialer) dialedAddrs() []string {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
|
|
||||||
return slices.Clone(d.dialed)
|
|
||||||
}
|
|
||||||
|
|
||||||
// startUpstream launches a fake upstream with the routes the fetch tests
|
|
||||||
// exercise and stops it when the test finishes.
|
|
||||||
func startUpstream(t *testing.T) *httptest.Server {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
mux := http.NewServeMux()
|
|
||||||
mux.HandleFunc("/image", func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
w.Header().Set("Content-Type", contentTypeJPEG)
|
|
||||||
_, _ = io.WriteString(w, imagePayload)
|
|
||||||
})
|
|
||||||
mux.HandleFunc("/status/500", func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
w.WriteHeader(http.StatusInternalServerError)
|
|
||||||
})
|
|
||||||
mux.HandleFunc("/html", func(w http.ResponseWriter, _ *http.Request) {
|
|
||||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
|
||||||
_, _ = io.WriteString(w, "<html></html>")
|
|
||||||
})
|
|
||||||
mux.HandleFunc("/redirect/private", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
http.Redirect(w, r, "http://169.254.169.254/latest/meta-data/", http.StatusFound)
|
|
||||||
})
|
|
||||||
mux.HandleFunc("/redirect/public", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
http.Redirect(w, r, "/image", http.StatusFound)
|
|
||||||
})
|
|
||||||
mux.HandleFunc("/redirect/chain", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
http.Redirect(w, r, "/redirect/hop", http.StatusFound)
|
|
||||||
})
|
|
||||||
mux.HandleFunc("/redirect/hop", func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
http.Redirect(w, r, "/image", http.StatusFound)
|
|
||||||
})
|
|
||||||
|
|
||||||
srv := httptest.NewServer(mux)
|
|
||||||
t.Cleanup(srv.Close)
|
|
||||||
|
|
||||||
return srv
|
|
||||||
}
|
|
||||||
|
|
||||||
// newServerFetcher builds a fetcher whose transport routes testPublicHost to
|
|
||||||
// srv, leaving the real SSRF validation and redirect checks in place.
|
|
||||||
func newServerFetcher(
|
|
||||||
t *testing.T,
|
|
||||||
srv *httptest.Server,
|
|
||||||
cfg *Config,
|
|
||||||
) (*HTTPFetcher, *recordingDialer) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
if cfg == nil {
|
|
||||||
cfg = DefaultConfig()
|
|
||||||
}
|
|
||||||
|
|
||||||
cfg.AllowHTTP = true
|
|
||||||
|
|
||||||
f := New(cfg)
|
|
||||||
|
|
||||||
transport, ok := f.client.Transport.(*http.Transport)
|
|
||||||
if !ok {
|
|
||||||
t.Fatalf("transport is %T, want *http.Transport", f.client.Transport)
|
|
||||||
}
|
|
||||||
|
|
||||||
dialer := &recordingDialer{target: srv.Listener.Addr().String()}
|
|
||||||
transport.DialContext = dialer.dialContext
|
|
||||||
|
|
||||||
return f, dialer
|
|
||||||
}
|
|
||||||
|
|
||||||
// testContext returns a context cancelled when the test ends, bounding any
|
|
||||||
// fetch that would otherwise block on a leaked semaphore slot.
|
|
||||||
func testContext(t *testing.T) context.Context {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
||||||
t.Cleanup(cancel)
|
|
||||||
|
|
||||||
return ctx
|
|
||||||
}
|
|
||||||
|
|
||||||
// fetchImage fetches path from the fake upstream and fails on error.
|
|
||||||
func fetchImage(t *testing.T, f *HTTPFetcher, path string) *FetchResult {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
res, err := f.Fetch(testContext(t), upstreamURL(path))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Fetch(%s) error = %v", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return res
|
|
||||||
}
|
|
||||||
|
|
||||||
// fetchExpectError fetches path and fails unless Fetch returns an error.
|
|
||||||
func fetchExpectError(t *testing.T, f *HTTPFetcher, path string) error {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
res, err := f.Fetch(testContext(t), upstreamURL(path))
|
|
||||||
if err == nil {
|
|
||||||
_ = res.Content.Close()
|
|
||||||
|
|
||||||
t.Fatalf("Fetch(%s) = nil error, want an error", path)
|
|
||||||
}
|
|
||||||
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// fetchBody fetches path and returns the fully read, closed response body.
|
|
||||||
func fetchBody(t *testing.T, f *HTTPFetcher, path string) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
res := fetchImage(t, f, path)
|
|
||||||
defer func() { _ = res.Content.Close() }()
|
|
||||||
|
|
||||||
data, err := io.ReadAll(res.Content)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read body: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return string(data)
|
|
||||||
}
|
|
||||||
|
|
||||||
// semLen reports how many per-host semaphore slots are currently held.
|
|
||||||
func semLen(f *HTTPFetcher, host string) int {
|
|
||||||
return len(f.getHostSemaphore(host))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchRedirectToPrivateIPBlocked(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
srv := startUpstream(t)
|
|
||||||
f, dialer := newServerFetcher(t, srv, nil)
|
|
||||||
|
|
||||||
_, err := f.Fetch(testContext(t), upstreamURL("/redirect/private"))
|
|
||||||
if !errors.Is(err, ErrSSRFBlocked) {
|
|
||||||
t.Fatalf("Fetch() error = %v, want ErrSSRFBlocked", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, addr := range dialer.dialedAddrs() {
|
|
||||||
if strings.Contains(addr, "169.254.169.254") {
|
|
||||||
t.Errorf("dialer connected to the private redirect target: %s", addr)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchRedirectToPublicSucceeds(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
srv := startUpstream(t)
|
|
||||||
f, _ := newServerFetcher(t, srv, nil)
|
|
||||||
|
|
||||||
if body := fetchBody(t, f, "/redirect/public"); body != imagePayload {
|
|
||||||
t.Errorf("body = %q, want %q", body, imagePayload)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchRedirectChainSucceeds(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
srv := startUpstream(t)
|
|
||||||
f, _ := newServerFetcher(t, srv, nil)
|
|
||||||
|
|
||||||
if body := fetchBody(t, f, "/redirect/chain"); body != imagePayload {
|
|
||||||
t.Errorf("body = %q, want %q", body, imagePayload)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchRejectsNon2xx(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
srv := startUpstream(t)
|
|
||||||
f, _ := newServerFetcher(t, srv, nil)
|
|
||||||
|
|
||||||
err := fetchExpectError(t, f, "/status/500")
|
|
||||||
if !errors.Is(err, ErrUpstreamError) {
|
|
||||||
t.Fatalf("Fetch() error = %v, want ErrUpstreamError", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchRejectsDisallowedContentType(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
srv := startUpstream(t)
|
|
||||||
f, _ := newServerFetcher(t, srv, nil)
|
|
||||||
|
|
||||||
err := fetchExpectError(t, f, "/html")
|
|
||||||
if !errors.Is(err, ErrInvalidContentType) {
|
|
||||||
t.Fatalf("Fetch() error = %v, want ErrInvalidContentType", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchMaxResponseSizeEnforced(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
srv := startUpstream(t)
|
|
||||||
|
|
||||||
cfg := DefaultConfig()
|
|
||||||
cfg.MaxResponseSize = 8
|
|
||||||
|
|
||||||
f, _ := newServerFetcher(t, srv, cfg)
|
|
||||||
|
|
||||||
res := fetchImage(t, f, "/image")
|
|
||||||
defer func() { _ = res.Content.Close() }()
|
|
||||||
|
|
||||||
data, err := io.ReadAll(res.Content)
|
|
||||||
if !errors.Is(err, ErrResponseTooLarge) {
|
|
||||||
t.Fatalf("read error = %v, want ErrResponseTooLarge", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if int64(len(data)) > cfg.MaxResponseSize {
|
|
||||||
t.Errorf("read %d bytes, exceeds limit %d", len(data), cfg.MaxResponseSize)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchSemaphoreReleasedOnError(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
srv := startUpstream(t)
|
|
||||||
|
|
||||||
cfg := DefaultConfig()
|
|
||||||
cfg.MaxConnectionsPerHost = 1
|
|
||||||
|
|
||||||
f, _ := newServerFetcher(t, srv, cfg)
|
|
||||||
|
|
||||||
err := fetchExpectError(t, f, "/status/500")
|
|
||||||
if !errors.Is(err, ErrUpstreamError) {
|
|
||||||
t.Fatalf("Fetch() error = %v, want ErrUpstreamError", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if held := semLen(f, testPublicHost); held != 0 {
|
|
||||||
t.Fatalf("semaphore slot leaked after error: %d held", held)
|
|
||||||
}
|
|
||||||
|
|
||||||
// One slot per host: this fetch proceeds only if the slot was released.
|
|
||||||
res := fetchImage(t, f, "/image")
|
|
||||||
_ = res.Content.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
// assertSlotReleasedByClose fetches an image over a one-slot host, hands the
|
|
||||||
// open result to consume, and asserts the slot is held before and freed after,
|
|
||||||
// then that a follow-up fetch can still acquire it.
|
|
||||||
func assertSlotReleasedByClose(
|
|
||||||
t *testing.T,
|
|
||||||
consume func(*testing.T, *FetchResult),
|
|
||||||
) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
srv := startUpstream(t)
|
|
||||||
|
|
||||||
cfg := DefaultConfig()
|
|
||||||
cfg.MaxConnectionsPerHost = 1
|
|
||||||
|
|
||||||
f, _ := newServerFetcher(t, srv, cfg)
|
|
||||||
|
|
||||||
res := fetchImage(t, f, "/image")
|
|
||||||
if held := semLen(f, testPublicHost); held != 1 {
|
|
||||||
t.Fatalf("slot not held while body is open: %d held", held)
|
|
||||||
}
|
|
||||||
|
|
||||||
consume(t, res)
|
|
||||||
|
|
||||||
if held := semLen(f, testPublicHost); held != 0 {
|
|
||||||
t.Fatalf("slot not released after close: %d held", held)
|
|
||||||
}
|
|
||||||
|
|
||||||
next := fetchImage(t, f, "/image")
|
|
||||||
_ = next.Content.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchSemaphoreReleasedOnBodyClose(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
assertSlotReleasedByClose(t, func(t *testing.T, res *FetchResult) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
_, err := io.ReadAll(res.Content)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read body: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = res.Content.Close()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("close body: %v", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFetchSemaphoreReleasedOnPartialReadClose(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
assertSlotReleasedByClose(t, func(t *testing.T, res *FetchResult) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
buf := make([]byte, 1)
|
|
||||||
|
|
||||||
_, err := res.Content.Read(buf)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("partial read: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = res.Content.Close()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("close body: %v", err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// The dial-time re-resolution in ssrfSafeDialer is what closes the DNS
|
|
||||||
// rebinding window: even if validateURL saw a public answer earlier, the
|
|
||||||
// dialer independently re-checks the address it is about to connect to. A full
|
|
||||||
// rebinding simulation (a resolver returning public, then private) would mean
|
|
||||||
// replacing the global net.DefaultResolver with a fake DNS server, which is
|
|
||||||
// heavyweight and unsafe to mutate under parallel -race tests. The property is
|
|
||||||
// proven directly here instead: the dialer rejects a private target outright,
|
|
||||||
// which is exactly the check that fires when a validated host later resolves
|
|
||||||
// to a private address.
|
|
||||||
func TestSSRFSafeDialerBlocksPrivateTarget(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
for _, addr := range []string{
|
|
||||||
"169.254.169.254:80", // link-local (cloud metadata)
|
|
||||||
"127.0.0.1:80", // loopback
|
|
||||||
"10.0.0.5:80", // RFC 1918 private
|
|
||||||
} {
|
|
||||||
t.Run(addr, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
_, err := ssrfSafeDialer(context.Background(), "tcp", addr)
|
|
||||||
if !errors.Is(err, ErrSSRFBlocked) {
|
|
||||||
t.Errorf("ssrfSafeDialer(%q) = %v, want ErrSSRFBlocked", addr, err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSSRFSafeDialerAllowsPublicTarget(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// A cancelled context makes the dial fail immediately without touching the
|
|
||||||
// network; the point is only that a public literal is not SSRF-blocked.
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
cancel()
|
|
||||||
|
|
||||||
_, err := ssrfSafeDialer(ctx, "tcp", testPublicHost+":80")
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("expected a dial error for an unreachable public target")
|
|
||||||
}
|
|
||||||
|
|
||||||
if errors.Is(err, ErrSSRFBlocked) {
|
|
||||||
t.Errorf("public target was SSRF-blocked: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+5
-1
@@ -17,7 +17,11 @@ run_with_cgo_deps() {
|
|||||||
main() {
|
main() {
|
||||||
cd "$ROOT"
|
cd "$ROOT"
|
||||||
echo "Running tests..."
|
echo "Running tests..."
|
||||||
run_with_cgo_deps "CGO_ENABLED=1 go test -timeout 30s -race -v ./..."
|
# Run without -v first for clean output on success; on failure rerun
|
||||||
|
# with -v for full diagnostics, then exit non-zero (REPO_POLICIES.md
|
||||||
|
# conditional-verbose-rerun pattern). The first run already proved the
|
||||||
|
# tests broken, so the build fails even if the rerun happens to pass.
|
||||||
|
run_with_cgo_deps "CGO_ENABLED=1 go test -timeout 30s -race -cover ./... || { echo '--- Rerunning with -v for details ---'; CGO_ENABLED=1 go test -timeout 30s -race -v ./...; exit 1; }"
|
||||||
}
|
}
|
||||||
|
|
||||||
main "$@"
|
main "$@"
|
||||||
|
|||||||
Reference in New Issue
Block a user