diff --git a/cmd/summer/main.go b/cmd/summer/main.go index 0a43fab..af01a6f 100644 --- a/cmd/summer/main.go +++ b/cmd/summer/main.go @@ -29,6 +29,8 @@ func toolCommands() []bonfire.Command { makePluginCommand(), addPluginCommand(), devCommand(), + parityRecordCommand(), + parityReplayCommand(), } } diff --git a/cmd/summer/parity.go b/cmd/summer/parity.go new file mode 100644 index 0000000..052926a --- /dev/null +++ b/cmd/summer/parity.go @@ -0,0 +1,95 @@ +package main + +import ( + "context" + "fmt" + "strings" + + "git.golem15.com/golem15/summercms/bonfire" + "git.golem15.com/golem15/summercms/tide" +) + +func parityRecordCommand() bonfire.Command { + return bonfire.Command{ + Name: "parity:record", + Description: "Record HTTP responses for a one-flow YAML spec", + Flags: []bonfire.Flag{ + {Name: "spec", Description: "YAML request spec path"}, + {Name: "target", Description: "Base URL of the HTTP backend"}, + {Name: "output", Description: "Destination fixture path"}, + }, + Run: runParityRecord, + } +} + +func parityReplayCommand() bonfire.Command { + return bonfire.Command{ + Name: "parity:replay", + Description: "Replay recorded fixtures against an HTTP backend", + Flags: []bonfire.Flag{ + {Name: "fixtures", Description: "Recorded YAML fixture path"}, + {Name: "target", Description: "Base URL of the HTTP backend"}, + }, + Run: runParityReplay, + } +} + +func runParityRecord(ctx context.Context, in bonfire.Input, out bonfire.Output) error { + specPath, err := requireFlag(in, "spec", "parity:record") + if err != nil { + return err + } + target, err := requireFlag(in, "target", "parity:record") + if err != nil { + return err + } + output, err := requireFlag(in, "output", "parity:record") + if err != nil { + return err + } + spec, err := tide.LoadFlow(specPath) + if err != nil { + return err + } + flow, err := tide.RecordFlow(ctx, spec, tide.RecordConfig{Target: target}) + if err != nil { + return err + } + if err := tide.SaveFlow(output, flow); err != nil { + return err + } + out.Success(fmt.Sprintf("recorded %s", output)) + return nil +} + +func runParityReplay(ctx context.Context, in bonfire.Input, out bonfire.Output) error { + path, err := requireFlag(in, "fixtures", "parity:replay") + if err != nil { + return err + } + target, err := requireFlag(in, "target", "parity:replay") + if err != nil { + return err + } + flow, err := tide.LoadFlow(path) + if err != nil { + return err + } + result, err := tide.ReplayFlow(ctx, flow, tide.ReplayConfig{Target: target}) + if err != nil { + return err + } + if result.OK { + out.Success("replay matched") + } + return nil +} + +func requireFlag(in bonfire.Input, name, cmd string) (string, error) { + v, ok := in.Flag(name) + v = strings.TrimSpace(v) + if !ok || v == "" { + return "", fmt.Errorf("%s requires --%s", cmd, name) + } + return v, nil +} diff --git a/cmd/summer/parity_test.go b/cmd/summer/parity_test.go new file mode 100644 index 0000000..a9ce9f5 --- /dev/null +++ b/cmd/summer/parity_test.go @@ -0,0 +1,120 @@ +package main + +import ( + "bytes" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + + "git.golem15.com/golem15/summercms/bonfire" +) + +func TestParityCommands(t *testing.T) { + names := commandNames() + for _, want := range []string{"parity:record", "parity:replay"} { + if !containsName(names, want) { + t.Fatalf("missing %s in %v", want, names) + } + } + + spec := filepath.Join("..", "..", "tide", "testdata", "one-route-spec.yaml") + outDir := t.TempDir() + fixture := filepath.Join(outDir, "sample.yaml") + + origJSON := httptest.NewServer(jsonHandler(`{"data":"ok"}`)) + t.Cleanup(origJSON.Close) + changedJSON := httptest.NewServer(jsonHandler(`{"data":"no"}`)) + t.Cleanup(changedJSON.Close) + origPlain := httptest.NewServer(plainHandler("hello")) + t.Cleanup(origPlain.Close) + changedPlain := httptest.NewServer(plainHandler("hallo")) + t.Cleanup(changedPlain.Close) + + if err := runParity("parity:record", "--spec", spec, "--target", origJSON.URL, "--output", fixture); err != nil { + t.Fatalf("record: %v", err) + } + raw, err := os.ReadFile(fixture) + if err != nil { + t.Fatal(err) + } + if !bytes.Contains(raw, []byte(`{"data":"ok"}`)) { + t.Fatalf("recorded fixture missing body:\n%s", raw) + } + + if err := runParity("parity:replay", "--fixtures", fixture, "--target", origJSON.URL); err != nil { + t.Fatalf("replay identical JSON: %v", err) + } + + err = runParity("parity:replay", "--fixtures", fixture, "--target", changedJSON.URL) + if err == nil { + t.Fatal("changed JSON must fail") + } + msg := err.Error() + if !strings.Contains(msg, "$.data") { + t.Fatalf("JSON mismatch missing $.data: %s", msg) + } + if !strings.Contains(msg, "ok") || !strings.Contains(msg, "no") { + t.Fatalf("JSON mismatch missing expected/actual: %s", msg) + } + + plainFixture := filepath.Join(outDir, "plain.yaml") + if err := runParity("parity:record", "--spec", spec, "--target", origPlain.URL, "--output", plainFixture); err != nil { + t.Fatalf("record plain: %v", err) + } + if err := runParity("parity:replay", "--fixtures", plainFixture, "--target", origPlain.URL); err != nil { + t.Fatalf("replay identical plain: %v", err) + } + err = runParity("parity:replay", "--fixtures", plainFixture, "--target", changedPlain.URL) + if err == nil { + t.Fatal("changed bytes must fail") + } + if !strings.Contains(err.Error(), "1") { + t.Fatalf("byte mismatch missing offset: %v", err) + } +} + +func commandNames() []string { + var names []string + for _, c := range toolCommands() { + names = append(names, c.Name) + } + return names +} + +func containsName(names []string, want string) bool { + for _, name := range names { + if name == want { + return true + } + } + return false +} + +func runParity(args ...string) error { + var buf bytes.Buffer + root, err := bonfire.NewRoot("summer", toolCommands(), &buf) + if err != nil { + return err + } + root.SetArgs(args) + return root.Execute() +} + +func jsonHandler(body string) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(body)) + }) +} + +func plainHandler(body string) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/plain") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(body)) + }) +} diff --git a/go.mod b/go.mod index 14f5927..ad9445e 100644 --- a/go.mod +++ b/go.mod @@ -6,6 +6,7 @@ toolchain go1.27.0 require ( github.com/fsnotify/fsnotify v1.10.1 + github.com/goccy/go-yaml v1.19.2 github.com/knadh/koanf/parsers/yaml v1.1.1 github.com/knadh/koanf/providers/confmap v1.0.1 github.com/knadh/koanf/providers/env/v2 v2.0.1 diff --git a/go.sum b/go.sum index 5df5ed1..5c30426 100644 --- a/go.sum +++ b/go.sum @@ -5,6 +5,8 @@ github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx5 github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo= github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9LvH92wZUgs= github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= +github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM= +github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/knadh/koanf/maps v0.1.2 h1:RBfmAW5CnZT+PJ1CVc1QSJKf4Xu9kxfQgYVQSu8hpbo= diff --git a/tide/diff.go b/tide/diff.go new file mode 100644 index 0000000..1f90872 --- /dev/null +++ b/tide/diff.go @@ -0,0 +1,220 @@ +package tide + +import ( + "bytes" + "encoding/json" + "fmt" + "mime" + "strconv" + "strings" + "unicode" + "unicode/utf8" +) + +func compareBodies(want, got Response) []Diff { + wantJSON := isJSONContentType(want.Headers) + gotJSON := isJSONContentType(got.Headers) + if wantJSON && gotJSON { + return diffJSON([]byte(want.Body), []byte(got.Body)) + } + return diffBytes([]byte(want.Body), []byte(got.Body)) +} + +func isJSONContentType(headers map[string]string) bool { + ct := headerValue(headers, "Content-Type") + if ct == "" { + return false + } + media, _, err := mime.ParseMediaType(ct) + if err != nil { + return strings.Contains(strings.ToLower(ct), "json") + } + return media == "application/json" || strings.HasSuffix(media, "+json") +} + +func headerValue(headers map[string]string, name string) string { + if v, ok := headers[name]; ok { + return v + } + for k, v := range headers { + if strings.EqualFold(k, name) { + return v + } + } + return "" +} + +func diffJSON(want, got []byte) []Diff { + wantVal, err := decodeJSON(want) + if err != nil { + return []Diff{{Path: "$", Expected: "valid JSON", Actual: err.Error()}} + } + gotVal, err := decodeJSON(got) + if err != nil { + return []Diff{{Path: "$", Expected: formatValue(wantVal), Actual: err.Error()}} + } + var diffs []Diff + compareValue("$", wantVal, gotVal, &diffs) + return diffs +} + +func decodeJSON(raw []byte) (any, error) { + dec := json.NewDecoder(bytes.NewReader(raw)) + dec.UseNumber() + var v any + if err := dec.Decode(&v); err != nil { + return nil, err + } + return v, nil +} + +func compareValue(path string, want, got any, diffs *[]Diff) { + switch w := want.(type) { + case map[string]any: + g, ok := got.(map[string]any) + if !ok { + *diffs = append(*diffs, Diff{Path: path, Expected: formatValue(want), Actual: formatValue(got)}) + return + } + for k, wv := range w { + gv, exists := g[k] + child := pathJoin(path, k) + if !exists { + *diffs = append(*diffs, Diff{Path: child, Expected: formatValue(wv), Actual: ""}) + continue + } + compareValue(child, wv, gv, diffs) + } + for k, gv := range g { + if _, exists := w[k]; !exists { + *diffs = append(*diffs, Diff{Path: pathJoin(path, k), Expected: "", Actual: formatValue(gv)}) + } + } + case []any: + g, ok := got.([]any) + if !ok { + *diffs = append(*diffs, Diff{Path: path, Expected: formatValue(want), Actual: formatValue(got)}) + return + } + if len(w) != len(g) { + *diffs = append(*diffs, Diff{ + Path: path, + Expected: fmt.Sprintf("array[%d]", len(w)), + Actual: fmt.Sprintf("array[%d]", len(g)), + }) + } + n := min(len(w), len(g)) + for i := 0; i < n; i++ { + compareValue(fmt.Sprintf("%s[%d]", path, i), w[i], g[i], diffs) + } + default: + if !scalarEqual(want, got) { + *diffs = append(*diffs, Diff{Path: path, Expected: formatValue(want), Actual: formatValue(got)}) + } + } +} + +func scalarEqual(want, got any) bool { + if want == nil || got == nil { + return want == nil && got == nil + } + switch w := want.(type) { + case json.Number: + g, ok := got.(json.Number) + return ok && w == g + case string: + g, ok := got.(string) + return ok && w == g + case bool: + g, ok := got.(bool) + return ok && w == g + default: + return false + } +} + +func formatValue(v any) string { + switch t := v.(type) { + case nil: + return "null" + case json.Number: + return "number " + string(t) + case string: + return strconv.Quote(t) + case bool: + return fmt.Sprintf("%v", t) + case map[string]any: + return "object" + case []any: + return fmt.Sprintf("array[%d]", len(t)) + default: + return fmt.Sprintf("%v", t) + } +} + +func pathJoin(parent, key string) string { + if parent == "$" { + return "$." + key + } + return parent + "." + key +} + +func diffBytes(want, got []byte) []Diff { + n := min(len(want), len(got)) + off := n + for i := 0; i < n; i++ { + if want[i] != got[i] { + off = i + break + } + } + if off == n && len(want) == len(got) { + return nil + } + return []Diff{{ + Path: fmt.Sprintf("body[%d]", off), + Expected: printableWindow(want, off), + Actual: printableWindow(got, off), + Offset: off, + Byte: true, + }} +} + +func printableWindow(b []byte, off int) string { + if len(b) == 0 { + return `""` + } + start := off - 8 + if start < 0 { + start = 0 + } + end := off + 8 + if end > len(b) { + end = len(b) + } + return quotePrintable(b[start:end]) +} + +func quotePrintable(b []byte) string { + var buf strings.Builder + buf.WriteByte('"') + for i := 0; i < len(b); { + r, size := utf8.DecodeRune(b[i:]) + if r == utf8.RuneError && size == 1 { + fmt.Fprintf(&buf, "\\x%02x", b[i]) + i++ + continue + } + if r == '\\' || r == '"' { + buf.WriteByte('\\') + buf.WriteRune(r) + } else if unicode.IsPrint(r) { + buf.WriteRune(r) + } else { + fmt.Fprintf(&buf, "\\u%04x", r) + } + i += size + } + buf.WriteByte('"') + return buf.String() +} diff --git a/tide/fixture.go b/tide/fixture.go new file mode 100644 index 0000000..b659ede --- /dev/null +++ b/tide/fixture.go @@ -0,0 +1,229 @@ +package tide + +import ( + "bytes" + "fmt" + "os" + "path/filepath" + "sort" + "strconv" + "strings" + + "github.com/goccy/go-yaml" + "github.com/goccy/go-yaml/token" +) + +func (b *Body) UnmarshalYAML(data []byte) error { + if b == nil { + return fmt.Errorf("tide: nil body") + } + var s string + if err := yaml.Unmarshal(data, &s); err != nil { + return err + } + *b = Body(s) + return nil +} + +// LoadFlow reads and validates a version-1 YAML flow from path. +func LoadFlow(path string) (Flow, error) { + raw, err := os.ReadFile(path) + if err != nil { + return Flow{}, fmt.Errorf("tide: read %s: %w", path, err) + } + return ParseFlow(raw) +} + +// ParseFlow decodes a version-1 YAML flow, rejecting unknown fields. +func ParseFlow(raw []byte) (Flow, error) { + var flow Flow + dec := yaml.NewDecoder(bytes.NewReader(raw), yaml.DisallowUnknownField()) + if err := dec.Decode(&flow); err != nil { + return Flow{}, fmt.Errorf("tide: parse flow: %w", err) + } + if err := validateFlow(flow); err != nil { + return Flow{}, err + } + return flow, nil +} + +// SaveFlow writes a validated flow atomically, syncing before rename. +func SaveFlow(path string, flow Flow) error { + if err := validateFlow(flow); err != nil { + return err + } + raw, err := marshalFlow(flow) + if err != nil { + return err + } + dir := filepath.Dir(path) + if dir != "" && dir != "." { + if err := os.MkdirAll(dir, 0o755); err != nil { + return fmt.Errorf("tide: create fixture dir: %w", err) + } + } + tmp, err := os.CreateTemp(dir, ".tide-*.tmp") + if err != nil { + return fmt.Errorf("tide: create temp fixture: %w", err) + } + tmpName := tmp.Name() + ok := false + defer func() { + if !ok { + _ = os.Remove(tmpName) + } + }() + if _, err := tmp.Write(raw); err != nil { + _ = tmp.Close() + return fmt.Errorf("tide: write fixture: %w", err) + } + if err := tmp.Sync(); err != nil { + _ = tmp.Close() + return fmt.Errorf("tide: sync fixture: %w", err) + } + if err := tmp.Close(); err != nil { + return fmt.Errorf("tide: close fixture: %w", err) + } + if err := os.Rename(tmpName, path); err != nil { + return fmt.Errorf("tide: commit fixture: %w", err) + } + ok = true + return nil +} + +func marshalFlow(flow Flow) ([]byte, error) { + var b strings.Builder + fmt.Fprintf(&b, "version: %d\n", flow.Version) + writeKV(&b, 0, "name", flow.Name) + if flow.Description != "" { + writeKV(&b, 0, "description", flow.Description) + } + if flow.SeedHook != "" { + writeKV(&b, 0, "seed_hook", flow.SeedHook) + } + b.WriteString("steps:\n") + for _, step := range flow.Steps { + writeStep(&b, step) + } + return []byte(b.String()), nil +} + +func writeStep(b *strings.Builder, step Step) { + fmt.Fprintf(b, " - id: %s\n", encodeScalar(step.ID)) + if step.RouteID != "" { + writeKV(b, 4, "route_id", step.RouteID) + } + b.WriteString(" request:\n") + writeRequest(b, 6, step.Request) + b.WriteString(" response:\n") + writeResponse(b, 6, step.Response) + writeCapture(b, 4, step.Capture) + writeNormalize(b, 4, step.Normalize) + writeHeaders(b, 4, step.Headers) +} + +func writeRequest(b *strings.Builder, indent int, req Request) { + writeKV(b, indent, "method", req.Method) + writeKV(b, indent, "path", req.Path) + if req.Query != "" { + writeKV(b, indent, "query", req.Query) + } + writeHeaders(b, indent, req.Headers) + writeBody(b, indent, string(req.Body)) +} + +func writeResponse(b *strings.Builder, indent int, resp Response) { + if resp.Status != 0 { + fmt.Fprintf(b, "%sstatus: %d\n", strings.Repeat(" ", indent), resp.Status) + } + writeHeaders(b, indent, resp.Headers) + writeBody(b, indent, string(resp.Body)) + if resp.BodyFile != "" { + writeKV(b, indent, "body_file", resp.BodyFile) + } + if resp.SHA256 != "" { + writeKV(b, indent, "sha256", resp.SHA256) + } +} + +func writeCapture(b *strings.Builder, indent int, rules []CaptureRule) { + if len(rules) == 0 { + return + } + pad := strings.Repeat(" ", indent) + fmt.Fprintf(b, "%scapture:\n", pad) + inner := strings.Repeat(" ", indent+2) + for _, rule := range rules { + fmt.Fprintf(b, "%s- path: %s\n", inner, encodeScalar(rule.Path)) + fmt.Fprintf(b, "%s as: %s\n", inner, encodeScalar(rule.As)) + } +} + +func writeNormalize(b *strings.Builder, indent int, rules []NormalizeRule) { + if len(rules) == 0 { + return + } + pad := strings.Repeat(" ", indent) + fmt.Fprintf(b, "%snormalize:\n", pad) + inner := strings.Repeat(" ", indent+2) + for _, rule := range rules { + fmt.Fprintf(b, "%s- path: %s\n", inner, encodeScalar(rule.Path)) + if rule.Disable { + fmt.Fprintf(b, "%s disable: true\n", inner) + } + } +} + +func writeHeaders(b *strings.Builder, indent int, headers map[string]string) { + if len(headers) == 0 { + return + } + pad := strings.Repeat(" ", indent) + fmt.Fprintf(b, "%sheaders:\n", pad) + keys := make([]string, 0, len(headers)) + for k := range headers { + keys = append(keys, k) + } + sort.Strings(keys) + inner := strings.Repeat(" ", indent+2) + for _, k := range keys { + fmt.Fprintf(b, "%s%s: %s\n", inner, k, encodeScalar(headers[k])) + } +} + +func writeBody(b *strings.Builder, indent int, body string) { + if body == "" { + return + } + pad := strings.Repeat(" ", indent) + header := token.LiteralBlockHeader(body) + if header == "" { + header = "|-" + } + fmt.Fprintf(b, "%sbody: %s\n", pad, header) + inner := strings.Repeat(" ", indent+2) + content := body + if strings.HasSuffix(content, "\n") { + content = content[:len(content)-1] + } + for _, line := range strings.Split(content, "\n") { + b.WriteString(inner) + b.WriteString(line) + b.WriteByte('\n') + } +} + +func writeKV(b *strings.Builder, indent int, key, value string) { + pad := strings.Repeat(" ", indent) + fmt.Fprintf(b, "%s%s: %s\n", pad, key, encodeScalar(value)) +} + +func encodeScalar(v string) string { + if v == "" { + return `""` + } + if token.IsNeedQuoted(v) || strings.ContainsAny(v, " \t") { + return strconv.Quote(v) + } + return v +} diff --git a/tide/flow.go b/tide/flow.go new file mode 100644 index 0000000..df23968 --- /dev/null +++ b/tide/flow.go @@ -0,0 +1,183 @@ +package tide + +import ( + "fmt" + "net/http" + "path/filepath" + "strings" +) + +const CurrentVersion = 1 + +// DefaultMaxBody is the default cap on recorded or replayed HTTP bodies. +const DefaultMaxBody = 8 << 20 + +// Flow is a versioned ordered list of HTTP steps. +type Flow struct { + Version int `yaml:"version"` + Name string `yaml:"name"` + Description string `yaml:"description,omitempty"` + SeedHook string `yaml:"seed_hook,omitempty"` + Steps []Step `yaml:"steps"` +} + +// Step is one request/response pair in a flow. +type Step struct { + ID string `yaml:"id"` + RouteID string `yaml:"route_id,omitempty"` + Request Request `yaml:"request"` + Response Response `yaml:"response"` + Capture []CaptureRule `yaml:"capture,omitempty"` + Normalize []NormalizeRule `yaml:"normalize,omitempty"` + Headers map[string]string `yaml:"headers,omitempty"` +} + +// Request is the outbound HTTP call for a step. +type Request struct { + Method string `yaml:"method"` + Path string `yaml:"path"` + Query string `yaml:"query,omitempty"` + Headers map[string]string `yaml:"headers,omitempty"` + Body Body `yaml:"body,omitempty"` +} + +// Response is the recorded or expected HTTP reply. +type Response struct { + Status int `yaml:"status,omitempty"` + Headers map[string]string `yaml:"headers,omitempty"` + Body Body `yaml:"body,omitempty"` + BodyFile string `yaml:"body_file,omitempty"` + SHA256 string `yaml:"sha256,omitempty"` +} + +// CaptureRule maps a JSON path in the response onto a variable name. +type CaptureRule struct { + Path string `yaml:"path"` + As string `yaml:"as"` +} + +// NormalizeRule names a per-step normalizer override. +type NormalizeRule struct { + Path string `yaml:"path,omitempty"` + Disable bool `yaml:"disable,omitempty"` +} + +// Body is verbatim request or response bytes stored as a YAML literal scalar. +type Body string + +// RecordConfig injects the HTTP target, client and body bound for recording. +type RecordConfig struct { + Target string + Client *http.Client + MaxBody int64 +} + +// ReplayConfig injects the HTTP target, client and body bound for replay. +type ReplayConfig struct { + Target string + Client *http.Client + MaxBody int64 +} + +// Result is the outcome of replaying a flow. +type Result struct { + OK bool + Steps []StepResult +} + +// StepResult is the outcome of one replayed step. +type StepResult struct { + ID string + OK bool + Diffs []Diff +} + +// Diff is one structural JSON or raw-byte mismatch. +type Diff struct { + Path string + Expected string + Actual string + Offset int + Byte bool +} + +// MismatchError is returned when replay finds one or more differences. +type MismatchError struct { + Result Result +} + +func (e *MismatchError) Error() string { + if e == nil { + return "tide: mismatch" + } + var b strings.Builder + for _, step := range e.Result.Steps { + for _, d := range step.Diffs { + if b.Len() > 0 { + b.WriteByte('\n') + } + if d.Byte { + fmt.Fprintf(&b, "step %s: body mismatch at byte %d: expected %s actual %s", step.ID, d.Offset, d.Expected, d.Actual) + continue + } + fmt.Fprintf(&b, "step %s: %s: expected %s actual %s", step.ID, d.Path, d.Expected, d.Actual) + } + } + if b.Len() == 0 { + return "tide: mismatch" + } + return b.String() +} + +func validateFlow(flow Flow) error { + if flow.Version != CurrentVersion { + return fmt.Errorf("tide: unsupported version %d (want %d)", flow.Version, CurrentVersion) + } + if strings.TrimSpace(flow.Name) == "" { + return fmt.Errorf("tide: flow name is required") + } + if len(flow.Steps) == 0 { + return fmt.Errorf("tide: flow %q has no steps", flow.Name) + } + seen := make(map[string]struct{}, len(flow.Steps)) + for i, step := range flow.Steps { + if strings.TrimSpace(step.ID) == "" { + return fmt.Errorf("tide: steps[%d] is missing id", i) + } + if _, dup := seen[step.ID]; dup { + return fmt.Errorf("tide: duplicate step id %q", step.ID) + } + seen[step.ID] = struct{}{} + if strings.TrimSpace(step.Request.Method) == "" { + return fmt.Errorf("tide: step %s is missing request method", step.ID) + } + if strings.TrimSpace(step.Request.Path) == "" { + return fmt.Errorf("tide: step %s is missing request path", step.ID) + } + if err := validateSidecar(step.Response.BodyFile); err != nil { + return fmt.Errorf("tide: step %s: %w", step.ID, err) + } + } + return nil +} + +func validateSidecar(path string) error { + if path == "" { + return nil + } + if filepath.IsAbs(path) { + return fmt.Errorf("body_file %q must be a relative path", path) + } + clean := filepath.ToSlash(filepath.Clean(path)) + if clean == ".." || strings.HasPrefix(clean, "../") { + return fmt.Errorf("body_file %q escapes the fixture directory", path) + } + return nil +} + +func maxBody(n int64) int64 { + if n <= 0 { + return DefaultMaxBody + } + return n +} diff --git a/tide/record.go b/tide/record.go new file mode 100644 index 0000000..6bd9802 --- /dev/null +++ b/tide/record.go @@ -0,0 +1,116 @@ +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, step := range spec.Steps { + resp, err := doStep(ctx, client, cfg.Target, step.Request, limit) + if err != nil { + return Flow{}, fmt.Errorf("tide: record step %s: %w", step.ID, err) + } + out.Steps[i].Response = resp + } + 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 + } + 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 { + ct := h.Get("Content-Type") + if ct == "" { + return nil + } + return map[string]string{"Content-Type": ct} +} + +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 +} diff --git a/tide/replay.go b/tide/replay.go new file mode 100644 index 0000000..194c894 --- /dev/null +++ b/tide/replay.go @@ -0,0 +1,48 @@ +package tide + +import ( + "context" + "fmt" + "strings" +) + +// ReplayFlow executes each recorded step against target and diffs responses. +func ReplayFlow(ctx context.Context, flow Flow, cfg ReplayConfig) (Result, error) { + if err := validateFlow(flow); err != nil { + return Result{}, err + } + if strings.TrimSpace(cfg.Target) == "" { + return Result{}, fmt.Errorf("tide: replay target is required") + } + client := cfg.Client + if client == nil { + client = defaultClient() + } + limit := maxBody(cfg.MaxBody) + result := Result{OK: true, Steps: make([]StepResult, 0, len(flow.Steps))} + for _, step := range flow.Steps { + got, err := doStep(ctx, client, cfg.Target, step.Request, limit) + if err != nil { + return result, fmt.Errorf("tide: replay step %s: %w", step.ID, err) + } + sr := StepResult{ID: step.ID, OK: true} + if step.Response.Status != 0 && got.Status != step.Response.Status { + sr.OK = false + sr.Diffs = append(sr.Diffs, Diff{ + Path: "status", + Expected: fmt.Sprintf("%d", step.Response.Status), + Actual: fmt.Sprintf("%d", got.Status), + }) + } + sr.Diffs = append(sr.Diffs, compareBodies(step.Response, got)...) + if len(sr.Diffs) > 0 { + sr.OK = false + result.OK = false + } + result.Steps = append(result.Steps, sr) + } + if !result.OK { + return result, &MismatchError{Result: result} + } + return result, nil +} diff --git a/tide/roundtrip_test.go b/tide/roundtrip_test.go new file mode 100644 index 0000000..a4417d6 --- /dev/null +++ b/tide/roundtrip_test.go @@ -0,0 +1,118 @@ +package tide + +import ( + "context" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestParityRoundTrip(t *testing.T) { + ctx := context.Background() + specPath := filepath.Join("testdata", "one-route-spec.yaml") + spec, err := LoadFlow(specPath) + if err != nil { + t.Fatal(err) + } + + t.Run("json record replay and scalar mismatch", func(t *testing.T) { + orig := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/sample" || r.Method != http.MethodGet { + http.NotFound(w, r) + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"data":"ok"}`)) + })) + t.Cleanup(orig.Close) + + changed := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"data":"no"}`)) + })) + t.Cleanup(changed.Close) + + recorded, err := RecordFlow(ctx, spec, RecordConfig{Target: orig.URL}) + if err != nil { + t.Fatal(err) + } + out := filepath.Join(t.TempDir(), "sample.yaml") + if err := SaveFlow(out, recorded); err != nil { + t.Fatal(err) + } + raw, err := os.ReadFile(out) + if err != nil { + t.Fatal(err) + } + text := string(raw) + if !strings.Contains(text, "version: 1") { + t.Fatalf("missing version:\n%s", text) + } + if !strings.Contains(text, "body: |") && !strings.Contains(text, "body: |-") { + t.Fatalf("body is not a literal block scalar:\n%s", text) + } + if !strings.Contains(text, `{"data":"ok"}`) { + t.Fatalf("missing recorded JSON body:\n%s", text) + } + if !strings.Contains(text, "application/json") { + t.Fatalf("missing Content-Type:\n%s", text) + } + + loaded, err := LoadFlow(out) + if err != nil { + t.Fatal(err) + } + if _, err := ReplayFlow(ctx, loaded, ReplayConfig{Target: orig.URL}); err != nil { + t.Fatalf("identical backend must replay: %v", err) + } + + _, err = ReplayFlow(ctx, loaded, ReplayConfig{Target: changed.URL}) + if err == nil { + t.Fatal("changed JSON must fail") + } + msg := err.Error() + if !strings.Contains(msg, "$.data") { + t.Fatalf("mismatch missing $.data: %s", msg) + } + if !strings.Contains(msg, "ok") || !strings.Contains(msg, "no") { + t.Fatalf("mismatch missing expected/actual: %s", msg) + } + }) + + t.Run("non-json byte offset mismatch", func(t *testing.T) { + orig := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/plain") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("hello")) + })) + t.Cleanup(orig.Close) + + changed := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/plain") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("hallo")) + })) + t.Cleanup(changed.Close) + + recorded, err := RecordFlow(ctx, spec, RecordConfig{Target: orig.URL}) + if err != nil { + t.Fatal(err) + } + if _, err := ReplayFlow(ctx, recorded, ReplayConfig{Target: orig.URL}); err != nil { + t.Fatalf("identical plain body must replay: %v", err) + } + _, err = ReplayFlow(ctx, recorded, ReplayConfig{Target: changed.URL}) + if err == nil { + t.Fatal("changed bytes must fail") + } + msg := err.Error() + if !strings.Contains(msg, "1") { + t.Fatalf("byte mismatch missing offset: %s", msg) + } + }) +} diff --git a/tide/testdata/one-route-spec.yaml b/tide/testdata/one-route-spec.yaml new file mode 100644 index 0000000..a6ca5ac --- /dev/null +++ b/tide/testdata/one-route-spec.yaml @@ -0,0 +1,9 @@ +version: 1 +name: one-route-sample +description: Record a single GET /sample request +steps: + - id: sample + route_id: GET /sample + request: + method: GET + path: /sample