From 93b7142059df9c640d96af46e374a7eacb88baa8 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Sat, 3 Oct 2026 19:39:58 +0200 Subject: [PATCH] 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 --- modules/fetchguard/client.go | 197 +++++++++++ modules/fetchguard/client_test.go | 143 ++++++++ modules/fetchguard/fetch.go | 55 ++-- .../testdata/upstream/post_json.upstream.yaml | 17 + modules/tide/upstream.go | 305 ++++++++++++++++++ modules/tide/upstream_test.go | 150 +++++++++ 6 files changed, 848 insertions(+), 19 deletions(-) create mode 100644 modules/fetchguard/client.go create mode 100644 modules/fetchguard/client_test.go create mode 100644 modules/tide/testdata/upstream/post_json.upstream.yaml create mode 100644 modules/tide/upstream.go create mode 100644 modules/tide/upstream_test.go diff --git a/modules/fetchguard/client.go b/modules/fetchguard/client.go new file mode 100644 index 0000000..502ae36 --- /dev/null +++ b/modules/fetchguard/client.go @@ -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() } diff --git a/modules/fetchguard/client_test.go b/modules/fetchguard/client_test.go new file mode 100644 index 0000000..6d3f117 --- /dev/null +++ b/modules/fetchguard/client_test.go @@ -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) } diff --git a/modules/fetchguard/fetch.go b/modules/fetchguard/fetch.go index b9c2a76..821749a 100644 --- a/modules/fetchguard/fetch.go +++ b/modules/fetchguard/fetch.go @@ -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 diff --git a/modules/tide/testdata/upstream/post_json.upstream.yaml b/modules/tide/testdata/upstream/post_json.upstream.yaml new file mode 100644 index 0000000..901ee07 --- /dev/null +++ b/modules/tide/testdata/upstream/post_json.upstream.yaml @@ -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"}' diff --git a/modules/tide/upstream.go b/modules/tide/upstream.go new file mode 100644 index 0000000..7fc6530 --- /dev/null +++ b/modules/tide/upstream.go @@ -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 .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() +} diff --git a/modules/tide/upstream_test.go b/modules/tide/upstream_test.go new file mode 100644 index 0000000..e743a27 --- /dev/null +++ b/modules/tide/upstream_test.go @@ -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) + } +}