bc57e9101744c6393d52f1a60c71c206d29347dd / internal/db/users_sessions_test.go · 4646 bytes · raw
package db
import (
"database/sql"
"errors"
"path/filepath"
"testing"
"time"
)
func openTestDB(t *testing.T) *sql.DB {
t.Helper()
database, err := Open(filepath.Join(t.TempDir(), "test.db"))
if err != nil {
t.Fatalf("Open: %v", err)
}
t.Cleanup(func() { database.Close() })
return database
}
func TestGetUserByName(t *testing.T) {
database := openTestDB(t)
id, err := CreateUser(database, "josie", "hash")
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
user, err := GetUserByName(database, "josie")
if err != nil {
t.Fatalf("GetUserByName: %v", err)
}
if user.ID != id || user.PasswordHash != "hash" {
t.Errorf("user = %+v, want id %d and stored hash", user, id)
}
if _, err := GetUserByName(database, "nobody"); !errors.Is(err, ErrNotFound) {
t.Errorf("missing user err = %v, want ErrNotFound", err)
}
if _, err := CreateUser(database, "josie", "other"); err == nil {
t.Error("duplicate username accepted")
}
}
func TestGetUserByID(t *testing.T) {
database := openTestDB(t)
id, err := CreateUser(database, "josie", "hash")
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
user, err := GetUserByID(database, id)
if err != nil {
t.Fatalf("GetUserByID: %v", err)
}
if user.Username != "josie" {
t.Errorf("username = %q, want josie", user.Username)
}
if _, err := GetUserByID(database, 999); !errors.Is(err, ErrNotFound) {
t.Errorf("missing user err = %v, want ErrNotFound", err)
}
}
func TestSessionLifecycle(t *testing.T) {
database := openTestDB(t)
id, err := CreateUser(database, "josie", "hash")
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
if err := CreateSession(database, "tok", id, time.Now().Add(time.Hour).Unix()); err != nil {
t.Fatalf("CreateSession: %v", err)
}
session, err := GetSession(database, "tok")
if err != nil {
t.Fatalf("GetSession: %v", err)
}
if session.UserID != id {
t.Errorf("session.UserID = %d, want %d", session.UserID, id)
}
if err := DeleteSession(database, "tok"); err != nil {
t.Fatalf("DeleteSession: %v", err)
}
if _, err := GetSession(database, "tok"); !errors.Is(err, ErrNotFound) {
t.Errorf("deleted session err = %v, want ErrNotFound", err)
}
}
func TestExpiredSessionNotReturned(t *testing.T) {
database := openTestDB(t)
id, err := CreateUser(database, "josie", "hash")
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
if err := CreateSession(database, "old", id, time.Now().Add(-time.Hour).Unix()); err != nil {
t.Fatalf("CreateSession: %v", err)
}
if _, err := GetSession(database, "old"); !errors.Is(err, ErrNotFound) {
t.Errorf("expired session err = %v, want ErrNotFound", err)
}
}
func TestSessionsCascadeWithUser(t *testing.T) {
database := openTestDB(t)
id, err := CreateUser(database, "josie", "hash")
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
if err := CreateSession(database, "tok", id, time.Now().Add(time.Hour).Unix()); err != nil {
t.Fatalf("CreateSession: %v", err)
}
if _, err := database.Exec(`DELETE FROM users WHERE id = ?`, id); err != nil {
t.Fatalf("delete user: %v", err)
}
if _, err := GetSession(database, "tok"); !errors.Is(err, ErrNotFound) {
t.Errorf("orphaned session err = %v, want ErrNotFound", err)
}
}
func TestUpdateUserPassword(t *testing.T) {
database := openTestDB(t)
id, err := CreateUser(database, "josie", "old-hash")
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
if err := UpdateUserPassword(database, id, "new-hash"); err != nil {
t.Fatalf("UpdateUserPassword: %v", err)
}
user, err := GetUserByID(database, id)
if err != nil {
t.Fatalf("GetUserByID: %v", err)
}
if user.PasswordHash != "new-hash" {
t.Errorf("PasswordHash = %q, want new-hash", user.PasswordHash)
}
if err := UpdateUserPassword(database, id+1, "x"); !errors.Is(err, ErrNotFound) {
t.Errorf("missing user err = %v, want ErrNotFound", err)
}
}
func TestDeleteOtherSessions(t *testing.T) {
database := openTestDB(t)
id, err := CreateUser(database, "josie", "hash")
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
for _, token := range []string{"keep", "drop-1", "drop-2"} {
if err := CreateSession(database, token, id, 9999999999); err != nil {
t.Fatalf("CreateSession %s: %v", token, err)
}
}
if err := DeleteOtherSessions(database, id, "keep"); err != nil {
t.Fatalf("DeleteOtherSessions: %v", err)
}
if _, err := GetSession(database, "keep"); err != nil {
t.Errorf("kept session gone: %v", err)
}
for _, token := range []string{"drop-1", "drop-2"} {
if _, err := GetSession(database, token); !errors.Is(err, ErrNotFound) {
t.Errorf("session %s err = %v, want ErrNotFound", token, err)
}
}
}