Files
summercms/modules/fetchguard/fetch_test.go
Jakub Zych 5e50b166ef refactor(10.2-01): nest framework packages under modules
- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
2026-09-28 02:21:02 +02:00

352 lines
11 KiB
Go

package fetchguard
import (
"errors"
"io"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"sync/atomic"
"testing"
"time"
"git.golem15.com/golem15/summercms/modules/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 TestDialControlRejectsUnsafeIPv6Transitions(t *testing.T) {
tests := []struct {
name string
ip string
}{
{name: "nat64 well-known loopback", ip: "64:ff9b::7f00:1"},
{name: "nat64 well-known rfc1918", ip: "64:ff9b::a00:1"},
{name: "nat64 well-known metadata", ip: "64:ff9b::a9fe:a9fe"},
{name: "nat64 local-use loopback", ip: "64:ff9b:1:7f00:0:100::"},
{name: "nat64 local-use rfc1918", ip: "64:ff9b:1:a00:0:100::"},
{name: "nat64 local-use metadata", ip: "64:ff9b:1:a9fe:a9:fe00::"},
{name: "6to4 loopback", ip: "2002:7f00:1::"},
{name: "6to4 rfc1918", ip: "2002:a00:1::"},
{name: "6to4 metadata", ip: "2002:a9fe:a9fe::"},
}
control := dialControl(Policy{Mode: PublicOnlyMode})
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := control("tcp6", "["+tt.ip+"]:443", nil)
if !errors.Is(err, errPrivateIP) {
t.Fatalf("dialControl(%s) error = %v, want errPrivateIP", tt.ip, err)
}
if got := mapTransportError(err).Reason; got != ReasonPrivateIP {
t.Fatalf("mapTransportError(%s) reason = %q, want %q", tt.ip, got, 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 past maxBytes+1, and under a loaded parallel
// test run the handler keeps flushing 64-byte chunks until the client's
// close propagates (observed ~80 KiB). The client must not have pulled an
// unbounded body first (the handler would hit 8 MiB), so 1 MiB is the
// meaningful ceiling.
if got := written.Load(); got > 1<<20 {
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
}
func TestDialControlRejectsZonedAndSpecialUse(t *testing.T) {
probes := []string{"[fe80::1%eth0]:443", "198.18.0.1:443", "192.0.0.1:443", "240.0.0.1:443"}
ctl := dialControl(Policy{Mode: PublicOnlyMode})
for _, a := range probes {
t.Run(a, func(t *testing.T) {
err := ctl("tcp", a, nil)
if !errors.Is(err, errPrivateIP) {
t.Fatalf("dialControl(%s) = %v, want errPrivateIP", a, err)
}
if got := mapTransportError(err).Reason; got != ReasonPrivateIP {
t.Fatalf("reason = %q, want %q", got, ReasonPrivateIP)
}
})
}
if err := ctl("tcp", "8.8.8.8:443", nil); err != nil {
t.Fatalf("public address rejected: %v", err)
}
if err := ctl("tcp", "[2606:4700:4700::1111]:443", nil); err != nil {
t.Fatalf("public v6 rejected: %v", err)
}
if err := ctl("tcp", "no-port", nil); err == nil {
t.Fatal("address without port accepted")
}
if err := ctl("tcp", "example.com:443", nil); err == nil || errors.Is(err, errPrivateIP) {
t.Fatalf("hostname dial target: err = %v, want unparseable error", err)
}
if err := dialControl(Policy{skipReservedCheck: true})("tcp", "127.0.0.1:1", nil); err != nil {
t.Fatalf("skipReservedCheck: %v", err)
}
}
func TestFetchPublicOnlyMapsSpecialUseToPrivateIP(t *testing.T) {
for _, u := range []string{"https://[fe80::1%25eth0]/", "https://198.18.0.1/", "https://192.0.0.1/", "https://240.0.0.1/"} {
t.Run(u, func(t *testing.T) {
_, err := Fetch(t.Context(), u, Policy{Mode: PublicOnlyMode, Timeout: 2 * time.Second, MaxBytes: 1024}, nil)
if got := reasonFrom(t, err); got != ReasonPrivateIP {
t.Fatalf("reason = %q, want %s (err %v)", got, ReasonPrivateIP, err)
}
})
}
}