| 27 | ) |
| 28 | |
| 29 | func TestCSRF(t *testing.T) { |
| 30 | a := assertions.New(t) |
| 31 | authKey := []byte("1234123412341234123412341234123412341234123412341234123412341234") |
| 32 | m := CSRF(authKey) |
| 33 | |
| 34 | t.Run("Protects non-idempotent methods when using a Session Token", func(t *testing.T) { |
| 35 | r := httptest.NewRequest(http.MethodPost, "/", nil) |
| 36 | r.Header.Set("Authorization", "Bearer "+auth.JoinToken(auth.SessionToken, "XXX", "YYY")) |
| 37 | rec := httptest.NewRecorder() |
| 38 | m(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 39 | })).ServeHTTP(rec, r) |
| 40 | res := rec.Result() |
| 41 | a.So(res.StatusCode, should.Equal, http.StatusForbidden) |
| 42 | }) |
| 43 | |
| 44 | t.Run("Allows non-idempotent methods when using an API Key", func(t *testing.T) { |
| 45 | r := httptest.NewRequest(http.MethodPost, "/", nil) |
| 46 | r.Header.Set("Authorization", "Bearer "+auth.JoinToken(auth.APIKey, "XXX", "YYY")) |
| 47 | rec := httptest.NewRecorder() |
| 48 | m(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 49 | })).ServeHTTP(rec, r) |
| 50 | res := rec.Result() |
| 51 | a.So(res.StatusCode, should.Equal, http.StatusOK) |
| 52 | }) |
| 53 | |
| 54 | t.Run("Allows access with valid CSRF token", func(t *testing.T) { |
| 55 | var csrfToken string |
| 56 | var r *http.Request |
| 57 | |
| 58 | // Obtain CSRF token. |
| 59 | r = httptest.NewRequest(http.MethodGet, "/", nil) |
| 60 | rec := httptest.NewRecorder() |
| 61 | m(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 62 | csrfToken = csrf.Token(r) |
| 63 | })).ServeHTTP(rec, r) |
| 64 | res := rec.Result() |
| 65 | |
| 66 | cookies := res.Cookies() |
| 67 | a.So(cookies, should.HaveLength, 1) |
| 68 | |
| 69 | // Make request |
| 70 | r = httptest.NewRequest(http.MethodPost, "/", nil) |
| 71 | r.Header.Set("X-CSRF-Token", csrfToken) |
| 72 | r.AddCookie(cookies[0]) |
| 73 | rec = httptest.NewRecorder() |
| 74 | m(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {})).ServeHTTP(rec, r) |
| 75 | res = rec.Result() |
| 76 | a.So(res.StatusCode, should.Equal, http.StatusOK) |
| 77 | }) |
| 78 | } |