josie / simplegit

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