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) } } }