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"
|
"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 {
|
type Result struct {
|
||||||
Body []byte
|
Body []byte
|
||||||
ContentType string
|
ContentType string
|
||||||
StatusCode int
|
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")
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
client := &http.Client{
|
client := newHTTPClient(newTransport(policy, timeout, false), timeout)
|
||||||
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,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -95,9 +82,39 @@ func Fetch(ctx context.Context, rawURL string, policy Policy, cfg *compass.Confi
|
|||||||
Body: data,
|
Body: data,
|
||||||
ContentType: resp.Header.Get("Content-Type"),
|
ContentType: resp.Header.Get("Content-Type"),
|
||||||
StatusCode: resp.StatusCode,
|
StatusCode: resp.StatusCode,
|
||||||
|
Header: resp.Header,
|
||||||
}, nil
|
}, 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) {
|
func resolveLimits(policy Policy, cfg *compass.Config) (int64, time.Duration, error) {
|
||||||
maxBytes := policy.MaxBytes
|
maxBytes := policy.MaxBytes
|
||||||
timeout := policy.Timeout
|
timeout := policy.Timeout
|
||||||
|
|||||||
17
modules/tide/testdata/upstream/post_json.upstream.yaml
vendored
Normal file
17
modules/tide/testdata/upstream/post_json.upstream.yaml
vendored
Normal file
@@ -0,0 +1,17 @@
|
|||||||
|
version: 1
|
||||||
|
exchanges:
|
||||||
|
- request:
|
||||||
|
method: POST
|
||||||
|
url: https://api.example.test/v1/things?mode=fast&lang=en
|
||||||
|
headers:
|
||||||
|
Accept: application/json
|
||||||
|
Authorization: Bearer {{secret:example-token}}
|
||||||
|
Content-Type: application/json
|
||||||
|
User-Agent: example-client/1.0
|
||||||
|
body: '{"name":"widget","count":2,"tags":["a","b"]}'
|
||||||
|
response:
|
||||||
|
status: 201
|
||||||
|
headers:
|
||||||
|
Content-Type: application/json
|
||||||
|
X-Request-Id: req-123
|
||||||
|
body: '{"id":7,"name":"widget"}'
|
||||||
305
modules/tide/upstream.go
Normal file
305
modules/tide/upstream.go
Normal file
@@ -0,0 +1,305 @@
|
|||||||
|
package tide
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/goccy/go-yaml"
|
||||||
|
)
|
||||||
|
|
||||||
|
// UpstreamSidecar is the upstream HTTP traffic a backend sent to outside
|
||||||
|
// services while one parity fixture was recorded. It lives next to the
|
||||||
|
// fixture as <fixture>.upstream.yaml (see UpstreamPath). Exchanges are
|
||||||
|
// consumed in recorded order across the whole flow.
|
||||||
|
type UpstreamSidecar struct {
|
||||||
|
Version int `yaml:"version"`
|
||||||
|
Exchanges []UpstreamExchange `yaml:"exchanges"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpstreamExchange is one outbound request and the vendor's answer.
|
||||||
|
type UpstreamExchange struct {
|
||||||
|
Request UpstreamRequest `yaml:"request"`
|
||||||
|
Response UpstreamResponse `yaml:"response"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpstreamRequest is the request the backend sent upstream. URL is absolute
|
||||||
|
// (scheme, host, path and query). Headers holds only the compared headers
|
||||||
|
// (UpstreamCompareHeaders); credential values are {{name}} placeholders.
|
||||||
|
type UpstreamRequest struct {
|
||||||
|
Method string `yaml:"method"`
|
||||||
|
URL string `yaml:"url"`
|
||||||
|
Headers map[string]string `yaml:"headers,omitempty"`
|
||||||
|
Body string `yaml:"body,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpstreamResponse is the recorded vendor answer the fake replays.
|
||||||
|
type UpstreamResponse struct {
|
||||||
|
Status int `yaml:"status"`
|
||||||
|
Headers map[string]string `yaml:"headers,omitempty"`
|
||||||
|
Body string `yaml:"body,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpstreamCompareHeaders are the request headers the fake asserts. A header
|
||||||
|
// is compared when either the recorded or the sent request carries it.
|
||||||
|
var UpstreamCompareHeaders = []string{
|
||||||
|
"User-Agent",
|
||||||
|
"Accept",
|
||||||
|
"Content-Type",
|
||||||
|
"Authorization",
|
||||||
|
"X-Api-Key",
|
||||||
|
"Anthropic-Version",
|
||||||
|
"Anthropic-Beta",
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpstreamPath maps a fixture path to its sidecar path: x.yaml becomes
|
||||||
|
// x.upstream.yaml.
|
||||||
|
func UpstreamPath(fixturePath string) string {
|
||||||
|
for _, ext := range []string{".yaml", ".yml"} {
|
||||||
|
if strings.HasSuffix(fixturePath, ext) {
|
||||||
|
return strings.TrimSuffix(fixturePath, ext) + ".upstream.yaml"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fixturePath + ".upstream.yaml"
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoadUpstream reads a version-1 sidecar, rejecting unknown fields. A missing
|
||||||
|
// file returns an error wrapping fs.ErrNotExist.
|
||||||
|
func LoadUpstream(path string) (UpstreamSidecar, error) {
|
||||||
|
raw, err := os.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return UpstreamSidecar{}, fmt.Errorf("tide: read upstream %s: %w", path, err)
|
||||||
|
}
|
||||||
|
var s UpstreamSidecar
|
||||||
|
dec := yaml.NewDecoder(bytes.NewReader(raw), yaml.DisallowUnknownField())
|
||||||
|
if err := dec.Decode(&s); err != nil {
|
||||||
|
return UpstreamSidecar{}, fmt.Errorf("tide: parse upstream %s: %w", path, err)
|
||||||
|
}
|
||||||
|
if err := validateUpstream(s); err != nil {
|
||||||
|
return UpstreamSidecar{}, fmt.Errorf("tide: upstream %s: %w", path, err)
|
||||||
|
}
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateUpstream(s UpstreamSidecar) error {
|
||||||
|
if s.Version != CurrentVersion {
|
||||||
|
return fmt.Errorf("version %d, want %d", s.Version, CurrentVersion)
|
||||||
|
}
|
||||||
|
for i, ex := range s.Exchanges {
|
||||||
|
if strings.TrimSpace(ex.Request.Method) == "" {
|
||||||
|
return fmt.Errorf("exchange %d: request method is required", i)
|
||||||
|
}
|
||||||
|
u, err := url.Parse(ex.Request.URL)
|
||||||
|
if err != nil || u.Scheme == "" || u.Host == "" {
|
||||||
|
return fmt.Errorf("exchange %d: request url %q must be absolute", i, ex.Request.URL)
|
||||||
|
}
|
||||||
|
if ex.Response.Status < 100 || ex.Response.Status > 999 {
|
||||||
|
return fmt.Errorf("exchange %d: response status %d out of range", i, ex.Response.Status)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpstreamFake is an http.RoundTripper that answers from a sidecar without
|
||||||
|
// dialing. Each request takes the next recorded exchange and is asserted
|
||||||
|
// against it: method, scheme, host, path, query (order-insensitive), the
|
||||||
|
// UpstreamCompareHeaders with {{name}} placeholders expanded from the store,
|
||||||
|
// and the body (JSON semantically, otherwise byte for byte). A mismatch fails
|
||||||
|
// the request and is remembered for Verify.
|
||||||
|
type UpstreamFake struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
store *Store
|
||||||
|
exchanges []UpstreamExchange
|
||||||
|
next int
|
||||||
|
errs []error
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewUpstreamFake returns a fake that replays s. store resolves {{name}}
|
||||||
|
// placeholders in the recorded requests and responses; it may be nil when the
|
||||||
|
// sidecar has none.
|
||||||
|
func NewUpstreamFake(s UpstreamSidecar, store *Store) *UpstreamFake {
|
||||||
|
return &UpstreamFake{store: store, exchanges: slices.Clone(s.Exchanges)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RoundTrip implements http.RoundTripper.
|
||||||
|
func (f *UpstreamFake) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
var body []byte
|
||||||
|
if req.Body != nil {
|
||||||
|
var err error
|
||||||
|
body, err = io.ReadAll(io.LimitReader(req.Body, DefaultMaxBody+1))
|
||||||
|
_ = req.Body.Close()
|
||||||
|
if err != nil {
|
||||||
|
return nil, f.fail(fmt.Errorf("tide: upstream %s %s: read body: %w", req.Method, req.URL, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
f.mu.Lock()
|
||||||
|
if f.next >= len(f.exchanges) {
|
||||||
|
f.mu.Unlock()
|
||||||
|
return nil, f.fail(fmt.Errorf("tide: upstream extra request %s %s", req.Method, redactURL(req.URL)))
|
||||||
|
}
|
||||||
|
idx := f.next
|
||||||
|
ex := f.exchanges[idx]
|
||||||
|
f.next++
|
||||||
|
f.mu.Unlock()
|
||||||
|
|
||||||
|
if problems := f.compare(ex.Request, req, body); len(problems) > 0 {
|
||||||
|
return nil, f.fail(fmt.Errorf("tide: upstream exchange %d %s %s: %s",
|
||||||
|
idx, req.Method, redactURL(req.URL), strings.Join(problems, "; ")))
|
||||||
|
}
|
||||||
|
resp, err := f.response(ex.Response, req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, f.fail(fmt.Errorf("tide: upstream exchange %d: %w", idx, err))
|
||||||
|
}
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify reports every mismatch, every extra request and every recorded
|
||||||
|
// exchange that was never requested. It returns nil when the backend sent
|
||||||
|
// exactly the recorded requests.
|
||||||
|
func (f *UpstreamFake) Verify() error {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
errs := slices.Clone(f.errs)
|
||||||
|
for i := f.next; i < len(f.exchanges); i++ {
|
||||||
|
r := f.exchanges[i].Request
|
||||||
|
errs = append(errs, fmt.Errorf("tide: upstream exchange %d unconsumed: %s %s", i, r.Method, r.URL))
|
||||||
|
}
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *UpstreamFake) fail(err error) error {
|
||||||
|
f.mu.Lock()
|
||||||
|
f.errs = append(f.errs, err)
|
||||||
|
f.mu.Unlock()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *UpstreamFake) compare(want UpstreamRequest, got *http.Request, body []byte) []string {
|
||||||
|
var problems []string
|
||||||
|
if !strings.EqualFold(want.Method, got.Method) {
|
||||||
|
problems = append(problems, fmt.Sprintf("method: want %s, got %s", want.Method, got.Method))
|
||||||
|
}
|
||||||
|
rawURL, err := f.store.Expand(want.URL)
|
||||||
|
if err != nil {
|
||||||
|
return append(problems, "url: "+err.Error())
|
||||||
|
}
|
||||||
|
wu, err := url.Parse(rawURL)
|
||||||
|
if err != nil {
|
||||||
|
return append(problems, "url: "+err.Error())
|
||||||
|
}
|
||||||
|
gu := got.URL
|
||||||
|
if !strings.EqualFold(wu.Scheme, gu.Scheme) {
|
||||||
|
problems = append(problems, fmt.Sprintf("scheme: want %s, got %s", wu.Scheme, gu.Scheme))
|
||||||
|
}
|
||||||
|
if !strings.EqualFold(wu.Host, gu.Host) {
|
||||||
|
problems = append(problems, fmt.Sprintf("host: want %s, got %s", wu.Host, gu.Host))
|
||||||
|
}
|
||||||
|
if wu.EscapedPath() != gu.EscapedPath() && wu.Path != gu.Path {
|
||||||
|
problems = append(problems, fmt.Sprintf("path: want %s, got %s", wu.EscapedPath(), gu.EscapedPath()))
|
||||||
|
}
|
||||||
|
problems = append(problems, compareQuery(wu.Query(), gu.Query())...)
|
||||||
|
|
||||||
|
for _, name := range UpstreamCompareHeaders {
|
||||||
|
wv, err := f.store.Expand(headerValue(want.Headers, name))
|
||||||
|
if err != nil {
|
||||||
|
problems = append(problems, "header "+name+": "+err.Error())
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
gv := got.Header.Get(name)
|
||||||
|
if wv == gv {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if isCredentialHeader(name) {
|
||||||
|
problems = append(problems, fmt.Sprintf("header %s: value differs", name))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
problems = append(problems, fmt.Sprintf("header %s: want %q, got %q", name, wv, gv))
|
||||||
|
}
|
||||||
|
|
||||||
|
wantBody, err := f.store.Expand(want.Body)
|
||||||
|
if err != nil {
|
||||||
|
return append(problems, "body: "+err.Error())
|
||||||
|
}
|
||||||
|
jsonBody := isJSONContentType(want.Headers) || isJSONContentType(map[string]string{"Content-Type": got.Header.Get("Content-Type")})
|
||||||
|
switch {
|
||||||
|
case jsonBody && (len(wantBody) > 0 || len(body) > 0):
|
||||||
|
if diffs := diffJSON([]byte(wantBody), body); len(diffs) > 0 {
|
||||||
|
d := diffs[0]
|
||||||
|
problems = append(problems, fmt.Sprintf("body %s: want %s, got %s", d.Path, d.Expected, d.Actual))
|
||||||
|
}
|
||||||
|
case wantBody != string(body):
|
||||||
|
problems = append(problems, fmt.Sprintf("body: want %d bytes, got %d bytes", len(wantBody), len(body)))
|
||||||
|
}
|
||||||
|
return problems
|
||||||
|
}
|
||||||
|
|
||||||
|
func compareQuery(want, got url.Values) []string {
|
||||||
|
var problems []string
|
||||||
|
keys := make([]string, 0, len(want)+len(got))
|
||||||
|
for k := range want {
|
||||||
|
keys = append(keys, k)
|
||||||
|
}
|
||||||
|
for k := range got {
|
||||||
|
if _, ok := want[k]; !ok {
|
||||||
|
keys = append(keys, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
slices.Sort(keys)
|
||||||
|
for _, k := range keys {
|
||||||
|
wv, gv := slices.Clone(want[k]), slices.Clone(got[k])
|
||||||
|
slices.Sort(wv)
|
||||||
|
slices.Sort(gv)
|
||||||
|
if !slices.Equal(wv, gv) {
|
||||||
|
problems = append(problems, fmt.Sprintf("query %s: want %q, got %q", k, wv, gv))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return problems
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *UpstreamFake) response(r UpstreamResponse, req *http.Request) (*http.Response, error) {
|
||||||
|
h := http.Header{}
|
||||||
|
for k, v := range r.Headers {
|
||||||
|
ev, err := f.store.Expand(v)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("response header %s: %w", k, err)
|
||||||
|
}
|
||||||
|
h.Set(k, ev)
|
||||||
|
}
|
||||||
|
body, err := f.store.Expand(r.Body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("response body: %w", err)
|
||||||
|
}
|
||||||
|
return &http.Response{
|
||||||
|
Status: fmt.Sprintf("%d %s", r.Status, http.StatusText(r.Status)),
|
||||||
|
StatusCode: r.Status,
|
||||||
|
Proto: "HTTP/1.1",
|
||||||
|
ProtoMajor: 1,
|
||||||
|
ProtoMinor: 1,
|
||||||
|
Header: h,
|
||||||
|
Body: io.NopCloser(strings.NewReader(body)),
|
||||||
|
ContentLength: int64(len(body)),
|
||||||
|
Request: req,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func isCredentialHeader(name string) bool {
|
||||||
|
return strings.EqualFold(name, "Authorization") || strings.EqualFold(name, "X-Api-Key")
|
||||||
|
}
|
||||||
|
|
||||||
|
// redactURL drops the query, which may carry a vendor token.
|
||||||
|
func redactURL(u *url.URL) string {
|
||||||
|
if u == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
c := *u
|
||||||
|
c.RawQuery = ""
|
||||||
|
c.User = nil
|
||||||
|
return c.String()
|
||||||
|
}
|
||||||
150
modules/tide/upstream_test.go
Normal file
150
modules/tide/upstream_test.go
Normal file
@@ -0,0 +1,150 @@
|
|||||||
|
package tide
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func upstreamTestSidecar() UpstreamSidecar {
|
||||||
|
return UpstreamSidecar{Version: 1, Exchanges: []UpstreamExchange{{
|
||||||
|
Request: UpstreamRequest{
|
||||||
|
Method: "POST",
|
||||||
|
URL: "https://api.example.test/v1/things?mode=fast&lang=en",
|
||||||
|
Headers: map[string]string{
|
||||||
|
"Authorization": "Bearer {{secret:example-token}}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
"User-Agent": "example-client/1.0",
|
||||||
|
},
|
||||||
|
Body: `{"name":"widget","count":2}`,
|
||||||
|
},
|
||||||
|
Response: UpstreamResponse{Status: 201, Body: `{"id":7}`},
|
||||||
|
}}}
|
||||||
|
}
|
||||||
|
|
||||||
|
func upstreamTestStore(t *testing.T) *Store {
|
||||||
|
t.Helper()
|
||||||
|
s, err := OpenStore("")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
s.Set("secret:example-token", "example-token-value")
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
type upstreamCall struct {
|
||||||
|
method, url, ua, auth, body string
|
||||||
|
}
|
||||||
|
|
||||||
|
func goodUpstreamCall() upstreamCall {
|
||||||
|
return upstreamCall{
|
||||||
|
method: "POST",
|
||||||
|
url: "https://api.example.test/v1/things?lang=en&mode=fast",
|
||||||
|
ua: "example-client/1.0",
|
||||||
|
auth: "Bearer example-token-value",
|
||||||
|
body: `{"count":2,"name":"widget"}`,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c upstreamCall) send(t *testing.T, f *UpstreamFake) error {
|
||||||
|
t.Helper()
|
||||||
|
req, err := http.NewRequest(c.method, c.url, bytes.NewReader([]byte(c.body)))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("User-Agent", c.ua)
|
||||||
|
req.Header.Set("Authorization", c.auth)
|
||||||
|
resp, err := f.RoundTrip(req)
|
||||||
|
if err == nil {
|
||||||
|
_ = resp.Body.Close()
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpstreamFakeRejectsMismatchedRequest(t *testing.T) {
|
||||||
|
ok := goodUpstreamCall()
|
||||||
|
t.Run("match", func(t *testing.T) {
|
||||||
|
f := NewUpstreamFake(upstreamTestSidecar(), upstreamTestStore(t))
|
||||||
|
if err := ok.send(t, f); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := f.Verify(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
edit func(*upstreamCall)
|
||||||
|
field string
|
||||||
|
}{
|
||||||
|
{"method", func(c *upstreamCall) { c.method = "PUT" }, "method"},
|
||||||
|
{"path", func(c *upstreamCall) { c.url = strings.Replace(c.url, "/v1/things", "/v1/other", 1) }, "path"},
|
||||||
|
{"query", func(c *upstreamCall) { c.url = strings.Replace(c.url, "mode=fast", "mode=slow", 1) }, "query mode"},
|
||||||
|
{"host", func(c *upstreamCall) { c.url = strings.Replace(c.url, "api.example.test", "api2.example.test", 1) }, "host"},
|
||||||
|
{"user agent", func(c *upstreamCall) { c.ua = "other/2.0" }, "header User-Agent"},
|
||||||
|
{"authorization", func(c *upstreamCall) { c.auth = "Bearer wrong-token" }, "header Authorization"},
|
||||||
|
{"json body", func(c *upstreamCall) { c.body = `{"count":3,"name":"widget"}` }, "body $.count"},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
f := NewUpstreamFake(upstreamTestSidecar(), upstreamTestStore(t))
|
||||||
|
c := ok
|
||||||
|
tc.edit(&c)
|
||||||
|
if err := c.send(t, f); err == nil || !strings.Contains(err.Error(), tc.field) {
|
||||||
|
t.Fatalf("RoundTrip err = %v, want it to name %q", err, tc.field)
|
||||||
|
}
|
||||||
|
err := f.Verify()
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tc.field) {
|
||||||
|
t.Fatalf("Verify = %v, want it to name %q", err, tc.field)
|
||||||
|
}
|
||||||
|
if strings.Contains(err.Error(), "wrong-token") || strings.Contains(err.Error(), "example-token-value") {
|
||||||
|
t.Fatalf("Verify leaks a credential: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
t.Run("extra request", func(t *testing.T) {
|
||||||
|
f := NewUpstreamFake(upstreamTestSidecar(), upstreamTestStore(t))
|
||||||
|
if err := ok.send(t, f); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := ok.send(t, f); err == nil {
|
||||||
|
t.Fatal("second request must fail")
|
||||||
|
}
|
||||||
|
if err := f.Verify(); err == nil || !strings.Contains(err.Error(), "extra request") {
|
||||||
|
t.Fatalf("Verify = %v, want extra request", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
t.Run("unconsumed exchange", func(t *testing.T) {
|
||||||
|
f := NewUpstreamFake(upstreamTestSidecar(), upstreamTestStore(t))
|
||||||
|
if err := f.Verify(); err == nil || !strings.Contains(err.Error(), "unconsumed") {
|
||||||
|
t.Fatalf("Verify = %v, want unconsumed", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
t.Run("unresolved placeholder", func(t *testing.T) {
|
||||||
|
f := NewUpstreamFake(upstreamTestSidecar(), nil)
|
||||||
|
if err := ok.send(t, f); err == nil || !strings.Contains(err.Error(), "header Authorization") {
|
||||||
|
t.Fatalf("err = %v, want unresolved Authorization placeholder", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadUpstreamAndPath(t *testing.T) {
|
||||||
|
if got := UpstreamPath("fixtures/routes/GET_x__ok.yaml"); got != "fixtures/routes/GET_x__ok.upstream.yaml" {
|
||||||
|
t.Fatalf("UpstreamPath = %q", got)
|
||||||
|
}
|
||||||
|
if got := UpstreamPath("a.yml"); got != "a.upstream.yaml" {
|
||||||
|
t.Fatalf("UpstreamPath(.yml) = %q", got)
|
||||||
|
}
|
||||||
|
if got := UpstreamPath("a"); got != "a.upstream.yaml" {
|
||||||
|
t.Fatalf("UpstreamPath(no ext) = %q", got)
|
||||||
|
}
|
||||||
|
s, err := LoadUpstream("testdata/upstream/post_json.upstream.yaml")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(s.Exchanges) != 1 || s.Exchanges[0].Response.Status != 201 {
|
||||||
|
t.Fatalf("sidecar = %+v", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user