287 lines
7.9 KiB
Go
287 lines
7.9 KiB
Go
package httpapi
|
|
|
|
import (
|
|
"io"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
|
|
"github.com/alexedwards/scs/v2"
|
|
"github.com/descrybe/descrybe-v2/apps/api/internal/auth"
|
|
"github.com/descrybe/descrybe-v2/apps/api/internal/config"
|
|
)
|
|
|
|
func testServerCSRF() *Server {
|
|
sm := scs.New()
|
|
sm.Cookie.Name = "descrybe_session"
|
|
return &Server{
|
|
Config: config.Config{
|
|
CSRFCookieName: "descrybe_csrf",
|
|
SessionSecure: false,
|
|
},
|
|
Sessions: sm,
|
|
Auth: &auth.Service{},
|
|
}
|
|
}
|
|
|
|
func testServerCSRFSecure(secure bool, appEnv string) *Server {
|
|
s := testServerCSRF()
|
|
s.Config.SessionSecure = secure
|
|
s.Config.AppEnv = appEnv
|
|
return s
|
|
}
|
|
|
|
func findCSRFCookie(cookies []*http.Cookie) *http.Cookie {
|
|
for _, c := range cookies {
|
|
if c.Name == "descrybe_csrf" && c.Value != "" {
|
|
return c
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func TestCSRFAllowsSafeMethodsWithoutHeader(t *testing.T) {
|
|
t.Parallel()
|
|
s := testServerCSRF()
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := httptest.NewRequest(http.MethodGet, "/api/auth/me", nil)
|
|
rec := httptest.NewRecorder()
|
|
h.ServeHTTP(rec, req)
|
|
if rec.Code != http.StatusNoContent {
|
|
t.Fatalf("GET status = %d, want 204", rec.Code)
|
|
}
|
|
found := false
|
|
for _, c := range rec.Result().Cookies() {
|
|
if c.Name == "descrybe_csrf" && c.Value != "" && !c.HttpOnly {
|
|
found = true
|
|
}
|
|
}
|
|
if !found {
|
|
t.Fatal("expected non-HttpOnly CSRF cookie on first GET")
|
|
}
|
|
hdr := rec.Header().Get("X-CSRF-Token")
|
|
if hdr == "" {
|
|
t.Fatal("expected X-CSRF-Token response header on GET (cross-origin SPA seed)")
|
|
}
|
|
cookie := findCSRFCookie(rec.Result().Cookies())
|
|
if cookie == nil || cookie.Value != hdr {
|
|
t.Fatalf("X-CSRF-Token header %q must match cookie value", hdr)
|
|
}
|
|
}
|
|
|
|
func TestCSRFRejectsPOSTWithoutToken(t *testing.T) {
|
|
t.Parallel()
|
|
s := testServerCSRF()
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusOK)
|
|
}))
|
|
|
|
for _, path := range []string{
|
|
"/api/auth/login",
|
|
"/api/auth/forgot-password",
|
|
"/api/auth/reset-password",
|
|
} {
|
|
req := httptest.NewRequest(http.MethodPost, path, nil)
|
|
rec := httptest.NewRecorder()
|
|
h.ServeHTTP(rec, req)
|
|
if rec.Code != http.StatusForbidden {
|
|
t.Fatalf("%s POST without CSRF status = %d, want 403", path, rec.Code)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestCSRFAcceptsMatchingHeader(t *testing.T) {
|
|
t.Parallel()
|
|
s := testServerCSRF()
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
_, _ = io.WriteString(w, "ok")
|
|
}))
|
|
|
|
getReq := httptest.NewRequest(http.MethodGet, "/healthz", nil)
|
|
getRec := httptest.NewRecorder()
|
|
h.ServeHTTP(getRec, getReq)
|
|
var token string
|
|
for _, c := range getRec.Result().Cookies() {
|
|
if c.Name == "descrybe_csrf" {
|
|
token = c.Value
|
|
}
|
|
}
|
|
if token == "" {
|
|
t.Fatal("missing CSRF cookie from GET")
|
|
}
|
|
|
|
postReq := httptest.NewRequest(http.MethodPost, "/api/auth/login", nil)
|
|
postReq.AddCookie(&http.Cookie{Name: "descrybe_csrf", Value: token})
|
|
postReq.Header.Set("X-CSRF-Token", token)
|
|
postRec := httptest.NewRecorder()
|
|
h.ServeHTTP(postRec, postReq)
|
|
if postRec.Code != http.StatusOK {
|
|
t.Fatalf("POST with CSRF status = %d, want 200", postRec.Code)
|
|
}
|
|
}
|
|
|
|
func TestCSRFRejectsMismatchedHeader(t *testing.T) {
|
|
t.Parallel()
|
|
s := testServerCSRF()
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusOK)
|
|
}))
|
|
|
|
req := httptest.NewRequest(http.MethodPost, "/api/auth/login", nil)
|
|
req.AddCookie(&http.Cookie{Name: "descrybe_csrf", Value: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"})
|
|
req.Header.Set("X-CSRF-Token", "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb")
|
|
rec := httptest.NewRecorder()
|
|
h.ServeHTTP(rec, req)
|
|
if rec.Code != http.StatusForbidden {
|
|
t.Fatalf("mismatched CSRF status = %d, want 403", rec.Code)
|
|
}
|
|
}
|
|
|
|
func TestCSRFCookieAttributesDev(t *testing.T) {
|
|
t.Parallel()
|
|
s := testServerCSRFSecure(false, "development")
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := httptest.NewRequest(http.MethodGet, "/api/auth/me", nil)
|
|
rec := httptest.NewRecorder()
|
|
h.ServeHTTP(rec, req)
|
|
c := findCSRFCookie(rec.Result().Cookies())
|
|
if c == nil {
|
|
t.Fatal("expected CSRF cookie")
|
|
}
|
|
if c.HttpOnly {
|
|
t.Fatal("CSRF cookie must not be HttpOnly (double-submit)")
|
|
}
|
|
if c.Secure {
|
|
t.Fatal("development without SessionSecure should not set Secure")
|
|
}
|
|
if c.SameSite != http.SameSiteLaxMode {
|
|
t.Fatalf("SameSite = %v, want Lax", c.SameSite)
|
|
}
|
|
if c.Path != "/" {
|
|
t.Fatalf("Path = %q, want /", c.Path)
|
|
}
|
|
if c.MaxAge != 7*24*60*60 {
|
|
t.Fatalf("MaxAge = %d, want 7d", c.MaxAge)
|
|
}
|
|
}
|
|
|
|
func TestCSRFCookieSecureWhenSessionSecure(t *testing.T) {
|
|
t.Parallel()
|
|
s := testServerCSRFSecure(true, "development")
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := httptest.NewRequest(http.MethodGet, "/api/auth/me", nil)
|
|
rec := httptest.NewRecorder()
|
|
h.ServeHTTP(rec, req)
|
|
c := findCSRFCookie(rec.Result().Cookies())
|
|
if c == nil {
|
|
t.Fatal("expected CSRF cookie")
|
|
}
|
|
if !c.Secure {
|
|
t.Fatal("SessionSecure=true should set Secure")
|
|
}
|
|
if c.HttpOnly {
|
|
t.Fatal("CSRF cookie must not be HttpOnly")
|
|
}
|
|
if c.SameSite != http.SameSiteLaxMode {
|
|
t.Fatalf("SameSite = %v, want Lax", c.SameSite)
|
|
}
|
|
}
|
|
|
|
func TestCSRFCookieSecureWhenProductionAppEnv(t *testing.T) {
|
|
t.Parallel()
|
|
// Defense in depth: APP_ENV=production forces Secure even if SessionSecure was left false.
|
|
s := testServerCSRFSecure(false, "production")
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := httptest.NewRequest(http.MethodGet, "/api/auth/me", nil)
|
|
rec := httptest.NewRecorder()
|
|
h.ServeHTTP(rec, req)
|
|
c := findCSRFCookie(rec.Result().Cookies())
|
|
if c == nil {
|
|
t.Fatal("expected CSRF cookie")
|
|
}
|
|
if !c.Secure {
|
|
t.Fatal("APP_ENV=production must set Secure via CookieSecure")
|
|
}
|
|
}
|
|
|
|
// Client-mint path: SPA sets descrybe_csrf locally; middleware must accept matching header+cookie
|
|
// without a prior server-issued Set-Cookie on this request.
|
|
func TestCSRFAcceptsClientMintedCookie(t *testing.T) {
|
|
t.Parallel()
|
|
s := testServerCSRF()
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
const token = "0123456789abcdef0123456789abcdef"
|
|
req := httptest.NewRequest(http.MethodPost, "/api/auth/login", nil)
|
|
req.AddCookie(&http.Cookie{Name: "descrybe_csrf", Value: token})
|
|
req.Header.Set("X-CSRF-Token", token)
|
|
rec := httptest.NewRecorder()
|
|
h.ServeHTTP(rec, req)
|
|
if rec.Code != http.StatusNoContent {
|
|
t.Fatalf("client-minted CSRF status = %d, want 204", rec.Code)
|
|
}
|
|
}
|
|
|
|
func TestCSRFExemptPathSegmentsOnly(t *testing.T) {
|
|
t.Parallel()
|
|
cases := []struct {
|
|
path string
|
|
exempt bool
|
|
}{
|
|
{"/api/v1", true},
|
|
{"/api/v1/products", true},
|
|
{"/api/v10", false},
|
|
{"/api/v1legacy", false},
|
|
{"/api/public", true},
|
|
{"/api/public/plans", true},
|
|
{"/api/publicish", false},
|
|
{"/api/webhooks", true},
|
|
{"/api/webhooks/stripe", true},
|
|
{"/api/webhooksx", false},
|
|
{"/api/auth/login", false},
|
|
{"/api/auth/forgot-password", false},
|
|
{"/api/auth/reset-password", false},
|
|
}
|
|
for _, tc := range cases {
|
|
if got := csrfExemptPath(tc.path); got != tc.exempt {
|
|
t.Fatalf("csrfExemptPath(%q) = %v, want %v", tc.path, got, tc.exempt)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestCSRFRequiresTokenOnV1LookalikePath(t *testing.T) {
|
|
t.Parallel()
|
|
s := testServerCSRF()
|
|
h := s.CSRF(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := httptest.NewRequest(http.MethodPost, "/api/v10/mutate", nil)
|
|
rec := httptest.NewRecorder()
|
|
h.ServeHTTP(rec, req)
|
|
if rec.Code != http.StatusForbidden {
|
|
t.Fatalf("status = %d, want 403 (lookalike must not skip CSRF)", rec.Code)
|
|
}
|
|
|
|
req2 := httptest.NewRequest(http.MethodPost, "/api/v1/products", nil)
|
|
rec2 := httptest.NewRecorder()
|
|
h.ServeHTTP(rec2, req2)
|
|
if rec2.Code != http.StatusNoContent {
|
|
t.Fatalf("v1 exempt status = %d, want 204", rec2.Code)
|
|
}
|
|
}
|