feat(14-01): guarded fetchguard client replayed through a tide upstream fake
- fetchguard.NewClient with Do, Send, Get and PostJSON over a capped, never-redirecting transport - WithTransport: code-only context seam for offline replay; Result gains Header - tide UpstreamSidecar, LoadUpstream, UpstreamPath and the asserting UpstreamFake
This commit is contained in:
197
modules/fetchguard/client.go
Normal file
197
modules/fetchguard/client.go
Normal file
@@ -0,0 +1,197 @@
|
||||
package fetchguard
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
)
|
||||
|
||||
// Client is a guarded outbound HTTP client for calling service APIs: any
|
||||
// method, JSON or multipart bodies, caller-supplied headers. It keeps the
|
||||
// guarantees of Fetch for the policy it was built with: URL validation
|
||||
// before any I/O, the dial-time private and reserved address check on every
|
||||
// new connection (except in TrustedMode), a response body cap and no
|
||||
// redirects. Status codes are returned, never judged.
|
||||
//
|
||||
// A Client is safe for concurrent use; build one per vendor and reuse it so
|
||||
// connections are kept alive.
|
||||
type Client struct {
|
||||
policy Policy
|
||||
maxBytes int64
|
||||
timeout time.Duration
|
||||
http *http.Client
|
||||
}
|
||||
|
||||
// NewClient resolves the policy's limits once (zero MaxBytes or Timeout fall
|
||||
// back to cfg, then to Defaults) and builds the client.
|
||||
func NewClient(policy Policy, cfg *compass.Config) (*Client, error) {
|
||||
maxBytes, timeout, err := resolveLimits(policy, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Client{
|
||||
policy: policy,
|
||||
maxBytes: maxBytes,
|
||||
timeout: timeout,
|
||||
http: newHTTPClient(newTransport(policy, timeout, true), timeout),
|
||||
}, nil
|
||||
}
|
||||
|
||||
type transportKey struct{}
|
||||
|
||||
// WithTransport returns a context whose requests a Client sends through rt
|
||||
// instead of its own network transport. It is the test and parity-replay
|
||||
// seam: a test hands it a fake http.RoundTripper (for example the tide
|
||||
// upstream fake) so the client's real request is asserted offline.
|
||||
//
|
||||
// The override is code-only. No Policy field, Client field, config key,
|
||||
// environment variable or request header can set it; only Go code holding
|
||||
// the context can. URL validation still runs before the override is
|
||||
// consulted, and redirects are still not followed.
|
||||
func WithTransport(ctx context.Context, rt http.RoundTripper) context.Context {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
return context.WithValue(ctx, transportKey{}, rt)
|
||||
}
|
||||
|
||||
func transportFrom(ctx context.Context) http.RoundTripper {
|
||||
if ctx == nil {
|
||||
return nil
|
||||
}
|
||||
rt, _ := ctx.Value(transportKey{}).(http.RoundTripper)
|
||||
return rt
|
||||
}
|
||||
|
||||
// Do validates req's URL against the client's policy and sends it. The
|
||||
// response body is capped at the policy's MaxBytes: reading past it fails
|
||||
// with an *Error whose Reason is ReasonTooLarge. The caller closes the body.
|
||||
// A 3xx response is returned as-is; redirects are never followed.
|
||||
//
|
||||
// A non-nil error is always *Error.
|
||||
func (c *Client) Do(req *http.Request) (*http.Response, error) {
|
||||
if req == nil || req.URL == nil {
|
||||
return nil, &Error{Reason: ReasonInvalidURL}
|
||||
}
|
||||
if err := c.check(req.URL); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hc := c.http
|
||||
if rt := transportFrom(req.Context()); rt != nil {
|
||||
hc = newHTTPClient(rt, c.timeout)
|
||||
}
|
||||
resp, err := hc.Do(req)
|
||||
if err != nil {
|
||||
return nil, mapTransportError(err)
|
||||
}
|
||||
resp.Body = &cappedBody{rc: resp.Body, max: c.maxBytes}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// Send sends req through Do and reads the whole capped body into a Result.
|
||||
func (c *Client) Send(req *http.Request) (*Result, error) {
|
||||
resp, err := c.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
data, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
var fe *Error
|
||||
if errors.As(err, &fe) {
|
||||
return nil, fe
|
||||
}
|
||||
return nil, &Error{Reason: ReasonNetworkError, Err: err}
|
||||
}
|
||||
return &Result{
|
||||
Body: data,
|
||||
ContentType: resp.Header.Get("Content-Type"),
|
||||
StatusCode: resp.StatusCode,
|
||||
Header: resp.Header,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// PostJSON marshals body as JSON and POSTs it to rawURL. Content-Type is set
|
||||
// to application/json first, then header is copied over it, so a caller can
|
||||
// override any header including Content-Type.
|
||||
func (c *Client) PostJSON(ctx context.Context, rawURL string, header http.Header, body any) (*Result, error) {
|
||||
return c.sendJSON(ctx, http.MethodPost, rawURL, header, body)
|
||||
}
|
||||
|
||||
// Get sends a GET to rawURL with header.
|
||||
func (c *Client) Get(ctx context.Context, rawURL string, header http.Header) (*Result, error) {
|
||||
return c.send(ctx, http.MethodGet, rawURL, nil, header, nil)
|
||||
}
|
||||
|
||||
func (c *Client) sendJSON(ctx context.Context, method, rawURL string, header http.Header, body any) (*Result, error) {
|
||||
raw, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return nil, &Error{Reason: ReasonInvalidURL, Err: err}
|
||||
}
|
||||
h := http.Header{}
|
||||
h.Set("Content-Type", "application/json")
|
||||
return c.send(ctx, method, rawURL, h, header, bytes.NewReader(raw))
|
||||
}
|
||||
|
||||
func (c *Client) send(ctx context.Context, method, rawURL string, base, header http.Header, body io.Reader) (*Result, error) {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, rawURL, body)
|
||||
if err != nil {
|
||||
return nil, &Error{Reason: ReasonInvalidURL, Err: err}
|
||||
}
|
||||
for k, vs := range base {
|
||||
req.Header[k] = append([]string(nil), vs...)
|
||||
}
|
||||
for k, vs := range header {
|
||||
req.Header[http.CanonicalHeaderKey(k)] = append([]string(nil), vs...)
|
||||
}
|
||||
return c.Send(req)
|
||||
}
|
||||
|
||||
// check applies the policy's URL rules before any I/O.
|
||||
func (c *Client) check(u *url.URL) error {
|
||||
if u.Scheme == "" || u.Host == "" {
|
||||
return &Error{Reason: ReasonInvalidURL}
|
||||
}
|
||||
if strings.ToLower(u.Scheme) != "https" {
|
||||
return &Error{Reason: ReasonScheme}
|
||||
}
|
||||
if c.policy.Mode == AllowHostsMode && !hostAllowed(u.Hostname(), c.policy.AllowHosts) {
|
||||
return &Error{Reason: ReasonInvalidURL}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// cappedBody fails with ReasonTooLarge once more than max bytes were read.
|
||||
type cappedBody struct {
|
||||
rc io.ReadCloser
|
||||
max int64
|
||||
n int64
|
||||
}
|
||||
|
||||
func (b *cappedBody) Read(p []byte) (int, error) {
|
||||
if b.n > b.max {
|
||||
return 0, &Error{Reason: ReasonTooLarge}
|
||||
}
|
||||
if room := b.max - b.n + 1; int64(len(p)) > room {
|
||||
p = p[:room]
|
||||
}
|
||||
n, err := b.rc.Read(p)
|
||||
b.n += int64(n)
|
||||
if b.n > b.max {
|
||||
return n - int(b.n-b.max), &Error{Reason: ReasonTooLarge}
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (b *cappedBody) Close() error { return b.rc.Close() }
|
||||
143
modules/fetchguard/client_test.go
Normal file
143
modules/fetchguard/client_test.go
Normal file
@@ -0,0 +1,143 @@
|
||||
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) }
|
||||
@@ -16,11 +16,15 @@ import (
|
||||
"git.golem15.com/golem15/summercms/modules/compass"
|
||||
)
|
||||
|
||||
// Result is a successful (including non-2xx, including 3xx) Fetch response.
|
||||
// Result is a successful (including non-2xx, including 3xx) Fetch or Client
|
||||
// response.
|
||||
type Result struct {
|
||||
Body []byte
|
||||
ContentType string
|
||||
StatusCode int
|
||||
// Header holds every response header, for callers that read rate-limit
|
||||
// or Retry-After values.
|
||||
Header http.Header
|
||||
}
|
||||
|
||||
var errPrivateIP = errors.New("private_ip")
|
||||
@@ -55,24 +59,7 @@ func Fetch(ctx context.Context, rawURL string, policy Policy, cfg *compass.Confi
|
||||
return nil, err
|
||||
}
|
||||
|
||||
client := &http.Client{
|
||||
Timeout: timeout,
|
||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
Transport: &http.Transport{
|
||||
// User-supplied URLs must not be forwarded through HTTP_PROXY:
|
||||
// the dial-time IP check would then see the proxy, not the target.
|
||||
Proxy: nil,
|
||||
DialContext: (&net.Dialer{
|
||||
Timeout: timeout,
|
||||
Control: dialControl(policy),
|
||||
}).DialContext,
|
||||
TLSClientConfig: policy.tlsConfig,
|
||||
DisableKeepAlives: true,
|
||||
ForceAttemptHTTP2: true,
|
||||
},
|
||||
}
|
||||
client := newHTTPClient(newTransport(policy, timeout, false), timeout)
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
|
||||
if err != nil {
|
||||
@@ -95,9 +82,39 @@ func Fetch(ctx context.Context, rawURL string, policy Policy, cfg *compass.Confi
|
||||
Body: data,
|
||||
ContentType: resp.Header.Get("Content-Type"),
|
||||
StatusCode: resp.StatusCode,
|
||||
Header: resp.Header,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// newTransport builds the guarded transport for policy. keepAlives is off for
|
||||
// one-shot Fetch calls and on for a reusable Client; the dial Control runs on
|
||||
// every new connection either way.
|
||||
func newTransport(policy Policy, timeout time.Duration, keepAlives bool) *http.Transport {
|
||||
return &http.Transport{
|
||||
// User-supplied URLs must not be forwarded through HTTP_PROXY:
|
||||
// the dial-time IP check would then see the proxy, not the target.
|
||||
Proxy: nil,
|
||||
DialContext: (&net.Dialer{
|
||||
Timeout: timeout,
|
||||
Control: dialControl(policy),
|
||||
}).DialContext,
|
||||
TLSClientConfig: policy.tlsConfig,
|
||||
DisableKeepAlives: !keepAlives,
|
||||
ForceAttemptHTTP2: true,
|
||||
}
|
||||
}
|
||||
|
||||
// newHTTPClient wraps rt in an http.Client that never follows redirects.
|
||||
func newHTTPClient(rt http.RoundTripper, timeout time.Duration) *http.Client {
|
||||
return &http.Client{
|
||||
Timeout: timeout,
|
||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
Transport: rt,
|
||||
}
|
||||
}
|
||||
|
||||
func resolveLimits(policy Policy, cfg *compass.Config) (int64, time.Duration, error) {
|
||||
maxBytes := policy.MaxBytes
|
||||
timeout := policy.Timeout
|
||||
|
||||
Reference in New Issue
Block a user