package cabana import ( "io" "net/http" "net/http/httptest" "strings" "testing" "git.golem15.com/golem15/summercms/modules/bouncer" "git.golem15.com/golem15/summercms/modules/pact" ) // TestPhase10CSRF walks every state-changing route service.mount registers. // Except login, a request carrying only the admin cookie and no // X-Requested-With header is refused with 403 forbidden before the handler // runs: the body spy is never read and the service, which has no database, // never answers with a database error. The same call with the header or with // a Bearer token reaches the handler. func TestPhase10CSRF(t *testing.T) { router := &handlerRouter{handlers: map[string]http.HandlerFunc{}} svc := phase09DeniedService() svc.mount(router) super := &bouncer.Principal{ID: 1, Backend: true, IsSuperuser: true} login := http.MethodPost + " " + adminAPI("/auth/login") unsafe := 0 for _, key := range router.order { method, path, _ := strings.Cut(key, " ") if method != http.MethodPost && method != http.MethodPut && method != http.MethodDelete { continue } if key == login { continue } unsafe++ handler := router.handlers[key] t.Run(key, func(t *testing.T) { refused, spy := csrfRequest(method, path, super) refused.AddCookie(&http.Cookie{Name: AdminCookieName, Value: "cookie-only-session"}) rec := httptest.NewRecorder() handler(rec, refused) if rec.Code != http.StatusForbidden { t.Fatalf("cookie-only status=%d body=%s", rec.Code, rec.Body.String()) } assertErrorCode(t, rec.Body.Bytes(), "forbidden") if spy.reads != 0 { t.Fatalf("refused request body was read %d times", spy.reads) } withHeader, _ := csrfRequest(method, path, super) withHeader.AddCookie(&http.Cookie{Name: AdminCookieName, Value: "cookie-only-session"}) withHeader.Header.Set("X-Requested-With", "XMLHttpRequest") rec = httptest.NewRecorder() handler(rec, withHeader) if rec.Code == http.StatusForbidden { t.Fatalf("request with X-Requested-With was refused: %s", rec.Body.String()) } withBearer, _ := csrfRequest(method, path, super) withBearer.Header.Set("Authorization", "Bearer not-a-real-token") rec = httptest.NewRecorder() handler(rec, withBearer) if rec.Code == http.StatusForbidden { t.Fatalf("Bearer request was refused: %s", rec.Body.String()) } }) } // refresh, logout, settings put, create, bulk-delete, widget action, // toolbar action, update, delete, link, unlink if unsafe != 11 { t.Fatalf("walked %d state-changing routes, want 11: %v", unsafe, router.order) } loginHandler := router.handlers[login] req, _ := csrfRequest(http.MethodPost, adminAPI("/auth/login"), nil) rec := httptest.NewRecorder() loginHandler(rec, req) if rec.Code == http.StatusForbidden { t.Fatalf("login without the header was refused: %s", rec.Body.String()) } } type readSpy struct { r io.Reader reads int } func (s *readSpy) Read(p []byte) (int, error) { s.reads++ return s.r.Read(p) } func csrfRequest(method, path string, principal *bouncer.Principal) (*http.Request, *readSpy) { spy := &readSpy{r: strings.NewReader(`{"ids":[1],"name":"csrf"}`)} req := httptest.NewRequest(method, path, spy) req.Header.Set("Content-Type", "application/json") req.SetPathValue("vendor", "acme") req.SetPathValue("plugin", "demo") req.SetPathValue("controller", "widgets") req.SetPathValue("id", "1") req.SetPathValue("name", "editors") req.SetPathValue("code", "demo") if principal != nil { req = req.WithContext(bouncer.WithUser(req.Context(), principal)) } return req, spy } // handlerRouter records the handler mounted for every route, so tests can // call exactly what service.mount registered. type handlerRouter struct { prefix string handlers map[string]http.HandlerFunc order []string } func (h *handlerRouter) Group(prefix string, middleware []string, fn func(pact.Router)) { h.GroupRaw(prefix, middleware, fn) } func (h *handlerRouter) GroupRaw(prefix string, _ []string, fn func(pact.Router)) { child := &handlerRouter{prefix: h.prefix + prefix, handlers: h.handlers} fn(child) h.order = append(h.order, child.order...) } func (h *handlerRouter) Get(path string, fn http.HandlerFunc, _ ...string) { h.add(http.MethodGet, path, fn) } func (h *handlerRouter) Post(path string, fn http.HandlerFunc, _ ...string) { h.add(http.MethodPost, path, fn) } func (h *handlerRouter) Put(path string, fn http.HandlerFunc, _ ...string) { h.add(http.MethodPut, path, fn) } func (h *handlerRouter) Patch(path string, fn http.HandlerFunc, _ ...string) { h.add(http.MethodPatch, path, fn) } func (h *handlerRouter) Delete(path string, fn http.HandlerFunc, _ ...string) { h.add(http.MethodDelete, path, fn) } func (h *handlerRouter) Where(string, string) {} func (h *handlerRouter) WhereIn(string, ...string) {} func (h *handlerRouter) add(method, path string, fn http.HandlerFunc) { key := method + " " + h.prefix + path h.handlers[key] = fn h.order = append(h.order, key) }