refactor(10.2-01): nest framework packages under modules
- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
This commit is contained in:
512
modules/tide/capture_test.go
Normal file
512
modules/tide/capture_test.go
Normal file
@@ -0,0 +1,512 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
const testJWT = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiJhbGljZSJ9.signaturehere123456"
|
||||
|
||||
func TestCaptureAndPlaceholderResolution(t *testing.T) {
|
||||
store, err := OpenStore("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
step := Step{
|
||||
ID: "login",
|
||||
Response: Response{
|
||||
Headers: map[string]string{"Content-Type": "application/json"},
|
||||
Body: Body(`{"token":"` + testJWT + `"}`),
|
||||
},
|
||||
Capture: []CaptureRule{{
|
||||
From: "response.json",
|
||||
Path: "$.token",
|
||||
As: "jwt:alice",
|
||||
Identity: "alice",
|
||||
Category: "jwt",
|
||||
}},
|
||||
}
|
||||
if err := CaptureStep(store, &step); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ScrubStep(store, &step); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(string(step.Response.Body), "{{jwt:alice}}") {
|
||||
t.Fatalf("body not scrubbed: %s", step.Response.Body)
|
||||
}
|
||||
if strings.Contains(string(step.Response.Body), testJWT) {
|
||||
t.Fatal("jwt left in body")
|
||||
}
|
||||
got, err := store.Expand("Bearer {{jwt:alice}}")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got != "Bearer "+testJWT {
|
||||
t.Fatalf("expand: %q", got)
|
||||
}
|
||||
if _, err := store.Expand("{{missing}}"); err == nil || !strings.Contains(err.Error(), "unresolved") {
|
||||
t.Fatalf("missing placeholder: %v", err)
|
||||
}
|
||||
|
||||
escaped := Step{
|
||||
ID: "consent",
|
||||
Response: Response{
|
||||
Body: Body(`{"data":{"redirect_to":"http:\/\/127.0.0.1:8424\/oauth\/callback?code=oauthCode99"}}`),
|
||||
},
|
||||
Capture: []CaptureRule{{
|
||||
From: "response.json",
|
||||
Path: "$.data.redirect_to",
|
||||
As: "oauth:redirect",
|
||||
Category: "oauth_code",
|
||||
}},
|
||||
}
|
||||
if err := CaptureStep(store, &escaped); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ScrubStep(store, &escaped); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(string(escaped.Response.Body), "oauthCode99") {
|
||||
t.Fatalf("php-escaped redirect still has code: %s", escaped.Response.Body)
|
||||
}
|
||||
if !strings.Contains(string(escaped.Response.Body), "{{oauth:redirect}}") {
|
||||
t.Fatalf("php-escaped redirect not placeholder: %s", escaped.Response.Body)
|
||||
}
|
||||
codeStep := Step{
|
||||
ID: "consent-code",
|
||||
Response: Response{
|
||||
Body: Body(`{"data":{"redirect_to":"http://127.0.0.1:8424/oauth/callback?code=oauthCodeFromJSON"}}`),
|
||||
},
|
||||
Capture: []CaptureRule{{
|
||||
From: "response.json.query",
|
||||
Path: "$.data.redirect_to",
|
||||
Name: "code",
|
||||
As: "oauth:code",
|
||||
Category: "oauth_code",
|
||||
}},
|
||||
}
|
||||
if err := CaptureStep(store, &codeStep); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
gotCode, ok := store.Get("oauth:code")
|
||||
if !ok || gotCode != "oauthCodeFromJSON" {
|
||||
t.Fatalf("json.query code: ok=%v val=%q", ok, gotCode)
|
||||
}
|
||||
|
||||
step2 := Step{
|
||||
ID: "pkce",
|
||||
Request: Request{
|
||||
Method: "POST",
|
||||
Path: "/authorize",
|
||||
Body: Body("code_verifier=pkceVerifierValue1&client_id=x"),
|
||||
},
|
||||
Response: Response{
|
||||
Status: 302,
|
||||
Headers: map[string]string{"Location": "/cb?code=oauthCode99"},
|
||||
},
|
||||
Capture: []CaptureRule{
|
||||
{From: "request.form", Name: "code_verifier", As: "pkce:alice", Category: "pkce"},
|
||||
{From: "response.location.query", Name: "code", As: "oauth:alice-code", Category: "oauth_code"},
|
||||
},
|
||||
}
|
||||
if err := CaptureStep(store, &step2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ScrubStep(store, &step2); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(string(step2.Request.Body), "{{pkce:alice}}") || strings.Contains(string(step2.Request.Body), "pkceVerifierValue1") {
|
||||
t.Fatalf("pkce not scrubbed: %s", step2.Request.Body)
|
||||
}
|
||||
if !strings.Contains(step2.Response.Headers["Location"], "{{oauth:alice-code}}") {
|
||||
t.Fatalf("code not scrubbed: %v", step2.Response.Headers)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScrubRejectsUnclassifiedAndPersistsVars(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
varsPath := filepath.Join(dir, "vars.yaml")
|
||||
store, err := OpenStore(varsPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
step := Step{
|
||||
ID: "login",
|
||||
Response: Response{
|
||||
Headers: map[string]string{"Content-Type": "application/json"},
|
||||
Body: Body(`{"token":"` + testJWT + `"}`),
|
||||
},
|
||||
Capture: []CaptureRule{{From: "response.json", Path: "$.token", As: "jwt:alice", Category: "jwt"}},
|
||||
}
|
||||
if err := CaptureStep(store, &step); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ScrubStep(store, &step); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := store.Save(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, err := os.ReadFile(varsPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(string(raw), testJWT) {
|
||||
t.Fatalf("vars missing secret:\n%s", raw)
|
||||
}
|
||||
st, err := os.Stat(varsPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if st.Mode().Perm() != 0o600 {
|
||||
t.Fatalf("mode %o", st.Mode().Perm())
|
||||
}
|
||||
unknown := Step{
|
||||
ID: "bad",
|
||||
Response: Response{Body: Body(`{"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiJib2IifQ.otherSignatureValue99"}`)},
|
||||
}
|
||||
if err := ScrubStep(store, &unknown); err == nil {
|
||||
t.Fatal("unclassified leftover jwt must fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyLoginTokenAndPKCESession(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
case r.Method == http.MethodPost && r.URL.Path == "/login":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"token":"` + testJWT + `"}`))
|
||||
case r.URL.Path == "/me":
|
||||
if r.Header.Get("Authorization") != "Bearer "+testJWT {
|
||||
http.Error(w, "no", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
case r.URL.Path == "/authorize":
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
if !strings.Contains(string(body), "code_verifier=pkceVerifierValue1") {
|
||||
http.Error(w, "missing verifier", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Location", "/cb?code=oauthCode99")
|
||||
w.WriteHeader(http.StatusFound)
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
|
||||
fixtures := t.TempDir()
|
||||
varsPath := filepath.Join(t.TempDir(), "vars.yaml")
|
||||
rules := mustParseRules(t, ""+
|
||||
"client: nuxt\n"+
|
||||
"keep_request_headers:\n"+
|
||||
" - Authorization\n"+
|
||||
" - Content-Type\n"+
|
||||
"keep_response_headers:\n"+
|
||||
" - Content-Type\n"+
|
||||
" - Location\n"+
|
||||
"routes:\n"+
|
||||
" - method: POST\n"+
|
||||
" path: /login\n"+
|
||||
" capture:\n"+
|
||||
" - from: response.json\n"+
|
||||
" path: $.token\n"+
|
||||
" as: jwt:alice\n"+
|
||||
" identity: alice\n"+
|
||||
" category: jwt\n"+
|
||||
" - method: GET\n"+
|
||||
" path: /me\n"+
|
||||
" - method: POST\n"+
|
||||
" path: /authorize\n"+
|
||||
" capture:\n"+
|
||||
" - from: request.form\n"+
|
||||
" name: code_verifier\n"+
|
||||
" as: pkce:alice\n"+
|
||||
" category: pkce\n"+
|
||||
" - from: response.location.query\n"+
|
||||
" name: code\n"+
|
||||
" as: oauth:alice-code\n"+
|
||||
" category: oauth_code\n")
|
||||
proxy, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: upstream.URL,
|
||||
Fixtures: fixtures,
|
||||
VarsPath: varsPath,
|
||||
Rules: rules,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
post(t, srv.URL+"/login", "sess", "application/json", `{}`)
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/me", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, "sess")
|
||||
req.Header.Set("Authorization", "Bearer "+testJWT)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("me: %d", resp.StatusCode)
|
||||
}
|
||||
post(t, srv.URL+"/authorize", "sess", "application/x-www-form-urlencoded", "code_verifier=pkceVerifierValue1")
|
||||
if err := proxy.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
raw, err := os.ReadFile(filepath.Join(fixtures, "nuxt", "sess.yaml"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
text := string(raw)
|
||||
if strings.Contains(text, testJWT) || strings.Contains(text, "pkceVerifierValue1") || strings.Contains(text, "oauthCode99") {
|
||||
t.Fatalf("secrets in fixture:\n%s", text)
|
||||
}
|
||||
if !strings.Contains(text, "{{jwt:alice}}") || !strings.Contains(text, "{{pkce:alice}}") || !strings.Contains(text, "{{oauth:alice-code}}") {
|
||||
t.Fatalf("placeholders missing:\n%s", text)
|
||||
}
|
||||
|
||||
flow, err := LoadFlow(filepath.Join(fixtures, "nuxt", "sess.yaml"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
store, err := OpenStore(varsPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := ReplayFlow(context.Background(), flow, ReplayConfig{Target: upstream.URL, Store: store}); err != nil {
|
||||
t.Fatalf("replay captured session: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func post(t *testing.T, url, session, ct, body string) {
|
||||
t.Helper()
|
||||
req, err := http.NewRequest(http.MethodPost, url, strings.NewReader(body))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, session)
|
||||
req.Header.Set("Content-Type", ct)
|
||||
client := &http.Client{
|
||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode >= 400 {
|
||||
t.Fatalf("%s: %d", url, resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplayExpandsUnquotedIDPlaceholdersAfterCapture(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/login":
|
||||
_, _ = w.Write([]byte(`{"token":"` + testJWT + `","user":{"id":1}}`))
|
||||
case "/genres":
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":1,"name":"Rock","album_count":0}]}`))
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
flow := Flow{
|
||||
Version: 1,
|
||||
Name: "seed-then-list",
|
||||
Steps: []Step{
|
||||
{
|
||||
ID: "login",
|
||||
Request: Request{Method: http.MethodPost, Path: "/login"},
|
||||
Response: Response{
|
||||
Status: 200,
|
||||
Headers: jsonCT(),
|
||||
Body: Body(`{"token":"{{jwt:alice}}","user":{"id":{{id:alice}}}}`),
|
||||
},
|
||||
Capture: []CaptureRule{
|
||||
{From: "response.json", Path: "$.token", As: "jwt:alice", Category: "jwt"},
|
||||
{From: "response.json", Path: "$.user.id", As: "id:alice"},
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "genres",
|
||||
Request: Request{Method: http.MethodGet, Path: "/genres"},
|
||||
Response: Response{
|
||||
Status: 200,
|
||||
Headers: jsonCT(),
|
||||
Body: Body(`{"data":[{"id":{{id:alice}},"name":"Rock","album_count":0}]}`),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
if _, err := ReplayFlow(context.Background(), flow, ReplayConfig{Target: srv.URL}); err != nil {
|
||||
t.Fatalf("replay with unquoted id placeholders: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScrubShortNumericIDsDoNotCorruptPaths(t *testing.T) {
|
||||
store, err := OpenStore("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
store.Set("id:alice", "1")
|
||||
store.Set("id:genre", "4")
|
||||
step := Step{
|
||||
ID: "genres",
|
||||
Request: Request{
|
||||
Method: http.MethodGet,
|
||||
Path: "/_fonoteka/api/v1/genres/1",
|
||||
Headers: map[string]string{"Authorization": "Bearer x"},
|
||||
},
|
||||
Response: Response{
|
||||
Status: 200,
|
||||
Body: Body(`{"data":[{"id":1,"name":"Rock","album_count":0},{"id":15,"name":"Latin"},{"id":4,"name":"Jazz","album_count":1}]}`),
|
||||
},
|
||||
}
|
||||
if err := ScrubStep(store, &step); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if step.Request.Path != "/_fonoteka/api/v1/genres/{{id:alice}}" {
|
||||
t.Fatalf("path id not scrubbed: %s", step.Request.Path)
|
||||
}
|
||||
body := string(step.Response.Body)
|
||||
if body != `{"data":[{"id":1,"name":"Rock","album_count":0},{"id":15,"name":"Latin"},{"id":4,"name":"Jazz","album_count":1}]}` {
|
||||
t.Fatalf("JSON body short ids must stay literal: %s", body)
|
||||
}
|
||||
if strings.Contains(string(step.Request.Path), "v{{id:alice}}") {
|
||||
t.Fatalf("substring replace leaked: %s", step.Request.Path)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScrubShortIDsLeavePaginationAndIPv4Literal(t *testing.T) {
|
||||
store, err := OpenStore("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
store.Set("id:album", "1")
|
||||
step := Step{
|
||||
ID: "page",
|
||||
Request: Request{Method: http.MethodGet, Path: "/albums/1"},
|
||||
Response: Response{
|
||||
Status: 200,
|
||||
Body: Body(`{"id":1,"total":1,"host":"127.0.0.1"}`),
|
||||
},
|
||||
}
|
||||
if err := ScrubStep(store, &step); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if step.Request.Path != "/albums/{{id:album}}" {
|
||||
t.Fatalf("path id not scrubbed: %s", step.Request.Path)
|
||||
}
|
||||
body := string(step.Response.Body)
|
||||
if body != `{"id":1,"total":1,"host":"127.0.0.1"}` {
|
||||
t.Fatalf("pagination/IPv4 rewritten: %s", body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScrubRejectsUnknownPasswordKeepsAllowlist(t *testing.T) {
|
||||
store, err := OpenStore("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
allowed := Step{
|
||||
ID: "login",
|
||||
Request: Request{Method: http.MethodPost, Path: "/login", Body: Body(`{"email":"alice@parity.test","password":"parity-alice-pass"}`)},
|
||||
Response: Response{Status: 200, Body: Body(`{"ok":true}`)},
|
||||
}
|
||||
if err := ScrubStep(store, &allowed); err != nil {
|
||||
t.Fatalf("allow-listed test password: %v", err)
|
||||
}
|
||||
leaked := Step{
|
||||
ID: "login",
|
||||
Request: Request{Method: http.MethodPost, Path: "/login", Body: Body(`{"email":"alice@parity.test","password":"hunter2-live"}`)},
|
||||
Response: Response{Status: 200, Body: Body(`{"ok":true}`)},
|
||||
}
|
||||
if err := ScrubStep(store, &leaked); err == nil || !strings.Contains(err.Error(), "password") {
|
||||
t.Fatalf("unknown password: %v", err)
|
||||
}
|
||||
opaque := Step{
|
||||
ID: "tok",
|
||||
Response: Response{Status: 200, Body: Body(`{"access_token":"not-a-jwt-but-secret"}`)},
|
||||
}
|
||||
if err := ScrubStep(store, &opaque); err == nil || !strings.Contains(err.Error(), "access_token") {
|
||||
t.Fatalf("opaque access_token: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMismatchErrorRedactsJWT(t *testing.T) {
|
||||
err := &MismatchError{Result: Result{Steps: []StepResult{{
|
||||
ID: "t",
|
||||
Diffs: []Diff{{
|
||||
Path: "$.token",
|
||||
Expected: testJWT,
|
||||
Actual: testJWT + "x",
|
||||
}},
|
||||
}}}}
|
||||
msg := err.Error()
|
||||
if strings.Contains(msg, "eyJ") {
|
||||
t.Fatalf("mismatch leaked jwt: %s", msg)
|
||||
}
|
||||
if !strings.Contains(msg, "<redacted-jwt>") {
|
||||
t.Fatalf("expected redaction: %s", msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestScrubFormFieldDespiteSubstringSecrets(t *testing.T) {
|
||||
store, err := OpenStore("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
code := "overlapSECRET99"
|
||||
verifier := "xx" + code + "yyPKCEverifierValue"
|
||||
store.Set("oauth:code", code)
|
||||
step := Step{
|
||||
ID: "20",
|
||||
Request: Request{
|
||||
Method: "POST",
|
||||
Path: "/oauth/mcp/token",
|
||||
Body: Body("grant_type=authorization_code&code=" + code + "&code_verifier=" + verifier),
|
||||
},
|
||||
Response: Response{Status: 200, Body: Body(`{"ok":true}`)},
|
||||
Capture: []CaptureRule{
|
||||
{From: "request.form", Name: "code_verifier", As: "pkce:mcp", Category: "pkce"},
|
||||
{From: "request.form", Name: "code", As: "oauth:code", Category: "oauth_code"},
|
||||
},
|
||||
}
|
||||
if err := CaptureStep(store, &step); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ScrubStep(store, &step); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := string(step.Request.Body)
|
||||
if strings.Contains(got, verifier) || strings.Contains(got, code) {
|
||||
t.Fatalf("form still live: %s", got)
|
||||
}
|
||||
if !strings.Contains(got, "code_verifier={{pkce:mcp}}") {
|
||||
t.Fatalf("verifier placeholder: %s", got)
|
||||
}
|
||||
if !strings.Contains(got, "code={{oauth:code}}") {
|
||||
t.Fatalf("code placeholder: %s", got)
|
||||
}
|
||||
}
|
||||
288
modules/tide/diff.go
Normal file
288
modules/tide/diff.go
Normal file
@@ -0,0 +1,288 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"mime"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
func compareBodies(want, got Response, step Step) []Diff {
|
||||
wantJSON := isJSONContentType(want.Headers)
|
||||
gotJSON := isJSONContentType(got.Headers)
|
||||
if wantJSON && gotJSON {
|
||||
wantRaw, wantDiffs := normalizeJSON([]byte(want.Body), step)
|
||||
gotRaw, gotDiffs := normalizeJSON([]byte(got.Body), step)
|
||||
diffs := append([]Diff{}, wantDiffs...)
|
||||
diffs = append(diffs, gotDiffs...)
|
||||
diffs = append(diffs, diffJSON(wantRaw, gotRaw)...)
|
||||
return diffs
|
||||
}
|
||||
return diffBytes([]byte(want.Body), []byte(got.Body))
|
||||
}
|
||||
|
||||
var globalCompareHeaders = []string{
|
||||
"Content-Type",
|
||||
"X-Total-Count",
|
||||
"Access-Control-Allow-Origin",
|
||||
"Access-Control-Allow-Credentials",
|
||||
"Access-Control-Allow-Headers",
|
||||
"Access-Control-Allow-Methods",
|
||||
"Access-Control-Expose-Headers",
|
||||
}
|
||||
|
||||
var extraCompareHeaders = []string{
|
||||
"Cache-Control",
|
||||
"Pragma",
|
||||
"WWW-Authenticate",
|
||||
"Content-Disposition",
|
||||
"Link",
|
||||
"Location",
|
||||
}
|
||||
|
||||
var neverCompareHeaders = []string{
|
||||
"Date",
|
||||
"Server",
|
||||
"X-Request-Id",
|
||||
"X-Request-ID",
|
||||
"X-Correlation-Id",
|
||||
}
|
||||
|
||||
func compareHeaders(want, got, extra map[string]string) []Diff {
|
||||
allow := map[string]bool{}
|
||||
for _, n := range globalCompareHeaders {
|
||||
allow[http.CanonicalHeaderKey(n)] = true
|
||||
}
|
||||
for _, n := range extraCompareHeaders {
|
||||
allow[http.CanonicalHeaderKey(n)] = true
|
||||
}
|
||||
for n := range extra {
|
||||
allow[http.CanonicalHeaderKey(n)] = true
|
||||
}
|
||||
for _, n := range neverCompareHeaders {
|
||||
delete(allow, http.CanonicalHeaderKey(n))
|
||||
}
|
||||
var diffs []Diff
|
||||
for k, wv := range want {
|
||||
ck := http.CanonicalHeaderKey(k)
|
||||
if !allow[ck] {
|
||||
continue
|
||||
}
|
||||
gv := headerValue(got, k)
|
||||
if gv == wv {
|
||||
continue
|
||||
}
|
||||
if gv == "" {
|
||||
gv = "<missing>"
|
||||
}
|
||||
diffs = append(diffs, Diff{Path: "header." + ck, Expected: wv, Actual: gv})
|
||||
}
|
||||
return diffs
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
if dec.More() {
|
||||
return nil, fmt.Errorf("trailing JSON after first value")
|
||||
}
|
||||
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: "<missing>"})
|
||||
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: "<missing>", 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()
|
||||
}
|
||||
189
modules/tide/diff_contract_test.go
Normal file
189
modules/tide/diff_contract_test.go
Normal file
@@ -0,0 +1,189 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDiffContract(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
baseline := `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`
|
||||
reordered := `{"meta":{"total":1,"per_page":15,"last_page":1,"current_page":1},"data":[{"deleted_at":null,"created_at":"2026-01-01T00:00:00+00:00","is_owner":true,"tracklist":[],"price":"1.5000","slug":"keep-me","id":9}]}`
|
||||
|
||||
spec := Flow{Version: 1, Name: "diff-contract", Steps: []Step{{
|
||||
ID: "get",
|
||||
Request: Request{Method: http.MethodGet, Path: "/item"},
|
||||
}}}
|
||||
orig := jsonServer(t, baseline)
|
||||
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("baseline must pass: %v", err)
|
||||
}
|
||||
if _, err := ReplayFlow(ctx, recorded, ReplayConfig{Target: jsonServer(t, reordered).URL}); err != nil {
|
||||
t.Fatalf("object key reorder must pass: %v", err)
|
||||
}
|
||||
|
||||
type mutation struct {
|
||||
name string
|
||||
live string
|
||||
path string
|
||||
}
|
||||
cases := []mutation{
|
||||
{"null vs array", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":null,"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "tracklist"},
|
||||
{"empty object vs array", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":{},"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "tracklist"},
|
||||
{"carbon Z vs +00:00", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00Z","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "created_at"},
|
||||
{"non-date text", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"yesterday","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "created_at"},
|
||||
{"date present-null vs absent", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00"}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "deleted_at"},
|
||||
{"tri-state true vs false", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":false,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "is_owner"},
|
||||
{"tri-state true vs null", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":null,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "is_owner"},
|
||||
{"missing meta envelope", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}]}`, "meta"},
|
||||
{"extra links envelope key", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1},"links":{}}`, "links"},
|
||||
{"conditional reservation key", `{"data":[{"id":1,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null,"reservation":{"id":2}}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "reservation"},
|
||||
{"money string vs number", `{"data":[{"id":1,"slug":"keep-me","price":1.5,"tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "price"},
|
||||
{"integer id vs string", `{"data":[{"id":"1","slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "id"},
|
||||
{"integer id vs fraction", `{"data":[{"id":1.5,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "id"},
|
||||
{"exact slug mismatch", `{"data":[{"id":1,"slug":"other","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-01-01T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`, "slug"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
mut := cloneRecorded(t, recorded)
|
||||
mut.Steps[0].Response.Body = Body(tc.live)
|
||||
_, err := ReplayFlow(ctx, mut, ReplayConfig{Target: orig.URL})
|
||||
if err == nil {
|
||||
t.Fatal("mutated fixture must fail")
|
||||
}
|
||||
assertPathMismatch(t, err, tc.path)
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("csv exact-byte mismatch", func(t *testing.T) {
|
||||
want := csvServer(t, "a,b\n1,2\n", `attachment; filename="albums.csv"`)
|
||||
got := csvServer(t, "a,b\n1,3\n", `attachment; filename="albums.csv"`)
|
||||
spec := Flow{Version: 1, Name: "csv", Steps: []Step{{
|
||||
ID: "export",
|
||||
Request: Request{Method: http.MethodGet, Path: "/export"},
|
||||
}}}
|
||||
rec, err := RecordFlow(ctx, spec, RecordConfig{Target: want.URL})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := ReplayFlow(ctx, rec, ReplayConfig{Target: want.URL}); err != nil {
|
||||
t.Fatalf("csv baseline must pass: %v", err)
|
||||
}
|
||||
_, err = ReplayFlow(ctx, rec, ReplayConfig{Target: got.URL})
|
||||
if err == nil {
|
||||
t.Fatal("csv byte mismatch must fail")
|
||||
}
|
||||
msg := err.Error()
|
||||
if !strings.Contains(msg, "byte") && !strings.Contains(msg, "body[") {
|
||||
t.Fatalf("csv mismatch missing byte diagnostic: %s", msg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("image exact-byte mismatch", func(t *testing.T) {
|
||||
want := binaryServer(t, "image/png", []byte{0x89, 'P', 'N', 'G', 1})
|
||||
got := binaryServer(t, "image/png", []byte{0x89, 'P', 'N', 'G', 2})
|
||||
spec := Flow{Version: 1, Name: "png", Steps: []Step{{
|
||||
ID: "cover",
|
||||
Request: Request{Method: http.MethodGet, Path: "/cover"},
|
||||
}}}
|
||||
rec, err := RecordFlow(ctx, spec, RecordConfig{Target: want.URL})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := ReplayFlow(ctx, rec, ReplayConfig{Target: want.URL}); err != nil {
|
||||
t.Fatalf("image baseline must pass: %v", err)
|
||||
}
|
||||
_, err = ReplayFlow(ctx, rec, ReplayConfig{Target: got.URL})
|
||||
if err == nil {
|
||||
t.Fatal("image byte mismatch must fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "byte") && !strings.Contains(err.Error(), "body[") {
|
||||
t.Fatalf("image mismatch missing byte diagnostic: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("per-step normalization disable", func(t *testing.T) {
|
||||
later := jsonServer(t, `{"data":[{"id":9,"slug":"keep-me","price":"1.5000","tracklist":[],"is_owner":true,"created_at":"2026-02-02T00:00:00+00:00","deleted_at":null}],"meta":{"current_page":1,"last_page":1,"per_page":15,"total":1}}`)
|
||||
if _, err := ReplayFlow(ctx, recorded, ReplayConfig{Target: later.URL}); err != nil {
|
||||
t.Fatalf("masked dates/ids must pass: %v", err)
|
||||
}
|
||||
disabled := cloneRecorded(t, recorded)
|
||||
disabled.Steps[0].Normalize = []NormalizeRule{{Path: "created_at", Disable: true}}
|
||||
_, err := ReplayFlow(ctx, disabled, ReplayConfig{Target: later.URL})
|
||||
if err == nil {
|
||||
t.Fatal("disabled date mask must fail")
|
||||
}
|
||||
assertPathMismatch(t, err, "created_at")
|
||||
})
|
||||
}
|
||||
|
||||
func cloneRecorded(t *testing.T, flow Flow) Flow {
|
||||
t.Helper()
|
||||
path := filepath.Join(t.TempDir(), "clone.yaml")
|
||||
if err := SaveFlow(path, flow); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cloned, err := LoadFlow(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return cloned
|
||||
}
|
||||
|
||||
func assertPathMismatch(t *testing.T, err error, path string) {
|
||||
t.Helper()
|
||||
msg := err.Error()
|
||||
if !strings.Contains(msg, path) {
|
||||
t.Fatalf("want path %s in %s", path, msg)
|
||||
}
|
||||
if !strings.Contains(msg, "expected") || !strings.Contains(msg, "actual") {
|
||||
t.Fatalf("want expected/actual diagnostic in %s", msg)
|
||||
}
|
||||
var mis *MismatchError
|
||||
if !errors.As(err, &mis) {
|
||||
return
|
||||
}
|
||||
found := false
|
||||
for _, step := range mis.Result.Steps {
|
||||
for _, d := range step.Diffs {
|
||||
if strings.Contains(d.Path, path) && d.Expected != d.Actual {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("result diffs missing path %s: %+v", path, mis.Result.Steps)
|
||||
}
|
||||
}
|
||||
|
||||
func csvServer(t *testing.T, body, disposition string) *httptest.Server {
|
||||
t.Helper()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/csv")
|
||||
w.Header().Set("Content-Disposition", disposition)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(body))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
return srv
|
||||
}
|
||||
|
||||
func binaryServer(t *testing.T, ct string, body []byte) *httptest.Server {
|
||||
t.Helper()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", ct)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(body)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
return srv
|
||||
}
|
||||
49
modules/tide/diff_test.go
Normal file
49
modules/tide/diff_test.go
Normal file
@@ -0,0 +1,49 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDiffParityClasses(t *testing.T) {
|
||||
step := Step{ID: "d"}
|
||||
ct := jsonCT()
|
||||
check := func(name, want, got, path string) {
|
||||
t.Helper()
|
||||
diffs := compareBodies(Response{Headers: ct, Body: Body(want)}, Response{Headers: ct, Body: Body(got)}, step)
|
||||
if len(diffs) == 0 {
|
||||
t.Fatalf("%s: expected mismatch", name)
|
||||
}
|
||||
found := false
|
||||
for _, d := range diffs {
|
||||
if strings.Contains(d.Path, path) && d.Expected != d.Actual {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("%s: want path %s in %+v", name, path, diffs)
|
||||
}
|
||||
}
|
||||
check("null vs array", `{"tracklist":[]}`, `{"tracklist":null}`, "tracklist")
|
||||
check("money string vs number", `{"price":"1.5000"}`, `{"price":1.5}`, "price")
|
||||
check("null vs absent", `{"deleted_at":null}`, `{}`, "deleted_at")
|
||||
check("tri-state bool", `{"is_owner":null}`, `{"is_owner":false}`, "is_owner")
|
||||
check("conditional key", `{"data":{"name":"a"}}`, `{"data":{"name":"a","extra":1}}`, "extra")
|
||||
check("date vs Z", `{"created_at":"2026-01-01T00:00:00+00:00"}`, `{"created_at":"2026-01-01T00:00:00Z"}`, "created_at")
|
||||
}
|
||||
|
||||
func TestDecodeJSONRejectsTrailingValue(t *testing.T) {
|
||||
if _, err := decodeJSON([]byte(`{"data":[]}`)); err != nil {
|
||||
t.Fatalf("single value: %v", err)
|
||||
}
|
||||
if _, err := decodeJSON([]byte(`{"data":[]}{"debug":true}`)); err == nil || !strings.Contains(err.Error(), "trailing JSON") {
|
||||
t.Fatalf("trailing value: %v", err)
|
||||
}
|
||||
want := Response{Headers: jsonCT(), Body: Body(`{"data":[]}{"debug":true}`)}
|
||||
got := Response{Headers: jsonCT(), Body: Body(`{"data":[]}`)}
|
||||
diffs := compareBodies(want, got, Step{ID: "trail"})
|
||||
if len(diffs) == 0 {
|
||||
t.Fatal("trailing JSON envelope must mismatch")
|
||||
}
|
||||
}
|
||||
276
modules/tide/fixture.go
Normal file
276
modules/tide/fixture.go
Normal file
@@ -0,0 +1,276 @@
|
||||
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 {
|
||||
return saveFlow(path, flow, false)
|
||||
}
|
||||
|
||||
// SaveFlowExclusive writes a validated flow and fails if path already exists.
|
||||
func SaveFlowExclusive(path string, flow Flow) error {
|
||||
return saveFlow(path, flow, true)
|
||||
}
|
||||
|
||||
func saveFlow(path string, flow Flow, exclusive bool) 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)
|
||||
}
|
||||
}
|
||||
if exclusive {
|
||||
f, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: create fixture: %w", err)
|
||||
}
|
||||
ok := false
|
||||
defer func() {
|
||||
_ = f.Close()
|
||||
if !ok {
|
||||
_ = os.Remove(path)
|
||||
}
|
||||
}()
|
||||
if _, err := f.Write(raw); err != nil {
|
||||
return fmt.Errorf("tide: write fixture: %w", err)
|
||||
}
|
||||
if err := f.Sync(); err != nil {
|
||||
return fmt.Errorf("tide: sync fixture: %w", err)
|
||||
}
|
||||
if err := f.Close(); err != nil {
|
||||
return fmt.Errorf("tide: close fixture: %w", err)
|
||||
}
|
||||
ok = true
|
||||
return nil
|
||||
}
|
||||
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- as: %s\n", inner, encodeScalar(rule.As))
|
||||
if rule.From != "" {
|
||||
fmt.Fprintf(b, "%s from: %s\n", inner, encodeScalar(rule.From))
|
||||
}
|
||||
if rule.Path != "" {
|
||||
fmt.Fprintf(b, "%s path: %s\n", inner, encodeScalar(rule.Path))
|
||||
}
|
||||
if rule.Name != "" {
|
||||
fmt.Fprintf(b, "%s name: %s\n", inner, encodeScalar(rule.Name))
|
||||
}
|
||||
if rule.Identity != "" {
|
||||
fmt.Fprintf(b, "%s identity: %s\n", inner, encodeScalar(rule.Identity))
|
||||
}
|
||||
if rule.Category != "" {
|
||||
fmt.Fprintf(b, "%s category: %s\n", inner, encodeScalar(rule.Category))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
235
modules/tide/flow.go
Normal file
235
modules/tide/flow.go
Normal file
@@ -0,0 +1,235 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"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 named source onto a variable name.
|
||||
type CaptureRule struct {
|
||||
From string `yaml:"from,omitempty"`
|
||||
Path string `yaml:"path,omitempty"`
|
||||
Name string `yaml:"name,omitempty"`
|
||||
As string `yaml:"as"`
|
||||
Identity string `yaml:"identity,omitempty"`
|
||||
Category string `yaml:"category,omitempty"`
|
||||
}
|
||||
|
||||
// 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
|
||||
Store *Store
|
||||
Rules Rules
|
||||
}
|
||||
|
||||
// ReplayConfig injects the HTTP target, client and body bound for replay.
|
||||
type ReplayConfig struct {
|
||||
Target string
|
||||
Client *http.Client
|
||||
MaxBody int64
|
||||
Store *Store
|
||||
BaseDir string
|
||||
}
|
||||
|
||||
// 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
|
||||
Skipped 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, redactSecrets(d.Expected), redactSecrets(d.Actual))
|
||||
continue
|
||||
}
|
||||
fmt.Fprintf(&b, "step %s: %s: expected %s actual %s", step.ID, d.Path, redactSecrets(d.Expected), redactSecrets(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 materializeSidecar(base string, resp *Response) error {
|
||||
if resp == nil || resp.BodyFile == "" {
|
||||
return nil
|
||||
}
|
||||
if err := validateSidecar(resp.BodyFile); err != nil {
|
||||
return err
|
||||
}
|
||||
path := resp.BodyFile
|
||||
if base != "" {
|
||||
path = filepath.Join(base, resp.BodyFile)
|
||||
}
|
||||
resolved, err := resolvePath(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: read body_file %s: %w", resp.BodyFile, err)
|
||||
}
|
||||
if base != "" {
|
||||
root, err := resolvePath(base)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: fixture dir: %w", err)
|
||||
}
|
||||
if resolved != root && !strings.HasPrefix(resolved, root+string(os.PathSeparator)) {
|
||||
return fmt.Errorf("tide: body_file %q escapes the fixture directory", resp.BodyFile)
|
||||
}
|
||||
}
|
||||
raw, err := os.ReadFile(resolved)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: read body_file %s: %w", resp.BodyFile, err)
|
||||
}
|
||||
sum := sha256.Sum256(raw)
|
||||
got := hex.EncodeToString(sum[:])
|
||||
if resp.SHA256 != "" && !strings.EqualFold(got, resp.SHA256) {
|
||||
return fmt.Errorf("tide: body_file %s digest mismatch", resp.BodyFile)
|
||||
}
|
||||
resp.Body = Body(raw)
|
||||
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)
|
||||
}
|
||||
if _, err := resolvePath(path); err != nil {
|
||||
return fmt.Errorf("body_file %q: %w", path, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func maxBody(n int64) int64 {
|
||||
if n <= 0 {
|
||||
return DefaultMaxBody
|
||||
}
|
||||
return n
|
||||
}
|
||||
364
modules/tide/flow_contract_test.go
Normal file
364
modules/tide/flow_contract_test.go
Normal file
@@ -0,0 +1,364 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestFlowContract(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("record replay baseline", func(t *testing.T) {
|
||||
srv := jsonServer(t, `{"ok":true}`)
|
||||
spec := Flow{Version: 1, Name: "baseline", Steps: []Step{{
|
||||
ID: "a",
|
||||
Request: Request{Method: http.MethodGet, Path: "/ok"},
|
||||
}}}
|
||||
rec, err := RecordFlow(ctx, spec, RecordConfig{Target: srv.URL})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if rec.Steps[0].Response.Status != http.StatusOK {
|
||||
t.Fatalf("status %d", rec.Steps[0].Response.Status)
|
||||
}
|
||||
if _, err := ReplayFlow(ctx, rec, ReplayConfig{Target: srv.URL}); err != nil {
|
||||
t.Fatalf("baseline replay: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("oauth and 401 headers plus ignored date server request id", func(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Pragma", "no-cache")
|
||||
w.Header().Set("WWW-Authenticate", `Bearer error="invalid_token", resource_metadata="http://127.0.0.1/.well-known/oauth-protected-resource"`)
|
||||
w.Header().Set("Server", "php")
|
||||
w.Header().Set("X-Request-Id", "live-req")
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
_, _ = w.Write([]byte(`{"error":"invalid_token"}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
spec := Flow{Version: 1, Name: "oauth-401", Steps: []Step{{
|
||||
ID: "token",
|
||||
Request: Request{Method: http.MethodGet, Path: "/token"},
|
||||
}}}
|
||||
rec, err := RecordFlow(ctx, spec, RecordConfig{Target: srv.URL})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rec.Steps[0].Response.Headers["Date"] = "Wed, 01 Jan 2020 00:00:00 GMT"
|
||||
rec.Steps[0].Response.Headers["Server"] = "old"
|
||||
rec.Steps[0].Response.Headers["X-Request-Id"] = "fixture-req"
|
||||
if _, err := ReplayFlow(ctx, rec, ReplayConfig{Target: srv.URL}); err != nil {
|
||||
t.Fatalf("Date/Server/request id must be ignored: %v", err)
|
||||
}
|
||||
|
||||
badCache := cloneRecorded(t, rec)
|
||||
badCache.Steps[0].Response.Headers["Cache-Control"] = "public"
|
||||
_, err = ReplayFlow(ctx, badCache, ReplayConfig{Target: srv.URL})
|
||||
if err == nil || !strings.Contains(strings.ToLower(err.Error()), "cache-control") {
|
||||
t.Fatalf("Cache-Control mismatch: %v", err)
|
||||
}
|
||||
|
||||
badPragma := cloneRecorded(t, rec)
|
||||
badPragma.Steps[0].Response.Headers["Pragma"] = "public"
|
||||
_, err = ReplayFlow(ctx, badPragma, ReplayConfig{Target: srv.URL})
|
||||
if err == nil || !strings.Contains(strings.ToLower(err.Error()), "pragma") {
|
||||
t.Fatalf("Pragma mismatch: %v", err)
|
||||
}
|
||||
|
||||
badWWW := cloneRecorded(t, rec)
|
||||
badWWW.Steps[0].Response.Headers["WWW-Authenticate"] = `Bearer error="other"`
|
||||
_, err = ReplayFlow(ctx, badWWW, ReplayConfig{Target: srv.URL})
|
||||
if err == nil || !strings.Contains(strings.ToLower(err.Error()), "www-authenticate") {
|
||||
t.Fatalf("WWW-Authenticate mismatch: %v", err)
|
||||
}
|
||||
|
||||
badStatus := cloneRecorded(t, rec)
|
||||
badStatus.Steps[0].Response.Status = http.StatusOK
|
||||
_, err = ReplayFlow(ctx, badStatus, ReplayConfig{Target: srv.URL})
|
||||
if err == nil || !strings.Contains(err.Error(), "status") {
|
||||
t.Fatalf("status mismatch: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("csv content-disposition mismatch", func(t *testing.T) {
|
||||
srv := csvServer(t, "a,b\n1,2\n", `attachment; filename="albums.csv"`)
|
||||
spec := Flow{Version: 1, Name: "csv-disp", Steps: []Step{{
|
||||
ID: "export",
|
||||
Request: Request{Method: http.MethodGet, Path: "/export"},
|
||||
}}}
|
||||
rec, err := RecordFlow(ctx, spec, RecordConfig{Target: srv.URL})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := ReplayFlow(ctx, rec, ReplayConfig{Target: srv.URL}); err != nil {
|
||||
t.Fatalf("csv header baseline: %v", err)
|
||||
}
|
||||
bad := cloneRecorded(t, rec)
|
||||
bad.Steps[0].Response.Headers["Content-Disposition"] = `attachment; filename="other.csv"`
|
||||
_, err = ReplayFlow(ctx, bad, ReplayConfig{Target: srv.URL})
|
||||
if err == nil || !strings.Contains(strings.ToLower(err.Error()), "content-disposition") {
|
||||
t.Fatalf("Content-Disposition mismatch: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("sidecar digest and path", func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
payload := []byte("hello-bin")
|
||||
sum := sha256.Sum256(payload)
|
||||
if err := os.WriteFile(filepath.Join(dir, "data.bin"), payload, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
unsafe := "version: 1\nname: bin\nsteps:\n - id: a\n request:\n method: GET\n path: /bin\n response:\n body_file: ../secret.bin\n"
|
||||
if err := os.WriteFile(filepath.Join(dir, "unsafe.yaml"), []byte(unsafe), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := LoadFlow(filepath.Join(dir, "unsafe.yaml")); err == nil || !strings.Contains(err.Error(), "body_file") {
|
||||
t.Fatalf("parent sidecar must fail: %v", err)
|
||||
}
|
||||
abs := "version: 1\nname: bin\nsteps:\n - id: a\n request:\n method: GET\n path: /bin\n response:\n body_file: /tmp/secret.bin\n"
|
||||
if err := os.WriteFile(filepath.Join(dir, "abs.yaml"), []byte(abs), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := LoadFlow(filepath.Join(dir, "abs.yaml")); err == nil || !strings.Contains(err.Error(), "body_file") {
|
||||
t.Fatalf("absolute sidecar must fail: %v", err)
|
||||
}
|
||||
|
||||
srv := binaryServer(t, "application/octet-stream", payload)
|
||||
good := "version: 1\nname: bin\nsteps:\n - id: a\n request:\n method: GET\n path: /bin\n response:\n status: 200\n headers:\n Content-Type: application/octet-stream\n body_file: data.bin\n sha256: " + hex.EncodeToString(sum[:]) + "\n"
|
||||
gp := filepath.Join(dir, "good.yaml")
|
||||
if err := os.WriteFile(gp, []byte(good), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
flow, err := LoadFlow(gp)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := ReplayFlow(ctx, flow, ReplayConfig{Target: srv.URL, BaseDir: dir}); err != nil {
|
||||
t.Fatalf("matching sidecar digest must pass: %v", err)
|
||||
}
|
||||
bad := strings.Replace(good, hex.EncodeToString(sum[:]), strings.Repeat("ab", 32), 1)
|
||||
bp := filepath.Join(dir, "bad-digest.yaml")
|
||||
if err := os.WriteFile(bp, []byte(bad), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
flow, err = LoadFlow(bp)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = ReplayFlow(ctx, flow, ReplayConfig{Target: srv.URL, BaseDir: dir})
|
||||
if err == nil || !strings.Contains(err.Error(), "digest") {
|
||||
t.Fatalf("digest mismatch: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("comparison failure continues later steps", func(t *testing.T) {
|
||||
var hits atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/one":
|
||||
_, _ = w.Write([]byte(`{"data":"no"}`))
|
||||
default:
|
||||
_, _ = w.Write([]byte(`{"data":"ok"}`))
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
flow := Flow{
|
||||
Version: 1,
|
||||
Name: "mismatch-continue",
|
||||
Steps: []Step{
|
||||
{ID: "a", Request: Request{Method: http.MethodGet, Path: "/one"}, Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"data":"yes"}`)}},
|
||||
{ID: "b", Request: Request{Method: http.MethodGet, Path: "/two"}, Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"data":"ok"}`)}},
|
||||
},
|
||||
}
|
||||
res, err := ReplayFlow(ctx, flow, ReplayConfig{Target: srv.URL})
|
||||
if err == nil {
|
||||
t.Fatal("mismatch must error")
|
||||
}
|
||||
if hits.Load() != 2 {
|
||||
t.Fatalf("later step must run, hits=%d", hits.Load())
|
||||
}
|
||||
if len(res.Steps) != 2 || res.Steps[1].Skipped || !res.Steps[1].OK {
|
||||
t.Fatalf("step b should pass: %+v", res.Steps)
|
||||
}
|
||||
assertPathMismatch(t, err, "$.data")
|
||||
})
|
||||
|
||||
t.Run("capture failure skips remainder", func(t *testing.T) {
|
||||
var hits atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":"ok"}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
flow := Flow{
|
||||
Version: 1,
|
||||
Name: "capture-skip",
|
||||
Steps: []Step{
|
||||
{
|
||||
ID: "a",
|
||||
Request: Request{Method: http.MethodGet, Path: "/one"},
|
||||
Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"data":"ok"}`)},
|
||||
Capture: []CaptureRule{{From: "response.json", Path: "$.token", As: "jwt:alice"}},
|
||||
},
|
||||
{ID: "b", Request: Request{Method: http.MethodGet, Path: "/two"}, Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"data":"ok"}`)}},
|
||||
},
|
||||
}
|
||||
res, err := ReplayFlow(ctx, flow, ReplayConfig{Target: srv.URL})
|
||||
if err == nil {
|
||||
t.Fatal("failed capture must error")
|
||||
}
|
||||
if hits.Load() != 1 {
|
||||
t.Fatalf("capture fail should skip rest, hits=%d", hits.Load())
|
||||
}
|
||||
if len(res.Steps) != 2 || !res.Steps[1].Skipped {
|
||||
t.Fatalf("step b should skip: %+v", res.Steps)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "capture") {
|
||||
t.Fatalf("capture diagnostic: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing variable fails before send", func(t *testing.T) {
|
||||
var hits atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits.Add(1)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
flow := Flow{Version: 1, Name: "missing-var", Steps: []Step{
|
||||
{ID: "a", Request: Request{Method: http.MethodGet, Path: "/items/{{missing}}"}, Response: Response{Status: 200}},
|
||||
{ID: "b", Request: Request{Method: http.MethodGet, Path: "/later"}, Response: Response{Status: 200}},
|
||||
}}
|
||||
res, err := ReplayFlow(ctx, flow, ReplayConfig{Target: srv.URL})
|
||||
if err == nil || !strings.Contains(err.Error(), "unresolved") && !strings.Contains(err.Error(), "placeholder") {
|
||||
t.Fatalf("missing variable: %v", err)
|
||||
}
|
||||
if hits.Load() != 0 {
|
||||
t.Fatalf("must not send unresolved placeholder, hits=%d", hits.Load())
|
||||
}
|
||||
if len(res.Steps) != 2 || !res.Steps[1].Skipped {
|
||||
t.Fatalf("remainder must skip: %+v", res.Steps)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unknown capture source", func(t *testing.T) {
|
||||
store, err := OpenStore("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv := jsonServer(t, `{"ok":true}`)
|
||||
spec := Flow{Version: 1, Name: "bad-capture", Steps: []Step{{
|
||||
ID: "a",
|
||||
Request: Request{Method: http.MethodGet, Path: "/ok"},
|
||||
Capture: []CaptureRule{{From: "response.unknown", Path: "$.ok", As: "x"}},
|
||||
}}}
|
||||
_, err = RecordFlow(ctx, spec, RecordConfig{Target: srv.URL, Store: store})
|
||||
if err == nil || !strings.Contains(err.Error(), "unknown") {
|
||||
t.Fatalf("unknown capture from: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("capture-by-reference mismatch", func(t *testing.T) {
|
||||
store, err := OpenStore("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
n := 0
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/create":
|
||||
_, _ = w.Write([]byte(`{"token":"shareTokValue99"}`))
|
||||
case "/show":
|
||||
n++
|
||||
if n == 1 {
|
||||
_, _ = w.Write([]byte(`{"token":"shareTokValue99"}`))
|
||||
} else {
|
||||
_, _ = w.Write([]byte(`{"token":"shareTokOther00"}`))
|
||||
}
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
spec := Flow{
|
||||
Version: 1,
|
||||
Name: "capture-ref",
|
||||
Steps: []Step{
|
||||
{
|
||||
ID: "create",
|
||||
Request: Request{Method: http.MethodGet, Path: "/create"},
|
||||
Capture: []CaptureRule{{From: "response.json", Path: "$.token", As: "share:item"}},
|
||||
},
|
||||
{ID: "show", Request: Request{Method: http.MethodGet, Path: "/show"}},
|
||||
},
|
||||
}
|
||||
rec, err := RecordFlow(ctx, spec, RecordConfig{Target: srv.URL, Store: store})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(string(rec.Steps[0].Response.Body), "{{share:item}}") {
|
||||
t.Fatalf("create not scrubbed: %s", rec.Steps[0].Response.Body)
|
||||
}
|
||||
if !strings.Contains(string(rec.Steps[1].Response.Body), "{{share:item}}") {
|
||||
t.Fatalf("show not scrubbed by reference: %s", rec.Steps[1].Response.Body)
|
||||
}
|
||||
replayStore, err := OpenStore("")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = ReplayFlow(ctx, rec, ReplayConfig{Target: srv.URL, Store: replayStore})
|
||||
if err == nil {
|
||||
t.Fatal("capture-by-reference mismatch must fail")
|
||||
}
|
||||
assertPathMismatch(t, err, "token")
|
||||
})
|
||||
}
|
||||
|
||||
func TestFlowContractTwoFlowsContinue(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
var hits atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"v":2}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
first := Flow{Version: 1, Name: "first", Steps: []Step{{
|
||||
ID: "a", Request: Request{Method: http.MethodGet, Path: "/a"},
|
||||
Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"v":1}`)},
|
||||
}}}
|
||||
second := Flow{Version: 1, Name: "second", Steps: []Step{{
|
||||
ID: "b", Request: Request{Method: http.MethodGet, Path: "/b"},
|
||||
Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"v":2}`)},
|
||||
}}}
|
||||
_, err := ReplayFlow(ctx, first, ReplayConfig{Target: srv.URL})
|
||||
if err == nil {
|
||||
t.Fatal("first flow must fail")
|
||||
}
|
||||
if _, err := ReplayFlow(ctx, second, ReplayConfig{Target: srv.URL}); err != nil {
|
||||
t.Fatalf("later flow must still run: %v", err)
|
||||
}
|
||||
if hits.Load() != 2 {
|
||||
t.Fatalf("both flows must execute, hits=%d", hits.Load())
|
||||
}
|
||||
var mis *MismatchError
|
||||
if !errors.As(err, &mis) || mis.Result.OK {
|
||||
t.Fatalf("first flow mismatch result: %v", err)
|
||||
}
|
||||
}
|
||||
80
modules/tide/flow_test.go
Normal file
80
modules/tide/flow_test.go
Normal file
@@ -0,0 +1,80 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestFlowMismatchContinuesAndCaptureSkipsRest(t *testing.T) {
|
||||
var hits atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/one":
|
||||
_, _ = w.Write([]byte(`{"data":"no"}`))
|
||||
default:
|
||||
_, _ = w.Write([]byte(`{"data":"ok"}`))
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
mismatch := Flow{
|
||||
Version: 1,
|
||||
Name: "mismatch-continue",
|
||||
Steps: []Step{
|
||||
{
|
||||
ID: "a",
|
||||
Request: Request{Method: http.MethodGet, Path: "/one"},
|
||||
Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"data":"yes"}`)},
|
||||
},
|
||||
{
|
||||
ID: "b",
|
||||
Request: Request{Method: http.MethodGet, Path: "/two"},
|
||||
Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"data":"ok"}`)},
|
||||
},
|
||||
},
|
||||
}
|
||||
res, err := ReplayFlow(context.Background(), mismatch, ReplayConfig{Target: srv.URL})
|
||||
if err == nil {
|
||||
t.Fatal("mismatch must error")
|
||||
}
|
||||
if hits.Load() != 2 {
|
||||
t.Fatalf("mismatch should continue, hits=%d", hits.Load())
|
||||
}
|
||||
if len(res.Steps) != 2 || res.Steps[1].Skipped || !res.Steps[1].OK {
|
||||
t.Fatalf("step b should run and pass: %+v", res.Steps)
|
||||
}
|
||||
|
||||
hits.Store(0)
|
||||
captureFail := Flow{
|
||||
Version: 1,
|
||||
Name: "capture-skip",
|
||||
Steps: []Step{
|
||||
{
|
||||
ID: "a",
|
||||
Request: Request{Method: http.MethodGet, Path: "/two"},
|
||||
Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"data":"ok"}`)},
|
||||
Capture: []CaptureRule{{From: "response.json", Path: "$.token", As: "jwt:alice"}},
|
||||
},
|
||||
{
|
||||
ID: "b",
|
||||
Request: Request{Method: http.MethodGet, Path: "/two"},
|
||||
Response: Response{Status: 200, Headers: jsonCT(), Body: Body(`{"data":"ok"}`)},
|
||||
},
|
||||
},
|
||||
}
|
||||
res, err = ReplayFlow(context.Background(), captureFail, ReplayConfig{Target: srv.URL})
|
||||
if err == nil {
|
||||
t.Fatal("failed capture must error")
|
||||
}
|
||||
if hits.Load() != 1 {
|
||||
t.Fatalf("capture fail should skip rest, hits=%d", hits.Load())
|
||||
}
|
||||
if len(res.Steps) != 2 || !res.Steps[1].Skipped {
|
||||
t.Fatalf("step b should skip: %+v", res.Steps)
|
||||
}
|
||||
}
|
||||
99
modules/tide/headers_test.go
Normal file
99
modules/tide/headers_test.go
Normal file
@@ -0,0 +1,99 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestHeadersAllowListIgnoresDate(t *testing.T) {
|
||||
want := Response{
|
||||
Status: 401,
|
||||
Headers: map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
"WWW-Authenticate": `Bearer error="invalid_token"`,
|
||||
"Cache-Control": "no-store",
|
||||
"Date": "Wed, 01 Jan 2020 00:00:00 GMT",
|
||||
},
|
||||
Body: Body(`{"ok":false}`),
|
||||
}
|
||||
got := Response{
|
||||
Status: 401,
|
||||
Headers: map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
"WWW-Authenticate": `Bearer error="invalid_token"`,
|
||||
"Cache-Control": "no-store",
|
||||
"Date": "Thu, 02 Jan 2020 00:00:00 GMT",
|
||||
"Server": "php",
|
||||
},
|
||||
Body: Body(`{"ok":false}`),
|
||||
}
|
||||
diffs := compareStep(Step{ID: "h", Response: want, Headers: map[string]string{"WWW-Authenticate": "", "Cache-Control": ""}}, got)
|
||||
if len(diffs) != 0 {
|
||||
t.Fatalf("Date/Server must be ignored: %+v", diffs)
|
||||
}
|
||||
|
||||
got.Headers["WWW-Authenticate"] = `Bearer error="other"`
|
||||
diffs = compareStep(Step{ID: "h", Response: want}, got)
|
||||
if len(diffs) == 0 {
|
||||
t.Fatal("WWW-Authenticate mismatch must fail")
|
||||
}
|
||||
if !strings.Contains(strings.ToLower(diffs[0].Path), "www-authenticate") {
|
||||
t.Fatalf("path %s", diffs[0].Path)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeadersCompareLocation(t *testing.T) {
|
||||
want := Response{
|
||||
Status: 302,
|
||||
Headers: map[string]string{"Location": "http://127.0.0.1:8424/oauth/callback?code={{oauth:code}}"},
|
||||
}
|
||||
got := Response{
|
||||
Status: 302,
|
||||
Headers: map[string]string{"Location": "http://127.0.0.1:8424/oauth/callback?code={{oauth:code}}"},
|
||||
}
|
||||
if diffs := compareHeaders(want.Headers, got.Headers, nil); len(diffs) != 0 {
|
||||
t.Fatalf("identical Location must pass: %+v", diffs)
|
||||
}
|
||||
got.Headers["Location"] = "http://evil.example/oauth/callback?code={{oauth:code}}"
|
||||
diffs := compareHeaders(want.Headers, got.Headers, nil)
|
||||
if len(diffs) == 0 {
|
||||
t.Fatal("Location host mismatch must fail")
|
||||
}
|
||||
if !strings.Contains(strings.ToLower(diffs[0].Path), "location") {
|
||||
t.Fatalf("path %s", diffs[0].Path)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeadersReplayAgainstServer(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Pragma", "no-cache")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
flow := Flow{
|
||||
Version: 1,
|
||||
Name: "headers",
|
||||
Steps: []Step{{
|
||||
ID: "h",
|
||||
Request: Request{Method: http.MethodGet, Path: "/"},
|
||||
Response: Response{
|
||||
Status: 200,
|
||||
Headers: map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
"Cache-Control": "no-store",
|
||||
"Pragma": "no-cache",
|
||||
},
|
||||
Body: Body(`{"ok":true}`),
|
||||
},
|
||||
}},
|
||||
}
|
||||
if _, err := ReplayFlow(context.Background(), flow, ReplayConfig{Target: srv.URL}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
573
modules/tide/manifest.go
Normal file
573
modules/tide/manifest.go
Normal file
@@ -0,0 +1,573 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/goccy/go-yaml"
|
||||
)
|
||||
|
||||
const MaxBatch = 15
|
||||
|
||||
const (
|
||||
StatusPending = "pending"
|
||||
StatusPorted = "ported"
|
||||
)
|
||||
|
||||
const (
|
||||
ModeAllowIncomplete = "allow-incomplete"
|
||||
ModeRequireRecorded = "require-recorded"
|
||||
)
|
||||
|
||||
// Manifest is a generic ordered list of route cases. It has no app-specific names.
|
||||
type Manifest struct {
|
||||
Version int `yaml:"version"`
|
||||
AuthGroups []string `yaml:"auth_groups"`
|
||||
Seed *Seed `yaml:"seed,omitempty"`
|
||||
Routes []Route `yaml:"routes"`
|
||||
}
|
||||
|
||||
// Seed names an optional bootstrap flow recorded before route cases.
|
||||
type Seed struct {
|
||||
Hook string `yaml:"hook,omitempty"`
|
||||
Spec string `yaml:"spec,omitempty"`
|
||||
Fixture string `yaml:"fixture,omitempty"`
|
||||
}
|
||||
|
||||
// Route is one manifest entry with identity, cases and fixture path.
|
||||
type Route struct {
|
||||
ID string `yaml:"id"`
|
||||
Method string `yaml:"method"`
|
||||
Path string `yaml:"path"`
|
||||
AuthGroup string `yaml:"auth_group"`
|
||||
Status string `yaml:"status"`
|
||||
Identities []string `yaml:"identities,omitempty"`
|
||||
Cases []RouteCase `yaml:"cases,omitempty"`
|
||||
Headers map[string]string `yaml:"headers,omitempty"`
|
||||
Normalize []NormalizeRule `yaml:"normalize,omitempty"`
|
||||
SeedHook string `yaml:"seed_hook,omitempty"`
|
||||
Fixture string `yaml:"fixture,omitempty"`
|
||||
}
|
||||
|
||||
// RouteCase is one recorded request/response expectation for a route.
|
||||
type RouteCase struct {
|
||||
ID string `yaml:"id"`
|
||||
Identity string `yaml:"identity,omitempty"`
|
||||
Status int `yaml:"status,omitempty"`
|
||||
Headers map[string]string `yaml:"headers,omitempty"`
|
||||
Fixture string `yaml:"fixture,omitempty"`
|
||||
Request *Request `yaml:"request,omitempty"`
|
||||
Capture []CaptureRule `yaml:"capture,omitempty"`
|
||||
}
|
||||
|
||||
// ManifestConfig drives recording or replay of a manifest.
|
||||
type ManifestConfig struct {
|
||||
Target string
|
||||
Fixtures string
|
||||
Store *Store
|
||||
Rules Rules
|
||||
Update bool
|
||||
Resume bool
|
||||
NextBatch int
|
||||
Mode string
|
||||
SelfCheck bool
|
||||
}
|
||||
|
||||
// LoadManifest reads a strict YAML manifest.
|
||||
func LoadManifest(path string) (Manifest, error) {
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return Manifest{}, fmt.Errorf("tide: read manifest %s: %w", path, err)
|
||||
}
|
||||
return ParseManifest(raw)
|
||||
}
|
||||
|
||||
// ParseManifest decodes a manifest, rejecting unknown fields.
|
||||
func ParseManifest(raw []byte) (Manifest, error) {
|
||||
var m Manifest
|
||||
dec := yaml.NewDecoder(bytes.NewReader(raw), yaml.DisallowUnknownField())
|
||||
if err := dec.Decode(&m); err != nil {
|
||||
return Manifest{}, fmt.Errorf("tide: parse manifest: %w", err)
|
||||
}
|
||||
if m.Version != CurrentVersion {
|
||||
return Manifest{}, fmt.Errorf("tide: unsupported manifest version %d", m.Version)
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// ValidateManifest checks identities in both modes; case/fixture completeness depends on mode.
|
||||
func ValidateManifest(m Manifest, fixtures, mode string) error {
|
||||
if len(m.AuthGroups) == 0 {
|
||||
return fmt.Errorf("tide: manifest auth_groups is required")
|
||||
}
|
||||
groups := make(map[string]struct{}, len(m.AuthGroups))
|
||||
for _, g := range m.AuthGroups {
|
||||
if strings.TrimSpace(g) == "" {
|
||||
return fmt.Errorf("tide: empty auth group")
|
||||
}
|
||||
groups[g] = struct{}{}
|
||||
}
|
||||
if len(m.Routes) == 0 {
|
||||
return fmt.Errorf("tide: manifest has no routes")
|
||||
}
|
||||
seen := map[string]struct{}{}
|
||||
for i, route := range m.Routes {
|
||||
if strings.TrimSpace(route.ID) == "" {
|
||||
return fmt.Errorf("tide: routes[%d] is missing id", i)
|
||||
}
|
||||
if _, dup := seen[route.ID]; dup {
|
||||
return fmt.Errorf("tide: duplicate route id %q", route.ID)
|
||||
}
|
||||
seen[route.ID] = struct{}{}
|
||||
if strings.TrimSpace(route.Method) == "" || strings.TrimSpace(route.Path) == "" {
|
||||
return fmt.Errorf("tide: route %s is missing method or path", route.ID)
|
||||
}
|
||||
if _, ok := groups[route.AuthGroup]; !ok {
|
||||
return fmt.Errorf("tide: route %s has unknown auth group %q", route.ID, route.AuthGroup)
|
||||
}
|
||||
if route.Status != StatusPending && route.Status != StatusPorted {
|
||||
return fmt.Errorf("tide: route %s has unknown status %q", route.ID, route.Status)
|
||||
}
|
||||
if err := validateFixturePath(route.Fixture); err != nil {
|
||||
return fmt.Errorf("tide: route %s: %w", route.ID, err)
|
||||
}
|
||||
for j, c := range route.Cases {
|
||||
if err := validateFixturePath(c.Fixture); err != nil {
|
||||
return fmt.Errorf("tide: route %s case[%d]: %w", route.ID, j, err)
|
||||
}
|
||||
if mode == ModeRequireRecorded {
|
||||
if err := requireCase(route, c); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
if mode == ModeRequireRecorded {
|
||||
if len(route.Cases) == 0 {
|
||||
return fmt.Errorf("tide: route %s is missing cases", route.ID)
|
||||
}
|
||||
for _, c := range routeCases(route) {
|
||||
path := caseFixturePath(fixtures, route, c)
|
||||
if _, err := LoadFlow(path); err != nil {
|
||||
return fmt.Errorf("tide: route %s missing valid fixture %s: %w", route.ID, path, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func requireCase(route Route, c RouteCase) error {
|
||||
if strings.TrimSpace(c.ID) == "" {
|
||||
return fmt.Errorf("tide: route %s has a case without id", route.ID)
|
||||
}
|
||||
if c.Status == 0 {
|
||||
return fmt.Errorf("tide: route %s case %s is missing status", route.ID, c.ID)
|
||||
}
|
||||
req := caseRequest(route, c)
|
||||
if strings.TrimSpace(req.Method) == "" || strings.TrimSpace(req.Path) == "" {
|
||||
return fmt.Errorf("tide: route %s case %s is missing request", route.ID, c.ID)
|
||||
}
|
||||
if caseFixturePath("", route, c) == "" {
|
||||
return fmt.Errorf("tide: route %s case %s is missing fixture", route.ID, c.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateFixturePath(p string) error {
|
||||
if p == "" {
|
||||
return nil
|
||||
}
|
||||
if filepath.IsAbs(p) {
|
||||
return fmt.Errorf("fixture %q must be a relative path", p)
|
||||
}
|
||||
clean := filepath.ToSlash(filepath.Clean(p))
|
||||
if clean == ".." || strings.HasPrefix(clean, "../") {
|
||||
return fmt.Errorf("fixture %q escapes the fixture directory", p)
|
||||
}
|
||||
if _, err := resolvePath(p); err != nil {
|
||||
return fmt.Errorf("fixture %q: %w", p, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func routeCases(route Route) []RouteCase {
|
||||
if len(route.Cases) > 0 {
|
||||
return route.Cases
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func caseRequest(route Route, c RouteCase) Request {
|
||||
if c.Request != nil {
|
||||
return *c.Request
|
||||
}
|
||||
return Request{Method: route.Method, Path: route.Path}
|
||||
}
|
||||
|
||||
func caseFixturePath(root string, route Route, c RouteCase) string {
|
||||
rel := c.Fixture
|
||||
if rel == "" {
|
||||
rel = route.Fixture
|
||||
}
|
||||
if rel == "" && c.ID != "" {
|
||||
rel = filepath.ToSlash(filepath.Join("routes", sanitizeFile(route.ID)+"__"+sanitizeFile(c.ID)+".yaml"))
|
||||
}
|
||||
if rel == "" && route.ID != "" {
|
||||
rel = filepath.ToSlash(filepath.Join("routes", sanitizeFile(route.ID)+".yaml"))
|
||||
}
|
||||
if root == "" {
|
||||
return rel
|
||||
}
|
||||
return filepath.Join(root, rel)
|
||||
}
|
||||
|
||||
func sanitizeFile(s string) string {
|
||||
s = strings.TrimSpace(s)
|
||||
repl := strings.NewReplacer("/", "_", "\\", "_", " ", "_")
|
||||
return repl.Replace(s)
|
||||
}
|
||||
|
||||
func fixtureHash(path string) (string, error) {
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if len(bytes.TrimSpace(raw)) == 0 {
|
||||
return "", fmt.Errorf("empty fixture")
|
||||
}
|
||||
sum := sha256.Sum256(raw)
|
||||
return hex.EncodeToString(sum[:]), nil
|
||||
}
|
||||
|
||||
func recordedCase(fixtures string, route Route, c RouteCase) (bool, error) {
|
||||
path := caseFixturePath(fixtures, route, c)
|
||||
if _, err := os.Stat(path); err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return false, nil
|
||||
}
|
||||
return false, err
|
||||
}
|
||||
if _, err := LoadFlow(path); err != nil {
|
||||
return false, fmt.Errorf("tide: existing fixture %s is invalid: %w", path, err)
|
||||
}
|
||||
if _, err := fixtureHash(path); err != nil {
|
||||
return false, fmt.Errorf("tide: existing fixture %s hash: %w", path, err)
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
// RecordManifest records seed then route cases. NextBatch>15 is refused.
|
||||
func RecordManifest(ctx context.Context, m Manifest, cfg ManifestConfig) (Coverage, error) {
|
||||
if cfg.NextBatch > MaxBatch {
|
||||
return Coverage{}, fmt.Errorf("tide: --next-batch %d exceeds %d", cfg.NextBatch, MaxBatch)
|
||||
}
|
||||
if err := ValidateManifest(m, cfg.Fixtures, ModeAllowIncomplete); err != nil {
|
||||
return Coverage{}, err
|
||||
}
|
||||
cov := newCoverage(m)
|
||||
if m.Seed != nil && strings.TrimSpace(m.Seed.Spec) != "" {
|
||||
if err := recordSeed(ctx, m.Seed, cfg); err != nil {
|
||||
return cov, err
|
||||
}
|
||||
}
|
||||
limit := cfg.NextBatch
|
||||
if limit <= 0 {
|
||||
limit = len(m.Routes)
|
||||
}
|
||||
taken := 0
|
||||
for _, route := range m.Routes {
|
||||
cases := route.Cases
|
||||
if len(cases) == 0 {
|
||||
cov.markUnrecorded(route)
|
||||
continue
|
||||
}
|
||||
allRecorded := true
|
||||
for _, c := range cases {
|
||||
ok, err := recordedCase(cfg.Fixtures, route, c)
|
||||
if err != nil {
|
||||
return cov, err
|
||||
}
|
||||
if !ok {
|
||||
allRecorded = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if allRecorded && !cfg.Update {
|
||||
cov.markRecorded(route)
|
||||
continue
|
||||
}
|
||||
if taken >= limit {
|
||||
cov.markUnrecorded(route)
|
||||
continue
|
||||
}
|
||||
if err := recordRoute(ctx, route, cfg); err != nil {
|
||||
return cov, err
|
||||
}
|
||||
cov.markRecorded(route)
|
||||
taken++
|
||||
}
|
||||
cov.ResumeRemaining = countUnrecorded(m, cfg.Fixtures)
|
||||
return cov, nil
|
||||
}
|
||||
|
||||
func recordSeed(ctx context.Context, seed *Seed, cfg ManifestConfig) error {
|
||||
dest := seed.Fixture
|
||||
if dest == "" {
|
||||
dest = seed.Spec
|
||||
}
|
||||
if cfg.Fixtures != "" && !filepath.IsAbs(dest) {
|
||||
dest = filepath.Join(cfg.Fixtures, dest)
|
||||
}
|
||||
if recorded, err := seedAlreadyRecorded(dest); err != nil {
|
||||
return err
|
||||
} else if recorded {
|
||||
return nil
|
||||
}
|
||||
spec, err := LoadFlow(resolveSeedSpec(seed.Spec, cfg.Fixtures))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
flow, err := RecordFlow(ctx, spec, RecordConfig{Target: cfg.Target, Store: cfg.Store, Rules: cfg.Rules})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := os.Stat(dest); err == nil {
|
||||
return SaveFlow(dest, flow)
|
||||
}
|
||||
return SaveFlowExclusive(dest, flow)
|
||||
}
|
||||
|
||||
func seedAlreadyRecorded(path string) (bool, error) {
|
||||
st, err := os.Stat(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return false, nil
|
||||
}
|
||||
return false, err
|
||||
}
|
||||
if st.IsDir() || st.Size() == 0 {
|
||||
return false, nil
|
||||
}
|
||||
flow, err := LoadFlow(path)
|
||||
if err != nil {
|
||||
return false, nil
|
||||
}
|
||||
for _, step := range flow.Steps {
|
||||
if step.Response.Status == 0 {
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
return len(flow.Steps) > 0, nil
|
||||
}
|
||||
|
||||
func resolveSeedSpec(spec, fixtures string) string {
|
||||
if spec == "" || filepath.IsAbs(spec) || fixtures == "" {
|
||||
return spec
|
||||
}
|
||||
joined := filepath.Join(fixtures, spec)
|
||||
if _, err := os.Stat(joined); err == nil {
|
||||
return joined
|
||||
}
|
||||
return spec
|
||||
}
|
||||
|
||||
func recordRoute(ctx context.Context, route Route, cfg ManifestConfig) error {
|
||||
for _, c := range route.Cases {
|
||||
dest := caseFixturePath(cfg.Fixtures, route, c)
|
||||
if _, err := os.Stat(dest); err == nil {
|
||||
if cfg.Resume && !cfg.Update {
|
||||
if _, err := LoadFlow(dest); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !cfg.Update {
|
||||
return fmt.Errorf("tide: fixture %s exists; pass --update to overwrite", dest)
|
||||
}
|
||||
}
|
||||
spec := Flow{
|
||||
Version: CurrentVersion,
|
||||
Name: route.ID + "/" + c.ID,
|
||||
SeedHook: route.SeedHook,
|
||||
Steps: []Step{{
|
||||
ID: c.ID,
|
||||
RouteID: route.ID,
|
||||
Request: caseRequest(route, c),
|
||||
Capture: c.Capture,
|
||||
Normalize: route.Normalize,
|
||||
Headers: mergeHeaders(route.Headers, c.Headers),
|
||||
}},
|
||||
}
|
||||
flow, err := RecordFlow(ctx, spec, RecordConfig{Target: cfg.Target, Store: cfg.Store, Rules: cfg.Rules})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if cfg.Update {
|
||||
if err := SaveFlow(dest, flow); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := SaveFlowExclusive(dest, flow); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func mergeHeaders(a, b map[string]string) map[string]string {
|
||||
if len(a) == 0 && len(b) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make(map[string]string)
|
||||
for k, v := range a {
|
||||
out[k] = v
|
||||
}
|
||||
for k, v := range b {
|
||||
out[k] = v
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func countUnrecorded(m Manifest, fixtures string) int {
|
||||
n := 0
|
||||
for _, route := range m.Routes {
|
||||
if len(route.Cases) == 0 {
|
||||
n++
|
||||
continue
|
||||
}
|
||||
for _, c := range route.Cases {
|
||||
ok, err := recordedCase(fixtures, route, c)
|
||||
if err != nil || !ok {
|
||||
n++
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// ReplayManifest replays recorded route fixtures and builds a coverage table.
|
||||
func ReplayManifest(ctx context.Context, m Manifest, cfg ManifestConfig) (Coverage, error) {
|
||||
if err := ValidateManifest(m, cfg.Fixtures, ModeAllowIncomplete); err != nil {
|
||||
return Coverage{}, err
|
||||
}
|
||||
cov := newCoverage(m)
|
||||
var fail error
|
||||
for _, route := range m.Routes {
|
||||
cases := route.Cases
|
||||
if len(cases) == 0 {
|
||||
cov.markUnrecorded(route)
|
||||
if cfg.Mode == ModeRequireRecorded {
|
||||
fail = firstErr(fail, fmt.Errorf("tide: unrecorded required route %s", route.ID))
|
||||
}
|
||||
continue
|
||||
}
|
||||
routeFail := false
|
||||
routeRecorded := true
|
||||
for _, c := range cases {
|
||||
path := caseFixturePath(cfg.Fixtures, route, c)
|
||||
flow, err := LoadFlow(path)
|
||||
if err != nil {
|
||||
routeRecorded = false
|
||||
if cfg.Mode == ModeRequireRecorded {
|
||||
fail = firstErr(fail, fmt.Errorf("tide: unrecorded required route %s", route.ID))
|
||||
}
|
||||
continue
|
||||
}
|
||||
res, err := ReplayFlow(ctx, flow, ReplayConfig{Target: cfg.Target, Store: cfg.Store, BaseDir: filepath.Dir(path)})
|
||||
if err != nil {
|
||||
var mis *MismatchError
|
||||
if !errors.As(err, &mis) {
|
||||
return cov, err
|
||||
}
|
||||
res = mis.Result
|
||||
}
|
||||
if !res.OK {
|
||||
routeFail = true
|
||||
cov.addDiffs(route.ID, res)
|
||||
if cfg.SelfCheck || route.Status == StatusPorted {
|
||||
fail = firstErr(fail, fmt.Errorf("tide: route %s failed", route.ID))
|
||||
}
|
||||
}
|
||||
}
|
||||
if !routeRecorded {
|
||||
cov.markUnrecorded(route)
|
||||
continue
|
||||
}
|
||||
if routeFail {
|
||||
cov.markFailing(route)
|
||||
continue
|
||||
}
|
||||
cov.markPassing(route)
|
||||
}
|
||||
return cov, fail
|
||||
}
|
||||
|
||||
func firstErr(cur, next error) error {
|
||||
if cur == nil {
|
||||
return next
|
||||
}
|
||||
return cur
|
||||
}
|
||||
|
||||
func newCoverage(m Manifest) Coverage {
|
||||
c := Coverage{Total: len(m.Routes), Rows: make([]CoverageRow, 0, len(m.Routes))}
|
||||
for _, route := range m.Routes {
|
||||
if route.Status == StatusPending {
|
||||
c.Pending++
|
||||
} else if route.Status == StatusPorted {
|
||||
c.Ported++
|
||||
}
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *Coverage) markRecorded(route Route) {
|
||||
c.Recorded++
|
||||
c.Rows = append(c.Rows, CoverageRow{ID: route.ID, Status: route.Status, Outcome: "recorded"})
|
||||
}
|
||||
|
||||
func (c *Coverage) markPassing(route Route) {
|
||||
c.Recorded++
|
||||
c.Passing++
|
||||
c.Rows = append(c.Rows, CoverageRow{ID: route.ID, Status: route.Status, Outcome: "passing"})
|
||||
}
|
||||
|
||||
func (c *Coverage) markFailing(route Route) {
|
||||
c.Recorded++
|
||||
c.Failing++
|
||||
c.Rows = append(c.Rows, CoverageRow{ID: route.ID, Status: route.Status, Outcome: "failing"})
|
||||
}
|
||||
|
||||
func (c *Coverage) markUnrecorded(route Route) {
|
||||
c.Unrecorded++
|
||||
c.Rows = append(c.Rows, CoverageRow{ID: route.ID, Status: route.Status, Outcome: "unrecorded"})
|
||||
}
|
||||
|
||||
func (c *Coverage) addDiffs(id string, res Result) {
|
||||
for _, sr := range res.Steps {
|
||||
for _, d := range sr.Diffs {
|
||||
c.Diffs = append(c.Diffs, fmt.Sprintf("%s %s: %s expected %s actual %s", id, sr.ID, d.Path, redactSecrets(d.Expected), redactSecrets(d.Actual)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func ParseNextBatch(s string) (int, error) {
|
||||
if strings.TrimSpace(s) == "" {
|
||||
return 0, nil
|
||||
}
|
||||
n, err := strconv.Atoi(s)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("tide: --next-batch: %w", err)
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
116
modules/tide/manifest_contract_test.go
Normal file
116
modules/tide/manifest_contract_test.go
Normal file
@@ -0,0 +1,116 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestManifestContract(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
fixtures := t.TempDir()
|
||||
if err := os.MkdirAll(filepath.Join(fixtures, "routes"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var hits atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/pass":
|
||||
_, _ = w.Write([]byte(`{"v":1}`))
|
||||
case "/fail":
|
||||
_, _ = w.Write([]byte(`{"v":2}`))
|
||||
default:
|
||||
_, _ = w.Write([]byte(`{"v":0}`))
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
writeRouteFixture(t, srv.URL, fixtures, "pass", "/pass", `{"v":1}`)
|
||||
writeRouteFixture(t, srv.URL, fixtures, "fail", "/fail", `{"v":1}`)
|
||||
|
||||
m := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Routes: []Route{
|
||||
routeEntry("GET /pass public", "/pass", StatusPorted, "routes/pass.yaml"),
|
||||
routeEntry("GET /fail public", "/fail", StatusPorted, "routes/fail.yaml"),
|
||||
{ID: "GET /none public", Method: http.MethodGet, Path: "/none", AuthGroup: "public", Status: StatusPending},
|
||||
},
|
||||
}
|
||||
before := hits.Load()
|
||||
cov, err := ReplayManifest(ctx, m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, Mode: ModeRequireRecorded})
|
||||
if err == nil {
|
||||
t.Fatal("ported fail plus unrecorded required must be nonzero")
|
||||
}
|
||||
if cov.Passing != 1 || cov.Failing != 1 || cov.Unrecorded != 1 || cov.Recorded != 2 {
|
||||
t.Fatalf("coverage %+v", cov)
|
||||
}
|
||||
if cov.Pending != 1 || cov.Ported != 2 {
|
||||
t.Fatalf("pending/ported counts %+v", cov)
|
||||
}
|
||||
if !strings.Contains(cov.SummaryLine(), "recorded 2/3") {
|
||||
t.Fatalf("summary %s", cov.SummaryLine())
|
||||
}
|
||||
if hits.Load()-before != 2 {
|
||||
t.Fatalf("comparison failure must still replay later recorded flows, extra hits=%d", hits.Load()-before)
|
||||
}
|
||||
foundFailPath := false
|
||||
for _, d := range cov.Diffs {
|
||||
if strings.Contains(d, "GET /fail public") && (strings.Contains(d, "$.v") || strings.Contains(d, ".v")) {
|
||||
foundFailPath = true
|
||||
}
|
||||
}
|
||||
if !foundFailPath {
|
||||
t.Fatalf("coverage diffs missing path for failing route: %v", cov.Diffs)
|
||||
}
|
||||
|
||||
t.Run("pending mismatch is measured not fatal", func(t *testing.T) {
|
||||
pending := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Routes: []Route{routeEntry("GET /fail public", "/fail", StatusPending, "routes/fail.yaml")},
|
||||
}
|
||||
cov, err := ReplayManifest(ctx, pending, ManifestConfig{Target: srv.URL, Fixtures: fixtures})
|
||||
if err != nil {
|
||||
t.Fatalf("pending mismatch must not fail: %v", err)
|
||||
}
|
||||
if cov.Failing != 1 || cov.Passing != 0 {
|
||||
t.Fatalf("pending failing %+v", cov)
|
||||
}
|
||||
if _, err := ReplayManifest(ctx, pending, ManifestConfig{Target: srv.URL, Fixtures: fixtures, SelfCheck: true}); err == nil {
|
||||
t.Fatal("self-check must fail pending mismatches")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("require-recorded reports unrecorded", func(t *testing.T) {
|
||||
empty := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Routes: []Route{{
|
||||
ID: "GET /ghost public", Method: http.MethodGet, Path: "/ghost",
|
||||
AuthGroup: "public", Status: StatusPending,
|
||||
}},
|
||||
}
|
||||
if err := ValidateManifest(empty, fixtures, ModeAllowIncomplete); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ValidateManifest(empty, fixtures, ModeRequireRecorded); err == nil {
|
||||
t.Fatal("missing cases must fail require-recorded")
|
||||
}
|
||||
cov, err := ReplayManifest(ctx, empty, ManifestConfig{Target: srv.URL, Fixtures: fixtures, Mode: ModeRequireRecorded})
|
||||
if err == nil {
|
||||
t.Fatal("unrecorded required must error")
|
||||
}
|
||||
if cov.Unrecorded != 1 || cov.Passing != 0 {
|
||||
t.Fatalf("unrecorded coverage %+v", cov)
|
||||
}
|
||||
})
|
||||
}
|
||||
475
modules/tide/manifest_test.go
Normal file
475
modules/tide/manifest_test.go
Normal file
@@ -0,0 +1,475 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestManifestValidationAndCoverage(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
fixtures := filepath.Join(dir, "fixtures")
|
||||
if err := os.MkdirAll(filepath.Join(fixtures, "routes"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
manPath := filepath.Join(dir, "manifest.yaml")
|
||||
manYAML := "" +
|
||||
"version: 1\n" +
|
||||
"auth_groups:\n" +
|
||||
" - public\n" +
|
||||
" - session\n" +
|
||||
"routes:\n" +
|
||||
" - id: GET /ok public\n" +
|
||||
" method: GET\n" +
|
||||
" path: /ok\n" +
|
||||
" auth_group: public\n" +
|
||||
" status: ported\n" +
|
||||
" cases:\n" +
|
||||
" - id: ok\n" +
|
||||
" status: 200\n" +
|
||||
" fixture: routes/ok.yaml\n" +
|
||||
" request:\n" +
|
||||
" method: GET\n" +
|
||||
" path: /ok\n" +
|
||||
" - id: GET /missing public\n" +
|
||||
" method: GET\n" +
|
||||
" path: /missing\n" +
|
||||
" auth_group: public\n" +
|
||||
" status: pending\n" +
|
||||
" cases:\n" +
|
||||
" - id: missing\n" +
|
||||
" status: 200\n" +
|
||||
" fixture: routes/missing.yaml\n" +
|
||||
" request:\n" +
|
||||
" method: GET\n" +
|
||||
" path: /missing\n"
|
||||
if err := os.WriteFile(manPath, []byte(manYAML), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
m, err := LoadManifest(manPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ValidateManifest(m, fixtures, ModeAllowIncomplete); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ValidateManifest(m, fixtures, ModeRequireRecorded); err == nil {
|
||||
t.Fatal("require-recorded must fail with missing fixtures")
|
||||
}
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
cov, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, Resume: true})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cov.Recorded != 2 || cov.Unrecorded != 0 {
|
||||
t.Fatalf("record coverage %+v", cov)
|
||||
}
|
||||
if err := ValidateManifest(m, fixtures, ModeRequireRecorded); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCoverageTwoRouteReport(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
fixtures := filepath.Join(dir, "fx")
|
||||
if err := os.MkdirAll(filepath.Join(fixtures, "routes"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch r.URL.Path {
|
||||
case "/pass":
|
||||
_, _ = w.Write([]byte(`{"v":1}`))
|
||||
default:
|
||||
_, _ = w.Write([]byte(`{"v":2}`))
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
writeRouteFixture(t, srv.URL, fixtures, "pass", "/pass", `{"v":1}`)
|
||||
writeRouteFixture(t, srv.URL, fixtures, "fail", "/fail", `{"v":1}`)
|
||||
|
||||
m := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Routes: []Route{
|
||||
routeEntry("GET /pass public", "/pass", "ported", "routes/pass.yaml"),
|
||||
routeEntry("GET /fail public", "/fail", "ported", "routes/fail.yaml"),
|
||||
{ID: "GET /none public", Method: "GET", Path: "/none", AuthGroup: "public", Status: StatusPending},
|
||||
},
|
||||
}
|
||||
cov, err := ReplayManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, Mode: ModeRequireRecorded})
|
||||
if err == nil {
|
||||
t.Fatal("ported fail + unrecorded required must be nonzero")
|
||||
}
|
||||
if cov.Passing != 1 || cov.Failing != 1 || cov.Unrecorded != 1 || cov.Recorded != 2 {
|
||||
t.Fatalf("coverage %+v", cov)
|
||||
}
|
||||
if !strings.Contains(cov.SummaryLine(), "passing 1") {
|
||||
t.Fatalf("summary %s", cov.SummaryLine())
|
||||
}
|
||||
}
|
||||
|
||||
func TestManifestResumeBatches(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
fixtures := filepath.Join(dir, "fx")
|
||||
hits := map[string]int{}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits[r.URL.Path]++
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
m := Manifest{Version: 1, AuthGroups: []string{"public"}}
|
||||
for i := 1; i <= 16; i++ {
|
||||
id := fmt.Sprintf("GET /r/%d public", i)
|
||||
path := fmt.Sprintf("/r/%d", i)
|
||||
m.Routes = append(m.Routes, Route{
|
||||
ID: id,
|
||||
Method: http.MethodGet,
|
||||
Path: path,
|
||||
AuthGroup: "public",
|
||||
Status: StatusPending,
|
||||
Cases: []RouteCase{{
|
||||
ID: "ok",
|
||||
Status: 200,
|
||||
Fixture: fmt.Sprintf("routes/r%d.yaml", i),
|
||||
Request: &Request{Method: http.MethodGet, Path: path},
|
||||
}},
|
||||
})
|
||||
}
|
||||
if _, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, NextBatch: 16}); err == nil {
|
||||
t.Fatal("batch >15 must fail")
|
||||
}
|
||||
cov, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, NextBatch: 15, Resume: true})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cov.Recorded != 15 || cov.ResumeRemaining != 1 {
|
||||
t.Fatalf("first batch %+v remaining %d", cov, cov.ResumeRemaining)
|
||||
}
|
||||
cov, err = RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, NextBatch: 15, Resume: true})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cov.Recorded != 16 || cov.ResumeRemaining != 0 {
|
||||
t.Fatalf("second batch %+v remaining %d", cov, cov.ResumeRemaining)
|
||||
}
|
||||
for i := 1; i <= 16; i++ {
|
||||
if hits[fmt.Sprintf("/r/%d", i)] != 1 {
|
||||
t.Fatalf("route %d recaptured: %d", i, hits[fmt.Sprintf("/r/%d", i)])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestManifestAllowIncomplete154(t *testing.T) {
|
||||
var b strings.Builder
|
||||
b.WriteString("version: 1\nauth_groups:\n - public\nroutes:\n")
|
||||
for i := 1; i <= 154; i++ {
|
||||
fmt.Fprintf(&b, " - id: GET /r/%d public\n method: GET\n path: /r/%d\n auth_group: public\n status: pending\n", i, i)
|
||||
}
|
||||
m, err := ParseManifest([]byte(b.String()))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(m.Routes) != 154 {
|
||||
t.Fatalf("routes %d", len(m.Routes))
|
||||
}
|
||||
if err := ValidateManifest(m, t.TempDir(), ModeAllowIncomplete); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ValidateManifest(m, t.TempDir(), ModeRequireRecorded); err == nil {
|
||||
t.Fatal("154 empty cases must fail require-recorded")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManifestSidecarRejectAndDigest(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
base := "version: 1\nname: bin\nsteps:\n - id: a\n request:\n method: GET\n path: /bin\n response:\n body_file: ../secret.bin\n"
|
||||
p := filepath.Join(dir, "bad.yaml")
|
||||
if err := os.WriteFile(p, []byte(base), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := LoadFlow(p); err == nil {
|
||||
t.Fatal("parent sidecar must fail")
|
||||
}
|
||||
good := "version: 1\nname: bin\nsteps:\n - id: a\n request:\n method: GET\n path: /bin\n response:\n status: 200\n headers:\n Content-Type: application/octet-stream\n body_file: data.bin\n sha256: deadbeef\n"
|
||||
gp := filepath.Join(dir, "good.yaml")
|
||||
if err := os.WriteFile(gp, []byte(good), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, "data.bin"), []byte("hello"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
flow, err := LoadFlow(gp)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/octet-stream")
|
||||
_, _ = w.Write([]byte("hello"))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
_, err = ReplayFlow(context.Background(), flow, ReplayConfig{Target: srv.URL, BaseDir: dir})
|
||||
if err == nil || !strings.Contains(err.Error(), "digest") {
|
||||
t.Fatalf("digest mismatch: %v", err)
|
||||
}
|
||||
sum := sha256.Sum256([]byte("hello"))
|
||||
good2 := strings.Replace(good, "deadbeef", hex.EncodeToString(sum[:]), 1)
|
||||
if err := os.WriteFile(gp, []byte(good2), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
flow, err = LoadFlow(gp)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := ReplayFlow(context.Background(), flow, ReplayConfig{Target: srv.URL, BaseDir: dir}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManifestUnknownGroupAndDuplicate(t *testing.T) {
|
||||
m, err := ParseManifest([]byte("version: 1\nauth_groups:\n - public\nroutes:\n - id: a\n method: GET\n path: /a\n auth_group: oauth\n status: pending\n"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ValidateManifest(m, "", ModeAllowIncomplete); err == nil || !strings.Contains(err.Error(), "auth group") {
|
||||
t.Fatalf("unknown group: %v", err)
|
||||
}
|
||||
m, err = ParseManifest([]byte("version: 1\nauth_groups:\n - public\nroutes:\n - id: a\n method: GET\n path: /a\n auth_group: public\n status: pending\n - id: a\n method: GET\n path: /b\n auth_group: public\n status: pending\n"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ValidateManifest(m, "", ModeAllowIncomplete); err == nil || !strings.Contains(err.Error(), "duplicate") {
|
||||
t.Fatalf("duplicate: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func routeEntry(id, path, status, fixture string) Route {
|
||||
return Route{
|
||||
ID: id,
|
||||
Method: http.MethodGet,
|
||||
Path: path,
|
||||
AuthGroup: "public",
|
||||
Status: status,
|
||||
Fixture: fixture,
|
||||
Cases: []RouteCase{{
|
||||
ID: "ok",
|
||||
Status: 200,
|
||||
Fixture: fixture,
|
||||
Request: &Request{Method: http.MethodGet, Path: path},
|
||||
}},
|
||||
}
|
||||
}
|
||||
|
||||
func writeRouteFixture(t *testing.T, target, fixtures, name, path, body string) {
|
||||
t.Helper()
|
||||
spec := Flow{Version: 1, Name: name, Steps: []Step{{
|
||||
ID: "ok",
|
||||
Request: Request{Method: http.MethodGet, Path: path},
|
||||
}}}
|
||||
orig := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(body))
|
||||
}))
|
||||
t.Cleanup(orig.Close)
|
||||
rec, err := RecordFlow(context.Background(), spec, RecordConfig{Target: orig.URL})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := SaveFlow(filepath.Join(fixtures, "routes", name+".yaml"), rec); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManifestNoOverwriteWithoutUpdate(t *testing.T) {
|
||||
fixtures := t.TempDir()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
m := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Routes: []Route{{
|
||||
ID: "GET /a public",
|
||||
Method: http.MethodGet,
|
||||
Path: "/a",
|
||||
AuthGroup: "public",
|
||||
Status: StatusPending,
|
||||
Cases: []RouteCase{
|
||||
{ID: "one", Status: 200, Fixture: "routes/one.yaml", Request: &Request{Method: http.MethodGet, Path: "/a"}},
|
||||
{ID: "two", Status: 200, Fixture: "routes/two.yaml", Request: &Request{Method: http.MethodGet, Path: "/a"}},
|
||||
},
|
||||
}},
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Join(fixtures, "routes"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
first := Flow{Version: 1, Name: "one", Steps: []Step{{ID: "one", Request: Request{Method: http.MethodGet, Path: "/a"}}}}
|
||||
rec, err := RecordFlow(context.Background(), first, RecordConfig{Target: srv.URL})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := SaveFlow(filepath.Join(fixtures, "routes", "one.yaml"), rec); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures}); err == nil || !strings.Contains(err.Error(), "--update") {
|
||||
t.Fatalf("must refuse overwrite: %v", err)
|
||||
}
|
||||
if _, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, Update: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCoveragePendingMismatchDoesNotFail(t *testing.T) {
|
||||
fixtures := t.TempDir()
|
||||
if err := os.MkdirAll(filepath.Join(fixtures, "routes"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
writeRouteFixture(t, "", fixtures, "pend", "/pend", `{"v":1}`)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"v":9}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
m := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Routes: []Route{routeEntry("GET /pend public", "/pend", StatusPending, "routes/pend.yaml")},
|
||||
}
|
||||
cov, err := ReplayManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures})
|
||||
if err != nil {
|
||||
t.Fatalf("pending mismatch must not fail: %v", err)
|
||||
}
|
||||
if cov.Failing != 1 {
|
||||
t.Fatalf("failing %d", cov.Failing)
|
||||
}
|
||||
if _, err := ReplayManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, SelfCheck: true}); err == nil {
|
||||
t.Fatal("self-check must fail pending mismatches")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordSeedResolvesRelativeSpec(t *testing.T) {
|
||||
fixtures := t.TempDir()
|
||||
if err := os.MkdirAll(filepath.Join(fixtures, "seed"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
spec := Flow{Version: 1, Name: "seed", Steps: []Step{{
|
||||
ID: "ok",
|
||||
Request: Request{Method: http.MethodGet, Path: "/seed"},
|
||||
}}}
|
||||
if err := SaveFlow(filepath.Join(fixtures, "seed", "bootstrap.yaml"), spec); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
hits := 0
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits++
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
m := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Seed: &Seed{Spec: "seed/bootstrap.yaml", Fixture: "seed/bootstrap.yaml"},
|
||||
Routes: []Route{{
|
||||
ID: "GET /later public", Method: http.MethodGet, Path: "/later",
|
||||
AuthGroup: "public", Status: StatusPending,
|
||||
}},
|
||||
}
|
||||
if _, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hits != 1 {
|
||||
t.Fatalf("seed hits %d", hits)
|
||||
}
|
||||
if _, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hits != 1 {
|
||||
t.Fatalf("existing seed must not re-hit PHP: %d", hits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordManifestUpdateRerecords(t *testing.T) {
|
||||
fixtures := t.TempDir()
|
||||
hits := 0
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
hits++
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"n":` + strconv.Itoa(hits) + `}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
m := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Routes: []Route{{
|
||||
ID: "GET /x public", Method: http.MethodGet, Path: "/x",
|
||||
AuthGroup: "public", Status: StatusPending,
|
||||
Cases: []RouteCase{{
|
||||
ID: "ok", Status: 200, Fixture: "routes/x.yaml",
|
||||
Request: &Request{Method: http.MethodGet, Path: "/x"},
|
||||
}},
|
||||
}},
|
||||
}
|
||||
if _, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, Resume: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, Resume: true, Update: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hits != 2 {
|
||||
t.Fatalf("update must recapture, hits=%d", hits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouteCaseCaptureUpdatesStore(t *testing.T) {
|
||||
fixtures := t.TempDir()
|
||||
vars := filepath.Join(t.TempDir(), "vars.yaml")
|
||||
store, err := OpenStore(vars)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"data":{"token":"shareTokValue99"}}`))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
m := Manifest{
|
||||
Version: 1,
|
||||
AuthGroups: []string{"public"},
|
||||
Routes: []Route{{
|
||||
ID: "POST /share public", Method: http.MethodPost, Path: "/share",
|
||||
AuthGroup: "public", Status: StatusPending,
|
||||
Cases: []RouteCase{{
|
||||
ID: "ok", Status: 200, Fixture: "routes/share.yaml",
|
||||
Request: &Request{Method: http.MethodPost, Path: "/share"},
|
||||
Capture: []CaptureRule{{From: "response.json", Path: "$.data.token", As: "share:collection"}},
|
||||
}},
|
||||
}},
|
||||
}
|
||||
if _, err := RecordManifest(context.Background(), m, ManifestConfig{Target: srv.URL, Fixtures: fixtures, Store: store}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, ok := store.Get("share:collection")
|
||||
if !ok || got != "shareTokValue99" {
|
||||
t.Fatalf("capture: ok=%v val=%q", ok, got)
|
||||
}
|
||||
}
|
||||
156
modules/tide/normalize.go
Normal file
156
modules/tide/normalize.go
Normal file
@@ -0,0 +1,156 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var carbonOffsetRe = regexp.MustCompile(`^\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}\+00:00$`)
|
||||
|
||||
const (
|
||||
maskDatetime = "<datetime>"
|
||||
maskID = "<id>"
|
||||
)
|
||||
|
||||
func normalizeJSON(raw []byte, step Step) ([]byte, []Diff) {
|
||||
if len(strings.TrimSpace(string(raw))) == 0 {
|
||||
return raw, nil
|
||||
}
|
||||
val, err := decodeJSON(raw)
|
||||
if err != nil {
|
||||
return raw, nil
|
||||
}
|
||||
var diffs []Diff
|
||||
masked := maskValue("$", val, step, &diffs)
|
||||
out, err := json.Marshal(masked)
|
||||
if err != nil {
|
||||
return raw, diffs
|
||||
}
|
||||
return out, diffs
|
||||
}
|
||||
|
||||
func maskValue(path string, val any, step Step, diffs *[]Diff) any {
|
||||
switch v := val.(type) {
|
||||
case map[string]any:
|
||||
out := make(map[string]any, len(v))
|
||||
for k, child := range v {
|
||||
out[k] = maskValue(pathJoin(path, k), child, step, diffs)
|
||||
}
|
||||
return out
|
||||
case []any:
|
||||
out := make([]any, len(v))
|
||||
for i, child := range v {
|
||||
out[i] = maskValue(fmt.Sprintf("%s[%d]", path, i), child, step, diffs)
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return maskLeaf(path, val, step, diffs)
|
||||
}
|
||||
}
|
||||
|
||||
func maskLeaf(path string, val any, step Step, diffs *[]Diff) any {
|
||||
key := lastPathKey(path)
|
||||
if key == "slug" || disabledPath(step, path, key) {
|
||||
return val
|
||||
}
|
||||
if key == "collection_key" || key == "client_id" {
|
||||
if val == nil {
|
||||
return nil
|
||||
}
|
||||
if _, ok := val.(string); !ok {
|
||||
*diffs = append(*diffs, Diff{Path: path, Expected: "string " + key, Actual: formatValue(val)})
|
||||
return val
|
||||
}
|
||||
return maskID
|
||||
}
|
||||
if strings.HasSuffix(key, "_issued_at") {
|
||||
return maskIDValue(path, val, diffs)
|
||||
}
|
||||
if isDateKey(key) {
|
||||
return maskDate(path, val, diffs)
|
||||
}
|
||||
if isIDKey(key) {
|
||||
return maskIDValue(path, val, diffs)
|
||||
}
|
||||
return val
|
||||
}
|
||||
|
||||
func maskDate(path string, val any, diffs *[]Diff) any {
|
||||
if val == nil {
|
||||
return nil
|
||||
}
|
||||
s, ok := val.(string)
|
||||
if !ok {
|
||||
*diffs = append(*diffs, Diff{Path: path, Expected: "Carbon +00:00 string or null", Actual: formatValue(val)})
|
||||
return val
|
||||
}
|
||||
if !carbonOffsetRe.MatchString(s) {
|
||||
*diffs = append(*diffs, Diff{Path: path, Expected: "Carbon +00:00", Actual: strconvQuote(s)})
|
||||
return val
|
||||
}
|
||||
return maskDatetime
|
||||
}
|
||||
|
||||
func maskIDValue(path string, val any, diffs *[]Diff) any {
|
||||
if val == nil {
|
||||
return nil
|
||||
}
|
||||
n, ok := val.(json.Number)
|
||||
if !ok {
|
||||
*diffs = append(*diffs, Diff{Path: path, Expected: "integer id", Actual: formatValue(val)})
|
||||
return val
|
||||
}
|
||||
if strings.Contains(string(n), ".") {
|
||||
*diffs = append(*diffs, Diff{Path: path, Expected: "integer id", Actual: "number " + string(n)})
|
||||
return val
|
||||
}
|
||||
return maskID
|
||||
}
|
||||
|
||||
func isDateKey(key string) bool {
|
||||
return strings.HasSuffix(key, "_at") || key == "checkpoint"
|
||||
}
|
||||
|
||||
func isIDKey(key string) bool {
|
||||
if key == "id" {
|
||||
return true
|
||||
}
|
||||
if strings.HasSuffix(key, "_at") {
|
||||
return false
|
||||
}
|
||||
// "_ids" covers plural raw-integer-array fields such as collection_ids:
|
||||
// each array element still reaches maskLeaf individually (maskValue
|
||||
// recurses into []any before calling maskLeaf), so this masks every
|
||||
// element the same way a singular "_id" scalar would be masked.
|
||||
return strings.HasSuffix(key, "_id") || strings.HasSuffix(key, "_ids")
|
||||
}
|
||||
|
||||
func lastPathKey(path string) string {
|
||||
path = strings.TrimPrefix(path, "$.")
|
||||
if i := strings.LastIndex(path, "."); i >= 0 {
|
||||
path = path[i+1:]
|
||||
}
|
||||
if i := strings.IndexByte(path, '['); i >= 0 {
|
||||
path = path[:i]
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
func disabledPath(step Step, jsonPath, key string) bool {
|
||||
for _, rule := range step.Normalize {
|
||||
if !rule.Disable {
|
||||
continue
|
||||
}
|
||||
p := strings.TrimSpace(rule.Path)
|
||||
if p == jsonPath || p == key || strings.TrimPrefix(p, "$.") == strings.TrimPrefix(jsonPath, "$.") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func strconvQuote(s string) string {
|
||||
return `"` + s + `"`
|
||||
}
|
||||
59
modules/tide/normalize_test.go
Normal file
59
modules/tide/normalize_test.go
Normal file
@@ -0,0 +1,59 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNormalizeDateIDAndDisable(t *testing.T) {
|
||||
step := Step{ID: "n"}
|
||||
want := []byte(`{"id":1,"created_at":"2026-01-01T00:00:00+00:00","slug":"keep-me"}`)
|
||||
gotOK := []byte(`{"id":9,"created_at":"2026-02-02T00:00:00+00:00","slug":"keep-me"}`)
|
||||
gotZ := []byte(`{"id":9,"created_at":"2026-02-02T00:00:00Z","slug":"keep-me"}`)
|
||||
gotStrID := []byte(`{"id":"9","created_at":"2026-02-02T00:00:00+00:00","slug":"keep-me"}`)
|
||||
gotSlug := []byte(`{"id":9,"created_at":"2026-02-02T00:00:00+00:00","slug":"other"}`)
|
||||
|
||||
if diffs := compareBodies(Response{Headers: jsonCT(), Body: Body(want)}, Response{Headers: jsonCT(), Body: Body(gotOK)}, step); len(diffs) != 0 {
|
||||
t.Fatalf("masked date/id should pass: %+v", diffs)
|
||||
}
|
||||
nullID := []byte(`{"id":1,"discogs_id":null,"created_at":"2026-01-01T00:00:00+00:00","slug":"keep-me"}`)
|
||||
if diffs := compareBodies(Response{Headers: jsonCT(), Body: Body(nullID)}, Response{Headers: jsonCT(), Body: Body(nullID)}, step); len(diffs) != 0 {
|
||||
t.Fatalf("null *_id must pass: %+v", diffs)
|
||||
}
|
||||
syncA := []byte(`{"collection_key":"aaa","checkpoint":"2026-01-01T00:00:00+00:00","total_estimate":0}`)
|
||||
syncB := []byte(`{"collection_key":"bbb","checkpoint":"2026-02-02T00:00:00+00:00","total_estimate":0}`)
|
||||
if diffs := compareBodies(Response{Headers: jsonCT(), Body: Body(syncA)}, Response{Headers: jsonCT(), Body: Body(syncB)}, step); len(diffs) != 0 {
|
||||
t.Fatalf("collection_key/checkpoint must mask: %+v", diffs)
|
||||
}
|
||||
zdiffs := compareBodies(Response{Headers: jsonCT(), Body: Body(want)}, Response{Headers: jsonCT(), Body: Body(gotZ)}, step)
|
||||
if len(zdiffs) == 0 {
|
||||
t.Fatal("Z date must fail")
|
||||
}
|
||||
if !strings.Contains(zdiffs[0].Path, "created_at") {
|
||||
t.Fatalf("Z path %s", zdiffs[0].Path)
|
||||
}
|
||||
idDiffs := compareBodies(Response{Headers: jsonCT(), Body: Body(want)}, Response{Headers: jsonCT(), Body: Body(gotStrID)}, step)
|
||||
if len(idDiffs) == 0 || !strings.Contains(idDiffs[0].Path, "id") {
|
||||
t.Fatalf("string id: %+v", idDiffs)
|
||||
}
|
||||
slugDiffs := compareBodies(Response{Headers: jsonCT(), Body: Body(want)}, Response{Headers: jsonCT(), Body: Body(gotSlug)}, step)
|
||||
if len(slugDiffs) == 0 || !strings.Contains(slugDiffs[0].Path, "slug") {
|
||||
t.Fatalf("slug must stay exact: %+v", slugDiffs)
|
||||
}
|
||||
|
||||
issuedA := []byte(`{"client_id":"aaa","client_id_issued_at":111}`)
|
||||
issuedB := []byte(`{"client_id":"bbb","client_id_issued_at":222}`)
|
||||
if diffs := compareBodies(Response{Headers: jsonCT(), Body: Body(issuedA)}, Response{Headers: jsonCT(), Body: Body(issuedB)}, step); len(diffs) != 0 {
|
||||
t.Fatalf("oauth client_id/issued_at must mask: %+v", diffs)
|
||||
}
|
||||
|
||||
disabled := Step{ID: "n", Normalize: []NormalizeRule{{Path: "created_at", Disable: true}}}
|
||||
dDiffs := compareBodies(Response{Headers: jsonCT(), Body: Body(want)}, Response{Headers: jsonCT(), Body: Body(gotOK)}, disabled)
|
||||
if len(dDiffs) == 0 {
|
||||
t.Fatal("disabled date mask must compare raw timestamps")
|
||||
}
|
||||
}
|
||||
|
||||
func jsonCT() map[string]string {
|
||||
return map[string]string{"Content-Type": "application/json"}
|
||||
}
|
||||
385
modules/tide/proxy.go
Normal file
385
modules/tide/proxy.go
Normal file
@@ -0,0 +1,385 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httputil"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultListen = "127.0.0.1:8422"
|
||||
DefaultUpstream = "http://127.0.0.1:8423"
|
||||
)
|
||||
|
||||
type captureKey struct{}
|
||||
|
||||
type captureState struct {
|
||||
session string
|
||||
method string
|
||||
path string
|
||||
query string
|
||||
reqBody []byte
|
||||
reqHeaders http.Header
|
||||
}
|
||||
|
||||
type sessionBuf struct {
|
||||
name string
|
||||
steps []Step
|
||||
started bool
|
||||
}
|
||||
|
||||
// ProxyConfig pins a loopback reverse proxy that records named sessions.
|
||||
type ProxyConfig struct {
|
||||
Listen string
|
||||
Upstream string
|
||||
Session string
|
||||
Rules Rules
|
||||
RulesPath string
|
||||
VarsPath string
|
||||
Fixtures string
|
||||
MaxBody int64
|
||||
Update bool
|
||||
}
|
||||
|
||||
// Proxy is a fixed-upstream recording reverse proxy.
|
||||
type Proxy struct {
|
||||
cfg ProxyConfig
|
||||
upstream *url.URL
|
||||
rp *httputil.ReverseProxy
|
||||
limit int64
|
||||
store *Store
|
||||
|
||||
mu sync.Mutex
|
||||
sessions map[string]*sessionBuf
|
||||
failed map[string]error
|
||||
}
|
||||
|
||||
// NewProxy validates loopback bind/upstream and builds a recording handler.
|
||||
func NewProxy(cfg ProxyConfig) (*Proxy, error) {
|
||||
if strings.TrimSpace(cfg.Listen) == "" {
|
||||
cfg.Listen = DefaultListen
|
||||
}
|
||||
if strings.TrimSpace(cfg.Upstream) == "" {
|
||||
cfg.Upstream = DefaultUpstream
|
||||
}
|
||||
if err := requireLoopbackAddr(cfg.Listen); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstream, err := parseLoopbackUpstream(cfg.Upstream)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if strings.TrimSpace(cfg.Fixtures) == "" {
|
||||
return nil, fmt.Errorf("tide: proxy fixtures directory is required")
|
||||
}
|
||||
if err := validateRules(cfg.Rules); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := varsOutsideFixtures(cfg.VarsPath, cfg.Fixtures); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
store, err := OpenStore(cfg.VarsPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p := &Proxy{
|
||||
cfg: cfg,
|
||||
upstream: upstream,
|
||||
limit: maxBody(cfg.MaxBody),
|
||||
store: store,
|
||||
sessions: make(map[string]*sessionBuf),
|
||||
failed: make(map[string]error),
|
||||
}
|
||||
p.rp = &httputil.ReverseProxy{
|
||||
Rewrite: func(pr *httputil.ProxyRequest) {
|
||||
pr.SetURL(p.upstream)
|
||||
pr.Out.Host = p.upstream.Host
|
||||
pr.Out.Header.Del(SessionHeader)
|
||||
},
|
||||
ModifyResponse: p.modifyResponse,
|
||||
ErrorHandler: func(w http.ResponseWriter, r *http.Request, err error) {
|
||||
http.Error(w, "tide: upstream error: "+err.Error(), http.StatusBadGateway)
|
||||
},
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
// Handler returns the recording reverse-proxy handler.
|
||||
func (p *Proxy) Handler() http.Handler {
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *Proxy) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
state, err := p.beginCapture(r)
|
||||
if err != nil {
|
||||
status := http.StatusBadRequest
|
||||
if isTruncated(err) {
|
||||
status = http.StatusRequestEntityTooLarge
|
||||
}
|
||||
http.Error(w, err.Error(), status)
|
||||
return
|
||||
}
|
||||
body := state.reqBody
|
||||
r = r.WithContext(context.WithValue(r.Context(), captureKey{}, state))
|
||||
r.Body = io.NopCloser(bytes.NewReader(body))
|
||||
r.ContentLength = int64(len(body))
|
||||
r.GetBody = func() (io.ReadCloser, error) {
|
||||
return io.NopCloser(bytes.NewReader(body)), nil
|
||||
}
|
||||
p.rp.ServeHTTP(w, r)
|
||||
}
|
||||
|
||||
// ListenAndServe binds the loopback listener until ctx is cancelled, then flushes.
|
||||
func (p *Proxy) ListenAndServe(ctx context.Context) error {
|
||||
ln, err := net.Listen("tcp", p.cfg.Listen)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: listen %s: %w", p.cfg.Listen, err)
|
||||
}
|
||||
addr := ln.Addr().String()
|
||||
if err := requireLoopbackAddr(addr); err != nil {
|
||||
_ = ln.Close()
|
||||
return err
|
||||
}
|
||||
srv := &http.Server{
|
||||
Handler: p,
|
||||
ReadHeaderTimeout: 10 * time.Second,
|
||||
}
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
errCh <- srv.Serve(ln)
|
||||
}()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
_ = srv.Shutdown(shutdownCtx)
|
||||
<-errCh
|
||||
return p.Flush()
|
||||
case err := <-errCh:
|
||||
flushErr := p.Flush()
|
||||
if err != nil && err != http.ErrServerClosed {
|
||||
if flushErr != nil {
|
||||
return fmt.Errorf("%v; flush: %w", err, flushErr)
|
||||
}
|
||||
return err
|
||||
}
|
||||
return flushErr
|
||||
}
|
||||
}
|
||||
|
||||
// Flush writes every complete captured session. Failed sessions leave no new fixture.
|
||||
func (p *Proxy) Flush() error {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
var first error
|
||||
for name, buf := range p.sessions {
|
||||
if err := p.failed[name]; err != nil {
|
||||
if first == nil {
|
||||
first = err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := p.writeSessionLocked(buf); err != nil {
|
||||
if first == nil {
|
||||
first = err
|
||||
}
|
||||
}
|
||||
}
|
||||
return first
|
||||
}
|
||||
|
||||
func (p *Proxy) beginCapture(r *http.Request) (*captureState, error) {
|
||||
session := strings.TrimSpace(r.Header.Get(SessionHeader))
|
||||
if session == "" {
|
||||
session = strings.TrimSpace(p.cfg.Session)
|
||||
}
|
||||
if err := validateSessionName(session); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
raw, err := readBounded(r.Body, p.limit)
|
||||
if err != nil {
|
||||
p.failSession(session, fmt.Errorf("tide: session %q: %w", session, err))
|
||||
return nil, err
|
||||
}
|
||||
return &captureState{
|
||||
session: session,
|
||||
method: r.Method,
|
||||
path: r.URL.Path,
|
||||
query: r.URL.RawQuery,
|
||||
reqBody: raw,
|
||||
reqHeaders: r.Header.Clone(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (p *Proxy) modifyResponse(resp *http.Response) error {
|
||||
state, _ := resp.Request.Context().Value(captureKey{}).(*captureState)
|
||||
if state == nil {
|
||||
return fmt.Errorf("tide: missing capture state")
|
||||
}
|
||||
raw, err := readBounded(resp.Body, p.limit)
|
||||
if err != nil {
|
||||
resp.Body.Close()
|
||||
p.failSession(state.session, fmt.Errorf("tide: session %q: %w", state.session, err))
|
||||
return err
|
||||
}
|
||||
resp.Body = io.NopCloser(bytes.NewReader(raw))
|
||||
resp.ContentLength = int64(len(raw))
|
||||
resp.Header.Set("Content-Length", fmt.Sprintf("%d", len(raw)))
|
||||
if err := p.recordStep(state, resp, raw); err != nil {
|
||||
p.failSession(state.session, err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *Proxy) recordStep(state *captureState, resp *http.Response, respBody []byte) error {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if err := p.failed[state.session]; err != nil {
|
||||
return err
|
||||
}
|
||||
buf, err := p.sessionLocked(state.session)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
route := p.cfg.Rules.Match(state.method, state.path)
|
||||
reqHeaders := filterHeaders(state.reqHeaders, p.cfg.Rules.requestHeaders(route))
|
||||
respHeaders := filterHeaders(resp.Header, p.cfg.Rules.responseHeaders(route))
|
||||
step := Step{
|
||||
ID: fmt.Sprintf("%d", len(buf.steps)+1),
|
||||
Request: Request{
|
||||
Method: state.method,
|
||||
Path: state.path,
|
||||
Query: state.query,
|
||||
Headers: reqHeaders,
|
||||
Body: Body(state.reqBody),
|
||||
},
|
||||
Response: Response{
|
||||
Status: resp.StatusCode,
|
||||
Headers: respHeaders,
|
||||
Body: Body(respBody),
|
||||
},
|
||||
}
|
||||
if route != nil {
|
||||
step.Capture = append([]CaptureRule(nil), route.Capture...)
|
||||
}
|
||||
if err := CaptureStep(p.store, &step); err != nil {
|
||||
return fmt.Errorf("tide: session %q: %w", state.session, err)
|
||||
}
|
||||
if err := ScrubStep(p.store, &step); err != nil {
|
||||
return fmt.Errorf("tide: session %q: %w", state.session, err)
|
||||
}
|
||||
if err := p.store.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
buf.steps = append(buf.steps, step)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *Proxy) sessionLocked(name string) (*sessionBuf, error) {
|
||||
if buf, ok := p.sessions[name]; ok {
|
||||
return buf, nil
|
||||
}
|
||||
dest := p.sessionPath(name)
|
||||
if !p.cfg.Update {
|
||||
if _, err := os.Stat(dest); err == nil {
|
||||
return nil, fmt.Errorf("tide: duplicate session %q", name)
|
||||
}
|
||||
}
|
||||
buf := &sessionBuf{name: name}
|
||||
p.sessions[name] = buf
|
||||
return buf, nil
|
||||
}
|
||||
|
||||
func (p *Proxy) writeSessionLocked(buf *sessionBuf) error {
|
||||
if buf == nil || len(buf.steps) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := p.failed[buf.name]; err != nil {
|
||||
return err
|
||||
}
|
||||
flow := Flow{
|
||||
Version: CurrentVersion,
|
||||
Name: buf.name,
|
||||
Steps: append([]Step(nil), buf.steps...),
|
||||
}
|
||||
dest := p.sessionPath(buf.name)
|
||||
if p.cfg.Update {
|
||||
return SaveFlow(dest, flow)
|
||||
}
|
||||
return SaveFlowExclusive(dest, flow)
|
||||
}
|
||||
|
||||
func (p *Proxy) sessionPath(name string) string {
|
||||
return filepath.Join(p.cfg.Fixtures, p.cfg.Rules.Client, name+".yaml")
|
||||
}
|
||||
|
||||
func (p *Proxy) failSession(name string, err error) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if _, ok := p.failed[name]; !ok {
|
||||
p.failed[name] = err
|
||||
}
|
||||
_ = os.Remove(p.sessionPath(name))
|
||||
}
|
||||
|
||||
func validateSessionName(name string) error {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return fmt.Errorf("tide: session name is required (X-Parity-Session or --session)")
|
||||
}
|
||||
if strings.ContainsAny(name, `/\`) || strings.Contains(name, "..") || name != filepath.Base(name) {
|
||||
return fmt.Errorf("tide: session name %q must not contain path separators", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func requireLoopbackAddr(hostport string) error {
|
||||
host, _, err := net.SplitHostPort(hostport)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: listen %q: %w", hostport, err)
|
||||
}
|
||||
if !isLoopbackHost(host) {
|
||||
return fmt.Errorf("tide: listen %q must be loopback", hostport)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseLoopbackUpstream(raw string) (*url.URL, error) {
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tide: upstream: %w", err)
|
||||
}
|
||||
if u.Scheme != "http" {
|
||||
return nil, fmt.Errorf("tide: upstream %q must be plain http on loopback", raw)
|
||||
}
|
||||
if u.Host == "" || !isLoopbackHost(u.Hostname()) {
|
||||
return nil, fmt.Errorf("tide: upstream %q must be loopback", raw)
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
func isLoopbackHost(host string) bool {
|
||||
if host == "" {
|
||||
return false
|
||||
}
|
||||
if strings.EqualFold(host, "localhost") {
|
||||
return true
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
return ip != nil && ip.IsLoopback()
|
||||
}
|
||||
|
||||
func isTruncated(err error) bool {
|
||||
return err != nil && strings.Contains(err.Error(), "exceeds")
|
||||
}
|
||||
427
modules/tide/proxy_security_test.go
Normal file
427
modules/tide/proxy_security_test.go
Normal file
@@ -0,0 +1,427 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestProxySecurity(t *testing.T) {
|
||||
t.Run("rejects non-loopback bind and client-selected upstream", func(t *testing.T) {
|
||||
fixtures := t.TempDir()
|
||||
rules := mustParseRules(t, testRulesYAML())
|
||||
if _, err := NewProxy(ProxyConfig{Listen: "0.0.0.0:8422", Upstream: DefaultUpstream, Fixtures: fixtures, Rules: rules}); err == nil || !strings.Contains(err.Error(), "loopback") {
|
||||
t.Fatalf("non-loopback listen: %v", err)
|
||||
}
|
||||
if _, err := NewProxy(ProxyConfig{Listen: DefaultListen, Upstream: "http://example.com", Fixtures: fixtures, Rules: rules}); err == nil || !strings.Contains(err.Error(), "loopback") {
|
||||
t.Fatalf("non-loopback upstream: %v", err)
|
||||
}
|
||||
if _, err := NewProxy(ProxyConfig{Listen: DefaultListen, Upstream: "https://127.0.0.1:8423", Fixtures: fixtures, Rules: rules}); err == nil || !strings.Contains(err.Error(), "http") {
|
||||
t.Fatalf("https upstream: %v", err)
|
||||
}
|
||||
|
||||
var seenHost []string
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
seenHost = append(seenHost, r.Host)
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML())
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Host = "evil.example"
|
||||
req.Header.Set(SessionHeader, "pin")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
t.Fatalf("status %d", resp.StatusCode)
|
||||
}
|
||||
wantHost := strings.TrimPrefix(upstream.URL, "http://")
|
||||
if len(seenHost) != 1 || seenHost[0] != wantHost {
|
||||
t.Fatalf("upstream host %v want %s", seenHost, wantHost)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("request and response size caps leave no fixture", func(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
_, _ = w.Write([]byte("hello-world"))
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
fixtures := t.TempDir()
|
||||
proxy, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: upstream.URL,
|
||||
Fixtures: fixtures,
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
MaxBody: 4,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
req, err := http.NewRequest(http.MethodPost, srv.URL+"/sample", strings.NewReader("12345"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, "reqcap")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusRequestEntityTooLarge {
|
||||
t.Fatalf("request overflow status %d", resp.StatusCode)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(fixtures, "nuxt", "reqcap.yaml")); !os.IsNotExist(err) {
|
||||
t.Fatalf("request overflow committed fixture: %v", err)
|
||||
}
|
||||
|
||||
get, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
get.Header.Set(SessionHeader, "respcap")
|
||||
gresp, err := http.DefaultClient.Do(get)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
gresp.Body.Close()
|
||||
if gresp.StatusCode == http.StatusOK {
|
||||
t.Fatal("response overflow recorded as success")
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(fixtures, "nuxt", "respcap.yaml")); !os.IsNotExist(err) {
|
||||
t.Fatalf("response overflow committed fixture: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("cookie redirect passthrough and concurrent sessions", func(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/go" {
|
||||
w.Header().Set("Set-Cookie", "sid=abc; Path=/")
|
||||
w.Header().Set("Location", "/landed")
|
||||
w.WriteHeader(http.StatusFound)
|
||||
_, _ = w.Write([]byte("redirect-body"))
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
_, _ = w.Write([]byte("cookie=" + r.Header.Get("Cookie") + " path=" + r.URL.Path))
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
fixtures := t.TempDir()
|
||||
proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML())
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
client := &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
}}
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/go", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, "redir")
|
||||
req.Header.Set("Cookie", "keep=1")
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusFound || resp.Header.Get("Location") != "/landed" {
|
||||
t.Fatalf("redirect %d %s", resp.StatusCode, resp.Header.Get("Location"))
|
||||
}
|
||||
if !strings.Contains(resp.Header.Get("Set-Cookie"), "sid=abc") || string(body) != "redirect-body" {
|
||||
t.Fatalf("cookie/body passthrough cookie=%q body=%q", resp.Header.Get("Set-Cookie"), body)
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
errCh := make(chan error, 2)
|
||||
for _, name := range []string{"alpha", "beta"} {
|
||||
wg.Add(1)
|
||||
go func(session string) {
|
||||
defer wg.Done()
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/"+session, nil)
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
req.Header.Set(SessionHeader, session)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
errCh <- fmt.Errorf("session %s status %d", session, resp.StatusCode)
|
||||
}
|
||||
}(name)
|
||||
}
|
||||
wg.Wait()
|
||||
close(errCh)
|
||||
for err := range errCh {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := proxy.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
alpha, err := LoadFlow(filepath.Join(fixtures, "nuxt", "alpha.yaml"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
beta, err := LoadFlow(filepath.Join(fixtures, "nuxt", "beta.yaml"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(alpha.Steps) != 1 || alpha.Steps[0].Request.Path != "/alpha" {
|
||||
t.Fatalf("alpha: %+v", alpha.Steps)
|
||||
}
|
||||
if len(beta.Steps) != 1 || beta.Steps[0].Request.Path != "/beta" {
|
||||
t.Fatalf("beta: %+v", beta.Steps)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("path traversal session and sidecar digest", func(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
fixtures := t.TempDir()
|
||||
proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML())
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, "../escape")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusOK {
|
||||
t.Fatal("traversal session must fail")
|
||||
}
|
||||
if entries, _ := os.ReadDir(filepath.Dir(fixtures)); len(entries) == 0 {
|
||||
t.Fatal("temp fixtures dir vanished")
|
||||
}
|
||||
unsafe := "version: 1\nname: bin\nsteps:\n - id: a\n request:\n method: GET\n path: /bin\n response:\n body_file: ../secret.bin\n"
|
||||
p := filepath.Join(fixtures, "unsafe.yaml")
|
||||
if err := os.WriteFile(p, []byte(unsafe), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := LoadFlow(p); err == nil {
|
||||
t.Fatal("parent sidecar must fail")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("scrubs credentials across header cookie query json and form", func(t *testing.T) {
|
||||
const (
|
||||
invTok = "inv_abcd1234xyz"
|
||||
secret = "superSecretValue99"
|
||||
code = "oauthCodeValue99"
|
||||
cookieV = "cookieSecretValue99"
|
||||
)
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
case r.URL.Path == "/login":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Header().Set("Set-Cookie", "auth_token="+testJWT)
|
||||
_, _ = w.Write([]byte(`{"token":"` + testJWT + `"}`))
|
||||
case r.URL.Path == "/token":
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
case r.URL.Path == "/authorize":
|
||||
w.Header().Set("Location", "/cb?code="+code)
|
||||
w.WriteHeader(http.StatusFound)
|
||||
default:
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
}
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
fixtures := t.TempDir()
|
||||
varsPath := filepath.Join(t.TempDir(), "vars.yaml")
|
||||
rules := mustParseRules(t, ""+
|
||||
"client: nuxt\n"+
|
||||
"keep_request_headers:\n"+
|
||||
" - Authorization\n"+
|
||||
" - Cookie\n"+
|
||||
" - Content-Type\n"+
|
||||
"keep_response_headers:\n"+
|
||||
" - Content-Type\n"+
|
||||
" - Location\n"+
|
||||
" - Set-Cookie\n"+
|
||||
"routes:\n"+
|
||||
" - method: POST\n"+
|
||||
" path: /login\n"+
|
||||
" capture:\n"+
|
||||
" - from: response.json\n"+
|
||||
" path: $.token\n"+
|
||||
" as: jwt:alice\n"+
|
||||
" category: jwt\n"+
|
||||
" - method: GET\n"+
|
||||
" path: /me\n"+
|
||||
" - method: GET\n"+
|
||||
" path: /api\n"+
|
||||
" capture:\n"+
|
||||
" - from: request.query\n"+
|
||||
" name: token\n"+
|
||||
" as: token:mcp-read\n"+
|
||||
" category: token\n"+
|
||||
" - method: POST\n"+
|
||||
" path: /token\n"+
|
||||
" capture:\n"+
|
||||
" - from: request.form\n"+
|
||||
" name: client_secret\n"+
|
||||
" as: oauth:secret\n"+
|
||||
" category: oauth_secret\n"+
|
||||
" - method: GET\n"+
|
||||
" path: /authorize\n"+
|
||||
" capture:\n"+
|
||||
" - from: response.location.query\n"+
|
||||
" name: code\n"+
|
||||
" as: oauth:code\n"+
|
||||
" category: oauth_code\n"+
|
||||
" - method: GET\n"+
|
||||
" path: /cookie\n"+
|
||||
" capture:\n"+
|
||||
" - from: request.header\n"+
|
||||
" name: Cookie\n"+
|
||||
" as: cookie:auth\n"+
|
||||
" category: cookie\n")
|
||||
proxy, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: upstream.URL,
|
||||
Fixtures: fixtures,
|
||||
VarsPath: varsPath,
|
||||
Rules: rules,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
post(t, srv.URL+"/login", "creds", "application/json", `{}`)
|
||||
me, err := http.NewRequest(http.MethodGet, srv.URL+"/me", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
me.Header.Set(SessionHeader, "creds")
|
||||
me.Header.Set("Authorization", "Bearer "+testJWT)
|
||||
meResp, err := http.DefaultClient.Do(me)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
meResp.Body.Close()
|
||||
|
||||
q, err := http.NewRequest(http.MethodGet, srv.URL+"/api?token="+invTok, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
q.Header.Set(SessionHeader, "creds")
|
||||
qResp, err := http.DefaultClient.Do(q)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
qResp.Body.Close()
|
||||
|
||||
post(t, srv.URL+"/token", "creds", "application/x-www-form-urlencoded", "client_secret="+secret)
|
||||
get, err := http.NewRequest(http.MethodGet, srv.URL+"/authorize", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
get.Header.Set(SessionHeader, "creds")
|
||||
client := &http.Client{CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
}}
|
||||
aResp, err := client.Do(get)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
aResp.Body.Close()
|
||||
|
||||
ck, err := http.NewRequest(http.MethodGet, srv.URL+"/cookie", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ck.Header.Set(SessionHeader, "creds")
|
||||
ck.Header.Set("Cookie", "auth_token="+cookieV)
|
||||
cResp, err := http.DefaultClient.Do(ck)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cResp.Body.Close()
|
||||
|
||||
if err := proxy.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
raw, err := os.ReadFile(filepath.Join(fixtures, "nuxt", "creds.yaml"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
text := string(raw)
|
||||
for _, secret := range []string{testJWT, invTok, secret, code, cookieV, "auth_token=" + cookieV} {
|
||||
if strings.Contains(text, secret) {
|
||||
t.Fatalf("secret %q leaked into YAML:\n%s", secret, text)
|
||||
}
|
||||
}
|
||||
for _, ph := range []string{"{{jwt:alice}}", "{{token:mcp-read}}", "{{oauth:secret}}", "{{oauth:code}}", "{{cookie:auth}}"} {
|
||||
if !strings.Contains(text, ph) {
|
||||
t.Fatalf("missing placeholder %s:\n%s", ph, text)
|
||||
}
|
||||
}
|
||||
|
||||
unclassified := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"token":"` + testJWT + `"}`))
|
||||
}))
|
||||
t.Cleanup(unclassified.Close)
|
||||
badFix := t.TempDir()
|
||||
badProxy, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: unclassified.URL,
|
||||
Fixtures: badFix,
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
badSrv := httptest.NewServer(badProxy.Handler())
|
||||
t.Cleanup(badSrv.Close)
|
||||
breq, err := http.NewRequest(http.MethodGet, badSrv.URL+"/sample", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
breq.Header.Set(SessionHeader, "leak")
|
||||
bresp, err := http.DefaultClient.Do(breq)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bresp.Body.Close()
|
||||
if _, err := os.Stat(filepath.Join(badFix, "nuxt", "leak.yaml")); !os.IsNotExist(err) {
|
||||
t.Fatalf("unclassified jwt fixture committed: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
419
modules/tide/proxy_test.go
Normal file
419
modules/tide/proxy_test.go
Normal file
@@ -0,0 +1,419 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestProxyNamedSessionsOrderedAndIsolated(t *testing.T) {
|
||||
var seenHost []string
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
seenHost = append(seenHost, r.Host)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte(`{"path":"` + r.URL.Path + `"}`))
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
|
||||
fixtures := t.TempDir()
|
||||
proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML())
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
doProxy(t, srv.URL, "alpha", "/one", "")
|
||||
doProxy(t, srv.URL, "alpha", "/two", "")
|
||||
doProxy(t, srv.URL, "beta", "/other", "")
|
||||
|
||||
if err := proxy.Flush(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
alpha, err := LoadFlow(filepath.Join(fixtures, "nuxt", "alpha.yaml"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(alpha.Steps) != 2 {
|
||||
t.Fatalf("alpha steps: %d", len(alpha.Steps))
|
||||
}
|
||||
if alpha.Steps[0].Request.Path != "/one" || alpha.Steps[1].Request.Path != "/two" {
|
||||
t.Fatalf("alpha order: %+v", alpha.Steps)
|
||||
}
|
||||
beta, err := LoadFlow(filepath.Join(fixtures, "nuxt", "beta.yaml"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(beta.Steps) != 1 || beta.Steps[0].Request.Path != "/other" {
|
||||
t.Fatalf("beta interleaved: %+v", beta.Steps)
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/evil", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Host = "evil.example"
|
||||
req.Header.Set(SessionHeader, "gamma")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if len(seenHost) < 4 {
|
||||
t.Fatalf("upstream hits: %d", len(seenHost))
|
||||
}
|
||||
wantHost := strings.TrimPrefix(upstream.URL, "http://")
|
||||
if seenHost[len(seenHost)-1] != wantHost {
|
||||
t.Fatalf("client Host leaked: last=%s want=%s all=%v", seenHost[len(seenHost)-1], wantHost, seenHost)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyForwardsCookieRedirectAndBody(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/go" {
|
||||
w.Header().Set("Set-Cookie", "sid=abc; Path=/")
|
||||
w.Header().Set("Location", "/landed")
|
||||
w.WriteHeader(http.StatusFound)
|
||||
_, _ = w.Write([]byte("redirect-body"))
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte("cookie=" + r.Header.Get("Cookie")))
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
|
||||
fixtures := t.TempDir()
|
||||
proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML())
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
client := &http.Client{
|
||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
}
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/go", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, "redir")
|
||||
req.Header.Set("Cookie", "keep=1")
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusFound {
|
||||
t.Fatalf("status %d", resp.StatusCode)
|
||||
}
|
||||
if loc := resp.Header.Get("Location"); loc != "/landed" {
|
||||
t.Fatalf("location %q", loc)
|
||||
}
|
||||
if !strings.Contains(resp.Header.Get("Set-Cookie"), "sid=abc") {
|
||||
t.Fatalf("set-cookie %q", resp.Header.Get("Set-Cookie"))
|
||||
}
|
||||
if string(body) != "redirect-body" {
|
||||
t.Fatalf("body %q", body)
|
||||
}
|
||||
|
||||
req2, err := http.NewRequest(http.MethodGet, srv.URL+"/echo", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req2.Header.Set(SessionHeader, "redir")
|
||||
req2.Header.Set("Cookie", "keep=1")
|
||||
resp2, err := client.Do(req2)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, _ := io.ReadAll(resp2.Body)
|
||||
resp2.Body.Close()
|
||||
if string(got) != "cookie=keep=1" {
|
||||
t.Fatalf("cookie not forwarded: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyRejectsNonLoopbackOverflowAndCredentials(t *testing.T) {
|
||||
if _, err := NewProxy(ProxyConfig{
|
||||
Listen: "0.0.0.0:8422",
|
||||
Upstream: "http://127.0.0.1:8423",
|
||||
Fixtures: t.TempDir(),
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
}); err == nil || !strings.Contains(err.Error(), "loopback") {
|
||||
t.Fatalf("non-loopback listen: %v", err)
|
||||
}
|
||||
if _, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:8422",
|
||||
Upstream: "http://example.com",
|
||||
Fixtures: t.TempDir(),
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
}); err == nil || !strings.Contains(err.Error(), "loopback") {
|
||||
t.Fatalf("non-loopback upstream: %v", err)
|
||||
}
|
||||
|
||||
upstream := 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-world"))
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
|
||||
fixtures := t.TempDir()
|
||||
cfg := ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: upstream.URL,
|
||||
Fixtures: fixtures,
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
MaxBody: 4,
|
||||
}
|
||||
proxy, err := NewProxy(cfg)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, "overflow")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusRequestEntityTooLarge && resp.StatusCode != http.StatusBadGateway && resp.StatusCode != http.StatusInternalServerError {
|
||||
if resp.StatusCode == http.StatusOK {
|
||||
t.Fatal("overflow recorded as success")
|
||||
}
|
||||
}
|
||||
overflowPath := filepath.Join(fixtures, "nuxt", "overflow.yaml")
|
||||
if _, statErr := os.Stat(overflowPath); !os.IsNotExist(statErr) {
|
||||
t.Fatalf("partial fixture committed: %v", statErr)
|
||||
}
|
||||
|
||||
bigUp := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxIn0.abcabcabcabcabcabcab"}`))
|
||||
}))
|
||||
t.Cleanup(bigUp.Close)
|
||||
secretFix := t.TempDir()
|
||||
secretProxy, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: bigUp.URL,
|
||||
Fixtures: secretFix,
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
secretSrv := httptest.NewServer(secretProxy.Handler())
|
||||
t.Cleanup(secretSrv.Close)
|
||||
sreq, err := http.NewRequest(http.MethodGet, secretSrv.URL+"/sample", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sreq.Header.Set(SessionHeader, "secret")
|
||||
sresp, err := http.DefaultClient.Do(sreq)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sresp.Body.Close()
|
||||
if _, statErr := os.Stat(filepath.Join(secretFix, "nuxt", "secret.yaml")); !os.IsNotExist(statErr) {
|
||||
t.Fatalf("credential fixture was committed: %v", statErr)
|
||||
}
|
||||
|
||||
reqBad, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
reqBad.Header.Set(SessionHeader, "alice/../bob")
|
||||
bresp, err := http.DefaultClient.Do(reqBad)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bresp.Body.Close()
|
||||
if bresp.StatusCode == http.StatusOK {
|
||||
t.Fatal("path-separator session must fail")
|
||||
}
|
||||
|
||||
inside := filepath.Join(fixtures, "vars.yaml")
|
||||
if _, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: upstream.URL,
|
||||
Fixtures: fixtures,
|
||||
VarsPath: inside,
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
}); err == nil || !strings.Contains(err.Error(), "outside") {
|
||||
t.Fatalf("vars inside fixtures: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(inside, []byte("{}\n"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
linkDir := t.TempDir()
|
||||
link := filepath.Join(linkDir, "vars-link.yaml")
|
||||
if err := os.Symlink(inside, link); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: upstream.URL,
|
||||
Fixtures: fixtures,
|
||||
VarsPath: link,
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
}); err == nil || !strings.Contains(err.Error(), "outside") {
|
||||
t.Fatalf("symlink vars into fixtures: %v", err)
|
||||
}
|
||||
outside := filepath.Join(t.TempDir(), "vars.yaml")
|
||||
p, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: upstream.URL,
|
||||
Fixtures: fixtures,
|
||||
VarsPath: outside,
|
||||
Rules: mustParseRules(t, testRulesYAML()),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = p
|
||||
st, err := os.Stat(outside)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if st.Mode().Perm() != 0o600 {
|
||||
t.Fatalf("vars mode %o", st.Mode().Perm())
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyDuplicateSessionName(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
fixtures := t.TempDir()
|
||||
if err := os.MkdirAll(filepath.Join(fixtures, "nuxt"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(fixtures, "nuxt", "taken.yaml"), []byte("version: 1\n"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML())
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/sample", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, "taken")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if resp.StatusCode == http.StatusOK {
|
||||
t.Fatal("duplicate session must fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyFailedSessionLeavesNoFixture(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if r.URL.Path == "/ok" {
|
||||
_, _ = w.Write([]byte(`{"ok":true}`))
|
||||
return
|
||||
}
|
||||
_, _ = w.Write([]byte(`{"token":"` + testJWT + `"}`))
|
||||
}))
|
||||
t.Cleanup(upstream.Close)
|
||||
fixtures := t.TempDir()
|
||||
proxy := newTestProxy(t, upstream.URL, fixtures, testRulesYAML())
|
||||
srv := httptest.NewServer(proxy.Handler())
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
doProxy(t, srv.URL, "partial", "/ok", "")
|
||||
req, err := http.NewRequest(http.MethodGet, srv.URL+"/leak", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, "partial")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp.Body.Close()
|
||||
if err := proxy.Flush(); err == nil {
|
||||
t.Fatal("failed session flush must surface the capture error")
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(fixtures, "nuxt", "partial.yaml")); !os.IsNotExist(err) {
|
||||
t.Fatalf("partial session fixture committed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func newTestProxy(t *testing.T, upstream, fixtures, rulesYAML string) *Proxy {
|
||||
t.Helper()
|
||||
p, err := NewProxy(ProxyConfig{
|
||||
Listen: "127.0.0.1:0",
|
||||
Upstream: upstream,
|
||||
Fixtures: fixtures,
|
||||
Rules: mustParseRules(t, rulesYAML),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func mustParseRules(t *testing.T, raw string) Rules {
|
||||
t.Helper()
|
||||
rules, err := ParseRules([]byte(raw))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return rules
|
||||
}
|
||||
|
||||
func testRulesYAML() string {
|
||||
return "" +
|
||||
"client: nuxt\n" +
|
||||
"keep_request_headers:\n" +
|
||||
" - Cookie\n" +
|
||||
" - Content-Type\n" +
|
||||
"keep_response_headers:\n" +
|
||||
" - Content-Type\n" +
|
||||
" - Location\n" +
|
||||
" - Set-Cookie\n" +
|
||||
"routes:\n" +
|
||||
" - method: GET\n" +
|
||||
" path: /sample\n" +
|
||||
" - method: GET\n" +
|
||||
" path: /*\n"
|
||||
}
|
||||
|
||||
func doProxy(t *testing.T, base, session, path, body string) {
|
||||
t.Helper()
|
||||
var rdr io.Reader
|
||||
if body != "" {
|
||||
rdr = strings.NewReader(body)
|
||||
}
|
||||
req, err := http.NewRequest(http.MethodGet, base+path, rdr)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set(SessionHeader, session)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
b, _ := io.ReadAll(resp.Body)
|
||||
t.Fatalf("%s %s: %d %s", session, path, resp.StatusCode, b)
|
||||
}
|
||||
}
|
||||
147
modules/tide/record.go
Normal file
147
modules/tide/record.go
Normal file
@@ -0,0 +1,147 @@
|
||||
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
|
||||
}
|
||||
104
modules/tide/replay.go
Normal file
104
modules/tide/replay.go
Normal file
@@ -0,0 +1,104 @@
|
||||
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)
|
||||
store := cfg.Store
|
||||
if store == nil {
|
||||
store = mustMemoryStore()
|
||||
}
|
||||
result := Result{OK: true, Steps: make([]StepResult, 0, len(flow.Steps))}
|
||||
skipRest := false
|
||||
for _, step := range flow.Steps {
|
||||
if skipRest {
|
||||
result.Steps = append(result.Steps, StepResult{ID: step.ID, Skipped: true})
|
||||
continue
|
||||
}
|
||||
sr := StepResult{ID: step.ID, OK: true}
|
||||
want := step
|
||||
if err := materializeSidecar(cfg.BaseDir, &want.Response); err != nil {
|
||||
return result, fmt.Errorf("tide: replay step %s: %w", step.ID, err)
|
||||
}
|
||||
req, err := expandRequest(step.Request, store)
|
||||
if err != nil {
|
||||
sr.OK = false
|
||||
result.OK = false
|
||||
sr.Diffs = append(sr.Diffs, Diff{Path: "request", Expected: "resolved placeholders", Actual: err.Error()})
|
||||
result.Steps = append(result.Steps, sr)
|
||||
skipRest = true
|
||||
continue
|
||||
}
|
||||
got, err := doStep(ctx, client, cfg.Target, req, limit)
|
||||
if err != nil {
|
||||
return result, fmt.Errorf("tide: replay step %s: %w", step.ID, err)
|
||||
}
|
||||
live := step
|
||||
live.Request = req
|
||||
live.Response = got
|
||||
if err := CaptureStep(store, &live); err != nil {
|
||||
sr.OK = false
|
||||
result.OK = false
|
||||
sr.Diffs = append(sr.Diffs, Diff{Path: "capture", Expected: "captured value", Actual: err.Error()})
|
||||
result.Steps = append(result.Steps, sr)
|
||||
skipRest = true
|
||||
continue
|
||||
}
|
||||
if err := store.Save(); err != nil {
|
||||
return result, fmt.Errorf("tide: replay step %s: %w", step.ID, err)
|
||||
}
|
||||
want.Response, err = expandResponse(want.Response, store)
|
||||
if err != nil {
|
||||
sr.OK = false
|
||||
result.OK = false
|
||||
sr.Diffs = append(sr.Diffs, Diff{Path: "expected", Expected: "resolved placeholders", Actual: err.Error()})
|
||||
result.Steps = append(result.Steps, sr)
|
||||
skipRest = true
|
||||
continue
|
||||
}
|
||||
sr.Diffs = append(sr.Diffs, compareStep(want, live.Response)...)
|
||||
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
|
||||
}
|
||||
|
||||
func mustMemoryStore() *Store {
|
||||
s, _ := OpenStore("")
|
||||
return s
|
||||
}
|
||||
|
||||
func compareStep(want Step, got Response) []Diff {
|
||||
var diffs []Diff
|
||||
if want.Response.Status != 0 && got.Status != want.Response.Status {
|
||||
diffs = append(diffs, Diff{
|
||||
Path: "status",
|
||||
Expected: fmt.Sprintf("%d", want.Response.Status),
|
||||
Actual: fmt.Sprintf("%d", got.Status),
|
||||
})
|
||||
}
|
||||
diffs = append(diffs, compareHeaders(want.Response.Headers, got.Headers, want.Headers)...)
|
||||
diffs = append(diffs, compareBodies(want.Response, got, want)...)
|
||||
return diffs
|
||||
}
|
||||
53
modules/tide/report.go
Normal file
53
modules/tide/report.go
Normal file
@@ -0,0 +1,53 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// Coverage is the recorded/passing/failing/unrecorded table over manifest routes.
|
||||
type Coverage struct {
|
||||
Total int
|
||||
Recorded int
|
||||
Passing int
|
||||
Failing int
|
||||
Unrecorded int
|
||||
Pending int
|
||||
Ported int
|
||||
ResumeRemaining int
|
||||
Rows []CoverageRow
|
||||
Diffs []string
|
||||
}
|
||||
|
||||
// CoverageRow is one manifest route's outcome.
|
||||
type CoverageRow struct {
|
||||
ID string
|
||||
Status string
|
||||
Outcome string
|
||||
}
|
||||
|
||||
// CoverageHeaders is the table header for CLI output.
|
||||
func CoverageHeaders() []string {
|
||||
return []string{"route", "status", "outcome"}
|
||||
}
|
||||
|
||||
// CoverageRows renders the coverage table.
|
||||
func (c Coverage) CoverageRows() [][]string {
|
||||
rows := make([][]string, 0, len(c.Rows))
|
||||
for _, r := range c.Rows {
|
||||
rows = append(rows, []string{r.ID, r.Status, r.Outcome})
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
// SummaryLine is a one-line recorded/passing/failing/unrecorded report.
|
||||
func (c Coverage) SummaryLine() string {
|
||||
return fmt.Sprintf("recorded %d/%d passing %d failing %d unrecorded %d", c.Recorded, c.Total, c.Passing, c.Failing, c.Unrecorded)
|
||||
}
|
||||
|
||||
func (c Coverage) ResumeLine() string {
|
||||
if c.ResumeRemaining <= 0 {
|
||||
return "manifest recording complete"
|
||||
}
|
||||
return "resume remaining " + strconv.Itoa(c.ResumeRemaining)
|
||||
}
|
||||
245
modules/tide/roundtrip_test.go
Normal file
245
modules/tide/roundtrip_test.go
Normal file
@@ -0,0 +1,245 @@
|
||||
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)
|
||||
}
|
||||
if !strings.Contains(msg, "hello") || !strings.Contains(msg, "hallo") {
|
||||
t.Fatalf("byte mismatch missing printable bytes: %s", msg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("json key order ignored missing keys and types", func(t *testing.T) {
|
||||
orig := jsonServer(t, `{"z":1,"a":"x"}`)
|
||||
reordered := jsonServer(t, `{"a":"x","z":1}`)
|
||||
missing := jsonServer(t, `{"a":"x"}`)
|
||||
wrongType := jsonServer(t, `{"z":"1","a":"x"}`)
|
||||
|
||||
recorded, err := RecordFlow(ctx, spec, RecordConfig{Target: orig.URL})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := ReplayFlow(ctx, recorded, ReplayConfig{Target: reordered.URL}); err != nil {
|
||||
t.Fatalf("reordered keys must pass: %v", err)
|
||||
}
|
||||
_, err = ReplayFlow(ctx, recorded, ReplayConfig{Target: missing.URL})
|
||||
if err == nil || !strings.Contains(err.Error(), "$.z") {
|
||||
t.Fatalf("missing key must fail at $.z, got %v", err)
|
||||
}
|
||||
_, err = ReplayFlow(ctx, recorded, ReplayConfig{Target: wrongType.URL})
|
||||
if err == nil || !strings.Contains(err.Error(), "$.z") {
|
||||
t.Fatalf("number vs string must fail at $.z, got %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "1") {
|
||||
t.Fatalf("type mismatch missing values: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("rejects unknown fields empty names duplicates and unsafe paths", func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
writeFlow := func(name, body string) string {
|
||||
t.Helper()
|
||||
path := filepath.Join(dir, name)
|
||||
if err := os.WriteFile(path, []byte(body), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
base := "version: 1\nname: sample\nsteps:\n - id: a\n request:\n method: GET\n path: /sample\n"
|
||||
if _, err := LoadFlow(writeFlow("unknown.yaml", base+"extra: true\n")); err == nil {
|
||||
t.Fatal("unknown field must fail")
|
||||
}
|
||||
if _, err := LoadFlow(writeFlow("empty-name.yaml", "version: 1\nname: \"\"\nsteps:\n - id: a\n request:\n method: GET\n path: /sample\n")); err == nil || !strings.Contains(err.Error(), "name") {
|
||||
t.Fatalf("empty name must fail, got %v", err)
|
||||
}
|
||||
dup := "version: 1\nname: sample\nsteps:\n - id: a\n request:\n method: GET\n path: /a\n - id: a\n request:\n method: GET\n path: /b\n"
|
||||
if _, err := LoadFlow(writeFlow("dup.yaml", dup)); err == nil || !strings.Contains(err.Error(), "duplicate") {
|
||||
t.Fatalf("duplicate step id must fail, got %v", err)
|
||||
}
|
||||
unsafe := base + " response:\n body_file: ../secret.bin\n"
|
||||
if _, err := LoadFlow(writeFlow("unsafe.yaml", unsafe)); err == nil || !strings.Contains(err.Error(), "body_file") {
|
||||
t.Fatalf("parent body_file must fail, got %v", err)
|
||||
}
|
||||
abs := base + " response:\n body_file: /tmp/secret.bin\n"
|
||||
if _, err := LoadFlow(writeFlow("abs.yaml", abs)); err == nil || !strings.Contains(err.Error(), "body_file") {
|
||||
t.Fatalf("absolute body_file must fail, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("oversized body fails before fixture commit", func(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
_, _ = w.Write([]byte("hello"))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
out := filepath.Join(t.TempDir(), "too-big.yaml")
|
||||
_, err := RecordFlow(ctx, spec, RecordConfig{Target: srv.URL, MaxBody: 4})
|
||||
if err == nil || !strings.Contains(err.Error(), "exceeds") {
|
||||
t.Fatalf("truncated body must fail, got %v", err)
|
||||
}
|
||||
if _, statErr := os.Stat(out); !os.IsNotExist(statErr) {
|
||||
t.Fatalf("fixture was committed after truncation: %v", statErr)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("injected client is used and save leaves no temp files", func(t *testing.T) {
|
||||
srv := jsonServer(t, `{"ok":true}`)
|
||||
trip := &countTransport{rt: srv.Client().Transport}
|
||||
client := &http.Client{Transport: trip}
|
||||
recorded, err := RecordFlow(ctx, spec, RecordConfig{Target: srv.URL, Client: client})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if trip.n == 0 {
|
||||
t.Fatal("injected client was not used")
|
||||
}
|
||||
dir := t.TempDir()
|
||||
out := filepath.Join(dir, "saved.yaml")
|
||||
if err := SaveFlow(out, recorded); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, e := range entries {
|
||||
if strings.Contains(e.Name(), ".tmp") {
|
||||
t.Fatalf("temp file left behind: %s", e.Name())
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func jsonServer(t *testing.T, body string) *httptest.Server {
|
||||
t.Helper()
|
||||
srv := 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(body))
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
return srv
|
||||
}
|
||||
|
||||
type countTransport struct {
|
||||
rt http.RoundTripper
|
||||
n int
|
||||
}
|
||||
|
||||
func (c *countTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
c.n++
|
||||
if c.rt == nil {
|
||||
return http.DefaultTransport.RoundTrip(req)
|
||||
}
|
||||
return c.rt.RoundTrip(req)
|
||||
}
|
||||
147
modules/tide/rules.go
Normal file
147
modules/tide/rules.go
Normal file
@@ -0,0 +1,147 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"path"
|
||||
"strings"
|
||||
|
||||
"github.com/goccy/go-yaml"
|
||||
)
|
||||
|
||||
const (
|
||||
ClientNuxt = "nuxt"
|
||||
ClientMCP = "mcp"
|
||||
|
||||
SessionHeader = "X-Parity-Session"
|
||||
)
|
||||
|
||||
// Rules is the committed, non-secret capture policy for a proxy session.
|
||||
type Rules struct {
|
||||
Client string `yaml:"client"`
|
||||
KeepRequestHeaders []string `yaml:"keep_request_headers"`
|
||||
KeepResponseHeaders []string `yaml:"keep_response_headers"`
|
||||
Routes []RouteRule `yaml:"routes"`
|
||||
}
|
||||
|
||||
// RouteRule matches a method and path pattern and names capture sources.
|
||||
type RouteRule struct {
|
||||
Method string `yaml:"method"`
|
||||
Path string `yaml:"path"`
|
||||
KeepRequestHeaders []string `yaml:"keep_request_headers,omitempty"`
|
||||
KeepResponseHeaders []string `yaml:"keep_response_headers,omitempty"`
|
||||
Capture []CaptureRule `yaml:"capture,omitempty"`
|
||||
}
|
||||
|
||||
// LoadRules reads a strict YAML rule file, rejecting unknown fields.
|
||||
func LoadRules(path string) (Rules, error) {
|
||||
raw, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return Rules{}, fmt.Errorf("tide: read rules %s: %w", path, err)
|
||||
}
|
||||
return ParseRules(raw)
|
||||
}
|
||||
|
||||
// ParseRules decodes rules YAML with unknown-field rejection.
|
||||
func ParseRules(raw []byte) (Rules, error) {
|
||||
var rules Rules
|
||||
dec := yaml.NewDecoder(bytes.NewReader(raw), yaml.DisallowUnknownField())
|
||||
if err := dec.Decode(&rules); err != nil {
|
||||
return Rules{}, fmt.Errorf("tide: parse rules: %w", err)
|
||||
}
|
||||
if err := validateRules(rules); err != nil {
|
||||
return Rules{}, err
|
||||
}
|
||||
return rules, nil
|
||||
}
|
||||
|
||||
func validateRules(rules Rules) error {
|
||||
switch rules.Client {
|
||||
case ClientNuxt, ClientMCP:
|
||||
default:
|
||||
return fmt.Errorf("tide: rules client must be %q or %q", ClientNuxt, ClientMCP)
|
||||
}
|
||||
for i, route := range rules.Routes {
|
||||
if strings.TrimSpace(route.Method) == "" {
|
||||
return fmt.Errorf("tide: rules routes[%d] is missing method", i)
|
||||
}
|
||||
if strings.TrimSpace(route.Path) == "" {
|
||||
return fmt.Errorf("tide: rules routes[%d] is missing path", i)
|
||||
}
|
||||
for j, cap := range route.Capture {
|
||||
if strings.TrimSpace(cap.As) == "" {
|
||||
return fmt.Errorf("tide: rules routes[%d].capture[%d] is missing as", i, j)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Match returns the first route rule for method and path, or nil.
|
||||
func (r Rules) Match(method, requestPath string) *RouteRule {
|
||||
for i := range r.Routes {
|
||||
route := &r.Routes[i]
|
||||
if !matchMethod(route.Method, method) {
|
||||
continue
|
||||
}
|
||||
if !matchPath(route.Path, requestPath) {
|
||||
continue
|
||||
}
|
||||
return route
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r Rules) requestHeaders(route *RouteRule) []string {
|
||||
if route != nil && len(route.KeepRequestHeaders) > 0 {
|
||||
return route.KeepRequestHeaders
|
||||
}
|
||||
return r.KeepRequestHeaders
|
||||
}
|
||||
|
||||
func (r Rules) responseHeaders(route *RouteRule) []string {
|
||||
if route != nil && len(route.KeepResponseHeaders) > 0 {
|
||||
return route.KeepResponseHeaders
|
||||
}
|
||||
return r.KeepResponseHeaders
|
||||
}
|
||||
|
||||
func matchMethod(pattern, method string) bool {
|
||||
pattern = strings.TrimSpace(pattern)
|
||||
if pattern == "" || pattern == "*" {
|
||||
return true
|
||||
}
|
||||
return strings.EqualFold(pattern, method)
|
||||
}
|
||||
|
||||
func matchPath(pattern, requestPath string) bool {
|
||||
if pattern == requestPath {
|
||||
return true
|
||||
}
|
||||
ok, err := path.Match(pattern, requestPath)
|
||||
return err == nil && ok
|
||||
}
|
||||
|
||||
func filterHeaders(h http.Header, keep []string) map[string]string {
|
||||
if len(keep) == 0 || h == nil {
|
||||
return nil
|
||||
}
|
||||
out := make(map[string]string)
|
||||
for _, name := range keep {
|
||||
vals := h.Values(name)
|
||||
if len(vals) == 0 {
|
||||
continue
|
||||
}
|
||||
if len(vals) == 1 {
|
||||
out[http.CanonicalHeaderKey(name)] = vals[0]
|
||||
continue
|
||||
}
|
||||
out[http.CanonicalHeaderKey(name)] = strings.Join(vals, "\n")
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return nil
|
||||
}
|
||||
return out
|
||||
}
|
||||
60
modules/tide/rules_test.go
Normal file
60
modules/tide/rules_test.go
Normal file
@@ -0,0 +1,60 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRulesLoadAndMatch(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "rules.yaml")
|
||||
raw := "" +
|
||||
"client: nuxt\n" +
|
||||
"keep_request_headers:\n" +
|
||||
" - Cookie\n" +
|
||||
"keep_response_headers:\n" +
|
||||
" - Location\n" +
|
||||
"routes:\n" +
|
||||
" - method: GET\n" +
|
||||
" path: /sample\n" +
|
||||
" capture:\n" +
|
||||
" - from: response.json\n" +
|
||||
" path: $.token\n" +
|
||||
" as: jwt:alice\n" +
|
||||
" category: jwt\n" +
|
||||
" - method: POST\n" +
|
||||
" path: /items/*\n"
|
||||
if err := os.WriteFile(path, []byte(raw), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rules, err := LoadRules(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if rules.Client != ClientNuxt {
|
||||
t.Fatalf("client: %s", rules.Client)
|
||||
}
|
||||
if got := rules.Match("GET", "/sample"); got == nil || len(got.Capture) != 1 || got.Capture[0].As != "jwt:alice" {
|
||||
t.Fatalf("GET /sample match: %+v", got)
|
||||
}
|
||||
if got := rules.Match("POST", "/items/9"); got == nil {
|
||||
t.Fatal("POST /items/9 should match wildcard")
|
||||
}
|
||||
if got := rules.Match("GET", "/other"); got != nil {
|
||||
t.Fatalf("unexpected match: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRulesRejectUnknownAndInvalid(t *testing.T) {
|
||||
if _, err := ParseRules([]byte("client: nuxt\nextra: true\n")); err == nil {
|
||||
t.Fatal("unknown field must fail")
|
||||
}
|
||||
if _, err := ParseRules([]byte("client: browser\n")); err == nil || !strings.Contains(err.Error(), "client") {
|
||||
t.Fatalf("invalid client: %v", err)
|
||||
}
|
||||
if _, err := ParseRules([]byte("client: mcp\nroutes:\n - method: GET\n path: /x\n capture:\n - path: $.a\n")); err == nil || !strings.Contains(err.Error(), "as") {
|
||||
t.Fatalf("capture missing as: %v", err)
|
||||
}
|
||||
}
|
||||
9
modules/tide/testdata/one-route-spec.yaml
vendored
Normal file
9
modules/tide/testdata/one-route-spec.yaml
vendored
Normal file
@@ -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
|
||||
738
modules/tide/variables.go
Normal file
738
modules/tide/variables.go
Normal file
@@ -0,0 +1,738 @@
|
||||
package tide
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/goccy/go-yaml"
|
||||
)
|
||||
|
||||
var (
|
||||
placeholderRe = regexp.MustCompile(`\{\{([^{}]+)\}\}`)
|
||||
jwtShapeRe = regexp.MustCompile(`eyJ[A-Za-z0-9_-]+\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]+`)
|
||||
invShapeRe = regexp.MustCompile(`inv_[A-Za-z0-9]{8,}`)
|
||||
cookieRe = regexp.MustCompile(`(?i)auth_token=([^;]+)`)
|
||||
secretFormRe = regexp.MustCompile(`(?i)client_secret=([^&\s]+)`)
|
||||
pkceFormRe = regexp.MustCompile(`(?i)code_verifier=([^&\s]+)`)
|
||||
passwordJSONRe = regexp.MustCompile(`(?i)"password"\s*:\s*"([^"]*)"`)
|
||||
passwordFormRe = regexp.MustCompile(`(?i)(?:^|&)password=([^&\s]*)`)
|
||||
accessTokenJSONRe = regexp.MustCompile(`(?i)"access_token"\s*:\s*"([^"]*)"`)
|
||||
secretJSONRe = regexp.MustCompile(`(?i)"client_secret"\s*:\s*"[^"]*"`)
|
||||
)
|
||||
|
||||
// allowedTestPasswords are documented onboarding secrets that remain in fixtures
|
||||
// as plaintext. Any other leftover password-shaped value is a capture leak.
|
||||
var allowedTestPasswords = map[string]struct{}{
|
||||
"parity-alice-pass": {},
|
||||
"parity-alice-next": {},
|
||||
}
|
||||
|
||||
// Store holds named capture values. When Path is set it is a mode-0600 private file.
|
||||
type Store struct {
|
||||
mu sync.Mutex
|
||||
path string
|
||||
vals map[string]string
|
||||
}
|
||||
|
||||
// OpenStore loads or creates a private variable map. Empty path is memory-only.
|
||||
func OpenStore(path string) (*Store, error) {
|
||||
s := &Store{path: path, vals: make(map[string]string)}
|
||||
if strings.TrimSpace(path) == "" {
|
||||
return s, nil
|
||||
}
|
||||
abs, err := filepath.Abs(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tide: vars path: %w", err)
|
||||
}
|
||||
s.path = abs
|
||||
st, err := os.Stat(abs)
|
||||
if err == nil {
|
||||
if st.IsDir() {
|
||||
return nil, fmt.Errorf("tide: vars %q is a directory", path)
|
||||
}
|
||||
raw, err := os.ReadFile(abs)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tide: read vars: %w", err)
|
||||
}
|
||||
if len(strings.TrimSpace(string(raw))) > 0 {
|
||||
if err := yaml.Unmarshal(raw, &s.vals); err != nil {
|
||||
return nil, fmt.Errorf("tide: parse vars: %w", err)
|
||||
}
|
||||
if s.vals == nil {
|
||||
s.vals = make(map[string]string)
|
||||
}
|
||||
}
|
||||
if err := os.Chmod(abs, 0o600); err != nil {
|
||||
return nil, fmt.Errorf("tide: chmod vars: %w", err)
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
if !os.IsNotExist(err) {
|
||||
return nil, fmt.Errorf("tide: stat vars: %w", err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(abs), 0o755); err != nil {
|
||||
return nil, fmt.Errorf("tide: create vars dir: %w", err)
|
||||
}
|
||||
if err := os.WriteFile(abs, []byte("{}\n"), 0o600); err != nil {
|
||||
return nil, fmt.Errorf("tide: create vars: %w", err)
|
||||
}
|
||||
_ = os.Chmod(abs, 0o600)
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// Save writes the map as YAML with mode 0600. Memory-only stores are a no-op.
|
||||
func (s *Store) Save() error {
|
||||
if s == nil || s.path == "" {
|
||||
return nil
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
keys := make([]string, 0, len(s.vals))
|
||||
for k := range s.vals {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
var b strings.Builder
|
||||
if len(keys) == 0 {
|
||||
b.WriteString("{}\n")
|
||||
}
|
||||
for _, k := range keys {
|
||||
fmt.Fprintf(&b, "%s: %s\n", strconv.Quote(k), strconv.Quote(s.vals[k]))
|
||||
}
|
||||
raw := []byte(b.String())
|
||||
tmp, err := os.CreateTemp(filepath.Dir(s.path), ".vars-*.tmp")
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: vars temp: %w", err)
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
if _, err := tmp.Write(raw); err != nil {
|
||||
_ = tmp.Close()
|
||||
_ = os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
_ = tmp.Chmod(0o600)
|
||||
if err := tmp.Close(); err != nil {
|
||||
_ = os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(tmpName, s.path); err != nil {
|
||||
_ = os.Remove(tmpName)
|
||||
return err
|
||||
}
|
||||
return os.Chmod(s.path, 0o600)
|
||||
}
|
||||
|
||||
// Get returns a stored value.
|
||||
func (s *Store) Get(name string) (string, bool) {
|
||||
if s == nil {
|
||||
return "", false
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
v, ok := s.vals[name]
|
||||
return v, ok
|
||||
}
|
||||
|
||||
// Set stores a named value.
|
||||
func (s *Store) Set(name, value string) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.vals[name] = value
|
||||
}
|
||||
|
||||
// Expand replaces {{name}} placeholders. Unresolved names fail before HTTP send.
|
||||
func (s *Store) Expand(text string) (string, error) {
|
||||
if !strings.Contains(text, "{{") {
|
||||
return text, nil
|
||||
}
|
||||
if s == nil {
|
||||
m := placeholderRe.FindStringSubmatch(text)
|
||||
if len(m) > 1 {
|
||||
return "", fmt.Errorf("tide: unresolved placeholder %q", m[1])
|
||||
}
|
||||
return "", fmt.Errorf("tide: unresolved placeholder")
|
||||
}
|
||||
var missing []string
|
||||
s.mu.Lock()
|
||||
out := placeholderRe.ReplaceAllStringFunc(text, func(m string) string {
|
||||
name := m[2 : len(m)-2]
|
||||
v, ok := s.vals[name]
|
||||
if !ok {
|
||||
missing = append(missing, name)
|
||||
return m
|
||||
}
|
||||
return v
|
||||
})
|
||||
s.mu.Unlock()
|
||||
if len(missing) > 0 {
|
||||
return "", fmt.Errorf("tide: unresolved placeholder %q", missing[0])
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func expandRequest(req Request, store *Store) (Request, error) {
|
||||
out := req
|
||||
var err error
|
||||
out.Path, err = store.Expand(req.Path)
|
||||
if err != nil {
|
||||
return Request{}, err
|
||||
}
|
||||
out.Query, err = store.Expand(req.Query)
|
||||
if err != nil {
|
||||
return Request{}, err
|
||||
}
|
||||
if req.Headers != nil {
|
||||
out.Headers = make(map[string]string, len(req.Headers))
|
||||
for k, v := range req.Headers {
|
||||
out.Headers[k], err = store.Expand(v)
|
||||
if err != nil {
|
||||
return Request{}, err
|
||||
}
|
||||
}
|
||||
}
|
||||
body, err := store.Expand(string(req.Body))
|
||||
if err != nil {
|
||||
return Request{}, err
|
||||
}
|
||||
out.Body = Body(body)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func expandResponse(resp Response, store *Store) (Response, error) {
|
||||
out := resp
|
||||
var err error
|
||||
if resp.Headers != nil {
|
||||
out.Headers = make(map[string]string, len(resp.Headers))
|
||||
for k, v := range resp.Headers {
|
||||
out.Headers[k], err = store.Expand(v)
|
||||
if err != nil {
|
||||
return Response{}, err
|
||||
}
|
||||
}
|
||||
}
|
||||
body, err := store.Expand(string(resp.Body))
|
||||
if err != nil {
|
||||
return Response{}, err
|
||||
}
|
||||
out.Body = Body(body)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// CaptureStep writes named values from the step into the store.
|
||||
func CaptureStep(store *Store, step *Step) error {
|
||||
if store == nil || step == nil || len(step.Capture) == 0 {
|
||||
return nil
|
||||
}
|
||||
for _, rule := range step.Capture {
|
||||
val, err := extractCapture(rule, step.Request, step.Response)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: capture %q: %w", rule.As, err)
|
||||
}
|
||||
if strings.TrimSpace(val) == "" {
|
||||
return fmt.Errorf("tide: capture %q was empty", rule.As)
|
||||
}
|
||||
store.Set(rule.As, val)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func extractCapture(rule CaptureRule, req Request, resp Response) (string, error) {
|
||||
from := strings.TrimSpace(rule.From)
|
||||
if from == "" {
|
||||
from = "response.json"
|
||||
}
|
||||
switch from {
|
||||
case "response.json":
|
||||
v, err := jsonPathValue([]byte(resp.Body), rule.Path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return scalarString(v), nil
|
||||
case "response.json.query":
|
||||
v, err := jsonPathValue([]byte(resp.Body), rule.Path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
raw := scalarString(v)
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
q := u.Query().Get(rule.Name)
|
||||
if q == "" {
|
||||
return "", fmt.Errorf("missing JSON URL query %s", rule.Name)
|
||||
}
|
||||
return q, nil
|
||||
case "response.header":
|
||||
v := headerValue(resp.Headers, rule.Name)
|
||||
if v == "" {
|
||||
return "", fmt.Errorf("missing response header %s", rule.Name)
|
||||
}
|
||||
return v, nil
|
||||
case "response.query":
|
||||
loc := headerValue(resp.Headers, "Location")
|
||||
if loc == "" {
|
||||
return "", fmt.Errorf("missing Location header")
|
||||
}
|
||||
u, err := url.Parse(loc)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
v := u.Query().Get(rule.Name)
|
||||
if v == "" {
|
||||
return "", fmt.Errorf("missing response query %s", rule.Name)
|
||||
}
|
||||
return v, nil
|
||||
case "response.location.query":
|
||||
loc := headerValue(resp.Headers, "Location")
|
||||
if loc == "" {
|
||||
return "", fmt.Errorf("missing Location header")
|
||||
}
|
||||
u, err := url.Parse(loc)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
v := u.Query().Get(rule.Name)
|
||||
if v == "" {
|
||||
return "", fmt.Errorf("missing Location query %s", rule.Name)
|
||||
}
|
||||
return v, nil
|
||||
case "request.form":
|
||||
v, err := formValue(string(req.Body), rule.Name)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if v == "" {
|
||||
return "", fmt.Errorf("missing form field %s", rule.Name)
|
||||
}
|
||||
return v, nil
|
||||
case "request.header":
|
||||
v := headerValue(req.Headers, rule.Name)
|
||||
if v == "" {
|
||||
return "", fmt.Errorf("missing request header %s", rule.Name)
|
||||
}
|
||||
return v, nil
|
||||
case "request.query":
|
||||
q, err := url.ParseQuery(req.Query)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
v := q.Get(rule.Name)
|
||||
if v == "" {
|
||||
return "", fmt.Errorf("missing request query %s", rule.Name)
|
||||
}
|
||||
return v, nil
|
||||
default:
|
||||
return "", fmt.Errorf("unknown capture from %q", from)
|
||||
}
|
||||
}
|
||||
|
||||
func formValue(body, name string) (string, error) {
|
||||
vals, err := url.ParseQuery(body)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return vals.Get(name), nil
|
||||
}
|
||||
|
||||
func jsonPathValue(raw []byte, path string) (any, error) {
|
||||
if strings.TrimSpace(path) == "" {
|
||||
return nil, fmt.Errorf("json path is required")
|
||||
}
|
||||
root, err := decodeJSON(raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
v, err := walkJSONPath(root, path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return v, nil
|
||||
}
|
||||
|
||||
func walkJSONPath(root any, path string) (any, error) {
|
||||
path = strings.TrimSpace(path)
|
||||
if path == "$" || path == "" {
|
||||
return root, nil
|
||||
}
|
||||
if !strings.HasPrefix(path, "$") {
|
||||
path = "$." + path
|
||||
}
|
||||
cur := root
|
||||
rest := strings.TrimPrefix(path, "$")
|
||||
for rest != "" {
|
||||
switch {
|
||||
case strings.HasPrefix(rest, "."):
|
||||
rest = rest[1:]
|
||||
name, next := splitPathSeg(rest)
|
||||
if name == "" {
|
||||
return nil, fmt.Errorf("invalid json path %s", path)
|
||||
}
|
||||
obj, ok := cur.(map[string]any)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("%s is not an object", path)
|
||||
}
|
||||
v, ok := obj[name]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("missing %s", "$."+name)
|
||||
}
|
||||
cur = v
|
||||
rest = next
|
||||
case strings.HasPrefix(rest, "["):
|
||||
end := strings.IndexByte(rest, ']')
|
||||
if end < 0 {
|
||||
return nil, fmt.Errorf("invalid json path %s", path)
|
||||
}
|
||||
idx, err := atoi(rest[1:end])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
arr, ok := cur.([]any)
|
||||
if !ok || idx < 0 || idx >= len(arr) {
|
||||
return nil, fmt.Errorf("missing %s[%d]", path, idx)
|
||||
}
|
||||
cur = arr[idx]
|
||||
rest = rest[end+1:]
|
||||
default:
|
||||
return nil, fmt.Errorf("invalid json path %s", path)
|
||||
}
|
||||
}
|
||||
return cur, nil
|
||||
}
|
||||
|
||||
func splitPathSeg(s string) (name, rest string) {
|
||||
i := 0
|
||||
for i < len(s) && s[i] != '.' && s[i] != '[' {
|
||||
i++
|
||||
}
|
||||
return s[:i], s[i:]
|
||||
}
|
||||
|
||||
func atoi(s string) (int, error) {
|
||||
n := 0
|
||||
if s == "" {
|
||||
return 0, fmt.Errorf("empty index")
|
||||
}
|
||||
for _, c := range s {
|
||||
if c < '0' || c > '9' {
|
||||
return 0, fmt.Errorf("invalid index %q", s)
|
||||
}
|
||||
n = n*10 + int(c-'0')
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func scalarString(v any) string {
|
||||
switch t := v.(type) {
|
||||
case nil:
|
||||
return ""
|
||||
case string:
|
||||
return t
|
||||
case json.Number:
|
||||
return string(t)
|
||||
case bool:
|
||||
return fmt.Sprintf("%v", t)
|
||||
default:
|
||||
return fmt.Sprintf("%v", t)
|
||||
}
|
||||
}
|
||||
|
||||
// ScrubStep replaces stored capture values with {{name}} in kept fields.
|
||||
func ScrubStep(store *Store, step *Step) error {
|
||||
if store == nil || step == nil {
|
||||
return nil
|
||||
}
|
||||
pairs := store.replacements()
|
||||
// Short numeric IDs belong in path/query/headers (e.g. /albums/1). JSON
|
||||
// bodies keep literal counts and IPv4; normalizeJSON already masks id/*_id.
|
||||
step.Request.Path = replaceAll(step.Request.Path, pairs, true)
|
||||
step.Request.Query = replaceAll(step.Request.Query, pairs, true)
|
||||
step.Request.Headers = scrubMap(step.Request.Headers, pairs, true)
|
||||
step.Request.Body = Body(replaceAll(string(step.Request.Body), pairs, false))
|
||||
for _, rule := range step.Capture {
|
||||
if strings.TrimSpace(rule.From) != "request.form" {
|
||||
continue
|
||||
}
|
||||
if rule.Name == "" || rule.As == "" {
|
||||
continue
|
||||
}
|
||||
step.Request.Body = Body(scrubFormField(string(step.Request.Body), rule.Name, rule.As))
|
||||
}
|
||||
step.Response.Headers = scrubMap(step.Response.Headers, pairs, true)
|
||||
step.Response.Body = Body(replaceAll(string(step.Response.Body), pairs, false))
|
||||
return rejectUnclassifiedCredentials(*step)
|
||||
}
|
||||
|
||||
func (s *Store) replacements() [][2]string {
|
||||
if s == nil {
|
||||
return nil
|
||||
}
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
keys := make([]string, 0, len(s.vals))
|
||||
for k, v := range s.vals {
|
||||
if v != "" {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
}
|
||||
sort.Slice(keys, func(i, j int) bool {
|
||||
return len(s.vals[keys[i]]) > len(s.vals[keys[j]])
|
||||
})
|
||||
out := make([][2]string, 0, len(keys))
|
||||
for _, k := range keys {
|
||||
out = append(out, [2]string{s.vals[k], "{{" + k + "}}"})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func scrubMap(in map[string]string, pairs [][2]string, allowShortNumeric bool) map[string]string {
|
||||
if in == nil {
|
||||
return nil
|
||||
}
|
||||
out := make(map[string]string, len(in))
|
||||
for k, v := range in {
|
||||
out[k] = replaceAll(v, pairs, allowShortNumeric)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func replaceAll(s string, pairs [][2]string, allowShortNumeric bool) string {
|
||||
for _, p := range pairs {
|
||||
if p[0] == "" {
|
||||
continue
|
||||
}
|
||||
if !allowShortNumeric && isAllDigits(p[0]) && len(p[0]) < 8 {
|
||||
continue
|
||||
}
|
||||
olds := []string{p[0]}
|
||||
if esc := phpJSONEscape(p[0]); esc != p[0] {
|
||||
olds = append(olds, esc)
|
||||
}
|
||||
for _, old := range olds {
|
||||
if len(old) >= 8 {
|
||||
s = strings.ReplaceAll(s, old, p[1])
|
||||
continue
|
||||
}
|
||||
s = replaceIsolated(s, old, p[1])
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func isAllDigits(s string) bool {
|
||||
if s == "" {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i] < '0' || s[i] > '9' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func phpJSONEscape(s string) string {
|
||||
return strings.ReplaceAll(s, "/", `\/`)
|
||||
}
|
||||
|
||||
func scrubFormField(body, name, as string) string {
|
||||
if name == "" || as == "" || body == "" {
|
||||
return body
|
||||
}
|
||||
re := regexp.MustCompile(`(?i)(^|&)(` + regexp.QuoteMeta(name) + `=)[^&]*`)
|
||||
return re.ReplaceAllString(body, `${1}${2}{{`+as+`}}`)
|
||||
}
|
||||
|
||||
func replaceIsolated(s, old, neu string) string {
|
||||
if old == "" || s == "" {
|
||||
return s
|
||||
}
|
||||
var b strings.Builder
|
||||
i := 0
|
||||
for i < len(s) {
|
||||
j := strings.Index(s[i:], old)
|
||||
if j < 0 {
|
||||
b.WriteString(s[i:])
|
||||
break
|
||||
}
|
||||
j += i
|
||||
leftOK := j == 0 || !isIdentByte(s[j-1])
|
||||
right := j + len(old)
|
||||
rightOK := right == len(s) || !isIdentByte(s[right])
|
||||
if leftOK && rightOK {
|
||||
b.WriteString(s[i:j])
|
||||
b.WriteString(neu)
|
||||
i = right
|
||||
continue
|
||||
}
|
||||
b.WriteString(s[i : j+len(old)])
|
||||
i = j + len(old)
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func isIdentByte(c byte) bool {
|
||||
return (c >= '0' && c <= '9') || (c >= 'A' && c <= 'Z') ||
|
||||
(c >= 'a' && c <= 'z') || c == '_' || c == '.'
|
||||
}
|
||||
|
||||
func rejectUnclassifiedCredentials(step Step) error {
|
||||
check := func(label, s string) error {
|
||||
if hit := remainingCredential(s); hit != "" {
|
||||
return fmt.Errorf("unclassified credential-shaped value (%s) in %s step %s", hit, label, step.ID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
for k, v := range step.Request.Headers {
|
||||
if err := check("request header "+k, v); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := check("request query", step.Request.Query); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := check("request body", string(step.Request.Body)); err != nil {
|
||||
return err
|
||||
}
|
||||
for k, v := range step.Response.Headers {
|
||||
if err := check("response header "+k, v); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return check("response body", string(step.Response.Body))
|
||||
}
|
||||
|
||||
func remainingCredential(s string) string {
|
||||
s = placeholderRe.ReplaceAllString(s, "")
|
||||
if s == "" {
|
||||
return ""
|
||||
}
|
||||
if jwtShapeRe.MatchString(s) {
|
||||
return "jwt"
|
||||
}
|
||||
if invShapeRe.MatchString(s) {
|
||||
return "token"
|
||||
}
|
||||
if cookieRe.MatchString(s) {
|
||||
return "cookie"
|
||||
}
|
||||
if secretFormRe.MatchString(s) {
|
||||
return "oauth_secret"
|
||||
}
|
||||
if pkceFormRe.MatchString(s) {
|
||||
return "pkce"
|
||||
}
|
||||
if leftoverPassword(s) {
|
||||
return "password"
|
||||
}
|
||||
if leftoverAccessToken(s) {
|
||||
return "access_token"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func leftoverPassword(s string) bool {
|
||||
for _, m := range passwordJSONRe.FindAllStringSubmatch(s, -1) {
|
||||
if m[1] != "" {
|
||||
if _, ok := allowedTestPasswords[m[1]]; !ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, m := range passwordFormRe.FindAllStringSubmatch(s, -1) {
|
||||
v, err := url.QueryUnescape(m[1])
|
||||
if err != nil {
|
||||
v = m[1]
|
||||
}
|
||||
if v != "" {
|
||||
if _, ok := allowedTestPasswords[v]; !ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func leftoverAccessToken(s string) bool {
|
||||
for _, m := range accessTokenJSONRe.FindAllStringSubmatch(s, -1) {
|
||||
if strings.TrimSpace(m[1]) != "" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func redactSecrets(s string) string {
|
||||
s = jwtShapeRe.ReplaceAllString(s, "<redacted-jwt>")
|
||||
s = invShapeRe.ReplaceAllString(s, "<redacted-inv>")
|
||||
s = secretFormRe.ReplaceAllString(s, "client_secret=<redacted>")
|
||||
s = secretJSONRe.ReplaceAllString(s, `"client_secret":"<redacted>"`)
|
||||
s = cookieRe.ReplaceAllString(s, "auth_token=<redacted>")
|
||||
return s
|
||||
}
|
||||
|
||||
func varsOutsideFixtures(varsPath, fixtures string) error {
|
||||
if varsPath == "" || fixtures == "" {
|
||||
return nil
|
||||
}
|
||||
absVars, err := resolvePath(varsPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: vars path: %w", err)
|
||||
}
|
||||
absFix, err := resolvePath(fixtures)
|
||||
if err != nil {
|
||||
return fmt.Errorf("tide: fixtures path: %w", err)
|
||||
}
|
||||
if absVars == absFix || strings.HasPrefix(absVars, absFix+string(os.PathSeparator)) {
|
||||
return fmt.Errorf("tide: vars file %q must be outside fixtures %q", varsPath, fixtures)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func resolvePath(path string) (string, error) {
|
||||
abs, err := filepath.Abs(path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if eval, err := filepath.EvalSymlinks(abs); err == nil {
|
||||
return eval, nil
|
||||
}
|
||||
parentEval, err := filepath.EvalSymlinks(filepath.Dir(abs))
|
||||
if err != nil {
|
||||
if _, statErr := os.Lstat(abs); statErr == nil {
|
||||
return "", fmt.Errorf("eval symlinks %s: %w", path, err)
|
||||
}
|
||||
return filepath.Clean(abs), nil
|
||||
}
|
||||
return filepath.Join(parentEval, filepath.Base(abs)), nil
|
||||
}
|
||||
|
||||
func mergeRouteCaptures(step *Step, rules Rules) {
|
||||
if step == nil || len(step.Capture) > 0 {
|
||||
return
|
||||
}
|
||||
route := rules.Match(step.Request.Method, step.Request.Path)
|
||||
if route != nil && len(route.Capture) > 0 {
|
||||
step.Capture = append([]CaptureRule(nil), route.Capture...)
|
||||
}
|
||||
}
|
||||
|
||||
func recordedResponseHeaders(h http.Header, rules Rules, method, path string) map[string]string {
|
||||
route := rules.Match(method, path)
|
||||
if keep := rules.responseHeaders(route); len(keep) > 0 {
|
||||
return filterHeaders(h, keep)
|
||||
}
|
||||
return keepResponseHeaders(h)
|
||||
}
|
||||
Reference in New Issue
Block a user