package web import ( "bytes" "context" "image" "image/png" "mime/multipart" "net/http" "net/http/httptest" "strings" "testing" "github.com/alexedwards/scs/v2/memstore" "github.com/google/uuid" "golang.org/x/crypto/bcrypt" "plumber" "plumber/internal/blob" "plumber/internal/pacific" "plumber/internal/store" ) func newTestServer(t *testing.T, cfg Config) (*Server, *store.Memory) { t.Helper() mem := store.NewMemory() return newTestServerStore(t, mem, cfg), mem } func newTestServerStore(t *testing.T, st store.Store, cfg Config) *Server { t.Helper() if cfg.Blob == nil { cfg.Blob = blob.Disabled{} } srv, err := New(st, memstore.New(), plumber.TemplateFS, plumber.StaticFS, cfg) if err != nil { t.Fatal(err) } return srv } func uniq(prefix string) string { return prefix + "_" + strings.ReplaceAll(uuid.NewString()[:8], "-", "") } func seedUser(t *testing.T, st store.Store, username, password string, role store.Role) *store.User { t.Helper() hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.MinCost) if err != nil { t.Fatal(err) } u := &store.User{ Username: username, PasswordHash: string(hash), Role: role, } if err := st.CreateUser(context.Background(), u); err != nil { t.Fatal(err) } return u } func sessionValue(cookies []*http.Cookie) string { for _, c := range cookies { if c.Name == "plumber_session" { return c.Value } } return "" } func loginUser(t *testing.T, h http.Handler, username, password string) []*http.Cookie { t.Helper() rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/login", nil)) pre := rec.Result().Cookies() preToken := sessionValue(pre) csrf := csrfFrom(rec.Body.String()) form := strings.NewReader("_csrf=" + csrf + "&username=" + username + "&password=" + password) req := httptest.NewRequest(http.MethodPost, "/login", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range pre { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("login %s: %d %s", username, rec.Code, rec.Body.String()) } post := mergeCookies(pre, rec.Result().Cookies()) postToken := sessionValue(post) if preToken == "" || postToken == "" || preToken == postToken { t.Fatalf("expected session token rotation on login; pre=%q post=%q", preToken, postToken) } // Old anonymous token must not unlock authenticated routes. req = httptest.NewRequest(http.MethodGet, "/profile", nil) for _, c := range pre { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("pre-auth cookie should not access profile, got %d", rec.Code) } return post } func registerUser(t *testing.T, h http.Handler, username, password string, setupSecret ...string) []*http.Cookie { t.Helper() rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/register", nil)) pre := rec.Result().Cookies() preToken := sessionValue(pre) csrf := csrfFrom(rec.Body.String()) form := "_csrf=" + csrf + "&username=" + username + "&password=" + password if len(setupSecret) > 0 && setupSecret[0] != "" { form += "&setup_secret=" + setupSecret[0] } req := httptest.NewRequest(http.MethodPost, "/register", strings.NewReader(form)) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range pre { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("register %s: %d %s", username, rec.Code, rec.Body.String()) } post := mergeCookies(pre, rec.Result().Cookies()) postToken := sessionValue(post) if preToken == "" || postToken == "" || preToken == postToken { t.Fatalf("expected session token rotation on register; pre=%q post=%q", preToken, postToken) } return post } func TestHomeEmptyAndViewport(t *testing.T) { srv, _ := newTestServer(t, Config{}) rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/", nil) srv.Handler().ServeHTTP(rec, req) if rec.Code != 200 { t.Fatalf("status %d: %s", rec.Code, rec.Body.String()) } body := rec.Body.String() if !strings.Contains(body, "width=device-width") { t.Fatal("missing mobile viewport") } if !strings.Contains(body, "not a substitute for a licensed plumber") { t.Fatal("missing disclaimer") } } func TestRegisterLoginAsk(t *testing.T) { srv, _ := newTestServer(t, Config{}) h := srv.Handler() name := uniq("ask") session := registerUser(t, h, name, "hunter22") rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/submit", nil) for _, c := range session { req.AddCookie(c) } h.ServeHTTP(rec, req) if rec.Code != 200 { t.Fatalf("submit form %d", rec.Code) } csrf := csrfFrom(rec.Body.String()) form := strings.NewReader("_csrf=" + csrf + "&title=Leaky+faucet&body=Drip+all+night.&city=Oakland") req = httptest.NewRequest(http.MethodPost, "/submit", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range session { req.AddCookie(c) } for _, c := range rec.Result().Cookies() { req.AddCookie(c) } rec2 := httptest.NewRecorder() h.ServeHTTP(rec2, req) if rec2.Code != http.StatusSeeOther { t.Fatalf("submit %d %s", rec2.Code, rec2.Body.String()) } } func TestAdminSetupSecretOnlyWhenNoAdmins(t *testing.T) { mem := store.NewMemory() secret := "one-time-admin-setup-secret" adminName := uniq("seed") srv := newTestServerStore(t, mem, Config{AdminSetupSecret: secret}) h := srv.Handler() plain := uniq("plain") registerUser(t, h, plain, "hunter22") uPlain, err := mem.UserByUsername(context.Background(), plain) if err != nil || uPlain.Admin() { t.Fatalf("register without setup secret must stay user: %+v %v", uPlain, err) } registerUser(t, h, adminName, "hunter22", secret) u, err := mem.UserByUsername(context.Background(), adminName) if err != nil || !u.Admin() { t.Fatalf("setup secret registrant should be admin: %+v %v", u, err) } later := uniq("later") srv2 := newTestServerStore(t, mem, Config{AdminSetupSecret: secret}) registerUser(t, srv2.Handler(), later, "hunter22", secret) u2, err := mem.UserByUsername(context.Background(), later) if err != nil { t.Fatal(err) } if u2.Admin() { t.Fatal("setup secret must not grant admin once an admin already exists") } } func TestAdminUsersPageAccessAndRoles(t *testing.T) { srv, mem := newTestServer(t, Config{}) h := srv.Handler() hubName := uniq("hub") bobName := uniq("bob") carolName := uniq("carol") seedUser(t, mem, hubName, "hunter22", store.RoleAdmin) adminCookies := loginUser(t, h, hubName, "hunter22") registerUser(t, h, bobName, "hunter22") rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/admin/users", nil) for _, c := range adminCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) if rec.Code != 200 { t.Fatalf("admin list %d", rec.Code) } if !strings.Contains(rec.Body.String(), bobName) { t.Fatal("missing bob on admin page") } bob, err := mem.UserByUsername(context.Background(), bobName) if err != nil { t.Fatal(err) } csrf := csrfFrom(rec.Body.String()) form := strings.NewReader("_csrf=" + csrf + "&role=admin") req = httptest.NewRequest(http.MethodPost, "/admin/users/"+bob.ID+"/role", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range adminCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("promote %d %s", rec.Code, rec.Body.String()) } bob, _ = mem.UserByUsername(context.Background(), bobName) if !bob.Admin() { t.Fatal("bob should be admin") } carolCookies := registerUser(t, h, carolName, "hunter22") rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/admin/users", nil) for _, c := range carolCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) if rec.Code != http.StatusForbidden { t.Fatalf("non-admin expected 403, got %d", rec.Code) } rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/admin/users", nil) for _, c := range adminCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) csrf = csrfFrom(rec.Body.String()) form = strings.NewReader("_csrf=" + csrf + "&role=user") req = httptest.NewRequest(http.MethodPost, "/admin/users/"+bob.ID+"/role", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range adminCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("demote bob %d", rec.Code) } hub, err := mem.UserByUsername(context.Background(), hubName) if err != nil { t.Fatal(err) } rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/admin/users", nil) for _, c := range adminCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) csrf = csrfFrom(rec.Body.String()) form = strings.NewReader("_csrf=" + csrf + "&role=user") req = httptest.NewRequest(http.MethodPost, "/admin/users/"+hub.ID+"/role", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range adminCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != 200 { t.Fatalf("demote last admin expected page with error, got %d", rec.Code) } if !strings.Contains(rec.Body.String(), "Cannot demote the last admin") { t.Fatalf("missing last-admin error: %s", rec.Body.String()) } hub, _ = mem.UserByUsername(context.Background(), hubName) if !hub.Admin() { t.Fatal("hub must remain admin") } } type fakeBlob struct { calls int last string } func (f *fakeBlob) Enabled() bool { return true } func (f *fakeBlob) Upload(_ context.Context, obj blob.FileUpload) (string, error) { f.calls++ f.last = obj.Key return "https://cdn.example.com/" + obj.Key, nil } func TestProfilePageAndState(t *testing.T) { srv, mem := newTestServer(t, Config{}) h := srv.Handler() name := uniq("alice") cookies := registerUser(t, h, name, "hunter22") rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/profile", nil) for _, c := range cookies { req.AddCookie(c) } h.ServeHTTP(rec, req) if rec.Code != 200 { t.Fatalf("profile %d", rec.Code) } if !strings.Contains(rec.Body.String(), "Your questions") { t.Fatal("expected user questions label") } if !strings.Contains(rec.Body.String(), "local plumbing codes") { t.Fatal("missing state helper copy") } csrf := csrfFrom(rec.Body.String()) var buf bytes.Buffer w := multipart.NewWriter(&buf) _ = w.WriteField("_csrf", csrf) _ = w.WriteField("state", "CA") _ = w.Close() req = httptest.NewRequest(http.MethodPost, "/profile", &buf) req.Header.Set("Content-Type", w.FormDataContentType()) for _, c := range cookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("save profile %d %s", rec.Code, rec.Body.String()) } u, err := mem.UserByUsername(context.Background(), name) if err != nil || u.State != "CA" { t.Fatalf("state not saved: %+v %v", u, err) } rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/profile", nil) for _, c := range cookies { req.AddCookie(c) } h.ServeHTTP(rec, req) csrf = csrfFrom(rec.Body.String()) buf.Reset() w = multipart.NewWriter(&buf) _ = w.WriteField("_csrf", csrf) _ = w.WriteField("state", "ZZ") _ = w.Close() req = httptest.NewRequest(http.MethodPost, "/profile", &buf) req.Header.Set("Content-Type", w.FormDataContentType()) for _, c := range cookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != 200 || !strings.Contains(rec.Body.String(), "valid US state") { t.Fatalf("expected invalid state error, got %d %s", rec.Code, rec.Body.String()) } } func TestProfileAdminAnsweredListAndAvatarUpload(t *testing.T) { fb := &fakeBlob{} srv, mem := newTestServer(t, Config{Blob: fb}) h := srv.Handler() hubName := uniq("hub") aliceName := uniq("alice") hub := seedUser(t, mem, hubName, "hunter22", store.RoleAdmin) alice := seedUser(t, mem, aliceName, "hunter22", store.RoleUser) adminCookies := loginUser(t, h, hubName, "hunter22") q := &store.RankedQuestion{ AuthorID: alice.ID, Title: "Drip", Body: "Under sink", City: "Oakland", } if err := mem.CreateQuestion(context.Background(), q); err != nil { t.Fatal(err) } ans := &store.Answer{ QuestionID: q.ID, AuthorID: hub.ID, Body: "Replace the cartridge.", } if err := mem.UpsertAnswer(context.Background(), ans); err != nil { t.Fatal(err) } rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/profile", nil) for _, c := range adminCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) if rec.Code != 200 { t.Fatalf("admin profile %d", rec.Code) } body := rec.Body.String() if !strings.Contains(body, "Questions you answered") || !strings.Contains(body, "Drip") { t.Fatalf("admin answered list missing: %s", body) } csrf := csrfFrom(body) var buf bytes.Buffer w := multipart.NewWriter(&buf) _ = w.WriteField("_csrf", csrf) _ = w.WriteField("state", "OR") part, err := w.CreateFormFile("avatar", "pic.png") if err != nil { t.Fatal(err) } img := image.NewRGBA(image.Rect(0, 0, 1, 1)) if err := png.Encode(part, img); err != nil { t.Fatal(err) } _ = w.Close() req = httptest.NewRequest(http.MethodPost, "/profile", &buf) req.Header.Set("Content-Type", w.FormDataContentType()) for _, c := range adminCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("avatar upload %d %s", rec.Code, rec.Body.String()) } if fb.calls != 1 { t.Fatalf("expected 1 upload, got %d", fb.calls) } hub, _ = mem.UserByUsername(context.Background(), hubName) if !strings.Contains(hub.AvatarURL, "cdn.example.com/avatars/") { t.Fatalf("avatar url %q", hub.AvatarURL) } } func TestMutationsVoteAnswerHideAndCSRF(t *testing.T) { srv, mem := newTestServer(t, Config{}) h := srv.Handler() adminName := uniq("admin") userName := uniq("user") admin := seedUser(t, mem, adminName, "hunter22", store.RoleAdmin) user := seedUser(t, mem, userName, "hunter22", store.RoleUser) adminCookies := loginUser(t, h, adminName, "hunter22") userCookies := loginUser(t, h, userName, "hunter22") q := &store.RankedQuestion{ AuthorID: user.ID, Title: "Pipe noise", Body: "Clanking", City: "SF", HuntDate: pacific.Today(), } if err := mem.CreateQuestion(context.Background(), q); err != nil { t.Fatal(err) } // Missing CSRF form := strings.NewReader("value=1&view=question") req := httptest.NewRequest(http.MethodPost, "/questions/"+q.ID+"/vote", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range userCookies { req.AddCookie(c) } rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusForbidden { t.Fatalf("missing csrf want 403, got %d", rec.Code) } // Anonymous HTMX vote → sign-in prompt rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/login", nil) h.ServeHTTP(rec, req) anon := rec.Result().Cookies() csrf := csrfFrom(rec.Body.String()) form = strings.NewReader("_csrf=" + csrf + "&value=1&view=question") req = httptest.NewRequest(http.MethodPost, "/questions/"+q.ID+"/vote", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") req.Header.Set("HX-Request", "true") for _, c := range anon { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != 200 || !strings.Contains(rec.Body.String(), "Sign in") { t.Fatalf("anon htmx vote: %d %s", rec.Code, rec.Body.String()) } // User vote + HTMX fragment rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/questions/"+q.ID, nil) for _, c := range userCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) csrf = csrfFrom(rec.Body.String()) form = strings.NewReader("_csrf=" + csrf + "&value=1&view=question") req = httptest.NewRequest(http.MethodPost, "/questions/"+q.ID+"/vote", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") req.Header.Set("HX-Request", "true") for _, c := range userCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != 200 { t.Fatalf("vote htmx %d %s", rec.Code, rec.Body.String()) } got, err := mem.GetQuestion(context.Background(), q.ID, user.ID) if err != nil || got.UserVote != 1 || got.Score != 1 { t.Fatalf("vote not applied: %+v %v", got, err) } // Non-admin answer rejected rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/questions/"+q.ID, nil) for _, c := range userCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) csrf = csrfFrom(rec.Body.String()) form = strings.NewReader("_csrf=" + csrf + "&body=Nope") req = httptest.NewRequest(http.MethodPost, "/questions/"+q.ID+"/answer", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range userCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusForbidden { t.Fatalf("non-admin answer want 403, got %d", rec.Code) } // Admin answer success (HTMX) rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/questions/"+q.ID, nil) for _, c := range adminCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) csrf = csrfFrom(rec.Body.String()) form = strings.NewReader("_csrf=" + csrf + "&body=Tighten+the+nuts.") req = httptest.NewRequest(http.MethodPost, "/questions/"+q.ID+"/answer", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") req.Header.Set("HX-Request", "true") for _, c := range adminCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != 200 || !strings.Contains(rec.Body.String(), "Tighten the nuts") { t.Fatalf("admin answer: %d %s", rec.Code, rec.Body.String()) } if _, err := mem.GetAnswer(context.Background(), q.ID); err != nil { t.Fatal(err) } // Hide invalid id rec = httptest.NewRecorder() req = httptest.NewRequest(http.MethodGet, "/questions/"+q.ID, nil) for _, c := range adminCookies { req.AddCookie(c) } h.ServeHTTP(rec, req) csrf = csrfFrom(rec.Body.String()) form = strings.NewReader("_csrf=" + csrf) req = httptest.NewRequest(http.MethodPost, "/questions/does-not-exist/hide", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range adminCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusNotFound { t.Fatalf("hide missing want 404, got %d", rec.Code) } // Admin hide success form = strings.NewReader("_csrf=" + csrf) req = httptest.NewRequest(http.MethodPost, "/questions/"+q.ID+"/hide", form) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") for _, c := range adminCookies { req.AddCookie(c) } rec = httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusSeeOther { t.Fatalf("hide %d %s", rec.Code, rec.Body.String()) } hidden, err := mem.GetQuestion(context.Background(), q.ID, admin.ID) if err != nil || !hidden.Hidden { t.Fatalf("question not hidden: %+v %v", hidden, err) } } func mergeCookies(sets ...[]*http.Cookie) []*http.Cookie { byName := map[string]*http.Cookie{} for _, set := range sets { for _, c := range set { byName[c.Name] = c } } out := make([]*http.Cookie, 0, len(byName)) for _, c := range byName { out = append(out, c) } return out } func csrfFrom(html string) string { const needle = `name="_csrf" value="` i := strings.Index(html, needle) if i < 0 { return "" } html = html[i+len(needle):] j := strings.Index(html, `"`) if j < 0 { return "" } return html[:j] }