From 9e745f7b8437fb0f8072e762731ccd6095f53a5a Mon Sep 17 00:00:00 2001 From: sneak Date: Mon, 21 Sep 2026 07:40:07 +0000 Subject: [PATCH] feat: CSRF protection on the login and URL-generator forms (closes #93) Both cookie-authenticated HTML form posts (POST / and POST /generate) now require a CSRF token via gorilla/csrf, the recorded default in GO_PACKAGE_DEFAULTS.md. The token cookie is independent of the session cookie, so it also covers the login POST, where no session exists yet (login CSRF). The token key is derived from the signing key with its own HKDF salt, so tokens survive restarts and reuse no other key material; gorilla/csrf supplies crypto/rand generation and constant-time compare. The form routes sit in a chi group behind the middleware; the hidden token field is rendered into login.html and generator.html. In local plaintext HTTP mode (debug) requests are marked plaintext so the library does not demand an https Referer or set a Secure cookie the browser would withhold; in production, behind the TLS-terminating proxy, it enforces its https Referer origin check. model: claude-opus-4-8 --- go.mod | 2 +- internal/handlers/auth.go | 40 +++++++----- internal/handlers/auth_csrf_internal_test.go | 35 +++++++---- internal/handlers/csrf.go | 65 ++++++++++++++++++++ internal/handlers/handlers.go | 31 ++++++---- internal/server/routes.go | 14 +++-- internal/templates/generator.html | 1 + internal/templates/login.html | 1 + 8 files changed, 144 insertions(+), 45 deletions(-) create mode 100644 internal/handlers/csrf.go 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 }}