diff --git a/Dockerfile b/Dockerfile index 3eb5782..27ab84b 100644 --- a/Dockerfile +++ b/Dockerfile @@ -67,8 +67,9 @@ RUN adduser -D -H -s /sbin/nologin pixad && \ mkdir -p /var/lib/pixa /etc/pixa && \ chown pixad:pixad /var/lib/pixa -# Copy default config (edit signing_key before use) -COPY config.example.yml /etc/pixa/config.yml +# Copy the image config; signing_key comes from PIXA_SIGNING_KEY. +# Mount a file over /etc/pixa/config.yml to override anything else. +COPY config.docker.yml /etc/pixa/config.yml USER pixad WORKDIR /var/lib/pixa diff --git a/README.md b/README.md index 8e36dc0..32fed7f 100644 --- a/README.md +++ b/README.md @@ -15,14 +15,25 @@ git clone https://git.eeqj.de/sneak/pixa.git cd pixa make build -# run with a config file -./bin/pixad --config config.example.yml +# run with a config file: copy the example and set a real signing key +# (the example placeholder is refused at startup), e.g. with +# openssl rand -base64 32 +cp config.example.yml config.yml +$EDITOR config.yml # replace the signing_key placeholder +./bin/pixad --config config.yml # or build and run via Docker make docker -docker run -p 8080:8080 pixad:latest +docker run -p 8080:8080 -e PIXA_SIGNING_KEY="$(openssl rand -base64 32)" pixa:latest ``` +A container is configured two ways. The signing key comes from the +`PIXA_SIGNING_KEY` environment variable, which the baked-in config +reads; if it is unset the container exits at startup naming the +variable. Everything else uses built-in defaults, so to change any +other setting mount your own file over `/etc/pixa/config.yml` (see +`config.example.yml` for the full set of keys). + ## Rationale Image-heavy web applications need a fast, caching reverse proxy that diff --git a/TODO.md b/TODO.md index 9421a23..dab3310 100644 --- a/TODO.md +++ b/TODO.md @@ -1,21 +1,27 @@ # Workflow -* branch (from `main`) +* branch per issue from `next` * do the work in Next Step * move Next Step to the top of Completed Steps * move the top item of Future Steps into Next Step * commit (`TODO.md` changes in the same commit as the work) -* merge to `main` if the branch is not protected, otherwise open a PR +* open a PR based on `next` +* an independent reviewer who did not write the change gates it +* the manager squash-merges the PR into `next` once review passes +* `next` stays green and mergeable to `main` at any time; only the owner + merges `next` into `main`, via the single milestone PR * push # Status -pre-1.0. No git tags exist. Recent work extracted the internal/magic, +pre-1.0. No git tags exist. The `1.0.0` milestone is in progress; work +lands on `next`, and `main` receives only the milestone PR that the +owner merges. `next` is at the canonical `golangci-lint` v2.12.2 config +and is green. Recent work extracted the internal/magic, internal/allowlist, internal/httpfetcher, and internal/signature -packages. The gosec findings from the 2026-07-06 survey are resolved -and `make check` is green on main. The disk cache is now size-bounded -with LRU eviction (`cache_max_bytes`), closing the unbounded disk -growth DoS vector. +packages. The gosec findings from the 2026-07-06 survey are resolved. +The disk cache is now size-bounded with LRU eviction +(`cache_max_bytes`), closing the unbounded disk growth DoS vector. # Next Step @@ -23,6 +29,14 @@ P1: implement blocked networks configuration to extend SSRF protection # Completed Steps +- 2026-09-21 http.Server hardening (closes #92): added + `HTTPReadHeaderTimeout` (10s, bounds the slowloris header dribble) and + `HTTPIdleTimeout` (120s, bounds keep-alive reuse) alongside the + existing timeouts and wired them onto the server; added a `LimitBody` + middleware capping the two form POST bodies (`POST /`, `POST /generate`) + at `MaxFormBytes` (1 MiB) and returning 413, applied ahead of the CSRF + middleware so an oversized body is refused as 413 rather than being read + as a missing CSRF token (403); left `WriteTimeout` at 60s unchanged - 2026-08-07 update golangci-lint to v2.12.2 with the canonical `.golangci.yml` (v2 schema, `default: all` minus six disabled linters, `lll` 88, tests included): bumped the pinned diff --git a/config.docker.yml b/config.docker.yml new file mode 100644 index 0000000..548b45a --- /dev/null +++ b/config.docker.yml @@ -0,0 +1,11 @@ +# Pixa configuration baked into the Docker image. +# +# The signing key is read from the PIXA_SIGNING_KEY environment +# variable; startup aborts naming it when it is unset. Every other key +# is omitted so its default applies. Operators who need more (an +# allowlist, metrics, and so on) mount their own file over +# /etc/pixa/config.yml. + +signing_key: "${ENV:PIXA_SIGNING_KEY}" +state_dir: /var/lib/pixa +port: 8080 diff --git a/go.mod b/go.mod index ce9b477..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 diff --git a/go.sum b/go.sum index 002a2b4..af3e229 100644 --- a/go.sum +++ b/go.sum @@ -175,6 +175,8 @@ github.com/googleapis/enterprise-certificate-proxy v0.3.6 h1:GW/XbdyBFQ8Qe+YAmFU github.com/googleapis/enterprise-certificate-proxy v0.3.6/go.mod h1:MkHOF77EYAE7qfSuSS9PU6g4Nt4e11cnsDUowfwewLA= github.com/googleapis/gax-go/v2 v2.14.2 h1:eBLnkZ9635krYIPD+ag1USrOAI0Nr0QYF3+/3GqO0k0= github.com/googleapis/gax-go/v2 v2.14.2/go.mod h1:ON64QhlJkhVtSqp4v1uaK92VyZ2gmvDQsweuyLV+8+w= +github.com/gorilla/csrf v1.7.3 h1:BHWt6FTLZAb2HtWT5KDBf6qgpZzvtbp9QWDRKZMXJC0= +github.com/gorilla/csrf v1.7.3/go.mod h1:F1Fj3KG23WYHE6gozCmBAezKookxbIvUJT+121wTuLk= github.com/gorilla/securecookie v1.1.2 h1:YCIWL56dvtr73r6715mJs5ZvhtnY73hBvEF8kXD8ePA= github.com/gorilla/securecookie v1.1.2/go.mod h1:NfCASbcHqRSY+3a8tlWJwsQap2VX5pwzwo4h3eOamfo= github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 h1:5ZPtiqj0JL5oKWmcsq4VMaAW5ukBEgSGXEN89zeH1Jo= diff --git a/internal/config/config.go b/internal/config/config.go index fb6e02d..e3ab570 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -44,6 +44,12 @@ const ( keyCacheMaxBytes = "cache_max_bytes" ) +// placeholderSigningKey is the dummy signing_key shipped in +// config.example.yml. It is 45 characters, so it passes the length +// check, but it is public in this repository and must be rejected at +// startup so no deployment ever signs URLs with it. +const placeholderSigningKey = "CHANGE_ME_generate_with_openssl_rand_base64_32" + // Static validation errors. Each use site attaches the offending key // and value by wrapping these with fmt.Errorf and %w. var ( @@ -61,6 +67,9 @@ var ( errPortOutOfRange = errors.New("outside the valid port range") errTooFewConnections = errors.New("must be at least 1") errValueTooShort = errors.New("value too short") + errPlaceholderKey = errors.New( + "is the placeholder from config.example.yml; " + + "generate a real key with: openssl rand -base64 32") errMustBeSetTogether = errors.New("must be set together") errMustNotBeNegative = errors.New("must not be negative") errOverflowsInt64 = errors.New("overflows a 64-bit integer") @@ -341,10 +350,10 @@ func (c *Config) ensureStateDirWritable() error { return nil } -// validate checks that all required configuration values are set and -// that every value is within its valid range. -func (c *Config) validate() error { - // The signing key value is never echoed in error messages. +// validateSigningKey checks that the signing key is present, long +// enough, and not the public placeholder from config.example.yml. The +// key value itself is never echoed in error messages. +func (c *Config) validateSigningKey() error { if c.SigningKey == "" { return fmt.Errorf("config key %q: %w", keySigningKey, errValueRequired) } @@ -356,6 +365,21 @@ func (c *Config) validate() error { keySigningKey, errValueTooShort, minKeyLength, len(c.SigningKey)) } + if c.SigningKey == placeholderSigningKey { + return fmt.Errorf("config key %q: %w", keySigningKey, errPlaceholderKey) + } + + return nil +} + +// validate checks that all required configuration values are set and +// that every value is within its valid range. +func (c *Config) validate() error { + err := c.validateSigningKey() + if err != nil { + return err + } + const maxPort = 65535 if c.Port < 1 || c.Port > maxPort { return fmt.Errorf("config key %q: value %d is %w 1-%d", diff --git a/internal/config/config_validation_internal_test.go b/internal/config/config_validation_internal_test.go index ca1573e..c12d27d 100644 --- a/internal/config/config_validation_internal_test.go +++ b/internal/config/config_validation_internal_test.go @@ -303,6 +303,11 @@ func invalidHostAndCredentialCases() []abortCase { yaml: "signing_key: short\n", wantErrSubstrings: []string{keySigningKey}, }, + { + name: "signing_key is the documented placeholder", + yaml: "signing_key: " + placeholderSigningKey + "\n", + wantErrSubstrings: []string{keySigningKey}, + }, { name: "signing_key missing", yaml: "port: 8080\n", 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 new file mode 100644 index 0000000..161253b --- /dev/null +++ b/internal/handlers/auth_csrf_internal_test.go @@ -0,0 +1,273 @@ +package handlers + +import ( + "context" + "log/slog" + "net/http" + "net/http/httptest" + "net/url" + "regexp" + "strings" + "testing" + + "github.com/go-chi/chi/v5" + + "sneak.berlin/go/pixa/internal/config" + "sneak.berlin/go/pixa/internal/encurl" + "sneak.berlin/go/pixa/internal/session" +) + +// testSigningKey is a throwaway signing key for the CSRF flow tests. It +// seeds the session manager, the encrypted-URL generator, and the CSRF +// 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( + `name="gorilla\.csrf\.Token" value="([^"]+)"`) + +// newCSRFTestRouter builds a router that mirrors the production wiring for +// the CSRF-protected UI routes (see server.SetupRoutes): the login and +// generator forms and their POST targets sit behind the real CSRF +// middleware. Requests are marked plaintext (Debug: true) so the flow runs +// over httptest's http transport without an https Referer. +func newCSRFTestRouter(t *testing.T) (*Handlers, http.Handler) { + t.Helper() + + cfg := &config.Config{SigningKey: testSigningKey, Debug: true} + + sessMgr, err := session.NewManager(testSigningKey) + if err != nil { + t.Fatalf("session.NewManager() error = %v", err) + } + + encGen, err := encurl.NewGenerator(testSigningKey) + if err != nil { + t.Fatalf("encurl.NewGenerator() error = %v", err) + } + + protect, err := newCSRFProtect(testSigningKey, cfg.Debug) + if err != nil { + t.Fatalf("newCSRFProtect() error = %v", err) + } + + h := &Handlers{ + log: slog.New(slog.DiscardHandler), + config: cfg, + sessMgr: sessMgr, + encGen: encGen, + csrfProtect: protect, + } + + r := chi.NewRouter() + r.Group(func(r chi.Router) { + r.Use(h.CSRF()) + r.Get("/", h.HandleRoot()) + r.Post("/", h.HandleRoot()) + r.Post("/generate", h.HandleGenerateURL()) + }) + + return h, r +} + +// csrfCredentials performs a GET that renders a form and returns the CSRF +// cookies the middleware set and the token embedded in the form. Passing +// the authenticated session cookie renders the generator form instead of +// the login form. +func csrfCredentials( + t *testing.T, srv http.Handler, reqCookies []*http.Cookie, +) ([]*http.Cookie, string) { + t.Helper() + + req := httptest.NewRequestWithContext( + context.Background(), http.MethodGet, "/", nil) + for _, c := range reqCookies { + req.AddCookie(c) + } + + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("GET / status = %d, want %d", rec.Code, http.StatusOK) + } + + match := csrfFieldPattern.FindStringSubmatch(rec.Body.String()) + if match == nil { + t.Fatalf("no CSRF token field found in rendered form") + } + + return rec.Result().Cookies(), match[1] +} + +// postForm submits form values with the given cookies and returns the +// recorder. +func postForm( + srv http.Handler, path string, + cookies []*http.Cookie, form url.Values, +) *httptest.ResponseRecorder { + 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 { + req.AddCookie(c) + } + + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, req) + + return rec +} + +// TestLoginPostRejectedWithoutToken verifies that POST / with no CSRF token +// is rejected. This is login CSRF: no session cookie exists yet, so the +// protection must rest on a token bound to a pre-session cookie. +func TestLoginPostRejectedWithoutToken(t *testing.T) { + t.Parallel() + + _, srv := newCSRFTestRouter(t) + + rec := postForm(srv, "/", nil, url.Values{loginKeyField: {testSigningKey}}) + + if rec.Code != http.StatusForbidden { + t.Errorf("POST / without token status = %d, want %d", + rec.Code, http.StatusForbidden) + } +} + +// TestLoginPostRejectedWithForeignToken verifies that a token that does not +// match the request's CSRF cookie is rejected: a token minted for one +// cookie cannot authorize a request carrying a different cookie. +func TestLoginPostRejectedWithForeignToken(t *testing.T) { + t.Parallel() + + _, srv := newCSRFTestRouter(t) + + cookiesA, _ := csrfCredentials(t, srv, nil) + _, tokenB := csrfCredentials(t, srv, nil) + + rec := postForm(srv, "/", cookiesA, url.Values{ + loginKeyField: {testSigningKey}, + csrfTokenField: {tokenB}, + }) + + if rec.Code != http.StatusForbidden { + t.Errorf("POST / with foreign token status = %d, want %d", + rec.Code, http.StatusForbidden) + } +} + +// TestLoginPostAcceptedWithValidToken verifies that POST / with a matching +// cookie and token succeeds: the login is processed and a session is +// established (303 redirect). +func TestLoginPostAcceptedWithValidToken(t *testing.T) { + t.Parallel() + + _, srv := newCSRFTestRouter(t) + + cookies, token := csrfCredentials(t, srv, nil) + + rec := postForm(srv, "/", cookies, url.Values{ + loginKeyField: {testSigningKey}, + csrfTokenField: {token}, + }) + + if rec.Code != http.StatusSeeOther { + t.Fatalf("POST / with valid token status = %d, want %d", + rec.Code, http.StatusSeeOther) + } + + var authed bool + + for _, c := range rec.Result().Cookies() { + if c.Name == session.CookieName && c.Value != "" { + authed = true + } + } + + if !authed { + t.Error("valid login did not set a session cookie") + } +} + +// TestGeneratePostRejectedWithoutToken verifies that POST /generate is +// rejected without a CSRF token even when the request carries a valid +// authenticated session. The session cookie is not sufficient; the policy +// requires a CSRF token on this cookie-authenticated form. +func TestGeneratePostRejectedWithoutToken(t *testing.T) { + t.Parallel() + + h, srv := newCSRFTestRouter(t) + + sessionCookie := newSessionCookie(t, h) + + rec := postForm(srv, "/generate", + []*http.Cookie{sessionCookie}, + url.Values{"url": {"https://example.com/a.jpg"}}) + + if rec.Code != http.StatusForbidden { + t.Errorf("POST /generate without token status = %d, want %d", + rec.Code, http.StatusForbidden) + } +} + +// TestGeneratePostAcceptedWithValidToken verifies that POST /generate +// succeeds with a valid session and a matching CSRF cookie and token. +func TestGeneratePostAcceptedWithValidToken(t *testing.T) { + t.Parallel() + + h, srv := newCSRFTestRouter(t) + + sessionCookie := newSessionCookie(t, h) + + cookies, token := csrfCredentials(t, srv, []*http.Cookie{sessionCookie}) + cookies = append(cookies, sessionCookie) + + rec := postForm(srv, "/generate", cookies, url.Values{ + "url": {"https://example.com/a.jpg"}, + "format": {"jpeg"}, + csrfTokenField: {token}, + }) + + if rec.Code != http.StatusOK { + t.Fatalf("POST /generate with valid token status = %d, want %d", + rec.Code, http.StatusOK) + } + + if !strings.Contains(rec.Body.String(), "/v1/e/") { + t.Error("generator response did not contain a generated URL") + } +} + +// newSessionCookie creates an authenticated session cookie via the +// handler's session manager. +func newSessionCookie(t *testing.T, h *Handlers) *http.Cookie { + t.Helper() + + rec := httptest.NewRecorder() + + err := h.sessMgr.CreateSession(rec) + if err != nil { + t.Fatalf("CreateSession() error = %v", err) + } + + for _, c := range rec.Result().Cookies() { + if c.Name == session.CookieName { + return c + } + } + + t.Fatalf("session manager did not set a %q cookie", session.CookieName) + + return nil +} diff --git a/internal/handlers/bodylimit.go b/internal/handlers/bodylimit.go new file mode 100644 index 0000000..a2c9adc --- /dev/null +++ b/internal/handlers/bodylimit.go @@ -0,0 +1,45 @@ +package handlers + +import ( + "errors" + "net/http" +) + +// MaxFormBytes bounds the request body accepted on the HTML form POST +// routes (POST / and POST /generate). The forms carry a handful of short +// fields, so 1 MiB is generous while making the bound explicit rather than +// resting on ParseForm's incidental 10 MB cap. +const MaxFormBytes = 1 << 20 // 1 MiB + +// LimitBody returns middleware that caps the request body on POST requests +// at maxBytes and rejects an oversized body with 413 Request Entity Too +// Large. +// +// It parses the form here, before the CSRF middleware reads the token from +// it. The CSRF middleware reads the token with PostFormValue, which +// swallows a parse error, so if the body were only capped there an +// oversized body would read as a missing token and be refused as 403. By +// parsing under the cap first, an oversized body is refused as 413. A +// successful parse is cached on the request, so the CSRF check and the +// handler reuse it rather than reading the body again. +func (s *Handlers) LimitBody(maxBytes int64) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPost { + r.Body = http.MaxBytesReader(w, r.Body, maxBytes) + + err := r.ParseForm() + + var tooLarge *http.MaxBytesError + if errors.As(err, &tooLarge) { + http.Error(w, "Request body too large", + http.StatusRequestEntityTooLarge) + + return + } + } + + next.ServeHTTP(w, r) + }) + } +} diff --git a/internal/handlers/bodylimit_internal_test.go b/internal/handlers/bodylimit_internal_test.go new file mode 100644 index 0000000..59e5fd7 --- /dev/null +++ b/internal/handlers/bodylimit_internal_test.go @@ -0,0 +1,177 @@ +package handlers + +import ( + "log/slog" + "net/http" + "net/url" + "strings" + "testing" + + "github.com/go-chi/chi/v5" + + "sneak.berlin/go/pixa/internal/config" + "sneak.berlin/go/pixa/internal/encurl" + "sneak.berlin/go/pixa/internal/session" +) + +// Form field names and a throwaway source image URL for the body-limit +// tests. +const ( + sourceURLField = "url" + testSourceURL = "https://example.com/a.jpg" +) + +// newBodyLimitTestRouter mirrors the production wiring for the form POST +// routes (see server.SetupRoutes): LimitBody sits in front of the CSRF +// middleware, which sits in front of the handlers. maxBytes is the body +// cap under test, so a test can trip the limit with a small body. +func newBodyLimitTestRouter( + t *testing.T, maxBytes int64, +) (*Handlers, http.Handler) { + t.Helper() + + cfg := &config.Config{SigningKey: testSigningKey, Debug: true} + + sessMgr, err := session.NewManager(testSigningKey) + if err != nil { + t.Fatalf("session.NewManager() error = %v", err) + } + + encGen, err := encurl.NewGenerator(testSigningKey) + if err != nil { + t.Fatalf("encurl.NewGenerator() error = %v", err) + } + + protect, err := newCSRFProtect(testSigningKey, cfg.Debug) + if err != nil { + t.Fatalf("newCSRFProtect() error = %v", err) + } + + h := &Handlers{ + log: slog.New(slog.DiscardHandler), + config: cfg, + sessMgr: sessMgr, + encGen: encGen, + csrfProtect: protect, + } + + r := chi.NewRouter() + r.Group(func(r chi.Router) { + r.Use(h.LimitBody(maxBytes)) + r.Use(h.CSRF()) + r.Get("/", h.HandleRoot()) + r.Post("/", h.HandleRoot()) + r.Post("/generate", h.HandleGenerateURL()) + }) + + return h, r +} + +// TestOversizedLoginPostRejectedBeforeCSRF is the core regression: an +// oversized POST / carrying an otherwise valid CSRF cookie and token must +// be rejected with 413. If the body limit ran after CSRF, the truncated +// body would read as a missing token and return 403; if it ran after the +// handler, a valid token would return 303. Getting 413 proves the limit +// fires before CSRF parses the form. +func TestOversizedLoginPostRejectedBeforeCSRF(t *testing.T) { + t.Parallel() + + _, srv := newBodyLimitTestRouter(t, 16) + + cookies, token := csrfCredentials(t, srv, nil) + + rec := postForm(srv, "/", cookies, url.Values{ + loginKeyField: {testSigningKey}, + csrfTokenField: {token}, + }) + + if rec.Code != http.StatusRequestEntityTooLarge { + t.Errorf("oversized POST / status = %d, want %d", + rec.Code, http.StatusRequestEntityTooLarge) + } +} + +// TestOversizedGeneratePostRejectedBeforeCSRF is the same regression for +// POST /generate, which also parses a form behind CSRF. +func TestOversizedGeneratePostRejectedBeforeCSRF(t *testing.T) { + t.Parallel() + + h, srv := newBodyLimitTestRouter(t, 16) + + sessionCookie := newSessionCookie(t, h) + + cookies, token := csrfCredentials(t, srv, []*http.Cookie{sessionCookie}) + cookies = append(cookies, sessionCookie) + + rec := postForm(srv, "/generate", cookies, url.Values{ + sourceURLField: {testSourceURL}, + csrfTokenField: {token}, + }) + + if rec.Code != http.StatusRequestEntityTooLarge { + t.Errorf("oversized POST /generate status = %d, want %d", + rec.Code, http.StatusRequestEntityTooLarge) + } +} + +// TestWithinLimitLoginPostSucceeds verifies the limit does not disturb a +// normal request: under the production cap, a valid login still parses and +// establishes a session (303). This guards against the body limit +// consuming or corrupting the form the CSRF check and handler depend on. +func TestWithinLimitLoginPostSucceeds(t *testing.T) { + t.Parallel() + + _, srv := newBodyLimitTestRouter(t, MaxFormBytes) + + cookies, token := csrfCredentials(t, srv, nil) + + rec := postForm(srv, "/", cookies, url.Values{ + loginKeyField: {testSigningKey}, + csrfTokenField: {token}, + }) + + if rec.Code != http.StatusSeeOther { + t.Fatalf("within-limit POST / status = %d, want %d", + rec.Code, http.StatusSeeOther) + } + + var authed bool + + for _, c := range rec.Result().Cookies() { + if c.Name == session.CookieName && c.Value != "" { + authed = true + } + } + + if !authed { + t.Error("within-limit valid login did not set a session cookie") + } +} + +// TestWithinLimitGeneratePostSucceeds is the same non-regression check for +// POST /generate. +func TestWithinLimitGeneratePostSucceeds(t *testing.T) { + t.Parallel() + + h, srv := newBodyLimitTestRouter(t, MaxFormBytes) + + sessionCookie := newSessionCookie(t, h) + + cookies, token := csrfCredentials(t, srv, []*http.Cookie{sessionCookie}) + cookies = append(cookies, sessionCookie) + + rec := postForm(srv, "/generate", cookies, url.Values{ + sourceURLField: {testSourceURL}, + "format": {"jpeg"}, + csrfTokenField: {token}, + }) + + if rec.Code != http.StatusOK { + t.Fatalf("within-limit POST /generate status = %d, want %d", + rec.Code, http.StatusOK) + } + + if !strings.Contains(rec.Body.String(), "/v1/e/") { + t.Error("within-limit generate response did not contain a generated URL") + } +} 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/httpfetcher/fetch_internal_test.go b/internal/httpfetcher/fetch_internal_test.go new file mode 100644 index 0000000..c0d6e9d --- /dev/null +++ b/internal/httpfetcher/fetch_internal_test.go @@ -0,0 +1,421 @@ +package httpfetcher + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "slices" + "strings" + "sync" + "testing" + "time" +) + +// testPublicHost is a TEST-NET-1 (RFC 5737) literal. isPrivateIP treats it as +// public, so validateURL and the redirect check accept it with no DNS lookup, +// while the recording dialer routes it to the local httptest server. The +// address is reserved for documentation and is never routed on the network. +const testPublicHost = "192.0.2.10" + +// imagePayload is the body served by the fake upstream's image route. +const imagePayload = "fake-jpeg-bytes" + +// errUnexpectedDial reports a dial to any host other than testPublicHost, which +// would mean SSRF protection let a forbidden target reach the transport. +var errUnexpectedDial = errors.New("unexpected dial target") + +// upstreamURL builds a fetch URL on the fake public host for the given path. +func upstreamURL(path string) string { + return "http://" + testPublicHost + path +} + +// recordingDialer records every address the transport asks it to dial and +// routes connections for testPublicHost to a real local server, so the SSRF +// checks run against a public-looking host while bytes go to httptest. +type recordingDialer struct { + target string + + mu sync.Mutex + dialed []string +} + +func (d *recordingDialer) dialContext( + ctx context.Context, + network, addr string, +) (net.Conn, error) { + d.mu.Lock() + d.dialed = append(d.dialed, addr) + d.mu.Unlock() + + host, _, err := net.SplitHostPort(addr) + if err != nil { + return nil, err + } + + if host != testPublicHost { + return nil, fmt.Errorf("%w: %s", errUnexpectedDial, addr) + } + + var dialer net.Dialer + + return dialer.DialContext(ctx, network, d.target) +} + +// dialedAddrs returns a copy of the addresses the dialer was asked to reach. +func (d *recordingDialer) dialedAddrs() []string { + d.mu.Lock() + defer d.mu.Unlock() + + return slices.Clone(d.dialed) +} + +// startUpstream launches a fake upstream with the routes the fetch tests +// exercise and stops it when the test finishes. +func startUpstream(t *testing.T) *httptest.Server { + t.Helper() + + mux := http.NewServeMux() + mux.HandleFunc("/image", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", contentTypeJPEG) + _, _ = io.WriteString(w, imagePayload) + }) + mux.HandleFunc("/status/500", func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + }) + mux.HandleFunc("/html", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + _, _ = io.WriteString(w, "") + }) + mux.HandleFunc("/redirect/private", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "http://169.254.169.254/latest/meta-data/", http.StatusFound) + }) + mux.HandleFunc("/redirect/public", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/image", http.StatusFound) + }) + mux.HandleFunc("/redirect/chain", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/redirect/hop", http.StatusFound) + }) + mux.HandleFunc("/redirect/hop", func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, "/image", http.StatusFound) + }) + + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + + return srv +} + +// newServerFetcher builds a fetcher whose transport routes testPublicHost to +// srv, leaving the real SSRF validation and redirect checks in place. +func newServerFetcher( + t *testing.T, + srv *httptest.Server, + cfg *Config, +) (*HTTPFetcher, *recordingDialer) { + t.Helper() + + if cfg == nil { + cfg = DefaultConfig() + } + + cfg.AllowHTTP = true + + f := New(cfg) + + transport, ok := f.client.Transport.(*http.Transport) + if !ok { + t.Fatalf("transport is %T, want *http.Transport", f.client.Transport) + } + + dialer := &recordingDialer{target: srv.Listener.Addr().String()} + transport.DialContext = dialer.dialContext + + return f, dialer +} + +// testContext returns a context cancelled when the test ends, bounding any +// fetch that would otherwise block on a leaked semaphore slot. +func testContext(t *testing.T) context.Context { + t.Helper() + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + t.Cleanup(cancel) + + return ctx +} + +// fetchImage fetches path from the fake upstream and fails on error. +func fetchImage(t *testing.T, f *HTTPFetcher, path string) *FetchResult { + t.Helper() + + res, err := f.Fetch(testContext(t), upstreamURL(path)) + if err != nil { + t.Fatalf("Fetch(%s) error = %v", path, err) + } + + return res +} + +// fetchExpectError fetches path and fails unless Fetch returns an error. +func fetchExpectError(t *testing.T, f *HTTPFetcher, path string) error { + t.Helper() + + res, err := f.Fetch(testContext(t), upstreamURL(path)) + if err == nil { + _ = res.Content.Close() + + t.Fatalf("Fetch(%s) = nil error, want an error", path) + } + + return err +} + +// fetchBody fetches path and returns the fully read, closed response body. +func fetchBody(t *testing.T, f *HTTPFetcher, path string) string { + t.Helper() + + res := fetchImage(t, f, path) + defer func() { _ = res.Content.Close() }() + + data, err := io.ReadAll(res.Content) + if err != nil { + t.Fatalf("read body: %v", err) + } + + return string(data) +} + +// semLen reports how many per-host semaphore slots are currently held. +func semLen(f *HTTPFetcher, host string) int { + return len(f.getHostSemaphore(host)) +} + +func TestFetchRedirectToPrivateIPBlocked(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + f, dialer := newServerFetcher(t, srv, nil) + + _, err := f.Fetch(testContext(t), upstreamURL("/redirect/private")) + if !errors.Is(err, ErrSSRFBlocked) { + t.Fatalf("Fetch() error = %v, want ErrSSRFBlocked", err) + } + + for _, addr := range dialer.dialedAddrs() { + if strings.Contains(addr, "169.254.169.254") { + t.Errorf("dialer connected to the private redirect target: %s", addr) + } + } +} + +func TestFetchRedirectToPublicSucceeds(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + f, _ := newServerFetcher(t, srv, nil) + + if body := fetchBody(t, f, "/redirect/public"); body != imagePayload { + t.Errorf("body = %q, want %q", body, imagePayload) + } +} + +func TestFetchRedirectChainSucceeds(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + f, _ := newServerFetcher(t, srv, nil) + + if body := fetchBody(t, f, "/redirect/chain"); body != imagePayload { + t.Errorf("body = %q, want %q", body, imagePayload) + } +} + +func TestFetchRejectsNon2xx(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + f, _ := newServerFetcher(t, srv, nil) + + err := fetchExpectError(t, f, "/status/500") + if !errors.Is(err, ErrUpstreamError) { + t.Fatalf("Fetch() error = %v, want ErrUpstreamError", err) + } +} + +func TestFetchRejectsDisallowedContentType(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + f, _ := newServerFetcher(t, srv, nil) + + err := fetchExpectError(t, f, "/html") + if !errors.Is(err, ErrInvalidContentType) { + t.Fatalf("Fetch() error = %v, want ErrInvalidContentType", err) + } +} + +func TestFetchMaxResponseSizeEnforced(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + + cfg := DefaultConfig() + cfg.MaxResponseSize = 8 + + f, _ := newServerFetcher(t, srv, cfg) + + res := fetchImage(t, f, "/image") + defer func() { _ = res.Content.Close() }() + + data, err := io.ReadAll(res.Content) + if !errors.Is(err, ErrResponseTooLarge) { + t.Fatalf("read error = %v, want ErrResponseTooLarge", err) + } + + if int64(len(data)) > cfg.MaxResponseSize { + t.Errorf("read %d bytes, exceeds limit %d", len(data), cfg.MaxResponseSize) + } +} + +func TestFetchSemaphoreReleasedOnError(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + + cfg := DefaultConfig() + cfg.MaxConnectionsPerHost = 1 + + f, _ := newServerFetcher(t, srv, cfg) + + err := fetchExpectError(t, f, "/status/500") + if !errors.Is(err, ErrUpstreamError) { + t.Fatalf("Fetch() error = %v, want ErrUpstreamError", err) + } + + if held := semLen(f, testPublicHost); held != 0 { + t.Fatalf("semaphore slot leaked after error: %d held", held) + } + + // One slot per host: this fetch proceeds only if the slot was released. + res := fetchImage(t, f, "/image") + _ = res.Content.Close() +} + +// assertSlotReleasedByClose fetches an image over a one-slot host, hands the +// open result to consume, and asserts the slot is held before and freed after, +// then that a follow-up fetch can still acquire it. +func assertSlotReleasedByClose( + t *testing.T, + consume func(*testing.T, *FetchResult), +) { + t.Helper() + + srv := startUpstream(t) + + cfg := DefaultConfig() + cfg.MaxConnectionsPerHost = 1 + + f, _ := newServerFetcher(t, srv, cfg) + + res := fetchImage(t, f, "/image") + if held := semLen(f, testPublicHost); held != 1 { + t.Fatalf("slot not held while body is open: %d held", held) + } + + consume(t, res) + + if held := semLen(f, testPublicHost); held != 0 { + t.Fatalf("slot not released after close: %d held", held) + } + + next := fetchImage(t, f, "/image") + _ = next.Content.Close() +} + +func TestFetchSemaphoreReleasedOnBodyClose(t *testing.T) { + t.Parallel() + + assertSlotReleasedByClose(t, func(t *testing.T, res *FetchResult) { + t.Helper() + + _, err := io.ReadAll(res.Content) + if err != nil { + t.Fatalf("read body: %v", err) + } + + err = res.Content.Close() + if err != nil { + t.Fatalf("close body: %v", err) + } + }) +} + +func TestFetchSemaphoreReleasedOnPartialReadClose(t *testing.T) { + t.Parallel() + + assertSlotReleasedByClose(t, func(t *testing.T, res *FetchResult) { + t.Helper() + + buf := make([]byte, 1) + + _, err := res.Content.Read(buf) + if err != nil { + t.Fatalf("partial read: %v", err) + } + + err = res.Content.Close() + if err != nil { + t.Fatalf("close body: %v", err) + } + }) +} + +// The dial-time re-resolution in ssrfSafeDialer is what closes the DNS +// rebinding window: even if validateURL saw a public answer earlier, the +// dialer independently re-checks the address it is about to connect to. A full +// rebinding simulation (a resolver returning public, then private) would mean +// replacing the global net.DefaultResolver with a fake DNS server, which is +// heavyweight and unsafe to mutate under parallel -race tests. The property is +// proven directly here instead: the dialer rejects a private target outright, +// which is exactly the check that fires when a validated host later resolves +// to a private address. +func TestSSRFSafeDialerBlocksPrivateTarget(t *testing.T) { + t.Parallel() + + for _, addr := range []string{ + "169.254.169.254:80", // link-local (cloud metadata) + "127.0.0.1:80", // loopback + "10.0.0.5:80", // RFC 1918 private + } { + t.Run(addr, func(t *testing.T) { + t.Parallel() + + _, err := ssrfSafeDialer(context.Background(), "tcp", addr) + if !errors.Is(err, ErrSSRFBlocked) { + t.Errorf("ssrfSafeDialer(%q) = %v, want ErrSSRFBlocked", addr, err) + } + }) + } +} + +func TestSSRFSafeDialerAllowsPublicTarget(t *testing.T) { + t.Parallel() + + // A cancelled context makes the dial fail immediately without touching the + // network; the point is only that a public literal is not SSRF-blocked. + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := ssrfSafeDialer(ctx, "tcp", testPublicHost+":80") + if err == nil { + t.Fatal("expected a dial error for an unreachable public target") + } + + if errors.Is(err, ErrSSRFBlocked) { + t.Errorf("public target was SSRF-blocked: %v", err) + } +} diff --git a/internal/middleware/middleware.go b/internal/middleware/middleware.go index 754ebd1..00d59f7 100644 --- a/internal/middleware/middleware.go +++ b/internal/middleware/middleware.go @@ -21,6 +21,33 @@ import ( // CORSMaxAgeSeconds is the max age for CORS preflight cache (24 hours). const CORSMaxAgeSeconds = 86400 +// HSTSValue is the Strict-Transport-Security header value: one year with +// includeSubDomains. Emitted unconditionally even though pixa listens plain +// HTTP behind a TLS-terminating proxy; browsers ignore an HSTS header received +// over plaintext (RFC 6797 section 8.1), so it never lies about the connection, +// and emitting it here avoids trusting a forwarded-proto header. +const HSTSValue = "max-age=31536000; includeSubDomains" + +// ContentSecurityPolicyValue is the Content-Security-Policy header value. +// default-src 'self' is the baseline and frame-ancestors 'none' is the primary +// clickjacking control. 'unsafe-inline' is required in script-src and style-src +// because the served templates carry inline onclick handlers (generator page) +// and the bundled Tailwind asset injects a runtime