test(06-04): add failing tests for SSRF-guarded Fetch
- Cover host allow-list, dotted-suffix bypass, dial-time private IP - Cover no-follow redirects, streaming byte cap, and D-14 config defaults
This commit is contained in:
276
fetchguard/fetch_test.go
Normal file
276
fetchguard/fetch_test.go
Normal file
@@ -0,0 +1,276 @@
|
||||
package fetchguard
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.golem15.com/golem15/summercms/compass"
|
||||
)
|
||||
|
||||
func TestFetchMalformedURL(t *testing.T) {
|
||||
_, err := Fetch(t.Context(), "not a url", Policy{Mode: PublicOnlyMode}, nil)
|
||||
if reasonFrom(t, err) != ReasonInvalidURL {
|
||||
t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonInvalidURL)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchHTTPSchemeRejectedWithoutIO(t *testing.T) {
|
||||
var hits atomic.Int64
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
|
||||
hits.Add(1)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
_, err := Fetch(t.Context(), srv.URL, Policy{Mode: PublicOnlyMode}, nil)
|
||||
if reasonFrom(t, err) != ReasonScheme {
|
||||
t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonScheme)
|
||||
}
|
||||
if hits.Load() != 0 {
|
||||
t.Fatal("http URL must not cause network I/O")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchAllowHostsRejectsUnknownHostBeforeDial(t *testing.T) {
|
||||
var hits atomic.Int64
|
||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
|
||||
hits.Add(1)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
_, err := Fetch(t.Context(), srv.URL, Policy{
|
||||
Mode: AllowHostsMode,
|
||||
AllowHosts: []string{"discogs.com"},
|
||||
}, nil)
|
||||
if reasonFrom(t, err) != ReasonInvalidURL {
|
||||
t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonInvalidURL)
|
||||
}
|
||||
if hits.Load() != 0 {
|
||||
t.Fatal("host allow-list miss must not dial")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchAllowHostsRejectsDottedSuffixBypass(t *testing.T) {
|
||||
_, err := Fetch(t.Context(), "https://evil-discogs.com/cover.jpg", Policy{
|
||||
Mode: AllowHostsMode,
|
||||
AllowHosts: []string{"discogs.com"},
|
||||
}, nil)
|
||||
if reasonFrom(t, err) != ReasonInvalidURL {
|
||||
t.Fatalf("evil-discogs.com must not match discogs.com, reason = %q", reasonFrom(t, err))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchPrivateIPBlockedInBothModes(t *testing.T) {
|
||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
|
||||
t.Error("handler must not run for a private dial")
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
t.Run("AllowHostsMode", func(t *testing.T) {
|
||||
_, err := Fetch(t.Context(), srv.URL, Policy{
|
||||
Mode: AllowHostsMode,
|
||||
AllowHosts: []string{"127.0.0.1"},
|
||||
Timeout: 2 * time.Second,
|
||||
MaxBytes: 1024,
|
||||
}, nil)
|
||||
if reasonFrom(t, err) != ReasonPrivateIP {
|
||||
t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonPrivateIP)
|
||||
}
|
||||
})
|
||||
t.Run("PublicOnlyMode", func(t *testing.T) {
|
||||
_, err := Fetch(t.Context(), srv.URL, Policy{
|
||||
Mode: PublicOnlyMode,
|
||||
Timeout: 2 * time.Second,
|
||||
MaxBytes: 1024,
|
||||
}, nil)
|
||||
if reasonFrom(t, err) != ReasonPrivateIP {
|
||||
t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonPrivateIP)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestFetchDoesNotFollowRedirect(t *testing.T) {
|
||||
var followed atomic.Bool
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/target", func(w http.ResponseWriter, r *http.Request) {
|
||||
followed.Store(true)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Location", "/target")
|
||||
w.WriteHeader(http.StatusFound)
|
||||
})
|
||||
srv := httptest.NewTLSServer(mux)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
res, err := Fetch(t.Context(), srv.URL, withTestLoopback(srv, Policy{
|
||||
Mode: PublicOnlyMode,
|
||||
MaxBytes: 1024,
|
||||
Timeout: 5 * time.Second,
|
||||
}), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Fetch: %v", err)
|
||||
}
|
||||
if res.StatusCode != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d (3xx returned, not followed)", res.StatusCode, http.StatusFound)
|
||||
}
|
||||
if followed.Load() {
|
||||
t.Fatal("redirect must not be followed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchTooLargeIsStreaming(t *testing.T) {
|
||||
const maxBytes int64 = 256
|
||||
var written atomic.Int64
|
||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/octet-stream")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
buf := make([]byte, 64)
|
||||
for written.Load() < 8<<20 {
|
||||
n, err := w.Write(buf)
|
||||
written.Add(int64(n))
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
_, err := Fetch(t.Context(), srv.URL, withTestLoopback(srv, Policy{
|
||||
Mode: PublicOnlyMode,
|
||||
MaxBytes: maxBytes,
|
||||
Timeout: 5 * time.Second,
|
||||
}), nil)
|
||||
if reasonFrom(t, err) != ReasonTooLarge {
|
||||
t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonTooLarge)
|
||||
}
|
||||
// TCP/HTTP buffering can write a little past maxBytes+1; the client must
|
||||
// not have pulled an unbounded body first (the handler would hit 8MiB).
|
||||
if got := written.Load(); got > 64<<10 {
|
||||
t.Fatalf("server wrote %d bytes, client appears to have buffered unbounded body", got)
|
||||
}
|
||||
if written.Load() < maxBytes+1 {
|
||||
t.Fatalf("server wrote %d bytes, want at least maxBytes+1 so the cap was hit", written.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchZeroPolicyUsesDefaults(t *testing.T) {
|
||||
max, timeout := Defaults()
|
||||
if max != 10*1024*1024 {
|
||||
t.Fatalf("Defaults maxBytes = %d, want 10MiB", max)
|
||||
}
|
||||
if timeout != 10*time.Second {
|
||||
t.Fatalf("Defaults timeout = %s, want 10s", timeout)
|
||||
}
|
||||
|
||||
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
io.WriteString(w, "ok")
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
res, err := Fetch(t.Context(), srv.URL, withTestLoopback(srv, Policy{Mode: PublicOnlyMode}), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Fetch with zero MaxBytes/Timeout and nil cfg: %v", err)
|
||||
}
|
||||
if string(res.Body) != "ok" {
|
||||
t.Fatalf("body = %q, want ok", res.Body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultsFromConfigExplicitZeroIsError(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "http.yaml"), []byte("fetch:\n max_bytes: 0\n timeout_seconds: 10\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg, err := compass.Open(compass.Options{Dir: dir, Environ: []string{}})
|
||||
if err != nil {
|
||||
t.Fatalf("Open: %v", err)
|
||||
}
|
||||
if _, ok := cfg.Lookup("http.fetch.max_bytes"); !ok {
|
||||
t.Fatal("expected http.fetch.max_bytes to be present in loaded YAML")
|
||||
}
|
||||
_, _, err = DefaultsFromConfig(cfg)
|
||||
if err == nil {
|
||||
t.Fatal("explicit http.fetch.max_bytes: 0 must be an error, not a silent fallback")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultsFromConfigExplicitNegativeIsError(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "http.yaml"), []byte("fetch:\n max_bytes: 10485760\n timeout_seconds: -1\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg, err := compass.Open(compass.Options{Dir: dir, Environ: []string{}})
|
||||
if err != nil {
|
||||
t.Fatalf("Open: %v", err)
|
||||
}
|
||||
_, _, err = DefaultsFromConfig(cfg)
|
||||
if err == nil {
|
||||
t.Fatal("explicit http.fetch.timeout_seconds: -1 must be an error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultsFromConfigAbsentFallsBack(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: t\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cfg, err := compass.Open(compass.Options{Dir: dir, Environ: []string{}})
|
||||
if err != nil {
|
||||
t.Fatalf("Open: %v", err)
|
||||
}
|
||||
max, timeout, err := DefaultsFromConfig(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("DefaultsFromConfig: %v", err)
|
||||
}
|
||||
wantMax, wantTO := Defaults()
|
||||
if max != wantMax || timeout != wantTO {
|
||||
t.Fatalf("got %d %s, want Defaults() %d %s", max, timeout, wantMax, wantTO)
|
||||
}
|
||||
}
|
||||
|
||||
func TestErrorUnwrap(t *testing.T) {
|
||||
inner := errors.New("dial tcp")
|
||||
err := &Error{Reason: ReasonNetworkError, Err: inner}
|
||||
if err.Error() == "" {
|
||||
t.Fatal("Error() must be non-empty")
|
||||
}
|
||||
if !errors.Is(err, inner) {
|
||||
t.Fatal("Unwrap must expose the inner error")
|
||||
}
|
||||
}
|
||||
|
||||
func reasonFrom(t *testing.T, err error) Reason {
|
||||
t.Helper()
|
||||
if err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
var fe *Error
|
||||
if !errors.As(err, &fe) {
|
||||
t.Fatalf("err = %v (%T), want *Error", err, err)
|
||||
}
|
||||
return fe.Reason
|
||||
}
|
||||
|
||||
// withTestLoopback trusts the httptest TLS cert and skips the reserved-IP
|
||||
// dial check so redirect/byte-cap tests can exercise a real listener on
|
||||
// 127.0.0.1. Production Policy values leave both fields unset.
|
||||
func withTestLoopback(srv *httptest.Server, p Policy) Policy {
|
||||
tr, ok := srv.Client().Transport.(*http.Transport)
|
||||
if !ok {
|
||||
panic("httptest client transport is not *http.Transport")
|
||||
}
|
||||
p.tlsConfig = tr.TLSClientConfig
|
||||
p.skipReservedCheck = true
|
||||
return p
|
||||
}
|
||||
Reference in New Issue
Block a user