diff --git a/backend/app/rest/api/middleware.go b/backend/app/rest/api/middleware.go new file mode 100644 index 00000000..8115b757 --- /dev/null +++ b/backend/app/rest/api/middleware.go @@ -0,0 +1,237 @@ +// Package api middleware: request-scoped HTTP middlewares used by the REST router. +package api + +import ( + "context" + "fmt" + "net/http" + "net/mail" + "regexp" + "strings" + "time" + + "github.com/didip/tollbooth/v8" + "github.com/didip/tollbooth/v8/limiter" + log "github.com/go-pkgz/lgr" + "github.com/umputun/remark42/backend/app/rest" + "github.com/umputun/remark42/backend/app/store" +) + +// timeout returns a middleware matching chi's middleware.Timeout: it sets a +// deadline on the request context and writes 504 Gateway Timeout if the +// deadline is exceeded. The 504 is sent once the downstream handler returns +// after observing the canceled context; a handler that ignores r.Context() +// is not aborted. +func timeout(d time.Duration) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ctx, cancel := context.WithTimeout(r.Context(), d) + defer func() { + cancel() + if ctx.Err() == context.DeadlineExceeded { + w.WriteHeader(http.StatusGatewayTimeout) + } + }() + next.ServeHTTP(w, r.WithContext(ctx)) + }) + } +} + +// rejectAnonUser is a middleware rejecting anonymous users +func rejectAnonUser(next http.Handler) http.Handler { + fn := func(w http.ResponseWriter, r *http.Request) { + user, err := rest.GetUserInfo(r) + if err != nil { + http.Error(w, "Unauthorized", http.StatusUnauthorized) + return + } + + if strings.HasPrefix(user.ID, "anonymous_") { + http.Error(w, "Access denied", http.StatusForbidden) + return + } + next.ServeHTTP(w, r) + } + return http.HandlerFunc(fn) +} + +// matchSiteID is a middleware rejecting users with mismatch between site param and and User.SiteID +func matchSiteID(next http.Handler) http.Handler { + fn := func(w http.ResponseWriter, r *http.Request) { + user, err := rest.GetUserInfo(r) + if err != nil { + http.Error(w, "Unauthorized", http.StatusUnauthorized) + return + } + + // skip for basic auth user + if user.Name == "admin" && user.ID == "admin" { + next.ServeHTTP(w, r) + return + } + + siteID := r.URL.Query().Get("site") + // require an explicit site so the user.SiteID check below cannot be bypassed + // by simply omitting the query parameter + if siteID == "" || user.SiteID != siteID { + http.Error(w, "Access denied", http.StatusForbidden) + return + } + next.ServeHTTP(w, r) + } + return http.HandlerFunc(fn) +} + +// cacheControl is a middleware setting cache expiration. Using url+version as etag +func cacheControl(expiration time.Duration, version string) func(http.Handler) http.Handler { + etag := func(r *http.Request, version string) string { + s := version + ":" + r.URL.String() + return store.EncodeID(s) + } + + return func(h http.Handler) http.Handler { + fn := func(w http.ResponseWriter, r *http.Request) { + e := `"` + etag(r, version) + `"` + w.Header().Set("Etag", e) + w.Header().Set("Cache-Control", fmt.Sprintf("max-age=%d, no-cache", int(expiration.Seconds()))) + + if match := r.Header.Get("If-None-Match"); match != "" { + if strings.Contains(match, e) { + w.WriteHeader(http.StatusNotModified) + return + } + } + h.ServeHTTP(w, r) + } + return http.HandlerFunc(fn) + } +} + +// apiCSPMiddleware overrides the global Content-Security-Policy on /api/v1 routes +// with a strict, default-deny policy. The global CSP (securityHeadersMiddleware) keeps +// 'self' 'unsafe-inline' for script-src/style-src because the widget HTML pages +// (/web/*.html) need inline bootstrap blocks. API responses serve JSON, XML/RSS, or +// images — none of those should ever execute scripts when rendered, so they get the +// strictest policy available as defense-in-depth against future trust-boundary bugs. +// +// Image-serving handlers (/api/v1/img, /api/v1/picture/{user}/{id}) re-apply the same +// rest.StrictImageCSP value at the handler level and additionally set Content-Disposition: +// inline; filename="image" (framing the response as a file rather than a renderable +// document) and X-Content-Type-Options: nosniff. The CSP re-apply is intentional belt-and- +// braces: if a future route refactor bypasses this middleware, the image handlers still +// emit the policy. +func apiCSPMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Security-Policy", rest.StrictImageCSP) + next.ServeHTTP(w, r) + }) +} + +// securityHeadersMiddleware sets security-related headers: +// - Content-Security-Policy: controls which resources the browser is allowed to load +// - Permissions-Policy: disables browser features (camera, mic, etc.) not needed by a comment widget +// - X-Content-Type-Options: prevents browsers from MIME-sniffing responses away from the declared type, +// stopping e.g. a user-uploaded image from being reinterpreted as executable HTML/JS +// - Referrer-Policy: controls how much URL information leaks in the Referer header on cross-origin +// requests; "strict-origin-when-cross-origin" sends only the origin (no path) to other domains +// and nothing at all on HTTPS→HTTP downgrades +func securityHeadersMiddleware(imageProxyEnabled bool, allowedAncestors []string) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + imgSrc := "*" + if imageProxyEnabled { + imgSrc = "'self'" + } + frameAncestors := "*" + if len(allowedAncestors) > 0 { + frameAncestors = strings.Join(allowedAncestors, " ") + } + // font-src is set to 'none' (no @font-face / no base64 fonts in the bundle). + w.Header().Set("Content-Security-Policy", fmt.Sprintf("default-src 'none'; base-uri 'none'; form-action 'none'; connect-src 'self'; frame-src 'self' mailto:; img-src %s; script-src 'self' 'unsafe-inline'; style-src 'self' 'unsafe-inline'; font-src 'none'; object-src 'none'; frame-ancestors %s;", imgSrc, frameAncestors)) + w.Header().Set("Permissions-Policy", "accelerometer=(), autoplay=(), camera=(), cross-origin-isolated=(), display-capture=(), encrypted-media=(), fullscreen=(), geolocation=(), gyroscope=(), keyboard-map=(), magnetometer=(), microphone=(), midi=(), payment=(), picture-in-picture=(), publickey-credentials-get=(), screen-wake-lock=(), sync-xhr=(), usb=(), xr-spatial-tracking=(), clipboard-read=(), clipboard-write=(), gamepad=(), hid=(), idle-detection=(), interest-cohort=(), serial=(), unload=(), window-management=()") + w.Header().Set("X-Content-Type-Options", "nosniff") + w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin") + next.ServeHTTP(w, r) + }) + } +} + +// subscribersOnly is a middleware rejecting non-paid_sub users +func subscribersOnly(enable bool) func(http.Handler) http.Handler { + return func(h http.Handler) http.Handler { + fn := func(w http.ResponseWriter, r *http.Request) { + if enable { + user, err := rest.GetUserInfo(r) + if err != nil { + http.Error(w, "Unauthorized", http.StatusUnauthorized) + return + } + if !user.PaidSub { + http.Error(w, "Access denied", http.StatusForbidden) + return + } + } + h.ServeHTTP(w, r) + } + return http.HandlerFunc(fn) + } +} + +// validEmailAuth is a middleware for auth endpoints for email method. +// it rejects login request if user, site or email are suspicious +func validEmailAuth() func(http.Handler) http.Handler { + + reUser := regexp.MustCompile(`^[\p{L}\d\s_]{4,64}$`) // matches ui side validation, adding min/max limitation + reSite := regexp.MustCompile(`^[a-zA-Z\d\s_.-]{1,64}$`) + + return func(h http.Handler) http.Handler { + fn := func(w http.ResponseWriter, r *http.Request) { + + if r.URL.Path != "/auth/email/login" { + // not email login, skip the check + h.ServeHTTP(w, r) + return + } + + if u := r.URL.Query().Get("user"); u != "" { + if !reUser.MatchString(u) { + log.Printf("[WARN] suspicious user rejected: %s", u) + http.Error(w, "Access denied", http.StatusForbidden) + return + } + } + + if a := r.URL.Query().Get("address"); a != "" { + if _, err := mail.ParseAddress(a); err != nil { + log.Printf("[WARN] suspicious address rejected: %s", a) + http.Error(w, "Access denied", http.StatusForbidden) + return + } + } + + if s := r.URL.Query().Get("site"); s != "" { + if !reSite.MatchString(s) { + log.Printf("[WARN] suspicious site rejected: %s", s) + http.Error(w, "Access denied", http.StatusForbidden) + return + } + } + + h.ServeHTTP(w, r) + } + return http.HandlerFunc(fn) + } +} + +// rateLimiter creates a rate limiting middleware with proper IP lookup configuration. +// tollbooth v8 requires explicit IP lookup method to be set. +// uses RemoteAddr which is set by rest.RealIP to the real client IP +// from X-Forwarded-For, X-Real-IP, or True-Client-IP headers. +func rateLimiter(maxReq float64) func(http.Handler) http.Handler { + lmt := tollbooth.NewLimiter(maxReq, nil) + lmt.SetIPLookup(limiter.IPLookup{ + Name: "RemoteAddr", + IndexFromRight: 0, + }) + return tollbooth.HTTPMiddleware(lmt) +} diff --git a/backend/app/rest/api/middleware_test.go b/backend/app/rest/api/middleware_test.go new file mode 100644 index 00000000..01150e29 --- /dev/null +++ b/backend/app/rest/api/middleware_test.go @@ -0,0 +1,266 @@ +package api + +import ( + "fmt" + "net/http" + "net/http/httptest" + "strconv" + "testing" + "time" + + "github.com/go-pkgz/auth/v2/token" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/umputun/remark42/backend/app/rest" + "github.com/umputun/remark42/backend/app/store" +) + +func TestTimeout(t *testing.T) { + t.Run("fast handler passes through and gets a deadline", func(t *testing.T) { + var gotDeadline bool + h := timeout(time.Second)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, gotDeadline = r.Context().Deadline() + w.WriteHeader(http.StatusCreated) + _, _ = w.Write([]byte("ok")) + })) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", http.NoBody)) + assert.True(t, gotDeadline, "request context should carry a deadline") + assert.Equal(t, http.StatusCreated, rec.Code) + assert.Equal(t, "ok", rec.Body.String()) + }) + + t.Run("deadline exceeded writes 504", func(t *testing.T) { + h := timeout(10 * time.Millisecond)(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + <-r.Context().Done() // honor the context: return only once the deadline fires + })) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", http.NoBody)) + assert.Equal(t, http.StatusGatewayTimeout, rec.Code) + }) +} + +func TestRest_rejectAnonUser(t *testing.T) { + ts := httptest.NewServer(fakeAuth(rejectAnonUser(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + fmt.Fprintln(w, "Hello") + })))) + defer ts.Close() + + resp, err := http.Get(ts.URL) + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode, "use not logged in") + + resp, err = http.Get(ts.URL + "?fake_id=anonymous_user123&fake_name=test") + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) + assert.Equal(t, http.StatusForbidden, resp.StatusCode, "anon rejected") + + resp, err = http.Get(ts.URL + "?fake_id=real_user123&fake_name=test") + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) + assert.Equal(t, http.StatusOK, resp.StatusCode, "real user") +} + +func TestRest_cacheControl(t *testing.T) { + tbl := []struct { + url string + version string + exp time.Duration + etag string + maxAge int + }{ + {"http://example.com/foo", "v1", time.Hour, "b433be1ea19edaee9dc92ca4b895b6bdf3c058cb", 3600}, + {"http://example.com/foo2", "v1", 10 * time.Hour, "6d8466aef3246c1057452561acddf7ad9d0d99e0", 36000}, + {"http://example.com/foo", "v2", time.Hour, "481700c52aab0dfbca99f3ffc2a4fbb27884c114", 3600}, + {"https://example.com/foo", "v2", time.Hour, "bebd4f1b87f474792c4e75e5affe31fbf67f5778", 3600}, + } + + for i, tt := range tbl { + t.Run(strconv.Itoa(i), func(t *testing.T) { + req := httptest.NewRequest("GET", tt.url, http.NoBody) + w := httptest.NewRecorder() + + h := cacheControl(tt.exp, tt.version)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + h.ServeHTTP(w, req) + resp := w.Result() + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.NoError(t, resp.Body.Close()) + t.Logf("%+v", resp.Header) + assert.Equal(t, `"`+tt.etag+`"`, resp.Header.Get("Etag")) + assert.Equal(t, `max-age=`+strconv.Itoa(int(tt.exp.Seconds()))+", no-cache", resp.Header.Get("Cache-Control")) + }) + } +} + +// TestRest_apiCSP locks in that /api/v1/* responses get a strict default-src 'none' +// override regardless of what the global CSP allows. The widget HTML pages +// (/web/*.html) still get the global CSP (with 'unsafe-inline' for bootstrap), +// so the test asserts the two policies diverge across origins. +func TestRest_apiCSP(t *testing.T) { + ts, _, teardown := startupT(t) + defer teardown() + client := http.Client{} + + // JSON API endpoint — must carry the strict policy + resp, err := client.Get(ts.URL + "/api/v1/config") + require.NoError(t, err) + defer resp.Body.Close() + csp := resp.Header.Get("Content-Security-Policy") + assert.Contains(t, csp, "default-src 'none'", + "API responses must override the global CSP with default-src 'none'; got %q", csp) + assert.Contains(t, csp, "sandbox", "API CSP must include sandbox; got %q", csp) + assert.NotContains(t, csp, "'unsafe-inline'", + "API CSP must not allow inline scripts/styles; got %q", csp) + + // RSS/XML endpoint — same strict policy, and the XML response itself must still be served + respRSS, err := client.Get(ts.URL + "/api/v1/rss/site?site=remark42") + require.NoError(t, err) + defer respRSS.Body.Close() + assert.Equal(t, http.StatusOK, respRSS.StatusCode, "RSS must still respond OK under strict CSP") + cspRSS := respRSS.Header.Get("Content-Security-Policy") + assert.Contains(t, cspRSS, "default-src 'none'", "RSS responses must carry the strict API CSP") + assert.Contains(t, cspRSS, "sandbox", "RSS CSP must include sandbox") + + // widget HTML — must keep the global CSP (unchanged, lax to support inline bootstrap) + resp2, err := client.Get(ts.URL + "/web/index.html") + require.NoError(t, err) + defer resp2.Body.Close() + csp2 := resp2.Header.Get("Content-Security-Policy") + assert.Contains(t, csp2, "'unsafe-inline'", + "widget HTML CSP must keep unsafe-inline for bootstrap; got %q", csp2) +} + +// check CSP, img-src should be 'self' with proxy enabled and * without it +func TestRest_securityHeaders(t *testing.T) { + ts, _, teardown := startupT(t) + + // with proxy disabled + client := http.Client{} + resp, err := client.Get(ts.URL + "/web/index.html") + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Contains(t, resp.Header.Get("Content-Security-Policy"), "img-src *;") + assert.Equal(t, "nosniff", resp.Header.Get("X-Content-Type-Options")) + assert.Equal(t, "strict-origin-when-cross-origin", resp.Header.Get("Referrer-Policy")) + teardown() + + // check CSP with proxy enabled + ts, _, teardown = startupT(t, func(srv *Rest) { + srv.ExternalImageProxy = true + }) + defer teardown() + resp, err = client.Get(ts.URL + "/web/index.html") + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Contains(t, resp.Header.Get("Content-Security-Policy"), "img-src 'self';") + assert.Equal(t, "nosniff", resp.Header.Get("X-Content-Type-Options")) + assert.Equal(t, "strict-origin-when-cross-origin", resp.Header.Get("Referrer-Policy")) +} + +func TestRest_subscribersOnly(t *testing.T) { + paidSubUser := &token.User{} + paidSubUser.SetPaidSub(true) + + tbl := []struct { + subsOnly bool + user token.User + setUser bool + status int + }{ + {true, token.User{}, false, http.StatusUnauthorized}, + {true, token.User{}, true, http.StatusForbidden}, + {false, token.User{}, false, http.StatusOK}, + {false, token.User{}, true, http.StatusOK}, + {true, *paidSubUser, true, http.StatusOK}, + } + + for i, tt := range tbl { + t.Run(strconv.Itoa(i), func(t *testing.T) { + req := httptest.NewRequest("GET", "http://example.com", http.NoBody) + if tt.setUser { + req = token.SetUserInfo(req, tt.user) + } + w := httptest.NewRecorder() + h := subscribersOnly(tt.subsOnly)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + h.ServeHTTP(w, req) + resp := w.Result() + assert.Equal(t, tt.status, resp.StatusCode) + assert.NoError(t, resp.Body.Close()) + }) + } +} + +func Test_validEmailAuth(t *testing.T) { + tbl := []struct { + req string + status int + }{ + {"/auth/email/login?site=remark42&address=umputun%example.com&user=someone", http.StatusOK}, + {"/auth/email/login?site=site-with-dash_and_underscore-and.dot&address=umputun%example.com&user=someone", http.StatusOK}, + {"/auth/email/login?site=remark42&address=umputun%example.com&user=someone+blah", http.StatusOK}, + {"/auth/email/login?site=remark42&address=umputun%example.com&user=Евгений+Умпутун", http.StatusOK}, + {"/auth/email/login?site=remark42&address=umputun%example.com&user=12", http.StatusForbidden}, + {"/auth/email/login?site=remark42&address=umputun%example.com&user=..blah+blah", http.StatusForbidden}, + {"/auth/email/login?site=remark42&address=umputun%example.com&user=someonelooong+loooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooong", http.StatusForbidden}, + {"/auth/twitter/login?site=remark42&address=umputun%example.com&user=..blah+blah", http.StatusOK}, + {"/auth/email/login?site=remark42&address=umputun%example.com", http.StatusOK}, + {"/auth/email/login?site=remark42&address=umputun+example.com&user=someone", http.StatusForbidden}, + {"/auth/email/login?site=bad!site&address=umputun%example.com&user=someone", http.StatusForbidden}, + {"/auth/email/login?site=loooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooongsite&address=umputun%example.com&user=someone", http.StatusForbidden}, + } + + for i, tt := range tbl { + t.Run(strconv.Itoa(i), func(t *testing.T) { + req := httptest.NewRequest("GET", "http://example.com"+tt.req, http.NoBody) + w := httptest.NewRecorder() + h := validEmailAuth()(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + h.ServeHTTP(w, req) + resp := w.Result() + assert.Equal(t, tt.status, resp.StatusCode) + assert.NoError(t, resp.Body.Close()) + }) + } +} + +// TestRest_matchSiteID reproduces the multi-tenant isolation gap in the matchSiteID +// middleware. Before the fix, the check `if siteID != "" && user.SiteID != siteID` +// silently allowed any authenticated request that omitted the ?site= query param. +// On admin and user-mutation routes this meant the cross-site check was bypassable +// just by dropping the parameter. The fix requires ?site= to be present and to match +// the user's bound site. +func TestRest_matchSiteID(t *testing.T) { + wrapped := matchSiteID(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("ok")) + })) + + cases := []struct { + name string + userSite string + query string + want int + }{ + {name: "matching site allowed", userSite: "site-a", query: "?site=site-a", want: http.StatusOK}, + {name: "mismatched site forbidden", userSite: "site-a", query: "?site=site-b", want: http.StatusForbidden}, + {name: "missing site param rejected", userSite: "site-a", query: "", want: http.StatusForbidden}, + {name: "empty site param rejected", userSite: "site-a", query: "?site=", want: http.StatusForbidden}, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + h := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + r = rest.SetUserInfo(r, store.User{ID: "u", Name: "u", SiteID: c.userSite}) + wrapped.ServeHTTP(w, r) + }) + ts := httptest.NewServer(h) + defer ts.Close() + resp, err := http.Get(ts.URL + c.query) + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) + assert.Equal(t, c.want, resp.StatusCode) + }) + } +} diff --git a/backend/app/rest/api/rest.go b/backend/app/rest/api/rest.go index 6cc46fc7..a052dfa0 100644 --- a/backend/app/rest/api/rest.go +++ b/backend/app/rest/api/rest.go @@ -8,15 +8,11 @@ import ( "fmt" "io/fs" "net/http" - "net/mail" "os" - "regexp" "strings" "sync" "time" - "github.com/didip/tollbooth/v8" - "github.com/didip/tollbooth/v8/limiter" "github.com/go-chi/chi/v5" "github.com/go-chi/cors" "github.com/go-pkgz/auth/v2" @@ -572,192 +568,6 @@ func URLKeyWithUser(r *http.Request) string { return key } -// rejectAnonUser is a middleware rejecting anonymous users -func rejectAnonUser(next http.Handler) http.Handler { - fn := func(w http.ResponseWriter, r *http.Request) { - user, err := rest.GetUserInfo(r) - if err != nil { - http.Error(w, "Unauthorized", http.StatusUnauthorized) - return - } - - if strings.HasPrefix(user.ID, "anonymous_") { - http.Error(w, "Access denied", http.StatusForbidden) - return - } - next.ServeHTTP(w, r) - } - return http.HandlerFunc(fn) -} - -// matchSiteID is a middleware rejecting users with mismatch between site param and and User.SiteID -func matchSiteID(next http.Handler) http.Handler { - fn := func(w http.ResponseWriter, r *http.Request) { - user, err := rest.GetUserInfo(r) - if err != nil { - http.Error(w, "Unauthorized", http.StatusUnauthorized) - return - } - - // skip for basic auth user - if user.Name == "admin" && user.ID == "admin" { - next.ServeHTTP(w, r) - return - } - - siteID := r.URL.Query().Get("site") - // require an explicit site so the user.SiteID check below cannot be bypassed - // by simply omitting the query parameter - if siteID == "" || user.SiteID != siteID { - http.Error(w, "Access denied", http.StatusForbidden) - return - } - next.ServeHTTP(w, r) - } - return http.HandlerFunc(fn) -} - -// cacheControl is a middleware setting cache expiration. Using url+version as etag -func cacheControl(expiration time.Duration, version string) func(http.Handler) http.Handler { - etag := func(r *http.Request, version string) string { - s := version + ":" + r.URL.String() - return store.EncodeID(s) - } - - return func(h http.Handler) http.Handler { - fn := func(w http.ResponseWriter, r *http.Request) { - e := `"` + etag(r, version) + `"` - w.Header().Set("Etag", e) - w.Header().Set("Cache-Control", fmt.Sprintf("max-age=%d, no-cache", int(expiration.Seconds()))) - - if match := r.Header.Get("If-None-Match"); match != "" { - if strings.Contains(match, e) { - w.WriteHeader(http.StatusNotModified) - return - } - } - h.ServeHTTP(w, r) - } - return http.HandlerFunc(fn) - } -} - -// apiCSPMiddleware overrides the global Content-Security-Policy on /api/v1 routes -// with a strict, default-deny policy. The global CSP (securityHeadersMiddleware) keeps -// 'self' 'unsafe-inline' for script-src/style-src because the widget HTML pages -// (/web/*.html) need inline bootstrap blocks. API responses serve JSON, XML/RSS, or -// images — none of those should ever execute scripts when rendered, so they get the -// strictest policy available as defense-in-depth against future trust-boundary bugs. -// -// Image-serving handlers (/api/v1/img, /api/v1/picture/{user}/{id}) re-apply the same -// rest.StrictImageCSP value at the handler level and additionally set Content-Disposition: -// inline; filename="image" (framing the response as a file rather than a renderable -// document) and X-Content-Type-Options: nosniff. The CSP re-apply is intentional belt-and- -// braces: if a future route refactor bypasses this middleware, the image handlers still -// emit the policy. -func apiCSPMiddleware(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Security-Policy", rest.StrictImageCSP) - next.ServeHTTP(w, r) - }) -} - -// securityHeadersMiddleware sets security-related headers: -// - Content-Security-Policy: controls which resources the browser is allowed to load -// - Permissions-Policy: disables browser features (camera, mic, etc.) not needed by a comment widget -// - X-Content-Type-Options: prevents browsers from MIME-sniffing responses away from the declared type, -// stopping e.g. a user-uploaded image from being reinterpreted as executable HTML/JS -// - Referrer-Policy: controls how much URL information leaks in the Referer header on cross-origin -// requests; "strict-origin-when-cross-origin" sends only the origin (no path) to other domains -// and nothing at all on HTTPS→HTTP downgrades -func securityHeadersMiddleware(imageProxyEnabled bool, allowedAncestors []string) func(http.Handler) http.Handler { - return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - imgSrc := "*" - if imageProxyEnabled { - imgSrc = "'self'" - } - frameAncestors := "*" - if len(allowedAncestors) > 0 { - frameAncestors = strings.Join(allowedAncestors, " ") - } - // font-src is set to 'none' (no @font-face / no base64 fonts in the bundle). - w.Header().Set("Content-Security-Policy", fmt.Sprintf("default-src 'none'; base-uri 'none'; form-action 'none'; connect-src 'self'; frame-src 'self' mailto:; img-src %s; script-src 'self' 'unsafe-inline'; style-src 'self' 'unsafe-inline'; font-src 'none'; object-src 'none'; frame-ancestors %s;", imgSrc, frameAncestors)) - w.Header().Set("Permissions-Policy", "accelerometer=(), autoplay=(), camera=(), cross-origin-isolated=(), display-capture=(), encrypted-media=(), fullscreen=(), geolocation=(), gyroscope=(), keyboard-map=(), magnetometer=(), microphone=(), midi=(), payment=(), picture-in-picture=(), publickey-credentials-get=(), screen-wake-lock=(), sync-xhr=(), usb=(), xr-spatial-tracking=(), clipboard-read=(), clipboard-write=(), gamepad=(), hid=(), idle-detection=(), interest-cohort=(), serial=(), unload=(), window-management=()") - w.Header().Set("X-Content-Type-Options", "nosniff") - w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin") - next.ServeHTTP(w, r) - }) - } -} - -// subscribersOnly is a middleware rejecting non-paid_sub users -func subscribersOnly(enable bool) func(http.Handler) http.Handler { - return func(h http.Handler) http.Handler { - fn := func(w http.ResponseWriter, r *http.Request) { - if enable { - user, err := rest.GetUserInfo(r) - if err != nil { - http.Error(w, "Unauthorized", http.StatusUnauthorized) - return - } - if !user.PaidSub { - http.Error(w, "Access denied", http.StatusForbidden) - return - } - } - h.ServeHTTP(w, r) - } - return http.HandlerFunc(fn) - } -} - -// validEmailAuth is a middleware for auth endpoints for email method. -// it rejects login request if user, site or email are suspicious -func validEmailAuth() func(http.Handler) http.Handler { - - reUser := regexp.MustCompile(`^[\p{L}\d\s_]{4,64}$`) // matches ui side validation, adding min/max limitation - reSite := regexp.MustCompile(`^[a-zA-Z\d\s_.-]{1,64}$`) - - return func(h http.Handler) http.Handler { - fn := func(w http.ResponseWriter, r *http.Request) { - - if r.URL.Path != "/auth/email/login" { - // not email login, skip the check - h.ServeHTTP(w, r) - return - } - - if u := r.URL.Query().Get("user"); u != "" { - if !reUser.MatchString(u) { - log.Printf("[WARN] suspicious user rejected: %s", u) - http.Error(w, "Access denied", http.StatusForbidden) - return - } - } - - if a := r.URL.Query().Get("address"); a != "" { - if _, err := mail.ParseAddress(a); err != nil { - log.Printf("[WARN] suspicious address rejected: %s", a) - http.Error(w, "Access denied", http.StatusForbidden) - return - } - } - - if s := r.URL.Query().Get("site"); s != "" { - if !reSite.MatchString(s) { - log.Printf("[WARN] suspicious site rejected: %s", s) - http.Error(w, "Access denied", http.StatusForbidden) - return - } - } - - h.ServeHTTP(w, r) - } - return http.HandlerFunc(fn) - } -} - func parseError(err error, defaultCode int) (code int) { code = defaultCode @@ -781,16 +591,3 @@ func parseError(err error, defaultCode int) (code int) { return code } - -// rateLimiter creates a rate limiting middleware with proper IP lookup configuration. -// tollbooth v8 requires explicit IP lookup method to be set. -// uses RemoteAddr which is set by rest.RealIP to the real client IP -// from X-Forwarded-For, X-Real-IP, or True-Client-IP headers. -func rateLimiter(maxReq float64) func(http.Handler) http.Handler { - lmt := tollbooth.NewLimiter(maxReq, nil) - lmt.SetIPLookup(limiter.IPLookup{ - Name: "RemoteAddr", - IndexFromRight: 0, - }) - return tollbooth.HTTPMiddleware(lmt) -} diff --git a/backend/app/rest/api/rest_test.go b/backend/app/rest/api/rest_test.go index 85121e28..ea293f90 100644 --- a/backend/app/rest/api/rest_test.go +++ b/backend/app/rest/api/rest_test.go @@ -193,28 +193,6 @@ func TestRest_RunAutocertModeHTTPOnly(t *testing.T) { srv.Shutdown() } -func TestRest_rejectAnonUser(t *testing.T) { - ts := httptest.NewServer(fakeAuth(rejectAnonUser(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - fmt.Fprintln(w, "Hello") - })))) - defer ts.Close() - - resp, err := http.Get(ts.URL) - require.NoError(t, err) - require.NoError(t, resp.Body.Close()) - assert.Equal(t, http.StatusUnauthorized, resp.StatusCode, "use not logged in") - - resp, err = http.Get(ts.URL + "?fake_id=anonymous_user123&fake_name=test") - require.NoError(t, err) - require.NoError(t, resp.Body.Close()) - assert.Equal(t, http.StatusForbidden, resp.StatusCode, "anon rejected") - - resp, err = http.Get(ts.URL + "?fake_id=real_user123&fake_name=test") - require.NoError(t, err) - require.NoError(t, resp.Body.Close()) - assert.Equal(t, http.StatusOK, resp.StatusCode, "real user") -} - func Test_URLKey(t *testing.T) { tbl := []struct { url string @@ -284,37 +262,6 @@ func TestRest_parseError(t *testing.T) { } } -func TestRest_cacheControl(t *testing.T) { - tbl := []struct { - url string - version string - exp time.Duration - etag string - maxAge int - }{ - {"http://example.com/foo", "v1", time.Hour, "b433be1ea19edaee9dc92ca4b895b6bdf3c058cb", 3600}, - {"http://example.com/foo2", "v1", 10 * time.Hour, "6d8466aef3246c1057452561acddf7ad9d0d99e0", 36000}, - {"http://example.com/foo", "v2", time.Hour, "481700c52aab0dfbca99f3ffc2a4fbb27884c114", 3600}, - {"https://example.com/foo", "v2", time.Hour, "bebd4f1b87f474792c4e75e5affe31fbf67f5778", 3600}, - } - - for i, tt := range tbl { - t.Run(strconv.Itoa(i), func(t *testing.T) { - req := httptest.NewRequest("GET", tt.url, http.NoBody) - w := httptest.NewRecorder() - - h := cacheControl(tt.exp, tt.version)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) - h.ServeHTTP(w, req) - resp := w.Result() - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.NoError(t, resp.Body.Close()) - t.Logf("%+v", resp.Header) - assert.Equal(t, `"`+tt.etag+`"`, resp.Header.Get("Etag")) - assert.Equal(t, `max-age=`+strconv.Itoa(int(tt.exp.Seconds()))+", no-cache", resp.Header.Get("Cache-Control")) - }) - } -} - func TestRest_frameAncestors(t *testing.T) { ts, _, teardown := startupT(t, func(o *Rest) { o.AllowedAncestors = []string{"'self'", "https://example.com"} @@ -341,138 +288,6 @@ func TestRest_frameAncestors(t *testing.T) { assert.Contains(t, resp.Header.Get("Content-Security-Policy"), "frame-ancestors *;") } -// TestRest_apiCSP locks in that /api/v1/* responses get a strict default-src 'none' -// override regardless of what the global CSP allows. The widget HTML pages -// (/web/*.html) still get the global CSP (with 'unsafe-inline' for bootstrap), -// so the test asserts the two policies diverge across origins. -func TestRest_apiCSP(t *testing.T) { - ts, _, teardown := startupT(t) - defer teardown() - client := http.Client{} - - // JSON API endpoint — must carry the strict policy - resp, err := client.Get(ts.URL + "/api/v1/config") - require.NoError(t, err) - defer resp.Body.Close() - csp := resp.Header.Get("Content-Security-Policy") - assert.Contains(t, csp, "default-src 'none'", - "API responses must override the global CSP with default-src 'none'; got %q", csp) - assert.Contains(t, csp, "sandbox", "API CSP must include sandbox; got %q", csp) - assert.NotContains(t, csp, "'unsafe-inline'", - "API CSP must not allow inline scripts/styles; got %q", csp) - - // RSS/XML endpoint — same strict policy, and the XML response itself must still be served - respRSS, err := client.Get(ts.URL + "/api/v1/rss/site?site=remark42") - require.NoError(t, err) - defer respRSS.Body.Close() - assert.Equal(t, http.StatusOK, respRSS.StatusCode, "RSS must still respond OK under strict CSP") - cspRSS := respRSS.Header.Get("Content-Security-Policy") - assert.Contains(t, cspRSS, "default-src 'none'", "RSS responses must carry the strict API CSP") - assert.Contains(t, cspRSS, "sandbox", "RSS CSP must include sandbox") - - // widget HTML — must keep the global CSP (unchanged, lax to support inline bootstrap) - resp2, err := client.Get(ts.URL + "/web/index.html") - require.NoError(t, err) - defer resp2.Body.Close() - csp2 := resp2.Header.Get("Content-Security-Policy") - assert.Contains(t, csp2, "'unsafe-inline'", - "widget HTML CSP must keep unsafe-inline for bootstrap; got %q", csp2) -} - -// check CSP, img-src should be 'self' with proxy enabled and * without it -func TestRest_securityHeaders(t *testing.T) { - ts, _, teardown := startupT(t) - - // with proxy disabled - client := http.Client{} - resp, err := client.Get(ts.URL + "/web/index.html") - require.NoError(t, err) - defer resp.Body.Close() - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Contains(t, resp.Header.Get("Content-Security-Policy"), "img-src *;") - assert.Equal(t, "nosniff", resp.Header.Get("X-Content-Type-Options")) - assert.Equal(t, "strict-origin-when-cross-origin", resp.Header.Get("Referrer-Policy")) - teardown() - - // check CSP with proxy enabled - ts, _, teardown = startupT(t, func(srv *Rest) { - srv.ExternalImageProxy = true - }) - defer teardown() - resp, err = client.Get(ts.URL + "/web/index.html") - require.NoError(t, err) - defer resp.Body.Close() - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Contains(t, resp.Header.Get("Content-Security-Policy"), "img-src 'self';") - assert.Equal(t, "nosniff", resp.Header.Get("X-Content-Type-Options")) - assert.Equal(t, "strict-origin-when-cross-origin", resp.Header.Get("Referrer-Policy")) -} - -func TestRest_subscribersOnly(t *testing.T) { - paidSubUser := &token.User{} - paidSubUser.SetPaidSub(true) - - tbl := []struct { - subsOnly bool - user token.User - setUser bool - status int - }{ - {true, token.User{}, false, http.StatusUnauthorized}, - {true, token.User{}, true, http.StatusForbidden}, - {false, token.User{}, false, http.StatusOK}, - {false, token.User{}, true, http.StatusOK}, - {true, *paidSubUser, true, http.StatusOK}, - } - - for i, tt := range tbl { - t.Run(strconv.Itoa(i), func(t *testing.T) { - req := httptest.NewRequest("GET", "http://example.com", http.NoBody) - if tt.setUser { - req = token.SetUserInfo(req, tt.user) - } - w := httptest.NewRecorder() - h := subscribersOnly(tt.subsOnly)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) - h.ServeHTTP(w, req) - resp := w.Result() - assert.Equal(t, tt.status, resp.StatusCode) - assert.NoError(t, resp.Body.Close()) - }) - } -} - -func Test_validEmailAuth(t *testing.T) { - tbl := []struct { - req string - status int - }{ - {"/auth/email/login?site=remark42&address=umputun%example.com&user=someone", http.StatusOK}, - {"/auth/email/login?site=site-with-dash_and_underscore-and.dot&address=umputun%example.com&user=someone", http.StatusOK}, - {"/auth/email/login?site=remark42&address=umputun%example.com&user=someone+blah", http.StatusOK}, - {"/auth/email/login?site=remark42&address=umputun%example.com&user=Евгений+Умпутун", http.StatusOK}, - {"/auth/email/login?site=remark42&address=umputun%example.com&user=12", http.StatusForbidden}, - {"/auth/email/login?site=remark42&address=umputun%example.com&user=..blah+blah", http.StatusForbidden}, - {"/auth/email/login?site=remark42&address=umputun%example.com&user=someonelooong+loooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooong", http.StatusForbidden}, - {"/auth/twitter/login?site=remark42&address=umputun%example.com&user=..blah+blah", http.StatusOK}, - {"/auth/email/login?site=remark42&address=umputun%example.com", http.StatusOK}, - {"/auth/email/login?site=remark42&address=umputun+example.com&user=someone", http.StatusForbidden}, - {"/auth/email/login?site=bad!site&address=umputun%example.com&user=someone", http.StatusForbidden}, - {"/auth/email/login?site=loooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooooongsite&address=umputun%example.com&user=someone", http.StatusForbidden}, - } - - for i, tt := range tbl { - t.Run(strconv.Itoa(i), func(t *testing.T) { - req := httptest.NewRequest("GET", "http://example.com"+tt.req, http.NoBody) - w := httptest.NewRecorder() - h := validEmailAuth()(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) - h.ServeHTTP(w, req) - resp := w.Result() - assert.Equal(t, tt.status, resp.StatusCode) - assert.NoError(t, resp.Body.Close()) - }) - } -} - // randomPath pick a file or folder name which is not in use for sure func randomPath(tempDir, basename, suffix string) (string, error) { for range 10 { @@ -738,43 +553,3 @@ func TestMain(m *testing.M) { goleak.IgnoreTopFunction("github.com/hashicorp/golang-lru/v2/expirable.NewLRU[...].func1"), ) } - -// TestRest_matchSiteID reproduces the multi-tenant isolation gap in the matchSiteID -// middleware. Before the fix, the check `if siteID != "" && user.SiteID != siteID` -// silently allowed any authenticated request that omitted the ?site= query param. -// On admin and user-mutation routes this meant the cross-site check was bypassable -// just by dropping the parameter. The fix requires ?site= to be present and to match -// the user's bound site. -func TestRest_matchSiteID(t *testing.T) { - wrapped := matchSiteID(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte("ok")) - })) - - cases := []struct { - name string - userSite string - query string - want int - }{ - {name: "matching site allowed", userSite: "site-a", query: "?site=site-a", want: http.StatusOK}, - {name: "mismatched site forbidden", userSite: "site-a", query: "?site=site-b", want: http.StatusForbidden}, - {name: "missing site param rejected", userSite: "site-a", query: "", want: http.StatusForbidden}, - {name: "empty site param rejected", userSite: "site-a", query: "?site=", want: http.StatusForbidden}, - } - - for _, c := range cases { - t.Run(c.name, func(t *testing.T) { - h := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - r = rest.SetUserInfo(r, store.User{ID: "u", Name: "u", SiteID: c.userSite}) - wrapped.ServeHTTP(w, r) - }) - ts := httptest.NewServer(h) - defer ts.Close() - resp, err := http.Get(ts.URL + c.query) - require.NoError(t, err) - require.NoError(t, resp.Body.Close()) - assert.Equal(t, c.want, resp.StatusCode) - }) - } -} diff --git a/backend/app/rest/api/ssl.go b/backend/app/rest/api/ssl.go index f163ec17..4e791c6e 100644 --- a/backend/app/rest/api/ssl.go +++ b/backend/app/rest/api/ssl.go @@ -1,7 +1,6 @@ package api import ( - "context" "crypto/tls" "fmt" "net/http" @@ -65,26 +64,6 @@ func (s *Rest) httpChallengeRouter(m *autocert.Manager) http.Handler { return router } -// timeout returns a middleware matching chi's middleware.Timeout: it sets a -// deadline on the request context and writes 504 Gateway Timeout if the -// deadline is exceeded. The 504 is sent once the downstream handler returns -// after observing the canceled context; a handler that ignores r.Context() -// is not aborted. -func timeout(d time.Duration) func(http.Handler) http.Handler { - return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - ctx, cancel := context.WithTimeout(r.Context(), d) - defer func() { - cancel() - if ctx.Err() == context.DeadlineExceeded { - w.WriteHeader(http.StatusGatewayTimeout) - } - }() - next.ServeHTTP(w, r.WithContext(ctx)) - }) - } -} - func (s *Rest) redirectHandler() http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { newURL, err := s.redirectURL(r) diff --git a/backend/app/rest/api/ssl_test.go b/backend/app/rest/api/ssl_test.go index cc8f8efd..d177c8ba 100644 --- a/backend/app/rest/api/ssl_test.go +++ b/backend/app/rest/api/ssl_test.go @@ -8,37 +8,11 @@ import ( "net/http/httptest" "os" "testing" - "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -func TestTimeout(t *testing.T) { - t.Run("fast handler passes through and gets a deadline", func(t *testing.T) { - var gotDeadline bool - h := timeout(time.Second)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - _, gotDeadline = r.Context().Deadline() - w.WriteHeader(http.StatusCreated) - _, _ = w.Write([]byte("ok")) - })) - rec := httptest.NewRecorder() - h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", http.NoBody)) - assert.True(t, gotDeadline, "request context should carry a deadline") - assert.Equal(t, http.StatusCreated, rec.Code) - assert.Equal(t, "ok", rec.Body.String()) - }) - - t.Run("deadline exceeded writes 504", func(t *testing.T) { - h := timeout(10*time.Millisecond)(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { - <-r.Context().Done() // honor the context: return only once the deadline fires - })) - rec := httptest.NewRecorder() - h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", http.NoBody)) - assert.Equal(t, http.StatusGatewayTimeout, rec.Code) - }) -} - func TestSSL_Redirect(t *testing.T) { rest := Rest{RemarkURL: "https://localhost:443"}