diff --git a/tide/diff.go b/tide/diff.go index 26263f8..7b636ec 100644 --- a/tide/diff.go +++ b/tide/diff.go @@ -42,6 +42,7 @@ var extraCompareHeaders = []string{ "WWW-Authenticate", "Content-Disposition", "Link", + "Location", } var neverCompareHeaders = []string{ diff --git a/tide/headers_test.go b/tide/headers_test.go index c0896b6..080d94d 100644 --- a/tide/headers_test.go +++ b/tide/headers_test.go @@ -45,6 +45,28 @@ func TestHeadersAllowListIgnoresDate(t *testing.T) { } } +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")