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) } req, err = prepareRequest(req, cfg.BaseDir) 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 }