Skip to content

Commit 42866fb

Browse files
fix(middleware): return a JSON error for rejected CSRF requests (#18)
CSRF now rejects an untrusted cross-origin request with a structured JSON 403, as CORS already does, instead of a `text/plain` body. - `CSRF` calls `CrossOriginProtection.Check` and returns `web.RespondError(ctx, w, errs.New(http.StatusForbidden, err))` on rejection. - The JSON `message` is the standard library's rejection reason. - `CSRF` still panics at construction on an invalid trusted origin. - `TestCSRF_Rejected` sends cross-origin POSTs with a cross-site `Sec-Fetch-Site` and with an untrusted `Origin`; each gets 403 with a JSON body. - `TestCSRF_TrustedOriginAllowed` checks that a POST from a trusted origin reaches the handler. Closes #17
1 parent c441bbd commit 42866fb

2 files changed

Lines changed: 82 additions & 22 deletions

File tree

‎web/middleware/csrf.go‎

Lines changed: 7 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -2,20 +2,21 @@ package middleware
22

33
import (
44
"context"
5-
"fmt"
65
"net/http"
76

7+
"github.com/adamwoolhether/httper/web"
8+
"github.com/adamwoolhether/httper/web/errs"
89
"github.com/adamwoolhether/httper/web/mux"
910
)
1011

1112
// CSRF uses the standard library CrossOriginProtection to prevent CSRF attacks.
13+
// It rejects an untrusted cross-origin request with a 403 JSON error.
1214
// Each trusted origin must be an exact scheme://host[:port] value, such as
1315
// "https://app.example.com:8443", with no path, query, or trailing slash.
1416
// Wildcards are not supported: an entry such as "https://*.example.com" is
1517
// accepted but matches no origin. CSRF panics if a trusted origin is invalid.
1618
func CSRF(allowedOrigins ...string) mux.Middleware {
1719
cop := http.NewCrossOriginProtection()
18-
cop.SetDenyHandler(errHandler())
1920
for _, origin := range allowedOrigins {
2021
if err := cop.AddTrustedOrigin(origin); err != nil {
2122
panic(err)
@@ -24,31 +25,15 @@ func CSRF(allowedOrigins ...string) mux.Middleware {
2425

2526
m := func(handler mux.Handler) mux.Handler {
2627
h := func(ctx context.Context, w http.ResponseWriter, r *http.Request) error {
27-
var err error
28+
if err := cop.Check(r); err != nil {
29+
return web.RespondError(ctx, w, errs.New(http.StatusForbidden, err))
30+
}
2831

29-
std := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
30-
ctx = r.Context()
31-
32-
err = handler(ctx, w, r)
33-
})
34-
35-
cop.Handler(std).ServeHTTP(w, r.WithContext(ctx))
36-
37-
return err
32+
return handler(ctx, w, r)
3833
}
3934

4035
return h
4136
}
4237

4338
return m
4439
}
45-
46-
func errHandler() http.HandlerFunc {
47-
f := func(w http.ResponseWriter, r *http.Request) {
48-
mux.SetStatusCode(r.Context(), http.StatusForbidden)
49-
50-
http.Error(w, fmt.Errorf("csrf: %s", http.StatusText(http.StatusForbidden)).Error(), http.StatusForbidden)
51-
}
52-
53-
return f
54-
}

‎web/middleware/csrf_test.go‎

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,86 @@
11
package middleware_test
22

33
import (
4+
"context"
5+
"encoding/json"
6+
"net/http"
7+
"net/http/httptest"
48
"testing"
59

610
"github.com/adamwoolhether/httper/web/middleware"
711
)
812

13+
func TestCSRF_Rejected(t *testing.T) {
14+
tests := map[string]http.Header{
15+
"cross-site fetch metadata": {"Sec-Fetch-Site": {"cross-site"}},
16+
"untrusted origin": {"Origin": {"https://evil.example.com"}},
17+
}
18+
19+
for name, header := range tests {
20+
t.Run(name, func(t *testing.T) {
21+
called := false
22+
handler := middleware.CSRF("https://app.example.com")(func(ctx context.Context, w http.ResponseWriter, r *http.Request) error {
23+
called = true
24+
return nil
25+
})
26+
27+
w := httptest.NewRecorder()
28+
r := httptest.NewRequest(http.MethodPost, "http://api.example.com/", nil)
29+
r.Header = header
30+
31+
if err := handler(r.Context(), w, r); err != nil {
32+
t.Fatalf("unexpected error: %v", err)
33+
}
34+
35+
if w.Code != http.StatusForbidden {
36+
t.Fatalf("status = %d, want %d", w.Code, http.StatusForbidden)
37+
}
38+
if called {
39+
t.Fatal("handler should not be called for a rejected request")
40+
}
41+
if ct := w.Header().Get("Content-Type"); ct != "application/json" {
42+
t.Fatalf("Content-Type = %q, want %q", ct, "application/json")
43+
}
44+
45+
var body struct {
46+
Code int `json:"code"`
47+
Message string `json:"message"`
48+
}
49+
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
50+
t.Fatalf("body %q should be JSON: %v", w.Body.String(), err)
51+
}
52+
if body.Code != http.StatusForbidden || body.Message == "" {
53+
t.Fatalf("body = %+v, want code %d and a message", body, http.StatusForbidden)
54+
}
55+
})
56+
}
57+
}
58+
59+
func TestCSRF_TrustedOriginAllowed(t *testing.T) {
60+
called := false
61+
handler := middleware.CSRF("https://app.example.com")(func(ctx context.Context, w http.ResponseWriter, r *http.Request) error {
62+
called = true
63+
w.WriteHeader(http.StatusOK)
64+
return nil
65+
})
66+
67+
w := httptest.NewRecorder()
68+
r := httptest.NewRequest(http.MethodPost, "http://api.example.com/", nil)
69+
r.Header.Set("Sec-Fetch-Site", "cross-site")
70+
r.Header.Set("Origin", "https://app.example.com")
71+
72+
if err := handler(r.Context(), w, r); err != nil {
73+
t.Fatalf("unexpected error: %v", err)
74+
}
75+
76+
if !called {
77+
t.Fatal("handler should be called for a trusted origin")
78+
}
79+
if w.Code != http.StatusOK {
80+
t.Fatalf("status = %d, want %d", w.Code, http.StatusOK)
81+
}
82+
}
83+
984
func TestCSRF_TrustedOrigin(t *testing.T) {
1085
tests := map[string]struct {
1186
origin string

0 commit comments

Comments
 (0)