From 405654b33f205b903caf04ff9d8ba13e6db7ee39 Mon Sep 17 00:00:00 2001 From: Serge Zaitsev Date: Tue, 25 Jan 2022 13:47:49 +0100 Subject: [PATCH] csrf checks for v8.3.5 (#234) --- pkg/api/http_server.go | 1 + pkg/macaron/binding.go | 18 ++++++++++++++---- pkg/middleware/csrf.go | 38 ++++++++++++++++++++++++++++++++++++++ 3 files changed, 53 insertions(+), 4 deletions(-) create mode 100644 pkg/middleware/csrf.go diff --git a/pkg/api/http_server.go b/pkg/api/http_server.go index f68d7b816b7..dc610567d3f 100644 --- a/pkg/api/http_server.go +++ b/pkg/api/http_server.go @@ -419,6 +419,7 @@ func (hs *HTTPServer) addMiddlewaresAndStaticRoutes() { } m.Use(middleware.Recovery(hs.Cfg)) + m.UseMiddleware(middleware.CSRF(hs.Cfg.LoginCookieName)) hs.mapStatic(m, hs.Cfg.StaticRootPath, "build", "public/build") hs.mapStatic(m, hs.Cfg.StaticRootPath, "", "public") diff --git a/pkg/macaron/binding.go b/pkg/macaron/binding.go index 0ea03925b19..fd37d2bf91e 100644 --- a/pkg/macaron/binding.go +++ b/pkg/macaron/binding.go @@ -2,8 +2,10 @@ package macaron import ( "encoding/json" + "errors" "fmt" "io" + "mime" "net/http" "reflect" ) @@ -11,9 +13,16 @@ import ( // Bind deserializes JSON payload from the request func Bind(req *http.Request, v interface{}) error { if req.Body != nil { - defer req.Body.Close() - err := json.NewDecoder(req.Body).Decode(v) - if err != nil && err != io.EOF { + m, _, err := mime.ParseMediaType(req.Header.Get("Content-type")) + if err != nil { + return err + } + if m != "application/json" { + return errors.New("bad content type") + } + defer func() { _ = req.Body.Close() }() + err = json.NewDecoder(req.Body).Decode(v) + if err != nil && !errors.Is(err, io.EOF) { return err } } @@ -29,7 +38,7 @@ func validate(obj interface{}) error { if validator, ok := obj.(Validator); ok { return validator.Validate() } - // Otherwise, use relfection to match `binding:"Required"` struct field tags. + // Otherwise, use reflection to match `binding:"Required"` struct field tags. // Resolve all pointers and interfaces, until we get a concrete type. t := reflect.TypeOf(obj) v := reflect.ValueOf(obj) @@ -69,6 +78,7 @@ func validate(obj interface{}) error { return err } } + default: // ignore } return nil } diff --git a/pkg/middleware/csrf.go b/pkg/middleware/csrf.go new file mode 100644 index 00000000000..190c3c758cf --- /dev/null +++ b/pkg/middleware/csrf.go @@ -0,0 +1,38 @@ +package middleware + +import ( + "net/http" + "net/url" + "strings" +) + +func CSRF(loginCookieName string) func(http.Handler) http.Handler { + // As per RFC 7231/4.2.2 these methods are idempotent: + // (GET is excluded because it may have side effects in some APIs) + safeMethods := []string{"HEAD", "OPTIONS", "TRACE"} + + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // If request has no login cookie - skip CSRF checks + if _, err := r.Cookie(loginCookieName); err == http.ErrNoCookie { + next.ServeHTTP(w, r) + return + } + // Skip CSRF checks for "safe" methods + for _, method := range safeMethods { + if r.Method == method { + next.ServeHTTP(w, r) + return + } + } + // Otherwise - verify that Origin matches the server origin + host := strings.Split(r.Host, ":")[0] + origin, err := url.Parse(r.Header.Get("Origin")) + if err != nil || (origin.String() != "" && origin.Hostname() != host) { + http.Error(w, "origin not allowed", http.StatusForbidden) + return + } + next.ServeHTTP(w, r) + }) + } +}