package db import ( "database/sql" "os" "path/filepath" "testing" ) func TestInitialize_CreatesDirectory(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "subdir", "test.db") db, err := Initialize(dbPath) if err != nil { t.Fatalf("expected no error, got: %v", err) } defer db.Close() // Verify directory was created if _, err := os.Stat(filepath.Dir(dbPath)); os.IsNotExist(err) { t.Fatal("expected directory to be created") } // Verify database file was created if _, err := os.Stat(dbPath); os.IsNotExist(err) { t.Fatal("expected database file to be created") } } func TestInitialize_OpensConnection(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "test.db") db, err := Initialize(dbPath) if err != nil { t.Fatalf("expected no error, got: %v", err) } defer db.Close() // Verify connection is alive if err := db.DB.Ping(); err != nil { t.Fatalf("expected ping to succeed, got: %v", err) } } func TestInitialize_ExistingFile(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "existing.db") // Create database once db1, err := Initialize(dbPath) if err != nil { t.Fatalf("first init failed: %v", err) } db1.Close() // Re-open existing database db2, err := Initialize(dbPath) if err != nil { t.Fatalf("second init failed: %v", err) } defer db2.Close() if err := db2.DB.Ping(); err != nil { t.Fatalf("expected ping to succeed, got: %v", err) } } func TestMigrate_CreatesTables(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "migrate-test.db") database, err := Initialize(dbPath) if err != nil { t.Fatalf("init failed: %v", err) } defer database.Close() if err := database.Migrate(); err != nil { t.Fatalf("migrate failed: %v", err) } // Verify tables exist expectedTables := []string{"users", "sessions", "audit_logs"} for _, table := range expectedTables { var count int row := database.DB.QueryRow( "SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name=?", table, ) if err := row.Scan(&count); err != nil { t.Fatalf("failed to check table %s: %v", table, err) } if count == 0 { t.Errorf("expected table %s to exist", table) } } } func TestMigrate_Idempotent(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "idempotent-test.db") database, err := Initialize(dbPath) if err != nil { t.Fatalf("init failed: %v", err) } defer database.Close() // Run migrations twice if err := database.Migrate(); err != nil { t.Fatalf("first migrate failed: %v", err) } if err := database.Migrate(); err != nil { t.Fatalf("second migrate should succeed (idempotent), got: %v", err) } } func TestMigrate_TableSchemas(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "schema-test.db") database, err := Initialize(dbPath) if err != nil { t.Fatalf("init failed: %v", err) } defer database.Close() database.Migrate() // Verify sessions table columns rows, err := database.DB.Query("PRAGMA table_info(sessions)") if err != nil { t.Fatalf("failed to get sessions schema: %v", err) } defer rows.Close() columns := map[string]bool{} for rows.Next() { var cid int var name, ctype string var notnull, pk int var dflt sql.NullString if err := rows.Scan(&cid, &name, &ctype, ¬null, &dflt, &pk); err != nil { t.Fatalf("failed to scan column: %v", err) } columns[name] = true _ = ctype } expectedCols := []string{"id", "user_id", "token_hash", "created_at", "expires_at"} for _, col := range expectedCols { if !columns[col] { t.Errorf("expected column %q in sessions table", col) } } } func TestMigrate_UsersTableSchema(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "users-schema-test.db") database, err := Initialize(dbPath) if err != nil { t.Fatalf("init failed: %v", err) } defer database.Close() database.Migrate() // Verify users table columns rows, err := database.DB.Query("PRAGMA table_info(users)") if err != nil { t.Fatalf("failed to get users schema: %v", err) } defer rows.Close() columns := map[string]string{} for rows.Next() { var cid int var name, ctype string var notnull, pk int var dflt sql.NullString if err := rows.Scan(&cid, &name, &ctype, ¬null, &dflt, &pk); err != nil { t.Fatalf("failed to scan column: %v", err) } columns[name] = ctype } expectedCols := []string{"id", "username", "display_name", "email", "groups", "password_hash", "disabled", "created_at", "updated_at"} for _, col := range expectedCols { if _, ok := columns[col]; !ok { t.Errorf("expected column %q in users table", col) } } // Verify username has UNIQUE constraint (SQLite creates an index for UNIQUE columns) var indexCount int database.DB.QueryRow("SELECT COUNT(*) FROM sqlite_master WHERE type='index' AND name LIKE 'sqlite_autoindex_users%' AND sql IS NULL").Scan(&indexCount) if indexCount == 0 { t.Error("expected UNIQUE constraint on username column") } // Spot-check specific types if columns["username"] != "TEXT" { t.Errorf("expected username type TEXT, got %s", columns["username"]) } if columns["disabled"] != "INTEGER" { t.Errorf("expected disabled type INTEGER, got %s", columns["disabled"]) } } func TestClose(t *testing.T) { tmpDir := t.TempDir() dbPath := filepath.Join(tmpDir, "close-test.db") database, err := Initialize(dbPath) if err != nil { t.Fatalf("init failed: %v", err) } if err := database.Close(); err != nil { t.Fatalf("close failed: %v", err) } // Ping should fail after close if err := database.DB.Ping(); err == nil { t.Fatal("expected ping to fail after close") } }