Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
eb75b92fb0 | ||
|
|
2d805125ee |
@@ -11,7 +11,6 @@ require (
|
|||||||
github.com/getsentry/sentry-go v0.40.0
|
github.com/getsentry/sentry-go v0.40.0
|
||||||
github.com/go-chi/chi/v5 v5.2.3
|
github.com/go-chi/chi/v5 v5.2.3
|
||||||
github.com/go-chi/cors v1.2.2
|
github.com/go-chi/cors v1.2.2
|
||||||
github.com/gorilla/csrf v1.7.3
|
|
||||||
github.com/gorilla/securecookie v1.1.2
|
github.com/gorilla/securecookie v1.1.2
|
||||||
github.com/prometheus/client_golang v1.23.2
|
github.com/prometheus/client_golang v1.23.2
|
||||||
github.com/slok/go-http-metrics v0.13.0
|
github.com/slok/go-http-metrics v0.13.0
|
||||||
|
|||||||
@@ -175,8 +175,6 @@ 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/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 h1:eBLnkZ9635krYIPD+ag1USrOAI0Nr0QYF3+/3GqO0k0=
|
||||||
github.com/googleapis/gax-go/v2 v2.14.2/go.mod h1:ON64QhlJkhVtSqp4v1uaK92VyZ2gmvDQsweuyLV+8+w=
|
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 h1:YCIWL56dvtr73r6715mJs5ZvhtnY73hBvEF8kXD8ePA=
|
||||||
github.com/gorilla/securecookie v1.1.2/go.mod h1:NfCASbcHqRSY+3a8tlWJwsQap2VX5pwzwo4h3eOamfo=
|
github.com/gorilla/securecookie v1.1.2/go.mod h1:NfCASbcHqRSY+3a8tlWJwsQap2VX5pwzwo4h3eOamfo=
|
||||||
github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 h1:5ZPtiqj0JL5oKWmcsq4VMaAW5ukBEgSGXEN89zeH1Jo=
|
github.com/grpc-ecosystem/grpc-gateway/v2 v2.26.3 h1:5ZPtiqj0JL5oKWmcsq4VMaAW5ukBEgSGXEN89zeH1Jo=
|
||||||
|
|||||||
+13
-23
@@ -2,7 +2,6 @@ package handlers
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/subtle"
|
"crypto/subtle"
|
||||||
"html/template"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -24,13 +23,13 @@ func (s *Handlers) HandleRoot() http.HandlerFunc {
|
|||||||
|
|
||||||
// Check if authenticated
|
// Check if authenticated
|
||||||
if s.sessMgr.IsAuthenticated(r) {
|
if s.sessMgr.IsAuthenticated(r) {
|
||||||
s.renderGenerator(w, r, nil)
|
s.renderGenerator(w, nil)
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Show login page
|
// Show login page
|
||||||
s.renderLogin(w, r, "")
|
s.renderLogin(w, "")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -38,7 +37,7 @@ func (s *Handlers) HandleRoot() http.HandlerFunc {
|
|||||||
func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) {
|
func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) {
|
||||||
err := r.ParseForm()
|
err := r.ParseForm()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.renderLogin(w, r, "Invalid form data")
|
s.renderLogin(w, "Invalid form data")
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -48,7 +47,7 @@ func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) {
|
|||||||
// Constant-time comparison to prevent timing attacks
|
// Constant-time comparison to prevent timing attacks
|
||||||
if subtle.ConstantTimeCompare([]byte(submittedKey), []byte(s.config.SigningKey)) != 1 {
|
if subtle.ConstantTimeCompare([]byte(submittedKey), []byte(s.config.SigningKey)) != 1 {
|
||||||
s.log.Warn("failed login attempt", "remote_addr", r.RemoteAddr)
|
s.log.Warn("failed login attempt", "remote_addr", r.RemoteAddr)
|
||||||
s.renderLogin(w, r, "Invalid signing key")
|
s.renderLogin(w, "Invalid signing key")
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -57,7 +56,7 @@ func (s *Handlers) handleLoginPost(w http.ResponseWriter, r *http.Request) {
|
|||||||
err = s.sessMgr.CreateSession(w)
|
err = s.sessMgr.CreateSession(w)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.log.Error("failed to create session", "error", err)
|
s.log.Error("failed to create session", "error", err)
|
||||||
s.renderLogin(w, r, "Failed to create session")
|
s.renderLogin(w, "Failed to create session")
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -88,7 +87,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc {
|
|||||||
|
|
||||||
err := r.ParseForm()
|
err := r.ParseForm()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.renderGenerator(w, r, &generatorData{Error: "Invalid form data"})
|
s.renderGenerator(w, &generatorData{Error: "Invalid form data"})
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -98,7 +97,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc {
|
|||||||
// Validate source URL
|
// Validate source URL
|
||||||
parsed, err := url.Parse(sourceURL)
|
parsed, err := url.Parse(sourceURL)
|
||||||
if err != nil || parsed.Host == "" {
|
if err != nil || parsed.Host == "" {
|
||||||
s.renderGeneratorWithForm(w, r, "Invalid source URL", r.Form)
|
s.renderGeneratorWithForm(w, "Invalid source URL", r.Form)
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -109,7 +108,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc {
|
|||||||
token, err := s.encGen.Generate(payload)
|
token, err := s.encGen.Generate(payload)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.log.Error("failed to generate encrypted URL", "error", err)
|
s.log.Error("failed to generate encrypted URL", "error", err)
|
||||||
s.renderGeneratorWithForm(w, r, "Failed to generate URL", r.Form)
|
s.renderGeneratorWithForm(w, "Failed to generate URL", r.Form)
|
||||||
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -122,7 +121,7 @@ func (s *Handlers) HandleGenerateURL() http.HandlerFunc {
|
|||||||
expiresAtStr = expiresAt.Format(time.RFC3339)
|
expiresAtStr = expiresAt.Format(time.RFC3339)
|
||||||
}
|
}
|
||||||
|
|
||||||
s.renderGenerator(w, r, &generatorData{
|
s.renderGenerator(w, &generatorData{
|
||||||
GeneratedURL: generatedURL,
|
GeneratedURL: generatedURL,
|
||||||
ExpiresAt: expiresAtStr,
|
ExpiresAt: expiresAtStr,
|
||||||
FormURL: sourceURL,
|
FormURL: sourceURL,
|
||||||
@@ -187,20 +186,15 @@ type generatorData struct {
|
|||||||
FormQuality string
|
FormQuality string
|
||||||
FormFit string
|
FormFit string
|
||||||
FormTTL string
|
FormTTL string
|
||||||
CSRFField template.HTML
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Handlers) renderLogin(
|
func (s *Handlers) renderLogin(w http.ResponseWriter, errorMsg string) {
|
||||||
w http.ResponseWriter, r *http.Request, errorMsg string,
|
|
||||||
) {
|
|
||||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||||
|
|
||||||
data := struct {
|
data := struct {
|
||||||
Error string
|
Error string
|
||||||
CSRFField template.HTML
|
|
||||||
}{
|
}{
|
||||||
Error: errorMsg,
|
Error: errorMsg,
|
||||||
CSRFField: csrfField(r),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
err := templates.Render(w, "login.html", data)
|
err := templates.Render(w, "login.html", data)
|
||||||
@@ -210,17 +204,13 @@ func (s *Handlers) renderLogin(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Handlers) renderGenerator(
|
func (s *Handlers) renderGenerator(w http.ResponseWriter, data *generatorData) {
|
||||||
w http.ResponseWriter, r *http.Request, data *generatorData,
|
|
||||||
) {
|
|
||||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||||
|
|
||||||
if data == nil {
|
if data == nil {
|
||||||
data = &generatorData{}
|
data = &generatorData{}
|
||||||
}
|
}
|
||||||
|
|
||||||
data.CSRFField = csrfField(r)
|
|
||||||
|
|
||||||
err := templates.Render(w, "generator.html", data)
|
err := templates.Render(w, "generator.html", data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.log.Error("failed to render generator template", "error", err)
|
s.log.Error("failed to render generator template", "error", err)
|
||||||
@@ -229,9 +219,9 @@ func (s *Handlers) renderGenerator(
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *Handlers) renderGeneratorWithForm(
|
func (s *Handlers) renderGeneratorWithForm(
|
||||||
w http.ResponseWriter, r *http.Request, errorMsg string, form url.Values,
|
w http.ResponseWriter, errorMsg string, form url.Values,
|
||||||
) {
|
) {
|
||||||
s.renderGenerator(w, r, &generatorData{
|
s.renderGenerator(w, &generatorData{
|
||||||
Error: errorMsg,
|
Error: errorMsg,
|
||||||
FormURL: form.Get("url"),
|
FormURL: form.Get("url"),
|
||||||
FormWidth: form.Get("width"),
|
FormWidth: form.Get("width"),
|
||||||
|
|||||||
@@ -1,273 +0,0 @@
|
|||||||
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
|
|
||||||
}
|
|
||||||
@@ -1,65 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
@@ -39,22 +39,15 @@ type Handlers struct {
|
|||||||
imgCache *imgcache.Cache
|
imgCache *imgcache.Cache
|
||||||
sessMgr *session.Manager
|
sessMgr *session.Manager
|
||||||
encGen *encurl.Generator
|
encGen *encurl.Generator
|
||||||
csrfProtect func(http.Handler) http.Handler
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// New creates a new Handlers instance.
|
// New creates a new Handlers instance.
|
||||||
func New(lc fx.Lifecycle, params Params) (*Handlers, error) {
|
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{
|
s := &Handlers{
|
||||||
log: params.Logger.Get(),
|
log: params.Logger.Get(),
|
||||||
hc: params.Healthcheck,
|
hc: params.Healthcheck,
|
||||||
db: params.Database,
|
db: params.Database,
|
||||||
config: params.Config,
|
config: params.Config,
|
||||||
csrfProtect: csrfProtect,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
lc.Append(fx.Hook{
|
lc.Append(fx.Hook{
|
||||||
|
|||||||
@@ -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, "<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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -44,17 +44,11 @@ func (s *Server) SetupRoutes() {
|
|||||||
// Static files (Tailwind CSS, etc.)
|
// Static files (Tailwind CSS, etc.)
|
||||||
s.router.Handle("/static/*", http.StripPrefix("/static/", static.Handler()))
|
s.router.Handle("/static/*", http.StripPrefix("/static/", static.Handler()))
|
||||||
|
|
||||||
// Login/generator UI. The form routes carry CSRF protection; the
|
// Login/generator UI
|
||||||
// token cookie is independent of the session cookie, so it also
|
s.router.Get("/", s.h.HandleRoot())
|
||||||
// covers the login POST, where no session exists yet.
|
s.router.Post("/", s.h.HandleRoot())
|
||||||
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.Get("/logout", s.h.HandleLogout())
|
||||||
|
s.router.Post("/generate", s.h.HandleGenerateURL())
|
||||||
|
|
||||||
// Main image proxy route
|
// Main image proxy route
|
||||||
// /v1/image/<host>/<path>/<width>x<height>.<format>
|
// /v1/image/<host>/<path>/<width>x<height>.<format>
|
||||||
|
|||||||
@@ -47,7 +47,6 @@
|
|||||||
{{end}}
|
{{end}}
|
||||||
|
|
||||||
<form method="POST" action="/generate" class="bg-white rounded-lg shadow-md p-6 space-y-4">
|
<form method="POST" action="/generate" class="bg-white rounded-lg shadow-md p-6 space-y-4">
|
||||||
{{ .CSRFField }}
|
|
||||||
<div>
|
<div>
|
||||||
<label for="url" class="block text-sm font-medium text-gray-700 mb-1">
|
<label for="url" class="block text-sm font-medium text-gray-700 mb-1">
|
||||||
Source URL
|
Source URL
|
||||||
|
|||||||
@@ -17,7 +17,6 @@
|
|||||||
{{end}}
|
{{end}}
|
||||||
|
|
||||||
<form method="POST" action="/" class="space-y-4">
|
<form method="POST" action="/" class="space-y-4">
|
||||||
{{ .CSRFField }}
|
|
||||||
<div>
|
<div>
|
||||||
<label for="key" class="block text-sm font-medium text-gray-700 mb-1">
|
<label for="key" class="block text-sm font-medium text-gray-700 mb-1">
|
||||||
Signing Key
|
Signing Key
|
||||||
|
|||||||
Reference in New Issue
Block a user