Author SHA1 Message Date
sneak 6a753f14c4 feat: CSRF protection on the login and URL-generator forms (closes #93)
check / check (push) Failing after 0s
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
2026-09-21 12:57:23 +00:00
sneak 5f88dc5cfa test: CSRF rejection/acceptance for POST / and POST /generate
Failing tests (TDD) for the two cookie-authenticated HTML form posts:
a POST without a CSRF token is rejected, a token that does not match the
request's CSRF cookie is rejected, and a matching cookie+token succeeds.
Login CSRF is covered specifically: the POST / cases carry no session,
so protection rests on a token bound to a pre-session cookie.

These reference production symbols not yet added (newCSRFProtect,
Handlers.CSRF, the csrfProtect field), so the package does not build
until the implementation lands.

model: claude-opus-4-8
2026-09-21 12:57:23 +00:00
clawbot 04b5db6fbf next -> main (1.0.0 milestone) (#105)
check / check (push) Failing after 0s
Accumulating milestone branch. One squashed commit per closed issue; `next` is kept green and mergeable to `main` at any time without notice.

Landed so far:

- `chore: update golangci-lint to v2.12.2 with canonical config` (#54) — canonical v2-schema `.golangci.yml`, pins bumped in `Dockerfile` and `script/bootstrap`, tree at `0 issues.`. Three behaviour deltas are recorded in that PR's body: `Cache.StoreVariant` takes a context, `MetadataStorage.Store` no longer leaks temp files on failure, and the `signing_key` too-short error text gained a `value too short:` prefix.

Sequencing for the milestone is tracked in #103.

Reviewed-on: #105
Co-authored-by: clawbot <clawbot@noreply.example.org>
2026-09-21 09:31:54 +02:00
10 changed files with 397 additions and 452 deletions
+1
View File
@@ -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
+2
View File
@@ -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=
+25 -15
View File
@@ -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"),
@@ -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
}
+65
View File
@@ -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)
}
+19 -12
View File
@@ -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{
-421
View File
@@ -1,421 +0,0 @@
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, "<html></html>")
})
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)
}
}
+10 -4
View File
@@ -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/<host>/<path>/<width>x<height>.<format>
+1
View File
@@ -47,6 +47,7 @@
{{end}}
<form method="POST" action="/generate" class="bg-white rounded-lg shadow-md p-6 space-y-4">
{{ .CSRFField }}
<div>
<label for="url" class="block text-sm font-medium text-gray-700 mb-1">
Source URL
+1
View File
@@ -17,6 +17,7 @@
{{end}}
<form method="POST" action="/" class="space-y-4">
{{ .CSRFField }}
<div>
<label for="key" class="block text-sm font-medium text-gray-700 mb-1">
Signing Key