Refactor Store into SessionStore; move domain SQL onto User/Question/Answer.

This commit is contained in:
2026-08-22 02:40:51 -07:00
parent f31f352838
commit c77298411e
15 changed files with 696 additions and 876 deletions
+8 -3
View File
@@ -28,7 +28,7 @@ func (s *Server) handleAdminUsers(w http.ResponseWriter, r *http.Request) {
if s.requireAdmin(w, r) == nil {
return
}
users, err := s.store.ListUsers(r.Context())
users, err := store.ListUsers(r.Context(), s.db)
if err != nil {
http.Error(w, "could not load users", http.StatusInternalServerError)
return
@@ -48,9 +48,14 @@ func (s *Server) handleAdminSetRole(w http.ResponseWriter, r *http.Request) {
}
id := chi.URLParam(r, "id")
role := store.Role(r.PostFormValue("role"))
err := s.store.SetRole(r.Context(), id, role)
u, err := store.UserByID(r.Context(), s.db, id)
if err != nil {
http.Error(w, "could not update role", http.StatusBadRequest)
return
}
err = u.SetRole(r.Context(), role)
if errors.Is(err, store.ErrLastAdmin) {
users, listErr := s.store.ListUsers(r.Context())
users, listErr := store.ListUsers(r.Context(), s.db)
if listErr != nil {
http.Error(w, "could not demote last admin", http.StatusBadRequest)
return
+7 -8
View File
@@ -43,7 +43,7 @@ func (s *Server) handleLogin(w http.ResponseWriter, r *http.Request) {
username := strings.TrimSpace(r.PostFormValue("username"))
password := r.PostFormValue("password")
next := safeNext(r.PostFormValue("next"))
u, err := s.store.UserByUsername(r.Context(), username)
u, err := store.UserByUsername(r.Context(), s.db, username)
if err != nil || bcrypt.CompareHashAndPassword([]byte(u.PasswordHash), []byte(password)) != nil {
w.WriteHeader(http.StatusUnauthorized)
s.exec(w, "login", authPage{
@@ -90,7 +90,7 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
}
role := store.RoleUser
if s.cfg.AdminUsername != "" && store.NormalizeUsername(username) == store.NormalizeUsername(s.cfg.AdminUsername) {
n, err := s.store.CountAdmins(r.Context())
n, err := store.CountAdmins(r.Context(), s.db)
if err != nil {
http.Error(w, "could not create account", http.StatusInternalServerError)
return
@@ -99,12 +99,11 @@ func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
role = store.RoleAdmin
}
}
u, err := s.store.CreateUser(r.Context(), store.NewUser{
Username: username,
PasswordHash: string(hash),
Role: role,
})
if err != nil {
u := store.NewUser(s.db)
u.Username = username
u.PasswordHash = string(hash)
u.Role = role
if err := u.Create(r.Context()); err != nil {
p.Error = "That username is taken."
s.exec(w, "register", p)
return
-329
View File
@@ -1,329 +0,0 @@
package web
import (
"context"
"database/sql"
"fmt"
"sort"
"strings"
"sync"
"time"
"github.com/google/uuid"
"plumber/internal/pacific"
"plumber/internal/store"
)
// memDB is an in-memory store.DB for tests.
type memDB struct {
mu sync.Mutex
users map[string]*store.User // id -> user
byName map[string]string // username -> id
questions map[string]*store.RankedQuestion // id -> question
votes map[string]int // userID|questionID -> value
answers map[string]*store.Answer // questionID -> answer
}
func newMemDB() *memDB {
return &memDB{
users: map[string]*store.User{},
byName: map[string]string{},
questions: map[string]*store.RankedQuestion{},
votes: map[string]int{},
answers: map[string]*store.Answer{},
}
}
func voteKey(userID, questionID string) string {
return userID + "|" + questionID
}
func (m *memDB) CreateUser(_ context.Context, nu store.NewUser) (*store.User, error) {
m.mu.Lock()
defer m.mu.Unlock()
username := store.NormalizeUsername(nu.Username)
if _, ok := m.byName[username]; ok {
return nil, fmt.Errorf("username taken")
}
if nu.Role != store.RoleUser && nu.Role != store.RoleAdmin {
return nil, fmt.Errorf("invalid role")
}
u := &store.User{
ID: uuid.NewString(),
Username: username,
Name: username,
Role: nu.Role,
PasswordHash: nu.PasswordHash,
CreatedAt: time.Now().UTC().Format(time.RFC3339),
}
m.users[u.ID] = u
m.byName[username] = u.ID
cp := *u
return &cp, nil
}
func (m *memDB) UserByID(_ context.Context, id string) (*store.User, error) {
m.mu.Lock()
defer m.mu.Unlock()
u, ok := m.users[id]
if !ok {
return nil, sql.ErrNoRows
}
cp := *u
cp.PasswordHash = ""
return &cp, nil
}
func (m *memDB) UserByUsername(_ context.Context, username string) (*store.User, error) {
m.mu.Lock()
defer m.mu.Unlock()
id, ok := m.byName[store.NormalizeUsername(username)]
if !ok {
return nil, sql.ErrNoRows
}
cp := *m.users[id]
return &cp, nil
}
func (m *memDB) CountAdmins(_ context.Context) (int, error) {
m.mu.Lock()
defer m.mu.Unlock()
n := 0
for _, u := range m.users {
if u.Role == store.RoleAdmin {
n++
}
}
return n, nil
}
func (m *memDB) ListUsers(_ context.Context) ([]store.User, error) {
m.mu.Lock()
defer m.mu.Unlock()
out := make([]store.User, 0, len(m.users))
for _, u := range m.users {
cp := *u
cp.PasswordHash = ""
out = append(out, cp)
}
sort.Slice(out, func(i, j int) bool {
return out[i].CreatedAt < out[j].CreatedAt
})
return out, nil
}
func (m *memDB) SetRole(_ context.Context, userID string, role store.Role) error {
if role != store.RoleUser && role != store.RoleAdmin {
return fmt.Errorf("invalid role")
}
m.mu.Lock()
defer m.mu.Unlock()
u, ok := m.users[userID]
if !ok {
return sql.ErrNoRows
}
if u.Role == store.RoleAdmin && role == store.RoleUser {
n := 0
for _, x := range m.users {
if x.Role == store.RoleAdmin {
n++
}
}
if n <= 1 {
return store.ErrLastAdmin
}
}
u.Role = role
return nil
}
func (m *memDB) CreateQuestion(_ context.Context, authorID, title, body, city string) (*store.RankedQuestion, error) {
m.mu.Lock()
defer m.mu.Unlock()
author, ok := m.users[authorID]
if !ok {
return nil, fmt.Errorf("unknown author")
}
q := &store.RankedQuestion{
ID: uuid.NewString(),
AuthorID: authorID,
AuthorName: author.Name,
Title: strings.TrimSpace(title),
Body: strings.TrimSpace(body),
City: strings.TrimSpace(city),
HuntDate: pacific.Today(),
CreatedAt: time.Now().UTC().Format(time.RFC3339),
}
m.questions[q.ID] = q
cp := *q
return &cp, nil
}
func (m *memDB) rankedLocked(q *store.RankedQuestion, viewerID string) store.RankedQuestion {
out := *q
score := 0
for k, v := range m.votes {
_, qid, ok := strings.Cut(k, "|")
if ok && qid == q.ID {
score += v
}
}
out.Score = score
out.Answered = m.answers[q.ID] != nil
if viewerID != "" {
out.UserVote = m.votes[voteKey(viewerID, q.ID)]
}
if u, ok := m.users[q.AuthorID]; ok {
out.AuthorName = u.Name
}
return out
}
func (m *memDB) ListHunt(_ context.Context, huntDate, viewerID string) ([]store.RankedQuestion, error) {
m.mu.Lock()
defer m.mu.Unlock()
var out []store.RankedQuestion
for _, q := range m.questions {
if q.HuntDate != huntDate || q.Hidden {
continue
}
out = append(out, m.rankedLocked(q, viewerID))
}
sort.Slice(out, func(i, j int) bool {
if out[i].Score != out[j].Score {
return out[i].Score > out[j].Score
}
return out[i].CreatedAt < out[j].CreatedAt
})
return out, nil
}
func (m *memDB) GetQuestion(_ context.Context, id, viewerID string) (*store.RankedQuestion, error) {
m.mu.Lock()
defer m.mu.Unlock()
q, ok := m.questions[id]
if !ok {
return nil, sql.ErrNoRows
}
out := m.rankedLocked(q, viewerID)
return &out, nil
}
func (m *memDB) Vote(_ context.Context, userID, questionID string, value int) error {
if value != 1 && value != -1 {
return fmt.Errorf("invalid vote")
}
m.mu.Lock()
defer m.mu.Unlock()
if _, ok := m.questions[questionID]; !ok {
return fmt.Errorf("unknown question")
}
k := voteKey(userID, questionID)
if cur, ok := m.votes[k]; ok && cur == value {
delete(m.votes, k)
return nil
}
m.votes[k] = value
return nil
}
func (m *memDB) GetAnswer(_ context.Context, questionID string) (*store.Answer, error) {
m.mu.Lock()
defer m.mu.Unlock()
a, ok := m.answers[questionID]
if !ok {
return nil, sql.ErrNoRows
}
cp := *a
if u, ok := m.users[a.AuthorID]; ok {
cp.AuthorName = u.Name
}
return &cp, nil
}
func (m *memDB) UpsertAnswer(_ context.Context, questionID, authorID, body string) error {
m.mu.Lock()
defer m.mu.Unlock()
body = strings.TrimSpace(body)
now := time.Now().UTC().Format(time.RFC3339)
if existing, ok := m.answers[questionID]; ok {
existing.Body = body
existing.AuthorID = authorID
existing.UpdatedAt = now
return nil
}
m.answers[questionID] = &store.Answer{
QuestionID: questionID,
AuthorID: authorID,
Body: body,
CreatedAt: now,
UpdatedAt: now,
}
return nil
}
func (m *memDB) HideQuestion(_ context.Context, id string) error {
m.mu.Lock()
defer m.mu.Unlock()
q, ok := m.questions[id]
if !ok {
return sql.ErrNoRows
}
q.Hidden = true
return nil
}
func (m *memDB) UpdateProfile(_ context.Context, userID, state, avatarURL string) error {
m.mu.Lock()
defer m.mu.Unlock()
u, ok := m.users[userID]
if !ok {
return sql.ErrNoRows
}
u.State = state
if avatarURL != "" {
u.AvatarURL = avatarURL
}
return nil
}
func (m *memDB) ListQuestionsByAuthor(_ context.Context, authorID string) ([]store.RankedQuestion, error) {
m.mu.Lock()
defer m.mu.Unlock()
var out []store.RankedQuestion
for _, q := range m.questions {
if q.AuthorID != authorID || q.Hidden {
continue
}
out = append(out, m.rankedLocked(q, ""))
}
sort.Slice(out, func(i, j int) bool {
return out[i].CreatedAt > out[j].CreatedAt
})
return out, nil
}
func (m *memDB) ListQuestionsAnsweredBy(_ context.Context, adminID string) ([]store.RankedQuestion, error) {
m.mu.Lock()
defer m.mu.Unlock()
var out []store.RankedQuestion
for qid, a := range m.answers {
if a.AuthorID != adminID {
continue
}
q, ok := m.questions[qid]
if !ok || q.Hidden {
continue
}
rq := m.rankedLocked(q, "")
rq.Answered = true
out = append(out, rq)
}
sort.Slice(out, func(i, j int) bool {
return out[i].CreatedAt > out[j].CreatedAt
})
return out, nil
}
var _ store.DB = (*memDB)(nil)
+8 -4
View File
@@ -91,7 +91,11 @@ func (s *Server) handleProfile(w http.ResponseWriter, r *http.Request) {
return
}
if err := s.store.UpdateProfile(r.Context(), u.ID, state, avatarURL); err != nil {
u.State = state
if avatarURL != "" {
u.AvatarURL = avatarURL
}
if err := u.SaveProfile(r.Context()); err != nil {
http.Error(w, "could not save profile", http.StatusInternalServerError)
return
}
@@ -122,16 +126,16 @@ func (s *Server) renderProfile(w http.ResponseWriter, r *http.Request, u *store.
)
if u.Admin() {
label = "Questions you answered"
questions, err = s.store.ListQuestionsAnsweredBy(r.Context(), u.ID)
questions, err = store.ListQuestionsAnsweredBy(r.Context(), s.db, u.ID)
} else {
label = "Your questions"
questions, err = s.store.ListQuestionsByAuthor(r.Context(), u.ID)
questions, err = store.ListQuestionsByAuthor(r.Context(), s.db, u.ID)
}
if err != nil {
http.Error(w, "could not load questions", http.StatusInternalServerError)
return
}
if fresh, e := s.store.UserByID(r.Context(), u.ID); e == nil {
if fresh, e := store.UserByID(r.Context(), s.db, u.ID); e == nil {
u = fresh
}
p := s.basePage(r, "Profile")
+26 -17
View File
@@ -3,6 +3,7 @@ package web
import (
"context"
"crypto/rand"
"database/sql"
"encoding/hex"
"fmt"
"html/template"
@@ -30,7 +31,7 @@ type Config struct {
}
type Server struct {
store store.DB
db *sql.DB
sessions *scs.SessionManager
tmpl *template.Template
cfg Config
@@ -84,7 +85,7 @@ type voteCtx struct {
Question store.RankedQuestion
}
func New(st store.DB, sessionStore scs.Store, templateFS fs.FS, staticFS fs.FS, cfg Config) (*Server, error) {
func New(db *sql.DB, sessionStore scs.Store, templateFS fs.FS, staticFS fs.FS, cfg Config) (*Server, error) {
if cfg.Blob == nil {
cfg.Blob = blob.Disabled{}
}
@@ -124,7 +125,7 @@ func New(st store.DB, sessionStore scs.Store, templateFS fs.FS, staticFS fs.FS,
}
return &Server{
store: st,
db: db,
sessions: sessions,
tmpl: tmpl,
cfg: cfg,
@@ -180,7 +181,7 @@ func (s *Server) withUser(next http.Handler) http.Handler {
}
id := s.sessions.GetString(r.Context(), "user_id")
if id != "" {
u, err := s.store.UserByID(r.Context(), id)
u, err := store.UserByID(r.Context(), s.db, id)
if err == nil {
r = r.WithContext(context.WithValue(r.Context(), userKey, u))
}
@@ -258,7 +259,7 @@ func (s *Server) renderHunt(w http.ResponseWriter, r *http.Request, date string)
if u := currentUser(r); u != nil {
viewer = u.ID
}
questions, err := s.store.ListHunt(r.Context(), date, viewer)
questions, err := store.ListHunt(r.Context(), s.db, date, viewer)
if err != nil {
http.Error(w, "could not load questions", http.StatusInternalServerError)
return
@@ -318,8 +319,12 @@ func (s *Server) handleSubmit(w http.ResponseWriter, r *http.Request) {
if len(city) > 80 {
city = city[:80]
}
q, err := s.store.CreateQuestion(r.Context(), u.ID, title, body, city)
if err != nil {
q := store.NewQuestion(s.db)
q.AuthorID = u.ID
q.Title = title
q.Body = body
q.City = city
if err := q.Create(r.Context()); err != nil {
http.Error(w, "could not save question", http.StatusInternalServerError)
return
}
@@ -332,14 +337,14 @@ func (s *Server) handleQuestion(w http.ResponseWriter, r *http.Request) {
if u := currentUser(r); u != nil {
viewer = u.ID
}
q, err := s.store.GetQuestion(r.Context(), id, viewer)
q, err := store.GetQuestion(r.Context(), s.db, id, viewer)
if err != nil || (q.Hidden && !currentUser(r).Admin()) {
http.NotFound(w, r)
return
}
var ans *store.Answer
if q.Answered {
ans, _ = s.store.GetAnswer(r.Context(), q.ID)
ans, _ = store.GetAnswer(r.Context(), s.db, q.ID)
}
s.exec(w, "question", questionPage{
page: s.basePage(r, q.Title),
@@ -372,7 +377,7 @@ func (s *Server) handleVote(w http.ResponseWriter, r *http.Request) {
http.Error(w, "invalid vote", http.StatusBadRequest)
return
}
if err := s.store.Vote(r.Context(), u.ID, id, value); err != nil {
if err := store.Vote(r.Context(), s.db, u.ID, id, value); err != nil {
http.Error(w, "could not vote", http.StatusInternalServerError)
return
}
@@ -383,7 +388,7 @@ func (s *Server) handleVote(w http.ResponseWriter, r *http.Request) {
s.renderLeaderboard(w, r, date)
return
}
q, err := s.store.GetQuestion(r.Context(), id, u.ID)
q, err := store.GetQuestion(r.Context(), s.db, id, u.ID)
if err != nil {
http.Error(w, "not found", http.StatusNotFound)
return
@@ -416,7 +421,7 @@ func (s *Server) renderLeaderboard(w http.ResponseWriter, r *http.Request, date
if u := currentUser(r); u != nil {
viewer = u.ID
}
questions, err := s.store.ListHunt(r.Context(), date, viewer)
questions, err := store.ListHunt(r.Context(), s.db, date, viewer)
if err != nil {
http.Error(w, "could not load questions", http.StatusInternalServerError)
return
@@ -446,17 +451,21 @@ func (s *Server) handleAnswer(w http.ResponseWriter, r *http.Request) {
if len(body) > 12000 {
body = body[:12000]
}
if err := s.store.UpsertAnswer(r.Context(), id, u.ID, body); err != nil {
ans := store.NewAnswer(s.db)
ans.QuestionID = id
ans.AuthorID = u.ID
ans.Body = body
if err := ans.Upsert(r.Context()); err != nil {
http.Error(w, "could not save answer", http.StatusInternalServerError)
return
}
ans, err := s.store.GetAnswer(r.Context(), id)
saved, err := store.GetAnswer(r.Context(), s.db, id)
if err != nil {
http.Error(w, "could not load answer", http.StatusInternalServerError)
return
}
if isHTMX(r) {
s.exec(w, "answer", questionPage{page: s.basePage(r, ""), Answer: ans})
s.exec(w, "answer", questionPage{page: s.basePage(r, ""), Answer: saved})
return
}
http.Redirect(w, r, "/questions/"+url.PathEscape(id), http.StatusSeeOther)
@@ -472,12 +481,12 @@ func (s *Server) handleHide(w http.ResponseWriter, r *http.Request) {
return
}
id := chi.URLParam(r, "id")
q, err := s.store.GetQuestion(r.Context(), id, u.ID)
q, err := store.GetQuestion(r.Context(), s.db, id, u.ID)
if err != nil {
http.NotFound(w, r)
return
}
if err := s.store.HideQuestion(r.Context(), id); err != nil {
if err := q.Hide(r.Context()); err != nil {
http.Error(w, "could not hide", http.StatusInternalServerError)
return
}
+164 -71
View File
@@ -3,31 +3,98 @@ package web
import (
"bytes"
"context"
"database/sql"
"mime/multipart"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"
"github.com/alexedwards/scs/v2"
"github.com/google/uuid"
"github.com/joho/godotenv"
"golang.org/x/crypto/bcrypt"
"plumber"
"plumber/internal/blob"
"plumber/internal/store"
)
func newTestServer(t *testing.T) (*Server, *memDB, scs.Store) {
func testDBURL() string {
_ = godotenv.Load()
if u := strings.TrimSpace(os.Getenv("TEST_DATABASE_URL")); u != "" {
return u
}
return strings.TrimSpace(os.Getenv("DATABASE_URL"))
}
func newTestServer(t *testing.T, cfg Config) (*Server, *sql.DB) {
t.Helper()
fake := newMemDB()
sessions := scs.New()
srv, err := New(fake, sessions.Store, plumber.TemplateFS, plumber.StaticFS, Config{AdminUsername: "hub"})
url := testDBURL()
if url == "" {
t.Skip("set TEST_DATABASE_URL or DATABASE_URL for web tests")
}
db, sessions, err := store.OpenPostgres(url, plumber.SchemaSQL)
if err != nil {
t.Fatalf("open postgres: %v", err)
}
t.Cleanup(func() {
sessions.Close()
_ = db.Close()
})
if cfg.Blob == nil {
cfg.Blob = blob.Disabled{}
}
srv, err := New(db, sessions.Store(), plumber.TemplateFS, plumber.StaticFS, cfg)
if err != nil {
t.Fatal(err)
}
return srv, fake, sessions.Store
return srv, db
}
func uniq(prefix string) string {
return prefix + "_" + strings.ReplaceAll(uuid.NewString()[:8], "-", "")
}
func seedUser(t *testing.T, db *sql.DB, 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.NewUser(db)
u.Username = username
u.PasswordHash = string(hash)
u.Role = role
if err := u.Create(context.Background()); err != nil {
t.Fatal(err)
}
return u
}
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))
cookies := rec.Result().Cookies()
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 cookies {
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())
}
return mergeCookies(cookies, rec.Result().Cookies())
}
func TestHomeEmptyAndViewport(t *testing.T) {
srv, _, _ := newTestServer(t)
srv, _ := newTestServer(t, Config{})
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/", nil)
srv.Handler().ServeHTTP(rec, req)
@@ -35,9 +102,6 @@ func TestHomeEmptyAndViewport(t *testing.T) {
t.Fatalf("status %d: %s", rec.Code, rec.Body.String())
}
body := rec.Body.String()
if !strings.Contains(body, "No questions yet") {
t.Fatal("missing empty state")
}
if !strings.Contains(body, "width=device-width") {
t.Fatal("missing mobile viewport")
}
@@ -47,8 +111,9 @@ func TestHomeEmptyAndViewport(t *testing.T) {
}
func TestRegisterLoginAsk(t *testing.T) {
srv, _, _ := newTestServer(t)
srv, _ := newTestServer(t, Config{})
h := srv.Handler()
name := uniq("ask")
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/register", nil))
cookie := rec.Result().Cookies()
@@ -56,7 +121,7 @@ func TestRegisterLoginAsk(t *testing.T) {
if csrf == "" {
t.Fatal("no csrf")
}
form := strings.NewReader("_csrf=" + csrf + "&username=hub&password=hunter22")
form := strings.NewReader("_csrf=" + csrf + "&username=" + name + "&password=hunter22")
req := httptest.NewRequest(http.MethodPost, "/register", form)
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
for _, c := range cookie {
@@ -95,23 +160,32 @@ func TestRegisterLoginAsk(t *testing.T) {
}
func TestSessionSurvivesServerRestart(t *testing.T) {
fake := newMemDB()
sessionStore := scs.New().Store
url := testDBURL()
if url == "" {
t.Skip("set TEST_DATABASE_URL or DATABASE_URL for web tests")
}
db, sessions, err := store.OpenPostgres(url, plumber.SchemaSQL)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
sessions.Close()
_ = db.Close()
})
sessionStore := sessions.Store()
srv1, err := New(fake, sessionStore, plumber.TemplateFS, plumber.StaticFS, Config{AdminUsername: "hub"})
srv1, err := New(db, sessionStore, plumber.TemplateFS, plumber.StaticFS, Config{})
if err != nil {
t.Fatal(err)
}
h1 := srv1.Handler()
name := uniq("sess")
rec := httptest.NewRecorder()
h1.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/register", nil))
preCookies := rec.Result().Cookies()
csrf := csrfFrom(rec.Body.String())
if csrf == "" {
t.Fatal("no csrf")
}
form := strings.NewReader("_csrf=" + csrf + "&username=hub&password=hunter22")
form := strings.NewReader("_csrf=" + csrf + "&username=" + name + "&password=hunter22")
req := httptest.NewRequest(http.MethodPost, "/register", form)
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
for _, c := range preCookies {
@@ -124,7 +198,7 @@ func TestSessionSurvivesServerRestart(t *testing.T) {
}
sessionCookies := mergeCookies(preCookies, rec.Result().Cookies())
srv2, err := New(fake, sessionStore, plumber.TemplateFS, plumber.StaticFS, Config{AdminUsername: "hub"})
srv2, err := New(db, sessionStore, plumber.TemplateFS, plumber.StaticFS, Config{})
if err != nil {
t.Fatal(err)
}
@@ -163,34 +237,46 @@ func registerUser(t *testing.T, h http.Handler, username, password string) []*ht
}
func TestAdminSeedOnlyWhenNoAdmins(t *testing.T) {
srv, fake, _ := newTestServer(t)
h := srv.Handler()
registerUser(t, h, "hub", "hunter22")
u, err := fake.UserByUsername(context.Background(), "hub")
if err != nil || !u.Admin() {
t.Fatalf("hub should be first admin: %+v %v", u, err)
}
registerUser(t, h, "hub2", "hunter22")
// Create another account that also matches AdminUsername after an admin exists — use a fresh server config with AdminUsername hub2 after hub exists
srv2, err := New(fake, scs.New().Store, plumber.TemplateFS, plumber.StaticFS, Config{AdminUsername: "lateradmin"})
srv, db := newTestServer(t, Config{})
n, err := store.CountAdmins(context.Background(), db)
if err != nil {
t.Fatal(err)
}
registerUser(t, srv2.Handler(), "lateradmin", "hunter22")
u2, err := fake.UserByUsername(context.Background(), "lateradmin")
if n > 0 {
t.Skip("admin already exists in database; bootstrap seed not exercised")
}
adminName := uniq("seed")
srv.cfg.AdminUsername = adminName
h := srv.Handler()
registerUser(t, h, adminName, "hunter22")
u, err := store.UserByUsername(context.Background(), db, adminName)
if err != nil || !u.Admin() {
t.Fatalf("first matching registrant should be admin: %+v %v", u, err)
}
later := uniq("later")
srv2, err := New(db, scs.New().Store, plumber.TemplateFS, plumber.StaticFS, Config{AdminUsername: later})
if err != nil {
t.Fatal(err)
}
registerUser(t, srv2.Handler(), later, "hunter22")
u2, err := store.UserByUsername(context.Background(), db, later)
if err != nil {
t.Fatal(err)
}
if u2.Admin() {
t.Fatal("lateradmin must stay user when an admin already exists")
t.Fatal("later admin username must stay user when an admin already exists")
}
}
func TestAdminUsersPageAccessAndRoles(t *testing.T) {
srv, fake, _ := newTestServer(t)
srv, db := newTestServer(t, Config{})
h := srv.Handler()
adminCookies := registerUser(t, h, "hub", "hunter22")
registerUser(t, h, "bob", "hunter22")
hubName := uniq("hub")
bobName := uniq("bob")
carolName := uniq("carol")
seedUser(t, db, 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)
@@ -201,11 +287,11 @@ func TestAdminUsersPageAccessAndRoles(t *testing.T) {
if rec.Code != 200 {
t.Fatalf("admin list %d", rec.Code)
}
if !strings.Contains(rec.Body.String(), "bob") {
if !strings.Contains(rec.Body.String(), bobName) {
t.Fatal("missing bob on admin page")
}
bob, err := fake.UserByUsername(context.Background(), "bob")
bob, err := store.UserByUsername(context.Background(), db, bobName)
if err != nil {
t.Fatal(err)
}
@@ -221,13 +307,12 @@ func TestAdminUsersPageAccessAndRoles(t *testing.T) {
if rec.Code != http.StatusSeeOther {
t.Fatalf("promote %d %s", rec.Code, rec.Body.String())
}
bob, _ = fake.UserByUsername(context.Background(), "bob")
bob, _ = store.UserByUsername(context.Background(), db, bobName)
if !bob.Admin() {
t.Fatal("bob should be admin")
}
// Non-admin forbidden
bobCookies := registerUser(t, h, "carol", "hunter22")
bobCookies := registerUser(t, h, carolName, "hunter22")
rec = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodGet, "/admin/users", nil)
for _, c := range bobCookies {
@@ -238,11 +323,7 @@ func TestAdminUsersPageAccessAndRoles(t *testing.T) {
t.Fatalf("non-admin expected 403, got %d", rec.Code)
}
// Demote last remaining admin after demoting bob first — leave only hub, then demote hub
hub, err := fake.UserByUsername(context.Background(), "hub")
if err != nil {
t.Fatal(err)
}
// Demote bob back to user
rec = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodGet, "/admin/users", nil)
for _, c := range adminCookies {
@@ -262,6 +343,18 @@ func TestAdminUsersPageAccessAndRoles(t *testing.T) {
t.Fatalf("demote bob %d", rec.Code)
}
admins, err := store.CountAdmins(context.Background(), db)
if err != nil {
t.Fatal(err)
}
if admins != 1 {
t.Skip("shared database has other admins; last-admin demote not isolated")
}
hub, err := store.UserByUsername(context.Background(), db, hubName)
if err != nil {
t.Fatal(err)
}
rec = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodGet, "/admin/users", nil)
for _, c := range adminCookies {
@@ -283,7 +376,7 @@ func TestAdminUsersPageAccessAndRoles(t *testing.T) {
if !strings.Contains(rec.Body.String(), "Cannot demote the last admin") {
t.Fatalf("missing last-admin error: %s", rec.Body.String())
}
hub, _ = fake.UserByUsername(context.Background(), "hub")
hub, _ = store.UserByUsername(context.Background(), db, hubName)
if !hub.Admin() {
t.Fatal("hub must remain admin")
}
@@ -303,9 +396,10 @@ func (f *fakeBlob) Upload(_ context.Context, obj blob.FileUpload) (string, error
}
func TestProfilePageAndState(t *testing.T) {
srv, fake, _ := newTestServer(t)
srv, db := newTestServer(t, Config{})
h := srv.Handler()
cookies := registerUser(t, h, "alice", "hunter22")
name := uniq("alice")
cookies := registerUser(t, h, name, "hunter22")
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/profile", nil)
@@ -339,12 +433,11 @@ func TestProfilePageAndState(t *testing.T) {
if rec.Code != http.StatusSeeOther {
t.Fatalf("save profile %d %s", rec.Code, rec.Body.String())
}
u, err := fake.UserByUsername(context.Background(), "alice")
u, err := store.UserByUsername(context.Background(), db, name)
if err != nil || u.State != "CA" {
t.Fatalf("state not saved: %+v %v", u, err)
}
// invalid state
rec = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodGet, "/profile", nil)
for _, c := range cookies {
@@ -370,27 +463,28 @@ func TestProfilePageAndState(t *testing.T) {
}
func TestProfileAdminAnsweredListAndAvatarUpload(t *testing.T) {
fake := newMemDB()
blob := &fakeBlob{}
sessions := scs.New()
srv, err := New(fake, sessions.Store, plumber.TemplateFS, plumber.StaticFS, Config{
AdminUsername: "hub",
Blob: blob,
})
if err != nil {
t.Fatal(err)
}
fb := &fakeBlob{}
srv, db := newTestServer(t, Config{Blob: fb})
h := srv.Handler()
adminCookies := registerUser(t, h, "hub", "hunter22")
userCookies := registerUser(t, h, "alice", "hunter22")
hubName := uniq("hub")
aliceName := uniq("alice")
hub := seedUser(t, db, hubName, "hunter22", store.RoleAdmin)
alice := seedUser(t, db, aliceName, "hunter22", store.RoleUser)
adminCookies := loginUser(t, h, hubName, "hunter22")
alice, _ := fake.UserByUsername(context.Background(), "alice")
hub, _ := fake.UserByUsername(context.Background(), "hub")
q, err := fake.CreateQuestion(context.Background(), alice.ID, "Drip", "Under sink", "Oakland")
if err != nil {
q := store.NewQuestion(db)
q.AuthorID = alice.ID
q.Title = "Drip"
q.Body = "Under sink"
q.City = "Oakland"
if err := q.Create(context.Background()); err != nil {
t.Fatal(err)
}
if err := fake.UpsertAnswer(context.Background(), q.ID, hub.ID, "Replace the cartridge."); err != nil {
ans := store.NewAnswer(db)
ans.QuestionID = q.ID
ans.AuthorID = hub.ID
ans.Body = "Replace the cartridge."
if err := ans.Upsert(context.Background()); err != nil {
t.Fatal(err)
}
@@ -429,14 +523,13 @@ func TestProfileAdminAnsweredListAndAvatarUpload(t *testing.T) {
if rec.Code != http.StatusSeeOther {
t.Fatalf("avatar upload %d %s", rec.Code, rec.Body.String())
}
if blob.calls != 1 {
t.Fatalf("expected 1 upload, got %d", blob.calls)
if fb.calls != 1 {
t.Fatalf("expected 1 upload, got %d", fb.calls)
}
hub, _ = fake.UserByUsername(context.Background(), "hub")
hub, _ = store.UserByUsername(context.Background(), db, hubName)
if !strings.Contains(hub.AvatarURL, "cdn.example.com/avatars/") {
t.Fatalf("avatar url %q", hub.AvatarURL)
}
_ = userCookies
}
func mergeCookies(sets ...[]*http.Cookie) []*http.Cookie {