feat(14-01): fetchguard client covers PUT, multipart, bearer and a trusted mode
- TrustedMode (declared after PublicOnlyMode) lifts the scheme, host and dial checks for Client only - PutJSON, PostMultipart with FormField/FormFile, Bearer - tests for modes, redirects, multipart order, body cap and the scheme guard - README, root modules row and outbound HTTP docs describe the client and its test seam
This commit is contained in:
@@ -2,9 +2,12 @@ package fetchguard_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -53,7 +56,7 @@ func TestClientPostJSONThroughUpstreamFake(t *testing.T) {
|
||||
}
|
||||
ctx := fetchguard.WithTransport(t.Context(), fake)
|
||||
header := http.Header{}
|
||||
header.Set("Authorization", fetchguardBearer("example-token-value"))
|
||||
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{
|
||||
@@ -81,8 +84,6 @@ func TestClientPostJSONThroughUpstreamFake(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
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]()} {
|
||||
@@ -141,3 +142,217 @@ func TestTransportSeamIsCodeOnly(t *testing.T) {
|
||||
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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user