package fetchguard_test import ( "errors" "net/http" "reflect" "strings" "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", fetchguardBearer("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 fetchguardBearer(token string) string { return "Bearer " + token } 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) }