fix(root): eliminate child shell leakage and terminal corruption on reload

This commit is contained in:
verse91
2026-04-11 23:54:11 +07:00
parent 8dff68601e
commit 11225df6c8
+128 -150
View File
@@ -11,7 +11,6 @@ import (
"strconv" "strconv"
"strings" "strings"
"syscall" "syscall"
"time"
"github.com/creack/pty" "github.com/creack/pty"
"github.com/spf13/cobra" "github.com/spf13/cobra"
@@ -42,7 +41,7 @@ It works exactly like coding editor suggestion menu drop down.`,
) )
func init() { 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() { func Execute() {
@@ -59,31 +58,42 @@ func Execute() {
} }
func detectShell() string { func detectShell() string {
// Look up to 5 levels to find the nearest shell
pid := os.Getppid() pid := os.Getppid()
for i := 0; i < 5 && pid > 1; i++ { for i := 0; i < 5 && pid > 1; i++ {
data, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid)) data, err := os.ReadFile(fmt.Sprintf("/proc/%d/comm", pid))
if err == nil { if err == nil {
comm := strings.ToLower(strings.TrimSpace(string(data))) comm := strings.ToLower(strings.TrimSpace(string(data)))
if strings.Contains(comm, "zsh") { return "zsh" } if strings.Contains(comm, "zsh") {
if strings.Contains(comm, "bash") { return "bash" } return "zsh"
if strings.Contains(comm, "fish") { return "fish" } }
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)) data, err = os.ReadFile(fmt.Sprintf("/proc/%d/stat", pid))
if err != nil { break } if err != nil {
break
}
fields := strings.Fields(string(data)) fields := strings.Fields(string(data))
if len(fields) > 3 { if len(fields) > 3 {
ppid, _ := strconv.Atoi(fields[3]) ppid, _ := strconv.Atoi(fields[3])
if ppid == pid || ppid <= 1 { break } if ppid == pid || ppid <= 1 {
break
}
pid = ppid pid = ppid
} else { break } } else {
break
}
} }
// Final fallback to system default
s := os.Getenv("SHELL") s := os.Getenv("SHELL")
if strings.Contains(s, "zsh") { return "zsh" } if strings.Contains(s, "zsh") {
return "zsh"
}
return "bash" return "bash"
} }
@@ -93,7 +103,6 @@ type procInfo struct {
comm string comm string
} }
// getActiveInnerShell scans the process tree to find the deepest shell running inside the PTY
func getActiveInnerShell(rootPid int, defaultShell string) string { func getActiveInnerShell(rootPid int, defaultShell string) string {
cmd := exec.Command("ps", "-e", "-o", "pid,ppid,comm") cmd := exec.Command("ps", "-e", "-o", "pid,ppid,comm")
out, err := cmd.Output() out, err := cmd.Output()
@@ -107,12 +116,10 @@ func getActiveInnerShell(rootPid int, defaultShell string) string {
for _, line := range lines { for _, line := range lines {
fields := strings.Fields(line) fields := strings.Fields(line)
if len(fields) >= 3 && fields[0] != "PID" { if len(fields) >= 3 && fields[0] != "PID" {
pid, err1 := strconv.Atoi(fields[0]) pid, _ := strconv.Atoi(fields[0])
ppid, err2 := strconv.Atoi(fields[1]) ppid, _ := strconv.Atoi(fields[1])
if err1 == nil && err2 == nil { comm := strings.ToLower(strings.Join(fields[2:], " "))
comm := strings.ToLower(strings.Join(fields[2:], " ")) childrenMap[ppid] = append(childrenMap[ppid], procInfo{pid, ppid, comm})
childrenMap[ppid] = append(childrenMap[ppid], procInfo{pid, ppid, comm})
}
} }
} }
@@ -123,39 +130,34 @@ func getActiveInnerShell(rootPid int, defaultShell string) string {
childShell := shell childShell := shell
if strings.Contains(child.comm, "zsh") { if strings.Contains(child.comm, "zsh") {
childShell = "zsh" childShell = "zsh"
} else if strings.Contains(child.comm, "bash") { }
if strings.Contains(child.comm, "bash") {
childShell = "bash" childShell = "bash"
} else if strings.Contains(child.comm, "fish") { }
if strings.Contains(child.comm, "fish") {
childShell = "fish" childShell = "fish"
} }
// recursively find even deeper shells
if deepest := findDeepest(child.pid, childShell); deepest != "" { if deepest := findDeepest(child.pid, childShell); deepest != "" {
shell = deepest shell = deepest
} }
} }
return shell return shell
} }
return findDeepest(rootPid, defaultShell) return findDeepest(rootPid, defaultShell)
} }
func runWrapper() { func runWrapper() {
r, w, err := os.Pipe() r, w, err := os.Pipe()
if err != nil { if err != nil {
fmt.Fprintln(os.Stderr, "Error creating pipe:", err)
return return
} }
var shellName string var shellName string
// 1. Priority: Shell from previous reload (inner deepest)
if active := os.Getenv("IRIS_ACTIVE_SHELL"); active != "" { if active := os.Getenv("IRIS_ACTIVE_SHELL"); active != "" {
shellName = active shellName = active
os.Unsetenv("IRIS_ACTIVE_SHELL") os.Unsetenv("IRIS_ACTIVE_SHELL")
// 2. Explicit flag
} else if shellFlag != "" { } else if shellFlag != "" {
shellName = shellFlag shellName = shellFlag
// 3. Dynamic detection (outer parent)
} else { } else {
shellName = detectShell() shellName = detectShell()
} }
@@ -164,18 +166,17 @@ func runWrapper() {
adapter := shell.Current adapter := shell.Current
c := exec.Command(adapter.GetShellPath()) c := exec.Command(adapter.GetShellPath())
c.ExtraFiles = make([]*os.File, 11) c.ExtraFiles = make([]*os.File, 11)
c.ExtraFiles[10] = w c.ExtraFiles[10] = w
c.Env = adapter.GetEnv(10, os.Getpid()) c.Env = adapter.GetEnv(10, os.Getpid())
ptmx, err := pty.Start(c) ptmx, err := pty.Start(c)
if err != nil { if err != nil {
fmt.Fprintln(os.Stderr, "Error starting pty:", err)
return return
} }
defer ptmx.Close() defer ptmx.Close()
_ = pty.InheritSize(os.Stdin, ptmx)
core.ShellPID = c.Process.Pid core.ShellPID = c.Process.Pid
oldState, err := term.MakeRaw(int(os.Stdin.Fd())) oldState, err := term.MakeRaw(int(os.Stdin.Fd()))
@@ -184,7 +185,6 @@ func runWrapper() {
} }
defer func() { _ = term.Restore(int(os.Stdin.Fd()), oldState) }() defer func() { _ = term.Restore(int(os.Stdin.Fd()), oldState) }()
// signal handling for resize and reload
sigCh := make(chan os.Signal, 2) sigCh := make(chan os.Signal, 2)
signal.Notify(sigCh, syscall.SIGWINCH, syscall.SIGUSR1) signal.Notify(sigCh, syscall.SIGWINCH, syscall.SIGUSR1)
go func() { go func() {
@@ -195,13 +195,20 @@ func runWrapper() {
case syscall.SIGUSR1: case syscall.SIGUSR1:
exe, _ := os.Executable() exe, _ := os.Executable()
os.Setenv("IRIS_RELOADED", "true") 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) innerShell := getActiveInnerShell(c.Process.Pid, shellName)
if innerShell != "" { if innerShell != "" {
os.Setenv("IRIS_ACTIVE_SHELL", 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 { if oldState != nil {
_ = term.Restore(int(os.Stdin.Fd()), oldState) _ = term.Restore(int(os.Stdin.Fd()), oldState)
} }
@@ -209,28 +216,26 @@ func runWrapper() {
} }
} }
}() }()
sigCh <- syscall.SIGWINCH
overlay := integration.NewOverlay() overlay := integration.NewOverlay()
// PTY -> Stdout // pty -> stdout
go func() { go func() {
buf := make([]byte, 4096) buf := make([]byte, 4096)
for { for {
n, err := ptmx.Read(buf) n, err := ptmx.Read(buf)
if err != nil { if err != nil {
if err == io.EOF { if err == io.EOF {
// Clean up the terminal and exit when the inner shell dies
_ = term.Restore(int(os.Stdin.Fd()), oldState) _ = term.Restore(int(os.Stdin.Fd()), oldState)
os.Exit(0) os.Exit(0)
} }
continue continue
} }
TermWrite(buf[:n]) os.Stdout.Write(buf[:n])
} }
}() }()
// Pipe IPC -> Logic -> Render // ipc pipe
go func() { go func() {
scanner := bufio.NewScanner(r) scanner := bufio.NewScanner(r)
scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) { scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) {
@@ -248,29 +253,29 @@ func runWrapper() {
for scanner.Scan() { for scanner.Scan() {
query := scanner.Text() query := scanner.Text()
results := mergeResults(query, "spec") // TODO: implement mode sync results := mergeResults(query, "spec")
if len(results) == 0 { if len(results) == 0 {
TermWrite([]byte(overlay.ClearAndDisable())) os.Stdout.Write([]byte(overlay.ClearAndDisable()))
continue continue
} }
os.Stdout.Write([]byte(overlay.Clear()))
TermWrite([]byte(overlay.Clear()))
overlay.UpdateItems(results) overlay.UpdateItems(results)
TermWrite([]byte(overlay.Render())) os.Stdout.Write([]byte(overlay.Render()))
} }
}() }()
// Stdin -> PTY var naiveBuffer string
buf := make([]byte, 1) mode := "spec"
var escSeq []byte
var naiveBuffer string // fallback naive tracker for bash
mode := "spec" // can be "spec" or "history"
renderOverlay := func() { renderOverlay := func() {
// don't render if shell is starting up
if isReload && naiveBuffer == "" {
return
}
results := mergeResults(naiveBuffer, mode) results := mergeResults(naiveBuffer, mode)
if len(results) == 0 { if len(results) == 0 {
TermWrite([]byte(overlay.ClearAndDisable())) os.Stdout.Write([]byte(overlay.ClearAndDisable()))
} else { } else {
var buf strings.Builder var buf strings.Builder
if overlay.Visible { if overlay.Visible {
@@ -278,109 +283,94 @@ func runWrapper() {
} }
overlay.UpdateItems(results) overlay.UpdateItems(results)
buf.WriteString(overlay.Render()) buf.WriteString(overlay.Render())
TermWrite([]byte(buf.String())) os.Stdout.Write([]byte(buf.String()))
} }
} }
for { for {
n, err := os.Stdin.Read(buf) inputSlice := make([]byte, 128)
n, err := os.Stdin.Read(inputSlice)
if err != nil { if err != nil {
break break
} }
if n > 0 { if n > 0 {
b := buf[0] shouldOverlayDraw := false
for i := 0; i < n; i++ {
b := inputSlice[i]
// escape sequence parsing intercepted := false
if b == '\033' { if overlay.Visible && (b == '\r' || b == 0x09) {
escSeq = append(escSeq, b) intercepted = true
continue selected := overlay.Items[overlay.Cursor].Cmd
} os.Stdout.Write([]byte(overlay.ClearAndDisable()))
if len(escSeq) > 0 { if mode == "spec" || b == 0x09 {
escSeq = append(escSeq, b) selected += " "
if len(escSeq) == 3 { }
seq := string(escSeq) naiveBuffer = selected
escSeq = nil mode = "spec"
ptmx.Write(adapter.PrepareSelectSequence(selected))
if b == '\r' {
ptmx.Write([]byte{'\r'})
naiveBuffer = ""
} else {
renderOverlay()
}
continue
}
if overlay.Visible { if !intercepted {
if seq == "\033[A" { // Up ptmx.Write([]byte{b})
overlay.Cursor--
if overlay.Cursor < 0 { // arrow key monitoring with tagged switch
overlay.Cursor = 0 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 { if shouldOverlayDraw {
renderOverlay() renderOverlay()
} }
@@ -388,19 +378,17 @@ func runWrapper() {
} }
} }
// mergeResults returns suggestions based on the active mode
func mergeResults(query string, mode string) []core.Suggestion { func mergeResults(query string, mode string) []core.Suggestion {
if query == "" { if query == "" {
return nil return nil
} }
if mode == "history" { if mode == "history" {
histResults, _ := integration.SearchHistory(query) histResults, _ := integration.SearchHistory(query)
cmdResults := []core.Suggestion{} cmdResults := []core.Suggestion{}
for _, h := range histResults { for _, h := range histResults {
cmdResults = append(cmdResults, core.Suggestion{ cmdResults = append(cmdResults, core.Suggestion{
Cmd: h.Cmd, Cmd: h.Cmd,
Desc: " history", // please dont remove space here Desc: " history",
Icon: fmt.Sprintf("%d", h.ID), Icon: fmt.Sprintf("%d", h.ID),
}) })
} }
@@ -409,15 +397,7 @@ func mergeResults(query string, mode string) []core.Suggestion {
} }
return cmdResults return cmdResults
} }
// Spec mode
if query == "" {
return nil
}
cmdResults := core.Lookup(query) cmdResults := core.Lookup(query)
// deduplicate by Cmd
seen := make(map[string]bool) seen := make(map[string]bool)
deduped := []core.Suggestion{} deduped := []core.Suggestion{}
for _, s := range cmdResults { for _, s := range cmdResults {
@@ -426,10 +406,8 @@ func mergeResults(query string, mode string) []core.Suggestion {
deduped = append(deduped, s) deduped = append(deduped, s)
} }
} }
if len(deduped) > 10 { if len(deduped) > 10 {
deduped = deduped[:10] deduped = deduped[:10]
} }
return deduped return deduped
} }