diff --git a/README.md b/README.md index c625996c..8dc3fc99 100644 --- a/README.md +++ b/README.md @@ -189,6 +189,7 @@ _this is the recommended way to run remark42_ | emoji | EMOJI | `false` | enable emoji support | | simple-view | SIMPLE_VIEW | `false` | minimized UI with basic info only | | proxy-cors | PROXY_CORS | `false` | disable internal CORS and delegate it to proxy | +| allowed-hosts | ALLOWED_HOSTS enable all | limit hosts/sources allowed to embed comments | | port | REMARK_PORT | `8080` | web server port | | web-root | REMARK_WEB_ROOT | `./web` | web server root directory | | update-limit | UPDATE_LIMIT | `0.5` | updates/sec limit | diff --git a/backend/app/cmd/server.go b/backend/app/cmd/server.go index 4430cd87..5dc2a57f 100644 --- a/backend/app/cmd/server.go +++ b/backend/app/cmd/server.go @@ -75,6 +75,7 @@ type ServerCommand struct { EnableEmoji bool `long:"emoji" env:"EMOJI" description:"enable emoji"` SimpleView bool `long:"simpler-view" env:"SIMPLE_VIEW" description:"minimal comment editor mode"` ProxyCORS bool `long:"proxy-cors" env:"PROXY_CORS" description:"disable internal CORS and delegate it to proxy"` + AllowedHosts []string `long:"allowed-hosts" env:"ALLOWED_HOSTS" description:"limit hosts/sources allowed to embed comments"` Auth struct { TTL struct { @@ -443,6 +444,7 @@ func (s *ServerCommand) newServerApp() (*serverApp, error) { AnonVote: s.AnonymousVote && s.RestrictVoteIP, SimpleView: s.SimpleView, ProxyCORS: s.ProxyCORS, + AllowedAncestors: s.AllowedHosts, SendJWTHeader: s.Auth.TTL.SendJWTHeader, } diff --git a/backend/app/rest/api/rest.go b/backend/app/rest/api/rest.go index 267661f5..33a4ac24 100644 --- a/backend/app/rest/api/rest.go +++ b/backend/app/rest/api/rest.go @@ -62,6 +62,7 @@ type Rest struct { SimpleView bool ProxyCORS bool SendJWTHeader bool + AllowedAncestors []string // sets Content-Security-Policy "frame-ancestors ..." SSLConfig SSLConfig httpsServer *http.Server @@ -201,6 +202,11 @@ func (s *Rest) routes() chi.Router { router.Use(corsMiddleware.Handler) } + if len(s.AllowedAncestors) > 0 { + log.Printf("[INFO] allowed from %+v only", s.AllowedAncestors) + router.Use(frameAncestors(s.AllowedAncestors)) + } + ipFn := func(ip string) string { return store.HashValue(ip, s.SharedSecret)[:12] } // logger uses it for anonymization logInfoWithBody := logger.New(logger.Log(log.Default()), logger.WithBody, logger.IPfn(ipFn), logger.Prefix("[INFO]")).Handler @@ -595,6 +601,22 @@ func cacheControl(expiration time.Duration, version string) func(http.Handler) h } } +// frameAncestors is a middleware setting Content-Security-Policy "frame-ancestors host1 host2 ..." +// prevents loading of comments widgets from any other origins. In case if the list of allowed empty, ignored. +func frameAncestors(hosts []string) func(http.Handler) http.Handler { + return func(h http.Handler) http.Handler { + fn := func(w http.ResponseWriter, r *http.Request) { + if len(hosts) == 0 { + h.ServeHTTP(w, r) + return + } + w.Header().Set("Content-Security-Policy", "frame-ancestors "+strings.Join(hosts, " ")+";") + h.ServeHTTP(w, r) + } + return http.HandlerFunc(fn) + } +} + func parseError(err error, defaultCode int) (code int) { code = defaultCode diff --git a/backend/app/rest/api/rest_test.go b/backend/app/rest/api/rest_test.go index c918174b..984ff8c4 100644 --- a/backend/app/rest/api/rest_test.go +++ b/backend/app/rest/api/rest_test.go @@ -333,6 +333,35 @@ func TestRest_cacheControl(t *testing.T) { } +func TestRest_frameAncestors(t *testing.T) { + + tbl := []struct { + hosts []string + header string + }{ + {[]string{"http://example.com"}, "frame-ancestors http://example.com;"}, + {[]string{}, ""}, + {[]string{"http://example.com", "http://example2.com"}, "frame-ancestors http://example.com http://example2.com;"}, + } + + for i, tt := range tbl { + tt := tt + t.Run(strconv.Itoa(i), func(t *testing.T) { + req := httptest.NewRequest("GET", "http://example.com", nil) + w := httptest.NewRecorder() + + h := frameAncestors(tt.hosts)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {})) + h.ServeHTTP(w, req) + resp := w.Result() + assert.Equal(t, http.StatusOK, resp.StatusCode) + t.Logf("%+v", resp.Header) + assert.Equal(t, tt.header, resp.Header.Get("Content-Security-Policy")) + + }) + } + +} + // randomPath pick a file or folder name which is not in use for sure func randomPath(tempDir, basename, suffix string) (string, error) { for i := 0; i < 10; i++ {