Resolve named placeholders from a private variable store, mask dates and ids after shape checks, and keep comparing independent steps. Co-authored-by: Cursor <cursoragent@cursor.com>
148 lines
3.7 KiB
Go
148 lines
3.7 KiB
Go
package tide
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"net/url"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
var errTruncated = errors.New("body truncated")
|
|
|
|
// RecordFlow executes each spec step against target and returns a complete flow.
|
|
func RecordFlow(ctx context.Context, spec Flow, cfg RecordConfig) (Flow, error) {
|
|
if err := validateFlow(spec); err != nil {
|
|
return Flow{}, err
|
|
}
|
|
if strings.TrimSpace(cfg.Target) == "" {
|
|
return Flow{}, fmt.Errorf("tide: record target is required")
|
|
}
|
|
client := cfg.Client
|
|
if client == nil {
|
|
client = defaultClient()
|
|
}
|
|
limit := maxBody(cfg.MaxBody)
|
|
out := spec
|
|
out.Version = CurrentVersion
|
|
out.Steps = make([]Step, len(spec.Steps))
|
|
copy(out.Steps, spec.Steps)
|
|
for i := range out.Steps {
|
|
step := &out.Steps[i]
|
|
mergeRouteCaptures(step, cfg.Rules)
|
|
req, err := expandRequest(step.Request, cfg.Store)
|
|
if err != nil {
|
|
return Flow{}, fmt.Errorf("tide: record step %s: %w", step.ID, err)
|
|
}
|
|
resp, err := doStep(ctx, client, cfg.Target, req, limit)
|
|
if err != nil {
|
|
return Flow{}, fmt.Errorf("tide: record step %s: %w", step.ID, err)
|
|
}
|
|
step.Response = resp
|
|
if err := CaptureStep(cfg.Store, step); err != nil {
|
|
return Flow{}, fmt.Errorf("tide: record step %s: %w", step.ID, err)
|
|
}
|
|
if err := ScrubStep(cfg.Store, step); err != nil {
|
|
return Flow{}, fmt.Errorf("tide: record step %s: %w", step.ID, err)
|
|
}
|
|
if err := cfg.Store.Save(); err != nil {
|
|
return Flow{}, err
|
|
}
|
|
}
|
|
if err := validateFlow(out); err != nil {
|
|
return Flow{}, err
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
func defaultClient() *http.Client {
|
|
return &http.Client{
|
|
Timeout: 30 * time.Second,
|
|
CheckRedirect: func(*http.Request, []*http.Request) error {
|
|
return http.ErrUseLastResponse
|
|
},
|
|
}
|
|
}
|
|
|
|
func doStep(ctx context.Context, client *http.Client, target string, req Request, limit int64) (Response, error) {
|
|
rawURL, err := joinURL(target, req.Path, req.Query)
|
|
if err != nil {
|
|
return Response{}, err
|
|
}
|
|
if int64(len(req.Body)) > limit {
|
|
return Response{}, fmt.Errorf("%w: request body exceeds %d bytes", errTruncated, limit)
|
|
}
|
|
var body io.Reader
|
|
if req.Body != "" {
|
|
body = strings.NewReader(string(req.Body))
|
|
}
|
|
httpReq, err := http.NewRequestWithContext(ctx, req.Method, rawURL, body)
|
|
if err != nil {
|
|
return Response{}, err
|
|
}
|
|
for k, v := range req.Headers {
|
|
httpReq.Header.Set(k, v)
|
|
}
|
|
httpResp, err := client.Do(httpReq)
|
|
if err != nil {
|
|
return Response{}, err
|
|
}
|
|
defer httpResp.Body.Close()
|
|
raw, err := readBounded(httpResp.Body, limit)
|
|
if err != nil {
|
|
return Response{}, err
|
|
}
|
|
return Response{
|
|
Status: httpResp.StatusCode,
|
|
Headers: keepResponseHeaders(httpResp.Header),
|
|
Body: Body(raw),
|
|
}, nil
|
|
}
|
|
|
|
func joinURL(target, path, rawQuery string) (string, error) {
|
|
base, err := url.Parse(target)
|
|
if err != nil {
|
|
return "", fmt.Errorf("target: %w", err)
|
|
}
|
|
if base.Scheme == "" || base.Host == "" {
|
|
return "", fmt.Errorf("target %q must be an absolute URL", target)
|
|
}
|
|
ref := &url.URL{Path: path, RawQuery: rawQuery}
|
|
return base.ResolveReference(ref).String(), nil
|
|
}
|
|
|
|
func keepResponseHeaders(h http.Header) map[string]string {
|
|
return filterHeaders(h, recordedHeaderNames)
|
|
}
|
|
|
|
var recordedHeaderNames = []string{
|
|
"Content-Type",
|
|
"Location",
|
|
"Set-Cookie",
|
|
"Cache-Control",
|
|
"Pragma",
|
|
"WWW-Authenticate",
|
|
"Content-Disposition",
|
|
"Access-Control-Allow-Origin",
|
|
"Access-Control-Allow-Credentials",
|
|
"Access-Control-Allow-Headers",
|
|
"Access-Control-Allow-Methods",
|
|
"Access-Control-Expose-Headers",
|
|
"X-Total-Count",
|
|
"Link",
|
|
}
|
|
|
|
func readBounded(r io.Reader, max int64) ([]byte, error) {
|
|
data, err := io.ReadAll(io.LimitReader(r, max+1))
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if int64(len(data)) > max {
|
|
return nil, fmt.Errorf("%w: exceeds %d bytes", errTruncated, max)
|
|
}
|
|
return data, nil
|
|
}
|