package fetchguard_test import ( "errors" "io" "net/http" "net/http/httptest" "reflect" "strings" "sync/atomic" "testing" "time" "git.golem15.com/golem15/summercms/modules/fetchguard" "git.golem15.com/golem15/summercms/modules/tide" ) // sidecarPath is shared with the tide package's own upstream tests. const sidecarPath = "../tide/testdata/upstream/post_json.upstream.yaml" func exampleStore(t *testing.T) *tide.Store { t.Helper() store, err := tide.OpenStore("") if err != nil { t.Fatal(err) } store.Set("secret:example-token", "example-token-value") return store } func reasonOf(t *testing.T, err error) fetchguard.Reason { t.Helper() if err == nil { t.Fatal("expected error") } var fe *fetchguard.Error if !errors.As(err, &fe) { t.Fatalf("err = %v (%T), want *fetchguard.Error", err, err) } return fe.Reason } func TestClientPostJSONThroughUpstreamFake(t *testing.T) { sidecar, err := tide.LoadUpstream(sidecarPath) if err != nil { t.Fatal(err) } fake := tide.NewUpstreamFake(sidecar, exampleStore(t)) client, err := fetchguard.NewClient(fetchguard.Policy{ Mode: fetchguard.AllowHostsMode, AllowHosts: []string{"api.example.test"}, Timeout: 5 * time.Second, }, nil) if err != nil { t.Fatal(err) } ctx := fetchguard.WithTransport(t.Context(), fake) header := http.Header{} header.Set("Authorization", fetchguard.Bearer("example-token-value")) header.Set("Accept", "application/json") header.Set("User-Agent", "example-client/1.0") res, err := client.PostJSON(ctx, "https://api.example.test/v1/things?lang=en&mode=fast", header, map[string]any{ "count": 2, "name": "widget", "tags": []string{"a", "b"}, }) if err != nil { t.Fatalf("PostJSON: %v", err) } if res.StatusCode != http.StatusCreated { t.Fatalf("status = %d, want 201", res.StatusCode) } if got := res.Header.Get("X-Request-Id"); got != "req-123" { t.Fatalf("X-Request-Id = %q", got) } if res.ContentType != "application/json" { t.Fatalf("content type = %q", res.ContentType) } if string(res.Body) != `{"id":7,"name":"widget"}` { t.Fatalf("body = %s", res.Body) } if err := fake.Verify(); err != nil { t.Fatalf("Verify: %v", err) } } func TestTransportSeamIsCodeOnly(t *testing.T) { rtType := reflect.TypeFor[http.RoundTripper]() for _, typ := range []reflect.Type{reflect.TypeFor[fetchguard.Policy](), reflect.TypeFor[fetchguard.Client]()} { for f := range typ.Fields() { if !f.IsExported() { continue } ft := f.Type if ft.Implements(rtType) || reflect.PointerTo(ft).Implements(rtType) || ft == rtType { t.Errorf("%s.%s (%s) exposes an http.RoundTripper", typ.Name(), f.Name, ft) } } } // The guard runs before the override is consulted: a host outside the // allow list never reaches the fake. var calls int stub := roundTripFunc(func(*http.Request) (*http.Response, error) { calls++ return &http.Response{StatusCode: 200, Body: http.NoBody, Header: http.Header{}}, nil }) client, err := fetchguard.NewClient(fetchguard.Policy{ Mode: fetchguard.AllowHostsMode, AllowHosts: []string{"api.example.test"}, Timeout: 2 * time.Second, }, nil) if err != nil { t.Fatal(err) } _, err = client.Get(fetchguard.WithTransport(t.Context(), stub), "https://other.example.test/x", nil) if reasonOf(t, err) != fetchguard.ReasonInvalidURL || calls != 0 { t.Fatalf("outside host: reason %v, calls %d", err, calls) } // Without the context override the request goes to the real transport, // whose dial guard refuses loopback. public, err := fetchguard.NewClient(fetchguard.Policy{Mode: fetchguard.PublicOnlyMode, Timeout: 2 * time.Second}, nil) if err != nil { t.Fatal(err) } _, err = public.Get(t.Context(), "https://127.0.0.1:1/", nil) if reasonOf(t, err) != fetchguard.ReasonPrivateIP { t.Fatalf("real transport: %v, want private_ip", err) } // A header naming a transport changes nothing. h := http.Header{"X-Transport": {"stub"}} _, err = public.Get(t.Context(), "https://127.0.0.1:1/", h) if reasonOf(t, err) != fetchguard.ReasonPrivateIP { t.Fatalf("header override: %v, want private_ip", err) } if !strings.Contains(err.Error(), "private_ip") { t.Fatalf("error text %q", err) } } type roundTripFunc func(*http.Request) (*http.Response, error) func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } func TestClientModes(t *testing.T) { var hits atomic.Int64 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { hits.Add(1) _, _ = io.WriteString(w, "ok") })) t.Cleanup(srv.Close) t.Run("AllowHostsMode refuses a host outside the list", func(t *testing.T) { c := mustClient(t, fetchguard.Policy{Mode: fetchguard.AllowHostsMode, AllowHosts: []string{"api.example.test"}}) _, err := c.Get(t.Context(), "https://other.example.test/", nil) if got := reasonOf(t, err); got != fetchguard.ReasonInvalidURL { t.Fatalf("reason = %s, want invalid_url", got) } }) t.Run("PublicOnlyMode refuses loopback at dial", func(t *testing.T) { c := mustClient(t, fetchguard.Policy{Mode: fetchguard.PublicOnlyMode}) _, err := c.Get(t.Context(), strings.Replace(srv.URL, "http://", "https://", 1), nil) if got := reasonOf(t, err); got != fetchguard.ReasonPrivateIP { t.Fatalf("reason = %s, want private_ip", got) } }) t.Run("TrustedMode reaches an http loopback endpoint", func(t *testing.T) { c := mustClient(t, fetchguard.Policy{Mode: fetchguard.TrustedMode}) res, err := c.Get(t.Context(), srv.URL+"/v1/models", nil) if err != nil { t.Fatalf("Get: %v", err) } if res.StatusCode != http.StatusOK || string(res.Body) != "ok" { t.Fatalf("result = %d %q", res.StatusCode, res.Body) } }) if hits.Load() != 1 { t.Fatalf("server hits = %d, want 1 (only the trusted call)", hits.Load()) } t.Run("TrustedMode still refuses other schemes", func(t *testing.T) { c := mustClient(t, fetchguard.Policy{Mode: fetchguard.TrustedMode}) _, err := c.Get(t.Context(), "ftp://127.0.0.1/x", nil) if got := reasonOf(t, err); got != fetchguard.ReasonScheme { t.Fatalf("reason = %s, want scheme", got) } }) t.Run("Fetch ignores TrustedMode", func(t *testing.T) { _, err := fetchguard.Fetch(t.Context(), strings.Replace(srv.URL, "http://", "https://", 1), fetchguard.Policy{Mode: fetchguard.TrustedMode, Timeout: 2 * time.Second, MaxBytes: 1024}, nil) if got := reasonOf(t, err); got != fetchguard.ReasonPrivateIP { t.Fatalf("reason = %s, want private_ip", got) } }) } func TestClientSchemeGuard(t *testing.T) { var hits atomic.Int64 srv := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { hits.Add(1) })) t.Cleanup(srv.Close) for _, p := range []fetchguard.Policy{ {Mode: fetchguard.AllowHostsMode, AllowHosts: []string{"127.0.0.1"}}, {Mode: fetchguard.PublicOnlyMode}, } { c := mustClient(t, p) _, err := c.PostJSON(t.Context(), srv.URL, nil, map[string]string{"a": "b"}) if got := reasonOf(t, err); got != fetchguard.ReasonScheme { t.Fatalf("mode %d: reason = %s, want scheme", p.Mode, got) } req, _ := http.NewRequestWithContext(t.Context(), http.MethodDelete, srv.URL, nil) if _, err := c.Do(req); reasonOf(t, err) != fetchguard.ReasonScheme { t.Fatalf("mode %d: Do reason = %v, want scheme", p.Mode, err) } } if hits.Load() != 0 { t.Fatal("http URL must not cause network I/O in a guarded mode") } c := mustClient(t, fetchguard.Policy{Mode: fetchguard.PublicOnlyMode}) if _, err := c.Get(t.Context(), "not a url", nil); reasonOf(t, err) != fetchguard.ReasonInvalidURL { t.Fatalf("malformed URL: %v", err) } if _, err := c.Do(nil); reasonOf(t, err) != fetchguard.ReasonInvalidURL { t.Fatalf("nil request: %v", err) } } func TestClientMultipart(t *testing.T) { type part struct{ name, filename, ctype, body string } var got []part var gotAuth, gotCT string srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { gotAuth = r.Header.Get("Authorization") gotCT = r.Header.Get("Content-Type") mr, err := r.MultipartReader() if err != nil { http.Error(w, err.Error(), 400) return } for { p, err := mr.NextPart() if err == io.EOF { break } if err != nil { http.Error(w, err.Error(), 400) return } b, _ := io.ReadAll(p) got = append(got, part{p.FormName(), p.FileName(), p.Header.Get("Content-Type"), string(b)}) } w.WriteHeader(http.StatusAccepted) })) t.Cleanup(srv.Close) c := mustClient(t, fetchguard.Policy{Mode: fetchguard.TrustedMode}) h := http.Header{} h.Set("Authorization", fetchguard.Bearer("tok")) h.Set("Content-Type", "application/json") // must not win over the boundary res, err := c.PostMultipart(t.Context(), srv.URL+"/upload", h, []fetchguard.FormField{{Name: "title", Value: "Hello"}, {Name: "kind", Value: "bug"}}, []fetchguard.FormFile{ {Field: "file", Filename: "shot.png", ContentType: "image/png", Body: strings.NewReader("\x89PNG")}, {Field: "extra", Filename: `a"b.txt`, Body: strings.NewReader("text")}, }) if err != nil { t.Fatalf("PostMultipart: %v", err) } if res.StatusCode != http.StatusAccepted { t.Fatalf("status = %d: %s", res.StatusCode, res.Body) } if gotAuth != "Bearer tok" || !strings.HasPrefix(gotCT, "multipart/form-data; boundary=") { t.Fatalf("auth %q, content type %q", gotAuth, gotCT) } want := []part{ {"title", "", "", "Hello"}, {"kind", "", "", "bug"}, {"file", "shot.png", "image/png", "\x89PNG"}, {"extra", `a"b.txt`, "application/octet-stream", "text"}, } if !reflect.DeepEqual(got, want) { t.Fatalf("parts = %+v\nwant %+v", got, want) } } func TestClientPutJSONHeaderOrder(t *testing.T) { var method, ct, body string srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { method, ct = r.Method, r.Header.Get("Content-Type") b, _ := io.ReadAll(r.Body) body = string(b) w.Header().Set("Retry-After", "3") w.WriteHeader(http.StatusTooManyRequests) })) t.Cleanup(srv.Close) c := mustClient(t, fetchguard.Policy{Mode: fetchguard.TrustedMode}) res, err := c.PutJSON(t.Context(), srv.URL, http.Header{"content-type": {"application/vnd.example+json"}}, []int{1, 2}) if err != nil { t.Fatal(err) } if method != http.MethodPut || ct != "application/vnd.example+json" || body != "[1,2]" { t.Fatalf("got %s %q %q", method, ct, body) } if res.StatusCode != http.StatusTooManyRequests || res.Header.Get("Retry-After") != "3" { t.Fatalf("status %d, Retry-After %q: the status is returned, not judged", res.StatusCode, res.Header.Get("Retry-After")) } if _, err := c.PostJSON(t.Context(), srv.URL, nil, func() {}); reasonOf(t, err) != fetchguard.ReasonInvalidURL { t.Fatalf("unmarshalable body: %v", err) } } func TestClientBodyCap(t *testing.T) { const maxBytes = 64 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { n := maxBytes if r.URL.Path == "/over" { n++ } _, _ = w.Write([]byte(strings.Repeat("x", n))) })) t.Cleanup(srv.Close) c := mustClient(t, fetchguard.Policy{Mode: fetchguard.TrustedMode, MaxBytes: maxBytes}) res, err := c.Get(t.Context(), srv.URL+"/exact", nil) if err != nil || len(res.Body) != maxBytes { t.Fatalf("exact: %v, %d bytes", err, len(res.Body)) } if _, err := c.Get(t.Context(), srv.URL+"/over", nil); reasonOf(t, err) != fetchguard.ReasonTooLarge { t.Fatalf("Send over cap: %v, want too_large", err) } req, _ := http.NewRequestWithContext(t.Context(), http.MethodGet, srv.URL+"/over", nil) resp, err := c.Do(req) if err != nil { t.Fatal(err) } defer resp.Body.Close() data, err := io.ReadAll(resp.Body) if reasonOf(t, err) != fetchguard.ReasonTooLarge { t.Fatalf("Do body over cap: %v, want too_large", err) } if len(data) != maxBytes { t.Fatalf("read %d bytes before the cap error, want %d", len(data), maxBytes) } if n, err := resp.Body.Read(make([]byte, 8)); n != 0 || reasonOf(t, err) != fetchguard.ReasonTooLarge { t.Fatalf("read after cap = %d, %v", n, err) } } func mustClient(t *testing.T, p fetchguard.Policy) *fetchguard.Client { t.Helper() if p.Timeout == 0 { p.Timeout = 5 * time.Second } c, err := fetchguard.NewClient(p, nil) if err != nil { t.Fatal(err) } return c }