package api import ( "context" "errors" "net/http" "net/http/httptest" "strings" "testing" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" "github.com/jackc/pgx/v5/pgxpool" "git.fabledsword.com/bvandeusen/minstrel/internal/auth" "git.fabledsword.com/bvandeusen/minstrel/internal/db/dbq" ) func sessionExists(t *testing.T, pool *pgxpool.Pool, id pgtype.UUID) bool { t.Helper() var exists bool if err := pool.QueryRow(context.Background(), `SELECT EXISTS (SELECT 1 FROM sessions WHERE id = $1)`, id, ).Scan(&exists); err != nil { t.Fatalf("exists check: %v", err) } return exists } // Sessions expire server-side on two limits, and an expired one stops // authenticating at once rather than when the GC sweep gets to it. func TestSessionExpiry_IdleAndAbsoluteLimits(t *testing.T) { _, pool := testHandlers(t) user := seedUser(t, pool, "alice", "hunter2", false) q := dbq.New(pool) mint := func(setup string) []byte { token, err := auth.MintSessionToken() if err != nil { t.Fatalf("mint: %v", err) } hash := auth.HashSessionToken(token) sess, err := q.InsertSession(context.Background(), dbq.InsertSessionParams{ UserID: user.ID, TokenHash: hash, UserAgent: "test", Ip: "192.0.2.1", }) if err != nil { t.Fatalf("insert: %v", err) } if setup != "" { if _, err := pool.Exec(context.Background(), `UPDATE sessions SET `+setup+` WHERE id = $1`, sess.ID); err != nil { t.Fatalf("age session: %v", err) } } return hash } fresh := mint("") idle := mint(`last_seen_at = now() - interval '31 days'`) old := mint(`created_at = now() - interval '366 days'`) if _, err := q.GetSessionByTokenHash(context.Background(), fresh); err != nil { t.Errorf("fresh session: %v, want found", err) } for name, hash := range map[string][]byte{"idle": idle, "absolute": old} { if _, err := q.GetSessionByTokenHash(context.Background(), hash); !errors.Is(err, pgx.ErrNoRows) { t.Errorf("%s-expired session: err = %v, want ErrNoRows", name, err) } } listed, err := q.ListSessionsForUser(context.Background(), user.ID) if err != nil { t.Fatalf("list: %v", err) } if len(listed) != 1 { t.Errorf("listed %d sessions, want only the live one", len(listed)) } swept, err := q.GcDeleteExpiredSessions(context.Background()) if err != nil { t.Fatalf("gc: %v", err) } if swept != 2 { t.Errorf("gc swept %d, want the 2 expired rows", swept) } } // Changing your password signs out every other device and keeps this one. func TestHandleChangePassword_RevokesOtherSessions(t *testing.T) { h, pool := testHandlers(t) user := seedUser(t, pool, "alice", "hunter2", false) current := seedSession(t, pool, user.ID, "192.0.2.1") other := seedSession(t, pool, user.ID, "198.51.100.7") body := strings.NewReader(`{"current_password":"hunter2","new_password":"correct-horse"}`) req := httptest.NewRequest(http.MethodPut, "/api/me/password", body) req.Header.Set("Content-Type", "application/json") req = withSession(req, user, current) w := httptest.NewRecorder() h.handleChangePassword(w, req) if w.Code != http.StatusNoContent { t.Fatalf("status = %d body = %s", w.Code, w.Body.String()) } if !sessionExists(t, pool, current) { t.Error("the device that changed the password was signed out") } if sessionExists(t, pool, other) { t.Error("another device's session survived the password change") } } // A reset by email ends every session the account has. func TestHandleResetPassword_RevokesAllSessions(t *testing.T) { h, pool := testHandlers(t) user := seedUser(t, pool, "alice", "hunter2", false) s1 := seedSession(t, pool, user.ID, "192.0.2.1") s2 := seedSession(t, pool, user.ID, "198.51.100.7") const token = "reset-token-for-session-lifecycle-test" if _, err := pool.Exec(context.Background(), `INSERT INTO password_resets (token, user_id, expires_at) VALUES ($1, $2, now() + interval '1 hour')`, token, user.ID, ); err != nil { t.Fatalf("seed reset: %v", err) } body := strings.NewReader(`{"token":"` + token + `","new_password":"correct-horse"}`) req := httptest.NewRequest(http.MethodPost, "/api/auth/reset-password", body) req.Header.Set("Content-Type", "application/json") w := httptest.NewRecorder() h.handleResetPassword(w, req) if w.Code != http.StatusNoContent { t.Fatalf("status = %d body = %s", w.Code, w.Body.String()) } if sessionExists(t, pool, s1) || sessionExists(t, pool, s2) { t.Error("a session survived the password reset") } }