diff --git a/go.mod b/go.mod index 8c31a7c..f3648e0 100644 --- a/go.mod +++ b/go.mod @@ -11,6 +11,7 @@ require ( github.com/getsentry/sentry-go v0.40.0 github.com/go-chi/chi/v5 v5.2.3 github.com/go-chi/cors v1.2.2 + github.com/gorilla/csrf v1.7.3 github.com/gorilla/securecookie v1.1.2 github.com/prometheus/client_golang v1.23.2 github.com/slok/go-http-metrics v0.13.0 @@ -71,7 +72,6 @@ require ( github.com/google/uuid v1.6.0 // indirect github.com/googleapis/enterprise-certificate-proxy v0.3.6 // indirect github.com/googleapis/gax-go/v2 v2.14.2 // indirect - github.com/gorilla/csrf v1.7.3 // indirect github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 // indirect github.com/hashicorp/consul/api v1.32.1 // indirect github.com/hashicorp/errwrap v1.1.0 // indirect diff --git a/internal/handlers/auth.go b/internal/handlers/auth.go index 977756a..5e74e1f 100644 --- a/internal/handlers/auth.go +++ b/internal/handlers/auth.go @@ -2,6 +2,7 @@ package handlers import ( "crypto/subtle" + "html/template" "net/http" "net/url" "strconv" @@ -23,13 +24,13 @@ func (s *Handlers) HandleRoot() http.HandlerFunc { // Check if authenticated if s.sessMgr.IsAuthenticated(r) { - s.renderGenerator(w, nil) + s.renderGenerator(w, r, nil) return } // Show login page - s.renderLogin(w, "") + s.renderLogin(w, r, "") } } @@ -37,7 +38,7 @@ func (s *Handlers) HandleRoot() http.HandlerFunc { func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) { err := r.ParseForm() if err != nil { - s.renderLogin(w, "Invalid form data") + s.renderLogin(w, r, "Invalid form data") return } @@ -47,7 +48,7 @@ func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) { // Constant-time comparison to prevent timing attacks if subtle.ConstantTimeCompare([]byte(submittedKey), []byte(s.config.SigningKey)) != 1 { s.log.Warn("failed login attempt", "remote_addr", r.RemoteAddr) - s.renderLogin(w, "Invalid signing key") + s.renderLogin(w, r, "Invalid signing key") return } @@ -56,7 +57,7 @@ func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) { err = s.sessMgr.CreateSession(w) if err != nil { s.log.Error("failed to create session", "error", err) - s.renderLogin(w, "Failed to create session") + s.renderLogin(w, r, "Failed to create session") return } @@ -87,7 +88,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc { err := r.ParseForm() if err != nil { - s.renderGenerator(w, &generatorData{Error: "Invalid form data"}) + s.renderGenerator(w, r, &generatorData{Error: "Invalid form data"}) return } @@ -97,7 +98,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc { // Validate source URL parsed, err := url.Parse(sourceURL) if err != nil || parsed.Host == "" { - s.renderGeneratorWithForm(w, "Invalid source URL", r.Form) + s.renderGeneratorWithForm(w, r, "Invalid source URL", r.Form) return } @@ -108,7 +109,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc { token, err := s.encGen.Generate(payload) if err != nil { s.log.Error("failed to generate encrypted URL", "error", err) - s.renderGeneratorWithForm(w, "Failed to generate URL", r.Form) + s.renderGeneratorWithForm(w, r, "Failed to generate URL", r.Form) return } @@ -121,7 +122,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc { expiresAtStr = expiresAt.Format(time.RFC3339) } - s.renderGenerator(w, &generatorData{ + s.renderGenerator(w, r, &generatorData{ GeneratedURL: generatedURL, ExpiresAt: expiresAtStr, FormURL: sourceURL, @@ -186,15 +187,20 @@ type generatorData struct { FormQuality string FormFit string FormTTL string + CSRFField template.HTML } -func (s *Handlers) renderLogin(w http.ResponseWriter, errorMsg string) { +func (s *Handlers) renderLogin( + w http.ResponseWriter, r *http.Request, errorMsg string, +) { w.Header().Set("Content-Type", "text/html; charset=utf-8") data := struct { - Error string + Error string + CSRFField template.HTML }{ - Error: errorMsg, + Error: errorMsg, + CSRFField: csrfField(r), } err := templates.Render(w, "login.html", data) @@ -204,13 +210,17 @@ func (s *Handlers) renderLogin(w http.ResponseWriter, errorMsg string) { } } -func (s *Handlers) renderGenerator(w http.ResponseWriter, data *generatorData) { +func (s *Handlers) renderGenerator( + w http.ResponseWriter, r *http.Request, data *generatorData, +) { w.Header().Set("Content-Type", "text/html; charset=utf-8") if data == nil { data = &generatorData{} } + data.CSRFField = csrfField(r) + err := templates.Render(w, "generator.html", data) if err != nil { s.log.Error("failed to render generator template", "error", err) @@ -219,9 +229,9 @@ func (s *Handlers) renderGenerator(w http.ResponseWriter, data *generatorData) { } func (s *Handlers) renderGeneratorWithForm( - w http.ResponseWriter, errorMsg string, form url.Values, + w http.ResponseWriter, r *http.Request, errorMsg string, form url.Values, ) { - s.renderGenerator(w, &generatorData{ + s.renderGenerator(w, r, &generatorData{ Error: errorMsg, FormURL: form.Get("url"), FormWidth: form.Get("width"), diff --git a/internal/handlers/auth_csrf_internal_test.go b/internal/handlers/auth_csrf_internal_test.go index 2db775a..161253b 100644 --- a/internal/handlers/auth_csrf_internal_test.go +++ b/internal/handlers/auth_csrf_internal_test.go @@ -1,7 +1,7 @@ package handlers import ( - "io" + "context" "log/slog" "net/http" "net/http/httptest" @@ -22,6 +22,13 @@ import ( // token key, exactly as the real signing key does in production. const testSigningKey = "test-signing-key-0123456789abcdef" +// Form field names used in the CSRF flow tests. +const ( + loginKeyField = "key" + // gorilla/csrf's default form field name, not a credential. + csrfTokenField = "gorilla.csrf.Token" //nolint:gosec // G101 false positive +) + // csrfFieldPattern extracts the token rendered by csrf.TemplateField into // the form. The field name is gorilla/csrf's default. var csrfFieldPattern = regexp.MustCompile( @@ -53,7 +60,7 @@ func newCSRFTestRouter(t *testing.T) (*Handlers, http.Handler) { } h := &Handlers{ - log: slog.New(slog.NewTextHandler(io.Discard, nil)), + log: slog.New(slog.DiscardHandler), config: cfg, sessMgr: sessMgr, encGen: encGen, @@ -80,7 +87,8 @@ func csrfCredentials( ) ([]*http.Cookie, string) { t.Helper() - req := httptest.NewRequest(http.MethodGet, "/", nil) + req := httptest.NewRequestWithContext( + context.Background(), http.MethodGet, "/", nil) for _, c := range reqCookies { req.AddCookie(c) } @@ -106,8 +114,9 @@ func postForm( srv http.Handler, path string, cookies []*http.Cookie, form url.Values, ) *httptest.ResponseRecorder { - req := httptest.NewRequest( - http.MethodPost, path, strings.NewReader(form.Encode())) + req := httptest.NewRequestWithContext( + context.Background(), http.MethodPost, path, + strings.NewReader(form.Encode())) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range cookies { @@ -128,7 +137,7 @@ func TestLoginPostRejectedWithoutToken(t *testing.T) { _, srv := newCSRFTestRouter(t) - rec := postForm(srv, "/", nil, url.Values{"key": {testSigningKey}}) + rec := postForm(srv, "/", nil, url.Values{loginKeyField: {testSigningKey}}) if rec.Code != http.StatusForbidden { t.Errorf("POST / without token status = %d, want %d", @@ -148,8 +157,8 @@ func TestLoginPostRejectedWithForeignToken(t *testing.T) { _, tokenB := csrfCredentials(t, srv, nil) rec := postForm(srv, "/", cookiesA, url.Values{ - "key": {testSigningKey}, - "gorilla.csrf.Token": {tokenB}, + loginKeyField: {testSigningKey}, + csrfTokenField: {tokenB}, }) if rec.Code != http.StatusForbidden { @@ -169,8 +178,8 @@ func TestLoginPostAcceptedWithValidToken(t *testing.T) { cookies, token := csrfCredentials(t, srv, nil) rec := postForm(srv, "/", cookies, url.Values{ - "key": {testSigningKey}, - "gorilla.csrf.Token": {token}, + loginKeyField: {testSigningKey}, + csrfTokenField: {token}, }) if rec.Code != http.StatusSeeOther { @@ -225,9 +234,9 @@ func TestGeneratePostAcceptedWithValidToken(t *testing.T) { cookies = append(cookies, sessionCookie) rec := postForm(srv, "/generate", cookies, url.Values{ - "url": {"https://example.com/a.jpg"}, - "format": {"jpeg"}, - "gorilla.csrf.Token": {token}, + "url": {"https://example.com/a.jpg"}, + "format": {"jpeg"}, + csrfTokenField: {token}, }) if rec.Code != http.StatusOK { diff --git a/internal/handlers/csrf.go b/internal/handlers/csrf.go new file mode 100644 index 0000000..166abdc --- /dev/null +++ b/internal/handlers/csrf.go @@ -0,0 +1,65 @@ +package handlers + +import ( + "html/template" + "net/http" + + "github.com/gorilla/csrf" + + "sneak.berlin/go/pixa/internal/seal" +) + +// csrfKeySalt provides domain separation for the CSRF authentication key, +// derived from the signing key so tokens survive restarts without extra +// configuration and never reuse the session or encrypted-URL key material. +const csrfKeySalt = "pixa-csrf-v1" + +// newCSRFProtect builds the CSRF-protection middleware for the +// state-mutating HTML form routes. The token lives in its own cookie, +// independent of the session cookie, so it also protects the login POST +// where no session exists yet (login CSRF). +// +// When plaintext is true (local HTTP development), requests are marked +// plaintext so the library neither demands an https Referer nor sets a +// Secure cookie the browser would withhold over http. In production the +// service runs behind a TLS-terminating proxy, so plaintext is false and +// the library enforces its https Referer origin check. +func newCSRFProtect( + signingKey string, plaintext bool, +) (func(http.Handler) http.Handler, error) { + key, err := seal.DeriveKey([]byte(signingKey), csrfKeySalt) + if err != nil { + return nil, err + } + + protect := csrf.Protect( + key[:], + csrf.Path("/"), + csrf.Secure(!plaintext), + csrf.SameSite(csrf.SameSiteStrictMode), + ) + + if !plaintext { + return protect, nil + } + + return func(next http.Handler) http.Handler { + protected := protect(next) + + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + protected.ServeHTTP(w, csrf.PlaintextHTTPRequest(r)) + }) + }, nil +} + +// CSRF returns the CSRF-protection middleware for the login and generator +// form routes. +func (s *Handlers) CSRF() func(http.Handler) http.Handler { + return s.csrfProtect +} + +// csrfField returns the hidden form input carrying the CSRF token for the +// given request, to be embedded verbatim in a rendered form. +func csrfField(r *http.Request) template.HTML { + return csrf.TemplateField(r) +} diff --git a/internal/handlers/handlers.go b/internal/handlers/handlers.go index 0398357..5323167 100644 --- a/internal/handlers/handlers.go +++ b/internal/handlers/handlers.go @@ -31,23 +31,30 @@ type Params struct { // Handlers provides HTTP request handlers. type Handlers struct { - log *slog.Logger - hc *healthcheck.Healthcheck - db *database.Database - config *config.Config - imgSvc *imgcache.Service - imgCache *imgcache.Cache - sessMgr *session.Manager - encGen *encurl.Generator + log *slog.Logger + hc *healthcheck.Healthcheck + db *database.Database + config *config.Config + imgSvc *imgcache.Service + imgCache *imgcache.Cache + sessMgr *session.Manager + encGen *encurl.Generator + csrfProtect func(http.Handler) http.Handler } // New creates a new Handlers instance. func New(lc fx.Lifecycle, params Params) (*Handlers, error) { + csrfProtect, err := newCSRFProtect(params.Config.SigningKey, params.Config.Debug) + if err != nil { + return nil, err + } + s := &Handlers{ - log: params.Logger.Get(), - hc: params.Healthcheck, - db: params.Database, - config: params.Config, + log: params.Logger.Get(), + hc: params.Healthcheck, + db: params.Database, + config: params.Config, + csrfProtect: csrfProtect, } lc.Append(fx.Hook{ diff --git a/internal/server/routes.go b/internal/server/routes.go index 4736dbc..1fe054e 100644 --- a/internal/server/routes.go +++ b/internal/server/routes.go @@ -44,11 +44,17 @@ func (s *Server) SetupRoutes() { // Static files (Tailwind CSS, etc.) s.router.Handle("/static/*", http.StripPrefix("/static/", static.Handler())) - // Login/generator UI - s.router.Get("/", s.h.HandleRoot()) - s.router.Post("/", s.h.HandleRoot()) + // Login/generator UI. The form routes carry CSRF protection; the + // token cookie is independent of the session cookie, so it also + // covers the login POST, where no session exists yet. + s.router.Group(func(r chi.Router) { + r.Use(s.h.CSRF()) + r.Get("/", s.h.HandleRoot()) + r.Post("/", s.h.HandleRoot()) + r.Post("/generate", s.h.HandleGenerateURL()) + }) + s.router.Get("/logout", s.h.HandleLogout()) - s.router.Post("/generate", s.h.HandleGenerateURL()) // Main image proxy route // /v1/image///x. diff --git a/internal/templates/generator.html b/internal/templates/generator.html index ede9951..d82424d 100644 --- a/internal/templates/generator.html +++ b/internal/templates/generator.html @@ -47,6 +47,7 @@ {{end}}
+ {{ .CSRFField }}