From 79615dd529d6f627df9e2a390dab5f601041716c Mon Sep 17 00:00:00 2001 From: Jason Ross Date: Thu, 11 Jun 2026 08:19:31 -0500 Subject: [PATCH] ensure zsh is installed first --- compat.go | 3 + helper_test.go | 4 +- main.go | 160 +++++++++++++++++++++++++++++++- post.go | 22 +++++ progress_test.go | 232 +++++++++++++++++++++++++++++++++++++++++++++++ testmain_test.go | 1 + 6 files changed, 420 insertions(+), 2 deletions(-) create mode 100644 progress_test.go diff --git a/compat.go b/compat.go index 7ce0f10..7474cf4 100644 --- a/compat.go +++ b/compat.go @@ -31,4 +31,7 @@ var ( // Filesystem paths osReleasePath = "/etc/os-release" passwdPath = "/etc/passwd" + + // Testing override + disableProgressTracking = false ) diff --git a/helper_test.go b/helper_test.go index 857f22f..5478f0b 100644 --- a/helper_test.go +++ b/helper_test.go @@ -23,13 +23,13 @@ func resetMocks() { osRemove = os.Remove osRemoveAll = os.RemoveAll osExit = os.Exit + progressConfigPath = progressConfigPathReal // Default to an empty reader so un-mocked tests don't hang waiting for user input. stdin = strings.NewReader("") readPassword = func() ([]byte, error) { return nil, errors.New("terminal blocked in test") } - // Reset global state variables to safe defaults isMacOS = false pkgMgr = "dnf" isRHELFamily = true @@ -37,6 +37,8 @@ func resetMocks() { osName = "linux" archName = "x86_64" + disableProgressTracking = true + osReleasePath = "/etc/os-release" passwdPath = "/etc/passwd" diff --git a/main.go b/main.go index 94e7ea3..d14f44e 100644 --- a/main.go +++ b/main.go @@ -25,6 +25,7 @@ import ( "fmt" "os" "path/filepath" + "strconv" "strings" "time" @@ -106,6 +107,55 @@ func runMain(args []string) { doFlatpak := (*only == "" || *only == "flatpak") && *gui && !isMacOS + step := 0 + if !disableProgressTracking { + step = readProgressStep() + if step >= 2 { + fmt.Println("\nBootstrap is already completed according to progress.config.") + fmt.Println("If you want to re-run, delete or reset progress.config.") + return + } + } + + originalSystemPkgs := systemPkgs + originalCustomPtrs := customPtrs + originalDoFlatpak := doFlatpak + originalFlatpakPkgs := flatpakPkgs + + if !disableProgressTracking && step == 0 { + // Minimum required system packages: zsh, git, curl + step1Sys := map[string]bool{"zsh": true, "git": true, "curl": true} + var keptSys []string + for _, p := range systemPkgs { + if step1Sys[p] { + keptSys = append(keptSys, p) + } + } + systemPkgs = keptSys + + // Minimum required custom packages: oh-my-zsh + var keptCust []*CustomPackage + for _, p := range customPtrs { + if strings.ToLower(p.Name) == "oh-my-zsh" { + keptCust = append(keptCust, p) + } + } + customPtrs = keptCust + + // No flatpaks in Step 1 + flatpakPkgs = nil + doFlatpak = false + } else if !disableProgressTracking && step == 1 { + // Exclude oh-my-zsh in Step 2 + var keptCust []*CustomPackage + for _, p := range customPtrs { + if strings.ToLower(p.Name) != "oh-my-zsh" { + keptCust = append(keptCust, p) + } + } + customPtrs = keptCust + } + sysCheck, flatCheck, custCheck := checkAllInParallel( *only == "" || *only == "system", systemPkgs, doFlatpak, flatpakPkgs, @@ -114,7 +164,52 @@ func runMain(args []string) { total := printCheckSummary(sysCheck, flatCheck, custCheck, *only) + if !disableProgressTracking && step == 0 { + // If we didn't need to install anything for step 1, and default shell is already zsh, + // we can skip step 1 and proceed to step 2 in the same run. + if total == 0 && isZshDefault() { + fmt.Println("\n[zsh/oh-my-zsh] Zsh and oh-my-zsh already installed and default shell is zsh. Proceeding to Step 2...") + if err := writeProgressStep(1); err != nil { + fmt.Fprintf(os.Stderr, "Failed to write progress: %v\n", err) + } + step = 1 + // Restore original packages for Step 2 + systemPkgs = originalSystemPkgs + flatpakPkgs = originalFlatpakPkgs + doFlatpak = originalDoFlatpak + + // Exclude oh-my-zsh + var keptCust []*CustomPackage + for _, p := range originalCustomPtrs { + if strings.ToLower(p.Name) != "oh-my-zsh" { + keptCust = append(keptCust, p) + } + } + customPtrs = keptCust + + // Re-run checks for Step 2 + sysCheck, flatCheck, custCheck = checkAllInParallel( + *only == "" || *only == "system", systemPkgs, + doFlatpak, flatpakPkgs, + *only == "" || *only == "custom", customPtrs, + ) + total = printCheckSummary(sysCheck, flatCheck, custCheck, *only) + } + } + if total == 0 { + if !disableProgressTracking && step == 0 { + // Zsh, git, curl and oh-my-zsh are installed, but default shell is not zsh. + ensureZshDefault() + if err := writeProgressStep(1); err != nil { + fmt.Fprintf(os.Stderr, "Failed to write progress: %v\n", err) + } + fmt.Println("\n[zsh] Default shell has been updated to zsh.") + fmt.Println("IMPORTANT: Please log out of your current session and log back in (or restart your terminal) for the shell change to take effect.") + fmt.Println("Once logged back in, please re-run the bootstrapper to complete the rest of the installation.") + osExit(0) + return + } fmt.Println("\nAll packages already installed.") writeRunLog() return @@ -129,9 +224,30 @@ func runMain(args []string) { checkSudo() promptGitHubToken() + if !disableProgressTracking && step == 0 { + // Run only Step 1: minimum required system packages, default shell, minimum custom packages (oh-my-zsh) + if *only == "" || *only == "system" { + installSystemPackages(sysCheck.toInstallRegular, sysCheck.toInstallSpecial) + ensureZshDefault() + } + if *only == "" || *only == "custom" { + installCustomPackages(custCheck.toInstall) + } + if err := writeProgressStep(1); err != nil { + fmt.Fprintf(os.Stderr, "Failed to write progress: %v\n", err) + } + fmt.Println("\n[zsh/oh-my-zsh] Step 1 of bootstrap completed successfully.") + fmt.Println("IMPORTANT: Please log out of your current session and log back in (or restart your terminal) so that zsh becomes your active shell.") + fmt.Println("Once logged back in, please re-run the bootstrapper to complete the rest of the installation.") + osExit(0) + return + } + if *only == "" || *only == "system" { installSystemPackages(sysCheck.toInstallRegular, sysCheck.toInstallSpecial) - ensureZshDefault() + if disableProgressTracking { + ensureZshDefault() + } } if doFlatpak { @@ -161,6 +277,12 @@ func runMain(args []string) { pyenvWG.Wait() } + if !disableProgressTracking { + if err := writeProgressStep(2); err != nil { + fmt.Fprintf(os.Stderr, "Failed to write progress: %v\n", err) + } + } + writeRunLog() printNotices() fmt.Println("\nDone.") @@ -240,3 +362,39 @@ func checkSudo() { return } } + +var progressConfigPathReal = func() string { + exe, err := os.Executable() + if err != nil { + return "progress.config" + } + return filepath.Join(filepath.Dir(exe), "progress.config") +} + +var progressConfigPath = progressConfigPathReal + +func readProgressStep() int { + path := progressConfigPath() + data, err := osReadFile(path) + if err != nil { + return 0 + } + lines := strings.Split(string(data), "\n") + for _, line := range lines { + line = strings.TrimSpace(line) + if strings.HasPrefix(line, "step=") { + stepStr := strings.TrimPrefix(line, "step=") + if step, err := strconv.Atoi(stepStr); err == nil { + return step + } + } + } + return 0 +} + +func writeProgressStep(step int) error { + path := progressConfigPath() + content := fmt.Sprintf("step=%d\n", step) + return osWriteFile(path, []byte(content), 0o644) +} + diff --git a/post.go b/post.go index c9593e3..e80426a 100644 --- a/post.go +++ b/post.go @@ -387,6 +387,28 @@ func userLoginShell(uid string) string { return "" } +func isZshDefault() bool { + if !hasCmd("zsh") { + return false + } + zshPath := "/bin/zsh" + if r, ok := probe([]string{"which", "zsh"}, 5*time.Second); ok && r.ExitCode == 0 { + if p := strings.TrimSpace(string(r.Stdout)); p != "" { + zshPath = p + } + } + username := invokingUser() + if username == "" { + return false + } + u, err := user.Lookup(username) + if err != nil { + return false + } + current := userLoginShell(u.Uid) + return current == zshPath +} + func cloneNvimConfig() { home, _ := os.UserHomeDir() configDir := filepath.Join(home, ".config", "nvim") diff --git a/progress_test.go b/progress_test.go new file mode 100644 index 0000000..9bc4e87 --- /dev/null +++ b/progress_test.go @@ -0,0 +1,232 @@ +package main + +import ( + "fmt" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func TestReadWriteProgressStep(t *testing.T) { + defer resetMocks() + + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, "progress.config") + + // Set up mock file operations + var fileContents []byte + osReadFile = func(name string) ([]byte, error) { + if name == configPath { + if fileContents == nil { + return nil, os.ErrNotExist + } + return fileContents, nil + } + return nil, os.ErrNotExist + } + osWriteFile = func(name string, data []byte, perm os.FileMode) error { + if name == configPath { + fileContents = data + return nil + } + return fmt.Errorf("unexpected file write to %s", name) + } + + // Override progressConfigPath + originalPathFunc := progressConfigPath + progressConfigPath = func() string { + return configPath + } + defer func() { + progressConfigPath = originalPathFunc + }() + + // 1. Initial state (no file) + disableProgressTracking = false + step := readProgressStep() + if step != 0 { + t.Errorf("expected step 0 initially, got %d", step) + } + + // 2. Write step 1 + err := writeProgressStep(1) + if err != nil { + t.Fatalf("failed to write step 1: %v", err) + } + step = readProgressStep() + if step != 1 { + t.Errorf("expected step 1 after writing, got %d", step) + } + + // 3. Write step 2 + err = writeProgressStep(2) + if err != nil { + t.Fatalf("failed to write step 2: %v", err) + } + step = readProgressStep() + if step != 2 { + t.Errorf("expected step 2 after writing, got %d", step) + } +} + +func TestRunMainStep2Exit(t *testing.T) { + defer resetMocks() + + // Mock file read to return step=2 + osReadFile = func(name string) ([]byte, error) { + if strings.HasSuffix(name, "progress.config") { + return []byte("step=2\n"), nil + } + return nil, os.ErrNotExist + } + + disableProgressTracking = false + + var exited bool + osExit = func(code int) { + exited = true + } + + runMain([]string{"bootstrap_environment"}) + + if exited { + t.Error("expected runMain to return normally rather than exit when step=2") + } +} + +func TestRunMainStep0OnlyZshGitCurl(t *testing.T) { + defer resetMocks() + + // Mock file read to return nothing (step=0) + var writtenData []byte + osReadFile = func(name string) ([]byte, error) { + if strings.HasSuffix(name, "progress.config") { + return nil, os.ErrNotExist + } + if strings.HasSuffix(name, "/etc/passwd") { + return []byte(fmt.Sprintf("%s:x:1000:1000::/home/user:/bin/bash\n", os.Getenv("USER"))), nil + } + return nil, os.ErrNotExist + } + osWriteFile = func(name string, data []byte, perm os.FileMode) error { + if strings.HasSuffix(name, "progress.config") { + writtenData = data + return nil + } + return nil + } + + hasCmd = func(name string) bool { + return name == "zsh" || name == "git" || name == "curl" + } + probe = func(argv []string, timeout time.Duration) (CmdResult, bool) { + if argv[0] == "which" && argv[1] == "zsh" { + return CmdResult{ExitCode: 0, Stdout: []byte("/bin/zsh\n")}, true + } + return CmdResult{ExitCode: 1}, true + } + + disableProgressTracking = false + + // Mock user selection: Accept + stdin = strings.NewReader("y\n") + + var runCmdCalls [][]string + runCmd = func(argv []string, opts CmdOpts) CmdResult { + runCmdCalls = append(runCmdCalls, argv) + return CmdResult{ExitCode: 0} + } + + var runShellCalls []string + runShell = func(cmd string, opts CmdOpts) CmdResult { + runShellCalls = append(runShellCalls, cmd) + return CmdResult{ExitCode: 0} + } + + var exited bool + osExit = func(code int) { + exited = true + } + + runMain([]string{"bootstrap_environment"}) + + if !exited { + t.Error("expected runMain to exit at step 1") + } + + // Verify step 1 was written to progress.config + if !strings.Contains(string(writtenData), "step=1") { + t.Errorf("expected step=1 written to progress.config, got: %q", string(writtenData)) + } +} + +func TestRunMainStep0TransitionToStep2(t *testing.T) { + defer resetMocks() + + // Zsh default is true, oh-my-zsh and zsh/git/curl already installed + // Under step=0, we should transition directly to step=1 and then execute step 2 in the same run. + var writtenData []byte + osReadFile = func(name string) ([]byte, error) { + if strings.HasSuffix(name, "progress.config") { + return nil, os.ErrNotExist + } + if strings.HasSuffix(name, "/etc/passwd") { + // passwd already says zsh is default shell + return []byte(fmt.Sprintf("%s:x:1000:1000::/home/user:/bin/zsh\n", os.Getenv("USER"))), nil + } + return nil, os.ErrNotExist + } + osWriteFile = func(name string, data []byte, perm os.FileMode) error { + if strings.HasSuffix(name, "progress.config") { + writtenData = data + return nil + } + return nil + } + + hasCmd = func(name string) bool { + return true // all installed + } + probe = func(argv []string, timeout time.Duration) (CmdResult, bool) { + if argv[0] == "which" && argv[1] == "zsh" { + return CmdResult{ExitCode: 0, Stdout: []byte("/bin/zsh\n")}, true + } + if len(argv) >= 3 && argv[1] == "install" && argv[2] == "--list" { + return CmdResult{ExitCode: 0, Stdout: []byte(" 3.10.0\n 3.11.0\n")}, true + } + return CmdResult{ExitCode: 0}, true + } + + disableProgressTracking = false + + // Mock user selection: Accept + stdin = strings.NewReader("y\n") + + runCmd = func(argv []string, opts CmdOpts) CmdResult { + return CmdResult{ExitCode: 0} + } + runShell = func(cmd string, opts CmdOpts) CmdResult { + return CmdResult{ExitCode: 0} + } + + var exited bool + osExit = func(code int) { + exited = true + } + + runMain([]string{"bootstrap_environment", "--only", "custom"}) + + // Since they are already installed, it will continue. + // Since we mock all installed, total items to install for Step 2 is 0. + // It should exit normally or print all packages installed. + if exited { + t.Errorf("did not expect exit during transition path. Issues: %v", issues) + } + + // Should have written step 1, then eventually step 2 + if !strings.Contains(string(writtenData), "step=2") { + t.Errorf("expected step=2 written to progress.config, got: %q", string(writtenData)) + } +} diff --git a/testmain_test.go b/testmain_test.go index 4145616..879daf3 100644 --- a/testmain_test.go +++ b/testmain_test.go @@ -12,5 +12,6 @@ import ( // (GitHub Actions auto-annotates "[ERROR]" lines as workflow errors). func TestMain(m *testing.M) { issueLogWriter = io.Discard + disableProgressTracking = true os.Exit(m.Run()) }