From 11225df6c839b5c287bf31dd5f7cc38085ea2251 Mon Sep 17 00:00:00 2001 From: verse91 Date: Sat, 11 Apr 2026 23:54:11 +0700 Subject: [PATCH] fix(root): eliminate child shell leakage and terminal corruption on reload --- root/root.go | 278 ++++++++++++++++++++++++--------------------------- 1 file changed, 128 insertions(+), 150 deletions(-) diff --git a/root/root.go b/root/root.go index 84a10aa..6784715 100644 --- a/root/root.go +++ b/root/root.go @@ -11,7 +11,6 @@ import ( "strconv" "strings" "syscall" - "time" "github.com/creack/pty" "github.com/spf13/cobra" @@ -42,7 +41,7 @@ It works exactly like coding editor suggestion menu drop down.`, ) func init() { - rootCmd.PersistentFlags().StringVarP(&shellFlag, "shell", "s", "", "Shell to use (bash, zsh, fish)") + rootCmd.PersistentFlags().StringVarP(&shellFlag, "shell", "s", "", "shell to use (bash, zsh, fish)") } func Execute() { @@ -59,31 +58,42 @@ func Execute() { } func detectShell() string { - // Look up to 5 levels to find the nearest shell pid := os.Getppid() for i := 0; i < 5 && pid > 1; i++ { data, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid)) if err == nil { comm := strings.ToLower(strings.TrimSpace(string(data))) - if strings.Contains(comm, "zsh") { return "zsh" } - if strings.Contains(comm, "bash") { return "bash" } - if strings.Contains(comm, "fish") { return "fish" } + if strings.Contains(comm, "zsh") { + return "zsh" + } + if strings.Contains(comm, "bash") { + return "bash" + } + if strings.Contains(comm, "fish") { + return "fish" + } } - // Move up to parent data, err = os.ReadFile(fmt.Sprintf("/proc/%d/stat", pid)) - if err != nil { break } + if err != nil { + break + } fields := strings.Fields(string(data)) if len(fields) > 3 { ppid, _ := strconv.Atoi(fields[3]) - if ppid == pid || ppid <= 1 { break } + if ppid == pid || ppid <= 1 { + break + } pid = ppid - } else { break } + } else { + break + } } - // Final fallback to system default s := os.Getenv("SHELL") - if strings.Contains(s, "zsh") { return "zsh" } + if strings.Contains(s, "zsh") { + return "zsh" + } return "bash" } @@ -93,7 +103,6 @@ type procInfo struct { comm string } -// getActiveInnerShell scans the process tree to find the deepest shell running inside the PTY func getActiveInnerShell(rootPid int, defaultShell string) string { cmd := exec.Command("ps", "-e", "-o", "pid,ppid,comm") out, err := cmd.Output() @@ -107,12 +116,10 @@ func getActiveInnerShell(rootPid int, defaultShell string) string { for _, line := range lines { fields := strings.Fields(line) if len(fields) >= 3 && fields[0] != "PID" { - pid, err1 := strconv.Atoi(fields[0]) - ppid, err2 := strconv.Atoi(fields[1]) - if err1 == nil && err2 == nil { - comm := strings.ToLower(strings.Join(fields[2:], " ")) - childrenMap[ppid] = append(childrenMap[ppid], procInfo{pid, ppid, comm}) - } + pid, _ := strconv.Atoi(fields[0]) + ppid, _ := strconv.Atoi(fields[1]) + comm := strings.ToLower(strings.Join(fields[2:], " ")) + childrenMap[ppid] = append(childrenMap[ppid], procInfo{pid, ppid, comm}) } } @@ -123,39 +130,34 @@ func getActiveInnerShell(rootPid int, defaultShell string) string { childShell := shell if strings.Contains(child.comm, "zsh") { childShell = "zsh" - } else if strings.Contains(child.comm, "bash") { + } + if strings.Contains(child.comm, "bash") { childShell = "bash" - } else if strings.Contains(child.comm, "fish") { + } + if strings.Contains(child.comm, "fish") { childShell = "fish" } - - // recursively find even deeper shells if deepest := findDeepest(child.pid, childShell); deepest != "" { shell = deepest } } return shell } - return findDeepest(rootPid, defaultShell) } func runWrapper() { r, w, err := os.Pipe() if err != nil { - fmt.Fprintln(os.Stderr, "Error creating pipe:", err) return } var shellName string - // 1. Priority: Shell from previous reload (inner deepest) if active := os.Getenv("IRIS_ACTIVE_SHELL"); active != "" { shellName = active os.Unsetenv("IRIS_ACTIVE_SHELL") - // 2. Explicit flag } else if shellFlag != "" { shellName = shellFlag - // 3. Dynamic detection (outer parent) } else { shellName = detectShell() } @@ -164,18 +166,17 @@ func runWrapper() { adapter := shell.Current c := exec.Command(adapter.GetShellPath()) - c.ExtraFiles = make([]*os.File, 11) c.ExtraFiles[10] = w c.Env = adapter.GetEnv(10, os.Getpid()) ptmx, err := pty.Start(c) if err != nil { - fmt.Fprintln(os.Stderr, "Error starting pty:", err) return } defer ptmx.Close() + _ = pty.InheritSize(os.Stdin, ptmx) core.ShellPID = c.Process.Pid oldState, err := term.MakeRaw(int(os.Stdin.Fd())) @@ -184,7 +185,6 @@ func runWrapper() { } defer func() { _ = term.Restore(int(os.Stdin.Fd()), oldState) }() - // signal handling for resize and reload sigCh := make(chan os.Signal, 2) signal.Notify(sigCh, syscall.SIGWINCH, syscall.SIGUSR1) go func() { @@ -195,13 +195,20 @@ func runWrapper() { case syscall.SIGUSR1: exe, _ := os.Executable() os.Setenv("IRIS_RELOADED", "true") - - // capture the actual shell running deep inside the PTY + + // capture the current shell state innerShell := getActiveInnerShell(c.Process.Pid, shellName) if innerShell != "" { os.Setenv("IRIS_ACTIVE_SHELL", innerShell) } + // kill the child shell before reloading + // this prevents multiple shells from fighting over the terminal + if c.Process != nil { + _ = syscall.Kill(c.Process.Pid, syscall.SIGKILL) + ptmx.Close() + } + if oldState != nil { _ = term.Restore(int(os.Stdin.Fd()), oldState) } @@ -209,28 +216,26 @@ func runWrapper() { } } }() - sigCh <- syscall.SIGWINCH overlay := integration.NewOverlay() - // PTY -> Stdout + // pty -> stdout go func() { buf := make([]byte, 4096) for { n, err := ptmx.Read(buf) if err != nil { if err == io.EOF { - // Clean up the terminal and exit when the inner shell dies _ = term.Restore(int(os.Stdin.Fd()), oldState) os.Exit(0) } continue } - TermWrite(buf[:n]) + os.Stdout.Write(buf[:n]) } }() - // Pipe IPC -> Logic -> Render + // ipc pipe go func() { scanner := bufio.NewScanner(r) scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) { @@ -248,29 +253,29 @@ func runWrapper() { for scanner.Scan() { query := scanner.Text() - results := mergeResults(query, "spec") // TODO: implement mode sync - + results := mergeResults(query, "spec") if len(results) == 0 { - TermWrite([]byte(overlay.ClearAndDisable())) + os.Stdout.Write([]byte(overlay.ClearAndDisable())) continue } - - TermWrite([]byte(overlay.Clear())) + os.Stdout.Write([]byte(overlay.Clear())) overlay.UpdateItems(results) - TermWrite([]byte(overlay.Render())) + os.Stdout.Write([]byte(overlay.Render())) } }() - // Stdin -> PTY - buf := make([]byte, 1) - var escSeq []byte - var naiveBuffer string // fallback naive tracker for bash - mode := "spec" // can be "spec" or "history" + var naiveBuffer string + mode := "spec" renderOverlay := func() { + // don't render if shell is starting up + if isReload && naiveBuffer == "" { + return + } + results := mergeResults(naiveBuffer, mode) if len(results) == 0 { - TermWrite([]byte(overlay.ClearAndDisable())) + os.Stdout.Write([]byte(overlay.ClearAndDisable())) } else { var buf strings.Builder if overlay.Visible { @@ -278,109 +283,94 @@ func runWrapper() { } overlay.UpdateItems(results) buf.WriteString(overlay.Render()) - TermWrite([]byte(buf.String())) + os.Stdout.Write([]byte(buf.String())) } } for { - n, err := os.Stdin.Read(buf) + inputSlice := make([]byte, 128) + n, err := os.Stdin.Read(inputSlice) if err != nil { break } if n > 0 { - b := buf[0] + shouldOverlayDraw := false + for i := 0; i < n; i++ { + b := inputSlice[i] - // escape sequence parsing - if b == '\033' { - escSeq = append(escSeq, b) - continue - } - if len(escSeq) > 0 { - escSeq = append(escSeq, b) - if len(escSeq) == 3 { - seq := string(escSeq) - escSeq = nil + intercepted := false + if overlay.Visible && (b == '\r' || b == 0x09) { + intercepted = true + selected := overlay.Items[overlay.Cursor].Cmd + os.Stdout.Write([]byte(overlay.ClearAndDisable())) + if mode == "spec" || b == 0x09 { + selected += " " + } + naiveBuffer = selected + mode = "spec" + ptmx.Write(adapter.PrepareSelectSequence(selected)) + if b == '\r' { + ptmx.Write([]byte{'\r'}) + naiveBuffer = "" + } else { + renderOverlay() + } + continue + } - if overlay.Visible { - if seq == "\033[A" { // Up - overlay.Cursor-- - if overlay.Cursor < 0 { - overlay.Cursor = 0 + if !intercepted { + ptmx.Write([]byte{b}) + + // arrow key monitoring with tagged switch + if b == '\033' && i+2 < n && inputSlice[i+1] == '[' { + if overlay.Visible { + switch inputSlice[i+2] { + case 'A': // up + overlay.Cursor-- + if overlay.Cursor < 0 { + overlay.Cursor = 0 + } + os.Stdout.Write([]byte(overlay.Render())) + case 'B': // down + overlay.Cursor++ + if overlay.Cursor >= len(overlay.Items) { + overlay.Cursor = len(overlay.Items) - 1 + } + os.Stdout.Write([]byte(overlay.Render())) } - TermWrite([]byte(overlay.Render())) - continue - } else if seq == "\033[B" { // Down - overlay.Cursor++ - if overlay.Cursor >= len(overlay.Items) { - overlay.Cursor = len(overlay.Items) - 1 - } - TermWrite([]byte(overlay.Render())) - continue } } - ptmx.Write([]byte(seq)) + + switch b { + case 0x12: // ctrl+r + if mode == "spec" { + mode = "history" + } else { + mode = "spec" + } + shouldOverlayDraw = true + case 0x09: // tab + if !overlay.Visible { + shouldOverlayDraw = true + } + case 127: // backspace + if len(naiveBuffer) > 0 { + naiveBuffer = naiveBuffer[:len(naiveBuffer)-1] + shouldOverlayDraw = true + } + case '\r', 0x03: // enter, ctrl+c + naiveBuffer = "" + mode = "spec" + os.Stdout.Write([]byte(overlay.ClearAndDisable())) + default: + if b >= 32 && b <= 126 { + naiveBuffer += string(b) + shouldOverlayDraw = true + } + } } - continue } - - // capture naive typing for bash test - shouldOverlayDraw := false - - // enter 0x0D, tab 0x09 - if overlay.Visible && (b == '\r' || b == 0x09) { - selected := overlay.Items[overlay.Cursor].Cmd - - TermWrite([]byte(overlay.ClearAndDisable())) - - // auto add space after tab for next suggestion in spec mode - if mode == "spec" || b == 0x09 { - selected += " " - } - - // sync our naive buffer with the selected result - naiveBuffer = selected - mode = "spec" // reset to spec mode after selection - - ptmx.Write(adapter.PrepareSelectSequence(selected)) - - if b == '\r' { - ptmx.Write([]byte{'\r'}) - naiveBuffer = "" - } else { - // if it was tab, we want to immediately show next suggestions - renderOverlay() - } - continue - } - - if b == 0x12 { // Ctrl+R to switch between showing commands mode and history mode - if mode == "spec" { - mode = "history" - } else { - mode = "spec" - } - shouldOverlayDraw = true - } else if b >= 32 && b <= 126 { - naiveBuffer += string(b) - shouldOverlayDraw = true - } else if b == 127 { // Backspace - if len(naiveBuffer) > 0 { - naiveBuffer = naiveBuffer[:len(naiveBuffer)-1] - shouldOverlayDraw = true - } - } else if b == '\r' || b == 0x03 { - naiveBuffer = "" - mode = "spec" // reset on clear - TermWrite([]byte(overlay.ClearAndDisable())) - } - - if b != 0x12 { - ptmx.Write([]byte{b}) - // small delay to let PTY echo arrive before overlay render - time.Sleep(3 * time.Millisecond) - } - if shouldOverlayDraw { renderOverlay() } @@ -388,19 +378,17 @@ func runWrapper() { } } -// mergeResults returns suggestions based on the active mode func mergeResults(query string, mode string) []core.Suggestion { if query == "" { return nil } - if mode == "history" { histResults, _ := integration.SearchHistory(query) cmdResults := []core.Suggestion{} for _, h := range histResults { cmdResults = append(cmdResults, core.Suggestion{ Cmd: h.Cmd, - Desc: " history", // please dont remove space here + Desc: " history", Icon: fmt.Sprintf("%d", h.ID), }) } @@ -409,15 +397,7 @@ func mergeResults(query string, mode string) []core.Suggestion { } return cmdResults } - - // Spec mode - if query == "" { - return nil - } - cmdResults := core.Lookup(query) - - // deduplicate by Cmd seen := make(map[string]bool) deduped := []core.Suggestion{} for _, s := range cmdResults { @@ -426,10 +406,8 @@ func mergeResults(query string, mode string) []core.Suggestion { deduped = append(deduped, s) } } - if len(deduped) > 10 { deduped = deduped[:10] } - return deduped }