diff --git a/cmd/server/main.go b/cmd/server/main.go index b09be88..aed2530 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -2,6 +2,7 @@ package main import ( "context" + "database/sql" "errors" "log" "net/http" @@ -22,29 +23,30 @@ import ( func main() { _ = godotenv.Load() - st := openStore() - defer st.Close() + db, sessions := openDB() + defer db.Close() + defer sessions.Close() uploader := blob.FromEnv() - handler := newHandler(st, uploader) + handler := newHandler(db, sessions, uploader) run(&http.Server{Addr: listenAddr(), Handler: handler}) } -func openStore() *store.Store { +func openDB() (*sql.DB, *store.SessionStore) { databaseURL := strings.TrimSpace(os.Getenv("DATABASE_URL")) if databaseURL == "" { log.Fatal("DATABASE_URL is required") } - st, err := store.OpenPostgres(databaseURL, plumber.SchemaSQL) + db, sessions, err := store.OpenPostgres(databaseURL, plumber.SchemaSQL) if err != nil { log.Fatalf("database: %v", err) } log.Printf("database: postgres") - return st + return db, sessions } -func newHandler(st *store.Store, uploader blob.Uploader) http.Handler { - srv, err := web.New(st, st.SessionStore(), plumber.TemplateFS, plumber.StaticFS, web.Config{ +func newHandler(db *sql.DB, sessions *store.SessionStore, uploader blob.Uploader) http.Handler { + srv, err := web.New(db, sessions.Store(), plumber.TemplateFS, plumber.StaticFS, web.Config{ AdminUsername: os.Getenv("ADMIN_USERNAME"), SecureCookie: os.Getenv("SECURE_COOKIE") == "1", Blob: uploader, diff --git a/internal/store/answer.go b/internal/store/answer.go new file mode 100644 index 0000000..571186c --- /dev/null +++ b/internal/store/answer.go @@ -0,0 +1,57 @@ +package store + +import ( + "context" + "database/sql" + "fmt" + "strings" + "time" +) + +// Answer is an admin reply to a question. +type Answer struct { + QuestionID string + AuthorID string + AuthorName string + Body string + CreatedAt string + UpdatedAt string + db *sql.DB +} + +// NewAnswer returns an Answer bound to db. +func NewAnswer(db *sql.DB) *Answer { + return &Answer{db: db} +} + +// Upsert inserts or updates the answer for QuestionID. +func (a *Answer) Upsert(ctx context.Context) error { + if a == nil || a.db == nil { + return fmt.Errorf("answer: no database") + } + a.Body = strings.TrimSpace(a.Body) + now := time.Now().UTC().Format(time.RFC3339) + if a.CreatedAt == "" { + a.CreatedAt = now + } + a.UpdatedAt = now + _, err := a.db.ExecContext(ctx, ` +INSERT INTO answers (question_id, author_id, body, created_at, updated_at) VALUES ($1, $2, $3, $4, $5) +ON CONFLICT (question_id) DO UPDATE SET body = excluded.body, author_id = excluded.author_id, updated_at = excluded.updated_at`, + a.QuestionID, a.AuthorID, a.Body, a.CreatedAt, a.UpdatedAt) + return err +} + +func GetAnswer(ctx context.Context, db *sql.DB, questionID string) (*Answer, error) { + var a Answer + err := db.QueryRowContext(ctx, ` +SELECT a.question_id, a.author_id, u.name, a.body, a.created_at, a.updated_at +FROM answers a +JOIN users u ON u.id = a.author_id +WHERE a.question_id = $1`, questionID).Scan(&a.QuestionID, &a.AuthorID, &a.AuthorName, &a.Body, &a.CreatedAt, &a.UpdatedAt) + if err != nil { + return nil, err + } + a.db = db + return &a, nil +} diff --git a/internal/store/db.go b/internal/store/db.go deleted file mode 100644 index 45514bc..0000000 --- a/internal/store/db.go +++ /dev/null @@ -1,48 +0,0 @@ -package store - -import ( - "context" - "errors" -) - -// ErrLastAdmin is returned when demoting the only remaining admin. -var ErrLastAdmin = errors.New("cannot demote the last admin") - -// Role is a user privilege level stored in users.role. -type Role string - -const ( - RoleUser Role = "user" - RoleAdmin Role = "admin" -) - -// NewUser is the input for CreateUser. -type NewUser struct { - Username string - PasswordHash string - Role Role -} - -// DB is the persistence API used by the web layer. -// Named DB to avoid colliding with scs.Store. -type DB interface { - CreateUser(ctx context.Context, user NewUser) (*User, error) - UserByID(ctx context.Context, id string) (*User, error) - UserByUsername(ctx context.Context, username string) (*User, error) - CountAdmins(ctx context.Context) (int, error) - ListUsers(ctx context.Context) ([]User, error) - SetRole(ctx context.Context, userID string, role Role) error - CreateQuestion(ctx context.Context, authorID, title, body, city string) (*RankedQuestion, error) - ListHunt(ctx context.Context, huntDate, viewerID string) ([]RankedQuestion, error) - GetQuestion(ctx context.Context, id, viewerID string) (*RankedQuestion, error) - Vote(ctx context.Context, userID, questionID string, value int) error - GetAnswer(ctx context.Context, questionID string) (*Answer, error) - UpsertAnswer(ctx context.Context, questionID, authorID, body string) error - HideQuestion(ctx context.Context, id string) error - UpdateProfile(ctx context.Context, userID, state, avatarURL string) error - ListQuestionsByAuthor(ctx context.Context, authorID string) ([]RankedQuestion, error) - ListQuestionsAnsweredBy(ctx context.Context, adminID string) ([]RankedQuestion, error) -} - -// Compile-time check: *Store implements DB. -var _ DB = (*Store)(nil) diff --git a/internal/store/postgres.go b/internal/store/postgres.go index 44b1069..cd5d24b 100644 --- a/internal/store/postgres.go +++ b/internal/store/postgres.go @@ -48,34 +48,33 @@ func postgresDSN(raw string) (string, error) { } // OpenPostgres connects to Postgres, applies schema/migrations, and starts session cleanup. -func OpenPostgres(databaseURL, schema string) (*Store, error) { +func OpenPostgres(databaseURL, schema string) (*sql.DB, *SessionStore, error) { dsn, err := postgresDSN(databaseURL) if err != nil { - return nil, err + return nil, nil, err } db, err := sql.Open("pgx", dsn) if err != nil { - return nil, err + return nil, nil, err } db.SetMaxOpenConns(20) db.SetMaxIdleConns(5) if err := db.Ping(); err != nil { _ = db.Close() - return nil, fmt.Errorf("postgres ping: %w", err) + return nil, nil, fmt.Errorf("postgres ping: %w", err) } if err := applySchema(db, schema); err != nil { _ = db.Close() - return nil, fmt.Errorf("apply schema: %w", err) + return nil, nil, fmt.Errorf("apply schema: %w", err) } if err := applySessionsSchema(db); err != nil { _ = db.Close() - return nil, fmt.Errorf("apply sessions schema: %w", err) + return nil, nil, fmt.Errorf("apply sessions schema: %w", err) } if err := migrateUserProfileColumns(db); err != nil { _ = db.Close() - return nil, fmt.Errorf("migrate profile columns: %w", err) + return nil, nil, fmt.Errorf("migrate profile columns: %w", err) } - st := &Store{db: db} - st.initSessionStore(5 * time.Minute) - return st, nil + sessions := NewSessionStore(db, 5*time.Minute) + return db, sessions, nil } diff --git a/internal/store/question.go b/internal/store/question.go new file mode 100644 index 0000000..81928a4 --- /dev/null +++ b/internal/store/question.go @@ -0,0 +1,168 @@ +package store + +import ( + "context" + "database/sql" + "fmt" + "strings" + "time" + + "github.com/google/uuid" + + "plumber/internal/pacific" +) + +// RankedQuestion is a question row with score / vote annotations for lists. +type RankedQuestion struct { + ID string + AuthorID string + AuthorName string + Title string + Body string + City string + HuntDate string + Hidden bool + CreatedAt string + Score int + Answered bool + UserVote int + db *sql.DB +} + +// NewQuestion returns a question bound to db (not yet inserted). +func NewQuestion(db *sql.DB) *RankedQuestion { + return &RankedQuestion{db: db} +} + +// Create inserts the question. Sets ID, HuntDate, and CreatedAt when empty. +func (q *RankedQuestion) Create(ctx context.Context) error { + if q == nil || q.db == nil { + return fmt.Errorf("question: no database") + } + q.Title = strings.TrimSpace(q.Title) + q.Body = strings.TrimSpace(q.Body) + q.City = strings.TrimSpace(q.City) + if q.ID == "" { + q.ID = uuid.NewString() + } + if q.HuntDate == "" { + q.HuntDate = pacific.Today() + } + if q.CreatedAt == "" { + q.CreatedAt = time.Now().UTC().Format(time.RFC3339) + } + _, err := q.db.ExecContext(ctx, `INSERT INTO questions (id, author_id, title, body, city, hunt_date, hidden, created_at) VALUES ($1, $2, $3, $4, $5, $6, 0, $7)`, + q.ID, q.AuthorID, q.Title, q.Body, q.City, q.HuntDate, q.CreatedAt) + return err +} + +// Hide marks the question hidden. +func (q *RankedQuestion) Hide(ctx context.Context) error { + if q == nil || q.db == nil { + return fmt.Errorf("question: no database") + } + _, err := q.db.ExecContext(ctx, `UPDATE questions SET hidden = 1 WHERE id = $1`, q.ID) + if err == nil { + q.Hidden = true + } + return err +} + +func ListHunt(ctx context.Context, db *sql.DB, huntDate, viewerID string) ([]RankedQuestion, error) { + rows, err := db.QueryContext(ctx, ` +SELECT q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, + COALESCE(SUM(v.value), 0) AS score, + CASE WHEN a.question_id IS NULL THEN 0 ELSE 1 END AS answered, + COALESCE((SELECT value FROM votes WHERE user_id = $1 AND question_id = q.id), 0) AS user_vote +FROM questions q +JOIN users u ON u.id = q.author_id +LEFT JOIN votes v ON v.question_id = q.id +LEFT JOIN answers a ON a.question_id = q.id +WHERE q.hunt_date = $2 AND q.hidden = 0 +GROUP BY q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, a.question_id +ORDER BY score DESC, q.created_at ASC`, viewerID, huntDate) + if err != nil { + return nil, err + } + defer rows.Close() + return scanRankedList(db, rows) +} + +func GetQuestion(ctx context.Context, db *sql.DB, id, viewerID string) (*RankedQuestion, error) { + row := db.QueryRowContext(ctx, ` +SELECT q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, + COALESCE((SELECT SUM(value) FROM votes WHERE question_id = q.id), 0) AS score, + CASE WHEN a.question_id IS NULL THEN 0 ELSE 1 END AS answered, + COALESCE((SELECT value FROM votes WHERE user_id = $1 AND question_id = q.id), 0) AS user_vote +FROM questions q +JOIN users u ON u.id = q.author_id +LEFT JOIN answers a ON a.question_id = q.id +WHERE q.id = $2`, viewerID, id) + q, err := scanRanked(db, row) + if err != nil { + return nil, err + } + return &q, nil +} + +func ListQuestionsByAuthor(ctx context.Context, db *sql.DB, authorID string) ([]RankedQuestion, error) { + rows, err := db.QueryContext(ctx, ` +SELECT q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, + COALESCE((SELECT SUM(value) FROM votes WHERE question_id = q.id), 0) AS score, + CASE WHEN a.question_id IS NULL THEN 0 ELSE 1 END AS answered, + 0 AS user_vote +FROM questions q +JOIN users u ON u.id = q.author_id +LEFT JOIN answers a ON a.question_id = q.id +WHERE q.author_id = $1 AND q.hidden = 0 +ORDER BY q.created_at DESC`, authorID) + if err != nil { + return nil, err + } + defer rows.Close() + return scanRankedList(db, rows) +} + +func ListQuestionsAnsweredBy(ctx context.Context, db *sql.DB, adminID string) ([]RankedQuestion, error) { + rows, err := db.QueryContext(ctx, ` +SELECT q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, + COALESCE((SELECT SUM(value) FROM votes WHERE question_id = q.id), 0) AS score, + 1 AS answered, + 0 AS user_vote +FROM answers ans +JOIN questions q ON q.id = ans.question_id +JOIN users u ON u.id = q.author_id +WHERE ans.author_id = $1 AND q.hidden = 0 +ORDER BY ans.updated_at DESC`, adminID) + if err != nil { + return nil, err + } + defer rows.Close() + return scanRankedList(db, rows) +} + +type scanned interface { + Scan(dest ...any) error +} + +func scanRanked(db *sql.DB, rows scanned) (RankedQuestion, error) { + var q RankedQuestion + var hidden, answered int + err := rows.Scan(&q.ID, &q.AuthorID, &q.AuthorName, &q.Title, &q.Body, &q.City, &q.HuntDate, &hidden, &q.CreatedAt, &q.Score, &answered, &q.UserVote) + q.Hidden = hidden != 0 + q.Answered = answered != 0 + q.db = db + return q, err +} + +func scanRankedList(db *sql.DB, rows *sql.Rows) ([]RankedQuestion, error) { + var out []RankedQuestion + for rows.Next() { + q, err := scanRanked(db, rows) + if err != nil { + return nil, err + } + out = append(out, q) + } + return out, rows.Err() +} diff --git a/internal/store/sessions.go b/internal/store/sessions.go index 33c507c..4d4edcf 100644 --- a/internal/store/sessions.go +++ b/internal/store/sessions.go @@ -26,13 +26,28 @@ type sessionStopper interface { StopCleanup() } -// SessionStore returns the scs store backed by this database. -func (s *Store) SessionStore() scs.Store { - return s.sessionStore +// SessionStore wraps scs Postgres session persistence and cleanup. +type SessionStore struct { + store scs.Store + stopper sessionStopper } -func (s *Store) initSessionStore(cleanupInterval time.Duration) { - ps := postgresstore.NewWithCleanupInterval(s.db, cleanupInterval) - s.sessionStore = ps - s.sessionStopper = ps +// NewSessionStore starts a postgresstore with the given cleanup interval. +func NewSessionStore(db *sql.DB, cleanupInterval time.Duration) *SessionStore { + ps := postgresstore.NewWithCleanupInterval(db, cleanupInterval) + return &SessionStore{store: ps, stopper: ps} +} + +// Store returns the scs.Store implementation. +func (s *SessionStore) Store() scs.Store { + return s.store +} + +// Close stops background session cleanup. +func (s *SessionStore) Close() { + if s == nil || s.stopper == nil { + return + } + s.stopper.StopCleanup() + s.stopper = nil } diff --git a/internal/store/store.go b/internal/store/store.go deleted file mode 100644 index 9c93af8..0000000 --- a/internal/store/store.go +++ /dev/null @@ -1,371 +0,0 @@ -package store - -import ( - "context" - "database/sql" - "fmt" - "strings" - "time" - - "github.com/alexedwards/scs/v2" - "github.com/google/uuid" - - "plumber/internal/pacific" -) - -type Store struct { - db *sql.DB - sessionStore scs.Store - sessionStopper sessionStopper -} - -type User struct { - ID string - Username string - Name string - Role Role - AvatarURL string - State string - CreatedAt string - PasswordHash string -} - -func (u *User) Admin() bool { - return u != nil && u.Role == RoleAdmin -} - -type RankedQuestion struct { - ID string - AuthorID string - AuthorName string - Title string - Body string - City string - HuntDate string - Hidden bool - CreatedAt string - Score int - Answered bool - UserVote int -} - -type Answer struct { - QuestionID string - AuthorID string - AuthorName string - Body string - CreatedAt string - UpdatedAt string -} - -func (s *Store) Close() error { - if s.sessionStopper != nil { - s.sessionStopper.StopCleanup() - s.sessionStopper = nil - } - return s.db.Close() -} - -func (s *Store) CreateUser(ctx context.Context, nu NewUser) (*User, error) { - if nu.Role != RoleUser && nu.Role != RoleAdmin { - return nil, fmt.Errorf("invalid role") - } - username := NormalizeUsername(nu.Username) - u := &User{ - ID: uuid.NewString(), - Username: username, - Name: username, - Role: nu.Role, - PasswordHash: nu.PasswordHash, - CreatedAt: time.Now().UTC().Format(time.RFC3339), - } - _, err := s.db.ExecContext(ctx, `INSERT INTO users (id, username, name, password_hash, role, avatar_url, state, created_at) VALUES ($1, $2, $3, $4, $5, '', '', $6)`, - u.ID, u.Username, u.Name, u.PasswordHash, string(u.Role), u.CreatedAt) - if err != nil { - return nil, err - } - return u, nil -} - -func (s *Store) CountAdmins(ctx context.Context) (int, error) { - var n int - err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = $1`, string(RoleAdmin)).Scan(&n) - return n, err -} - -func (s *Store) ListUsers(ctx context.Context) ([]User, error) { - rows, err := s.db.QueryContext(ctx, `SELECT id, username, name, role, avatar_url, state, created_at FROM users ORDER BY created_at ASC`) - if err != nil { - return nil, err - } - defer rows.Close() - var out []User - for rows.Next() { - var u User - var role string - if err := rows.Scan(&u.ID, &u.Username, &u.Name, &role, &u.AvatarURL, &u.State, &u.CreatedAt); err != nil { - return nil, err - } - u.Role = Role(role) - out = append(out, u) - } - return out, rows.Err() -} - -func (s *Store) SetRole(ctx context.Context, userID string, role Role) error { - if role != RoleUser && role != RoleAdmin { - return fmt.Errorf("invalid role") - } - tx, err := s.db.BeginTx(ctx, nil) - if err != nil { - return err - } - defer tx.Rollback() - - var current string - err = tx.QueryRowContext(ctx, `SELECT role FROM users WHERE id = $1`, userID).Scan(¤t) - if err != nil { - return err - } - if Role(current) == RoleAdmin && role == RoleUser { - var n int - if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = $1`, string(RoleAdmin)).Scan(&n); err != nil { - return err - } - if n <= 1 { - return ErrLastAdmin - } - } - res, err := tx.ExecContext(ctx, `UPDATE users SET role = $1 WHERE id = $2`, string(role), userID) - if err != nil { - return err - } - aff, err := res.RowsAffected() - if err != nil { - return err - } - if aff == 0 { - return sql.ErrNoRows - } - return tx.Commit() -} - -func (s *Store) UserByID(ctx context.Context, id string) (*User, error) { - return scanUser(s.db.QueryRowContext(ctx, `SELECT id, username, name, role, avatar_url, state, created_at FROM users WHERE id = $1`, id), false) -} - -func (s *Store) UserByUsername(ctx context.Context, username string) (*User, error) { - return scanUser(s.db.QueryRowContext(ctx, `SELECT id, username, name, role, avatar_url, state, created_at, password_hash FROM users WHERE username = $1`, NormalizeUsername(username)), true) -} - -func scanUser(row *sql.Row, withSecrets bool) (*User, error) { - var u User - var role string - var err error - if withSecrets { - err = row.Scan(&u.ID, &u.Username, &u.Name, &role, &u.AvatarURL, &u.State, &u.CreatedAt, &u.PasswordHash) - } else { - err = row.Scan(&u.ID, &u.Username, &u.Name, &role, &u.AvatarURL, &u.State, &u.CreatedAt) - } - if err != nil { - return nil, err - } - u.Role = Role(role) - return &u, nil -} - -func NormalizeUsername(s string) string { - return strings.ToLower(strings.TrimSpace(s)) -} - -func (s *Store) CreateQuestion(ctx context.Context, authorID, title, body, city string) (*RankedQuestion, error) { - q := &RankedQuestion{ - ID: uuid.NewString(), - AuthorID: authorID, - Title: strings.TrimSpace(title), - Body: strings.TrimSpace(body), - City: strings.TrimSpace(city), - HuntDate: pacific.Today(), - CreatedAt: time.Now().UTC().Format(time.RFC3339), - } - _, err := s.db.ExecContext(ctx, `INSERT INTO questions (id, author_id, title, body, city, hunt_date, hidden, created_at) VALUES ($1, $2, $3, $4, $5, $6, 0, $7)`, - q.ID, q.AuthorID, q.Title, q.Body, q.City, q.HuntDate, q.CreatedAt) - if err != nil { - return nil, err - } - return q, nil -} - -func (s *Store) ListHunt(ctx context.Context, huntDate, viewerID string) ([]RankedQuestion, error) { - rows, err := s.db.QueryContext(ctx, ` -SELECT q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, - COALESCE(SUM(v.value), 0) AS score, - CASE WHEN a.question_id IS NULL THEN 0 ELSE 1 END AS answered, - COALESCE((SELECT value FROM votes WHERE user_id = $1 AND question_id = q.id), 0) AS user_vote -FROM questions q -JOIN users u ON u.id = q.author_id -LEFT JOIN votes v ON v.question_id = q.id -LEFT JOIN answers a ON a.question_id = q.id -WHERE q.hunt_date = $2 AND q.hidden = 0 -GROUP BY q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, a.question_id -ORDER BY score DESC, q.created_at ASC`, viewerID, huntDate) - if err != nil { - return nil, err - } - defer rows.Close() - var out []RankedQuestion - for rows.Next() { - q, err := scanRanked(rows) - if err != nil { - return nil, err - } - out = append(out, q) - } - return out, rows.Err() -} - -func (s *Store) GetQuestion(ctx context.Context, id, viewerID string) (*RankedQuestion, error) { - row := s.db.QueryRowContext(ctx, ` -SELECT q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, - COALESCE((SELECT SUM(value) FROM votes WHERE question_id = q.id), 0) AS score, - CASE WHEN a.question_id IS NULL THEN 0 ELSE 1 END AS answered, - COALESCE((SELECT value FROM votes WHERE user_id = $1 AND question_id = q.id), 0) AS user_vote -FROM questions q -JOIN users u ON u.id = q.author_id -LEFT JOIN answers a ON a.question_id = q.id -WHERE q.id = $2`, viewerID, id) - q, err := scanRankedRow(row) - if err != nil { - return nil, err - } - return &q, nil -} - -type scanned interface { - Scan(dest ...any) error -} - -func scanRanked(rows scanned) (RankedQuestion, error) { - var q RankedQuestion - var hidden, answered int - err := rows.Scan(&q.ID, &q.AuthorID, &q.AuthorName, &q.Title, &q.Body, &q.City, &q.HuntDate, &hidden, &q.CreatedAt, &q.Score, &answered, &q.UserVote) - q.Hidden = hidden != 0 - q.Answered = answered != 0 - return q, err -} - -func scanRankedRow(row *sql.Row) (RankedQuestion, error) { - return scanRanked(row) -} - -func (s *Store) Vote(ctx context.Context, userID, questionID string, value int) error { - if value != 1 && value != -1 { - return fmt.Errorf("invalid vote") - } - tx, err := s.db.BeginTx(ctx, nil) - if err != nil { - return err - } - defer tx.Rollback() - var current sql.NullInt64 - err = tx.QueryRowContext(ctx, `SELECT value FROM votes WHERE user_id = $1 AND question_id = $2`, userID, questionID).Scan(¤t) - if err != nil && err != sql.ErrNoRows { - return err - } - if err == nil && current.Valid && int(current.Int64) == value { - _, err = tx.ExecContext(ctx, `DELETE FROM votes WHERE user_id = $1 AND question_id = $2`, userID, questionID) - } else { - _, err = tx.ExecContext(ctx, `INSERT INTO votes (user_id, question_id, value) VALUES ($1, $2, $3) -ON CONFLICT (user_id, question_id) DO UPDATE SET value = excluded.value`, userID, questionID, value) - } - if err != nil { - return err - } - return tx.Commit() -} - -func (s *Store) GetAnswer(ctx context.Context, questionID string) (*Answer, error) { - var a Answer - err := s.db.QueryRowContext(ctx, ` -SELECT a.question_id, a.author_id, u.name, a.body, a.created_at, a.updated_at -FROM answers a -JOIN users u ON u.id = a.author_id -WHERE a.question_id = $1`, questionID).Scan(&a.QuestionID, &a.AuthorID, &a.AuthorName, &a.Body, &a.CreatedAt, &a.UpdatedAt) - if err != nil { - return nil, err - } - return &a, nil -} - -func (s *Store) UpsertAnswer(ctx context.Context, questionID, authorID, body string) error { - body = strings.TrimSpace(body) - now := time.Now().UTC().Format(time.RFC3339) - _, err := s.db.ExecContext(ctx, ` -INSERT INTO answers (question_id, author_id, body, created_at, updated_at) VALUES ($1, $2, $3, $4, $5) -ON CONFLICT (question_id) DO UPDATE SET body = excluded.body, author_id = excluded.author_id, updated_at = excluded.updated_at`, - questionID, authorID, body, now, now) - return err -} - -func (s *Store) HideQuestion(ctx context.Context, id string) error { - _, err := s.db.ExecContext(ctx, `UPDATE questions SET hidden = 1 WHERE id = $1`, id) - return err -} - -func (s *Store) UpdateProfile(ctx context.Context, userID, state, avatarURL string) error { - state = strings.TrimSpace(state) - if avatarURL == "" { - _, err := s.db.ExecContext(ctx, `UPDATE users SET state = $1 WHERE id = $2`, state, userID) - return err - } - _, err := s.db.ExecContext(ctx, `UPDATE users SET state = $1, avatar_url = $2 WHERE id = $3`, state, avatarURL, userID) - return err -} - -func (s *Store) ListQuestionsByAuthor(ctx context.Context, authorID string) ([]RankedQuestion, error) { - rows, err := s.db.QueryContext(ctx, ` -SELECT q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, - COALESCE((SELECT SUM(value) FROM votes WHERE question_id = q.id), 0) AS score, - CASE WHEN a.question_id IS NULL THEN 0 ELSE 1 END AS answered, - 0 AS user_vote -FROM questions q -JOIN users u ON u.id = q.author_id -LEFT JOIN answers a ON a.question_id = q.id -WHERE q.author_id = $1 AND q.hidden = 0 -ORDER BY q.created_at DESC`, authorID) - if err != nil { - return nil, err - } - defer rows.Close() - return scanRankedList(rows) -} - -func (s *Store) ListQuestionsAnsweredBy(ctx context.Context, adminID string) ([]RankedQuestion, error) { - rows, err := s.db.QueryContext(ctx, ` -SELECT q.id, q.author_id, u.name, q.title, q.body, q.city, q.hunt_date, q.hidden, q.created_at, - COALESCE((SELECT SUM(value) FROM votes WHERE question_id = q.id), 0) AS score, - 1 AS answered, - 0 AS user_vote -FROM answers ans -JOIN questions q ON q.id = ans.question_id -JOIN users u ON u.id = q.author_id -WHERE ans.author_id = $1 AND q.hidden = 0 -ORDER BY ans.updated_at DESC`, adminID) - if err != nil { - return nil, err - } - defer rows.Close() - return scanRankedList(rows) -} - -func scanRankedList(rows *sql.Rows) ([]RankedQuestion, error) { - var out []RankedQuestion - for rows.Next() { - q, err := scanRanked(rows) - if err != nil { - return nil, err - } - out = append(out, q) - } - return out, rows.Err() -} diff --git a/internal/store/user.go b/internal/store/user.go new file mode 100644 index 0000000..21e57a3 --- /dev/null +++ b/internal/store/user.go @@ -0,0 +1,183 @@ +package store + +import ( + "context" + "database/sql" + "errors" + "fmt" + "strings" + "time" + + "github.com/google/uuid" +) + +// ErrLastAdmin is returned when demoting the only remaining admin. +var ErrLastAdmin = errors.New("cannot demote the last admin") + +// Role is a user privilege level stored in users.role. +type Role string + +const ( + RoleUser Role = "user" + RoleAdmin Role = "admin" +) + +// User is an account row. Methods run SQL against db. +type User struct { + ID string + Username string + Name string + Role Role + AvatarURL string + State string + CreatedAt string + PasswordHash string + db *sql.DB +} + +// NewUser returns a User bound to db (not yet inserted). +func NewUser(db *sql.DB) *User { + return &User{db: db} +} + +func (u *User) Admin() bool { + return u != nil && u.Role == RoleAdmin +} + +func NormalizeUsername(s string) string { + return strings.ToLower(strings.TrimSpace(s)) +} + +// Create inserts the user. Sets ID, Name, and CreatedAt when empty. +func (u *User) Create(ctx context.Context) error { + if u == nil || u.db == nil { + return fmt.Errorf("user: no database") + } + if u.Role != RoleUser && u.Role != RoleAdmin { + return fmt.Errorf("invalid role") + } + u.Username = NormalizeUsername(u.Username) + if u.ID == "" { + u.ID = uuid.NewString() + } + if u.Name == "" { + u.Name = u.Username + } + if u.CreatedAt == "" { + u.CreatedAt = time.Now().UTC().Format(time.RFC3339) + } + _, err := u.db.ExecContext(ctx, `INSERT INTO users (id, username, name, password_hash, role, avatar_url, state, created_at) VALUES ($1, $2, $3, $4, $5, '', '', $6)`, + u.ID, u.Username, u.Name, u.PasswordHash, string(u.Role), u.CreatedAt) + return err +} + +// SetRole updates this user's role (last-admin safe). +func (u *User) SetRole(ctx context.Context, role Role) error { + if u == nil || u.db == nil { + return fmt.Errorf("user: no database") + } + if role != RoleUser && role != RoleAdmin { + return fmt.Errorf("invalid role") + } + tx, err := u.db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + + var current string + err = tx.QueryRowContext(ctx, `SELECT role FROM users WHERE id = $1`, u.ID).Scan(¤t) + if err != nil { + return err + } + if Role(current) == RoleAdmin && role == RoleUser { + var n int + if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = $1`, string(RoleAdmin)).Scan(&n); err != nil { + return err + } + if n <= 1 { + return ErrLastAdmin + } + } + res, err := tx.ExecContext(ctx, `UPDATE users SET role = $1 WHERE id = $2`, string(role), u.ID) + if err != nil { + return err + } + aff, err := res.RowsAffected() + if err != nil { + return err + } + if aff == 0 { + return sql.ErrNoRows + } + if err := tx.Commit(); err != nil { + return err + } + u.Role = role + return nil +} + +// SaveProfile writes State and optionally AvatarURL. +func (u *User) SaveProfile(ctx context.Context) error { + if u == nil || u.db == nil { + return fmt.Errorf("user: no database") + } + u.State = strings.TrimSpace(u.State) + if u.AvatarURL == "" { + _, err := u.db.ExecContext(ctx, `UPDATE users SET state = $1 WHERE id = $2`, u.State, u.ID) + return err + } + _, err := u.db.ExecContext(ctx, `UPDATE users SET state = $1, avatar_url = $2 WHERE id = $3`, u.State, u.AvatarURL, u.ID) + return err +} + +func CountAdmins(ctx context.Context, db *sql.DB) (int, error) { + var n int + err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = $1`, string(RoleAdmin)).Scan(&n) + return n, err +} + +func ListUsers(ctx context.Context, db *sql.DB) ([]User, error) { + rows, err := db.QueryContext(ctx, `SELECT id, username, name, role, avatar_url, state, created_at FROM users ORDER BY created_at ASC`) + if err != nil { + return nil, err + } + defer rows.Close() + var out []User + for rows.Next() { + var u User + var role string + if err := rows.Scan(&u.ID, &u.Username, &u.Name, &role, &u.AvatarURL, &u.State, &u.CreatedAt); err != nil { + return nil, err + } + u.Role = Role(role) + u.db = db + out = append(out, u) + } + return out, rows.Err() +} + +func UserByID(ctx context.Context, db *sql.DB, id string) (*User, error) { + return scanUser(db, db.QueryRowContext(ctx, `SELECT id, username, name, role, avatar_url, state, created_at FROM users WHERE id = $1`, id), false) +} + +func UserByUsername(ctx context.Context, db *sql.DB, username string) (*User, error) { + return scanUser(db, db.QueryRowContext(ctx, `SELECT id, username, name, role, avatar_url, state, created_at, password_hash FROM users WHERE username = $1`, NormalizeUsername(username)), true) +} + +func scanUser(db *sql.DB, row *sql.Row, withSecrets bool) (*User, error) { + var u User + var role string + var err error + if withSecrets { + err = row.Scan(&u.ID, &u.Username, &u.Name, &role, &u.AvatarURL, &u.State, &u.CreatedAt, &u.PasswordHash) + } else { + err = row.Scan(&u.ID, &u.Username, &u.Name, &role, &u.AvatarURL, &u.State, &u.CreatedAt) + } + if err != nil { + return nil, err + } + u.Role = Role(role) + u.db = db + return &u, nil +} diff --git a/internal/store/vote.go b/internal/store/vote.go new file mode 100644 index 0000000..d3c07b6 --- /dev/null +++ b/internal/store/vote.go @@ -0,0 +1,34 @@ +package store + +import ( + "context" + "database/sql" + "fmt" +) + +// Vote toggles or sets a user's vote on a question (value must be 1 or -1). +func Vote(ctx context.Context, db *sql.DB, userID, questionID string, value int) error { + if value != 1 && value != -1 { + return fmt.Errorf("invalid vote") + } + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + var current sql.NullInt64 + err = tx.QueryRowContext(ctx, `SELECT value FROM votes WHERE user_id = $1 AND question_id = $2`, userID, questionID).Scan(¤t) + if err != nil && err != sql.ErrNoRows { + return err + } + if err == nil && current.Valid && int(current.Int64) == value { + _, err = tx.ExecContext(ctx, `DELETE FROM votes WHERE user_id = $1 AND question_id = $2`, userID, questionID) + } else { + _, err = tx.ExecContext(ctx, `INSERT INTO votes (user_id, question_id, value) VALUES ($1, $2, $3) +ON CONFLICT (user_id, question_id) DO UPDATE SET value = excluded.value`, userID, questionID, value) + } + if err != nil { + return err + } + return tx.Commit() +} diff --git a/internal/web/admin.go b/internal/web/admin.go index e7e9e70..b44c966 100644 --- a/internal/web/admin.go +++ b/internal/web/admin.go @@ -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 diff --git a/internal/web/auth.go b/internal/web/auth.go index 65cc384..5bf63d3 100644 --- a/internal/web/auth.go +++ b/internal/web/auth.go @@ -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 diff --git a/internal/web/memstore_test.go b/internal/web/memstore_test.go deleted file mode 100644 index 8476ce5..0000000 --- a/internal/web/memstore_test.go +++ /dev/null @@ -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) diff --git a/internal/web/profile.go b/internal/web/profile.go index 0ec071d..fb9a750 100644 --- a/internal/web/profile.go +++ b/internal/web/profile.go @@ -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") diff --git a/internal/web/server.go b/internal/web/server.go index 154515c..a6dd92d 100644 --- a/internal/web/server.go +++ b/internal/web/server.go @@ -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 } diff --git a/internal/web/server_test.go b/internal/web/server_test.go index 130157c..d6d9fb7 100644 --- a/internal/web/server_test.go +++ b/internal/web/server_test.go @@ -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 {