diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 5b78716..3686e15 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,6 +7,7 @@ jobs: test: runs-on: ${{ matrix.os }} strategy: + fail-fast: false # one OS failing must not hide the others' results matrix: os: [ubuntu-latest, macos-latest, windows-latest] steps: @@ -16,6 +17,10 @@ jobs: go-version: "1.26" - run: go vet ./... - run: go test ./... + # Show which service manager each OS really used (launchd on macOS, + # breakaway/WMI on Windows), not just that the suite passed. + - name: Service manager and Windows console, verbose + run: go test -count=1 -v -run "TestSessionRunsUnderTheRealServiceManager|TestWindowsSessionOnConPTY|TestDaemonStartsWithoutTheForkOverride|TestWMIStartsAProcess|TestChooseSupervisor" ./internal/server - if: runner.os == 'Linux' run: go test -race ./... - run: go build ./cmd/hqsh diff --git a/README.md b/README.md index d38beec..e2494a4 100644 --- a/README.md +++ b/README.md @@ -70,8 +70,14 @@ running. The daemon is started by the OS service manager where there is one: | --- | --- | --- | | Linux with systemd | a user unit, `hqsh-.service` | yes, with lingering on | | macOS | a launchd job, `sh.hqterm.hqsh.`, in `user/` | yes | +| Windows 10 1809+ / Server 2019+ | a process outside the ssh connection's job (breakaway, else WMI `Win32_Process.Create`) on a ConPTY | yes | | anything else | a detached process (setsid) | unless the OS kills it | +On Windows the session's shell is PowerShell 7 (`pwsh`), else Windows +PowerShell, else `cmd.exe`; `HQSH_SHELL` picks another on any OS. The host +needs the OpenSSH server (Settings → Optional features), and hqsh.exe on the +PATH or in `~\.local\bin`. + `hqsh server setup` shows which applies and turns lingering on (`loginctl enable-linger`) where systemd needs it; `--check` only reports. Without lingering, logind stops a user's units at their last logout, so hqsh falls @@ -105,8 +111,8 @@ tested end to end (a real shell over a pipe in CI, and over ssh by hand): reconnect after the connection dies, replay of exactly the missed output, detach and re-attach, exit status, Kitty and iTerm2 image escapes passed through untouched. 0.2.0 adds shared attach (several clients on one -session), `--steal` and `--read-only`. The server side runs on Linux and macOS; the client also -builds for Windows. Not yet: local echo prediction, a WebSocket bridge so a +session), `--steal` and `--read-only`. The server side runs on Linux (systemd), macOS (launchd) and Windows (ConPTY); the client +runs on all three. Not yet: local echo prediction, a WebSocket bridge so a phone/PWA can attach. Works with any modern terminal: Kitty, Ghostty, WezTerm, Rio, iTerm2, diff --git a/go.mod b/go.mod index ef656c2..22bd592 100644 --- a/go.mod +++ b/go.mod @@ -7,4 +7,4 @@ require ( golang.org/x/term v0.46.0 ) -require golang.org/x/sys v0.48.0 // indirect +require golang.org/x/sys v0.48.0 diff --git a/internal/server/console.go b/internal/server/console.go new file mode 100644 index 0000000..69dd828 --- /dev/null +++ b/internal/server/console.go @@ -0,0 +1,20 @@ +package server + +import "io" + +// console is the shell's terminal as the daemon sees it: a PTY on Unix, a +// pseudo console (ConPTY) on Windows. Reads are the shell's output, writes +// are keystrokes. +type console interface { + io.ReadWriteCloser + Resize(cols, rows uint16) error +} + +// The platform files provide: +// +// startShell(session, term string, cols, rows uint16) (console, func() int, error) +// starts the user's shell on a new console; the func waits for it to +// exit and returns its status (128+signal for a signal on Unix). +// lockSession(f *os.File) error exclusive, non-blocking: one daemon per session +// chmodSocket(path string) error the socket is the user's alone +// isDead(err error) bool a dial error meaning no daemon is behind the socket diff --git a/internal/server/console_unix.go b/internal/server/console_unix.go new file mode 100644 index 0000000..223105e --- /dev/null +++ b/internal/server/console_unix.go @@ -0,0 +1,113 @@ +//go:build !windows + +package server + +import ( + "errors" + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" + "syscall" + + "github.com/creack/pty" +) + +type unixConsole struct{ *os.File } + +func (c unixConsole) Resize(cols, rows uint16) error { + return pty.Setsize(c.File, &pty.Winsize{Cols: cols, Rows: rows}) +} + +// startShell runs the user's login shell ($HQSH_SHELL, else $SHELL, else +// /bin/sh) in their home directory on a new PTY. +func startShell(session, term string, cols, rows uint16) (console, func() int, error) { + shell := os.Getenv("HQSH_SHELL") + if shell == "" { + shell = os.Getenv("SHELL") + } + if shell == "" { + shell = "/bin/sh" + } + cmd := exec.Command(shell, "-l") + if home, err := os.UserHomeDir(); err == nil { + cmd.Dir = home + } + cmd.Env = shellEnv(os.Environ(), usableTerm(term), session) + ptmx, err := pty.StartWithSize(cmd, &pty.Winsize{Cols: cols, Rows: rows}) + if err != nil { + return nil, nil, err + } + return unixConsole{ptmx}, func() int { + _ = cmd.Wait() + return exitCode(cmd.ProcessState) + }, nil +} + +func lockSession(f *os.File) error { + return syscall.Flock(int(f.Fd()), syscall.LOCK_EX|syscall.LOCK_NB) +} + +func chmodSocket(path string) error { return os.Chmod(path, 0o600) } + +func isDead(err error) bool { + return errors.Is(err, syscall.ECONNREFUSED) || errors.Is(err, syscall.ENOENT) +} + +func exitCode(ps *os.ProcessState) int { + if ps == nil { + return 1 + } + if ws, ok := ps.Sys().(syscall.WaitStatus); ok && ws.Signaled() { + return 128 + int(ws.Signal()) + } + return ps.ExitCode() +} + +func shellEnv(env []string, term, session string) []string { + out := make([]string, 0, len(env)+2) + for _, kv := range env { + if strings.HasPrefix(kv, "TERM=") || strings.HasPrefix(kv, "HQSH_SESSION=") { + continue + } + out = append(out, kv) + } + return append(out, "TERM="+term, "HQSH_SESSION="+session) +} + +// usableTerm keeps the client's TERM when this host has a terminfo entry for +// it (xterm-kitty often is missing), else falls back to xterm-256color. +func usableTerm(term string) string { + const fallback = "xterm-256color" + if term == "" || strings.ContainsAny(term, "/\\\x00") || strings.HasPrefix(term, ".") { + return fallback + } + var dirs []string + if t := os.Getenv("TERMINFO"); t != "" { + dirs = append(dirs, t) + } + if home, err := os.UserHomeDir(); err == nil { + dirs = append(dirs, filepath.Join(home, ".terminfo")) + } + if td := os.Getenv("TERMINFO_DIRS"); td != "" { + dirs = append(dirs, filepath.SplitList(td)...) + } + dirs = append(dirs, "/etc/terminfo", "/lib/terminfo", "/usr/share/terminfo", "/usr/lib/terminfo", "/usr/share/lib/terminfo", "/opt/homebrew/share/terminfo", "/usr/local/share/terminfo") + anyDir := false + for _, dir := range dirs { + if st, err := os.Stat(dir); err != nil || !st.IsDir() { + continue + } + anyDir = true + for _, sub := range []string{term[:1], fmt.Sprintf("%x", term[0])} { + if _, err := os.Stat(filepath.Join(dir, sub, term)); err == nil { + return term + } + } + } + if !anyDir { + return term // no terminfo database to check against; trust the client + } + return fallback +} diff --git a/internal/server/console_windows.go b/internal/server/console_windows.go new file mode 100644 index 0000000..027946c --- /dev/null +++ b/internal/server/console_windows.go @@ -0,0 +1,170 @@ +//go:build windows + +package server + +import ( + "errors" + "os" + "os/exec" + "strings" + "sync" + "syscall" + "unicode/utf16" + "unsafe" + + "golang.org/x/sys/windows" +) + +// conPTY is a Windows pseudo console (Windows 10 1809 and later): the shell +// writes VT sequences to it like to a Unix PTY, and keystrokes go in as VT. +type conPTY struct { + hpc windows.Handle + in *os.File // our end of the shell's input + out *os.File // our end of the shell's output + once sync.Once +} + +func (c *conPTY) Read(p []byte) (int, error) { return c.out.Read(p) } +func (c *conPTY) Write(p []byte) (int, error) { return c.in.Write(p) } + +func (c *conPTY) Resize(cols, rows uint16) error { + return windows.ResizePseudoConsole(c.hpc, windows.Coord{X: int16(cols), Y: int16(rows)}) +} + +// Close ends the pseudo console (which ends the shell, if it is still +// there) and our pipe ends. Closing the console is also what lets the last +// read return EOF. +func (c *conPTY) Close() error { + c.once.Do(func() { + windows.ClosePseudoConsole(c.hpc) + c.in.Close() + c.out.Close() + }) + return nil +} + +// windowsShell picks the shell: $HQSH_SHELL, else PowerShell 7 (pwsh), +// else Windows PowerShell, else %ComSpec% (cmd.exe). +func windowsShell() string { + if s := os.Getenv("HQSH_SHELL"); s != "" { + return s + } + for _, s := range []string{"pwsh.exe", "powershell.exe"} { + if p, err := exec.LookPath(s); err == nil { + return p + } + } + if s := os.Getenv("ComSpec"); s != "" { + return s + } + return `C:\Windows\System32\cmd.exe` +} + +func startShell(session, term string, cols, rows uint16) (console, func() int, error) { + var inR, inW, outR, outW windows.Handle + if err := windows.CreatePipe(&inR, &inW, nil, 0); err != nil { + return nil, nil, err + } + if err := windows.CreatePipe(&outR, &outW, nil, 0); err != nil { + windows.CloseHandle(inR) + windows.CloseHandle(inW) + return nil, nil, err + } + var hpc windows.Handle + if err := windows.CreatePseudoConsole(windows.Coord{X: int16(cols), Y: int16(rows)}, inR, outW, 0, &hpc); err != nil { + for _, h := range []windows.Handle{inR, inW, outR, outW} { + windows.CloseHandle(h) + } + return nil, nil, err + } + // The pseudo console holds its own copies of these. + windows.CloseHandle(inR) + windows.CloseHandle(outW) + con := &conPTY{hpc: hpc, in: os.NewFile(uintptr(inW), "conpty-in"), out: os.NewFile(uintptr(outR), "conpty-out")} + + attrs, err := windows.NewProcThreadAttributeList(1) + if err != nil { + con.Close() + return nil, nil, err + } + defer attrs.Delete() + // The attribute's value is the HPCON itself, not a pointer to it. + if err := attrs.Update(windows.PROC_THREAD_ATTRIBUTE_PSEUDOCONSOLE, *(*unsafe.Pointer)(unsafe.Pointer(&hpc)), unsafe.Sizeof(hpc)); err != nil { + con.Close() + return nil, nil, err + } + si := windows.StartupInfoEx{ProcThreadAttributeList: attrs.List()} + si.Cb = uint32(unsafe.Sizeof(si)) + // No inherited std handles: the console is the shell's only terminal. + si.Flags = windows.STARTF_USESTDHANDLES + + cmdline, err := windows.UTF16PtrFromString(windows.EscapeArg(windowsShell())) + if err != nil { + con.Close() + return nil, nil, err + } + var dir *uint16 + if home, err := os.UserHomeDir(); err == nil { + dir, _ = windows.UTF16PtrFromString(home) + } + env := envBlock(shellEnvWindows(os.Environ(), session)) + var pi windows.ProcessInformation + if err := windows.CreateProcess(nil, cmdline, nil, nil, false, + windows.EXTENDED_STARTUPINFO_PRESENT|windows.CREATE_UNICODE_ENVIRONMENT, + &env[0], dir, &si.StartupInfo, &pi); err != nil { + con.Close() + return nil, nil, err + } + windows.CloseHandle(pi.Thread) + return con, func() int { + defer windows.CloseHandle(pi.Process) + if _, err := windows.WaitForSingleObject(pi.Process, windows.INFINITE); err != nil { + return 1 + } + var code uint32 + if err := windows.GetExitCodeProcess(pi.Process, &code); err != nil { + return 1 + } + return int(code) + }, nil +} + +// shellEnvWindows marks the session; TERM means nothing to Windows programs. +func shellEnvWindows(env []string, session string) []string { + out := make([]string, 0, len(env)+1) + for _, kv := range env { + if !strings.HasPrefix(strings.ToUpper(kv), "HQSH_SESSION=") { + out = append(out, kv) + } + } + return append(out, "HQSH_SESSION="+session) +} + +// envBlock is CreateProcess's environment: NUL-separated UTF-16, NUL-NUL ended. +func envBlock(env []string) []uint16 { + var b []uint16 + for _, kv := range env { + if strings.ContainsRune(kv, 0) { + continue + } + b = append(b, utf16.Encode([]rune(kv))...) + b = append(b, 0) + } + if len(b) == 0 { + b = append(b, 0) + } + return append(b, 0) +} + +func lockSession(f *os.File) error { + return windows.LockFileEx(windows.Handle(f.Fd()), windows.LOCKFILE_EXCLUSIVE_LOCK|windows.LOCKFILE_FAIL_IMMEDIATELY, 0, 1, 0, &windows.Overlapped{}) +} + +// chmodSocket: Windows has no mode bits for it; the socket lives in the +// user's profile, which only they (and admins) can open. +func chmodSocket(path string) error { return nil } + +func isDead(err error) bool { + return errors.Is(err, windows.WSAECONNREFUSED) || errors.Is(err, syscall.ENOENT) || + errors.Is(err, windows.ERROR_FILE_NOT_FOUND) || errors.Is(err, os.ErrNotExist) +} diff --git a/internal/server/server_unix.go b/internal/server/daemon.go similarity index 84% rename from internal/server/server_unix.go rename to internal/server/daemon.go index a8c1c53..fd77900 100644 --- a/internal/server/server_unix.go +++ b/internal/server/daemon.go @@ -1,5 +1,3 @@ -//go:build !windows - package server import ( @@ -13,11 +11,8 @@ import ( "sort" "strings" "sync" - "syscall" "time" - "github.com/creack/pty" - "github.com/profullstack/hqsh/internal/proto" "github.com/profullstack/hqsh/internal/ring" ) @@ -135,10 +130,6 @@ func List() ([]SessionInfo, error) { return out, nil } -func isDead(err error) bool { - return errors.Is(err, syscall.ECONNREFUSED) || errors.Is(err, syscall.ENOENT) -} - // status asks a daemon whether clients are attached, and how many. func status(path string) (attached bool, clients int, err error) { c, err := net.DialTimeout("unix", path, time.Second) @@ -170,7 +161,7 @@ type daemon struct { mu sync.Mutex // guards everything below clients map[*member]struct{} - ptmx *os.File + con console // the shell's terminal: a PTY, or a ConPTY on Windows started bool ln net.Listener cols, rows uint16 // the PTY's size: the smallest among the clients @@ -253,7 +244,7 @@ func Daemon(session string) error { return err } defer lock.Close() - if err := syscall.Flock(int(lock.Fd()), syscall.LOCK_EX|syscall.LOCK_NB); err != nil { + if err := lockSession(lock); err != nil { return nil // another daemon has it } d := &daemon{session: session, path: path, ring: ring.New(DefaultBuffer), @@ -284,7 +275,7 @@ func (d *daemon) listen() error { if err != nil { return err } - if err := os.Chmod(d.path, 0o600); err != nil { + if err := chmodSocket(d.path); err != nil { ln.Close() return err } @@ -407,7 +398,7 @@ func (d *daemon) serve(c net.Conn) { continue // watching only } d.mu.Lock() - p := d.ptmx + p := d.con d.mu.Unlock() if p != nil { if _, err := p.Write(f.Payload); err != nil { @@ -527,23 +518,23 @@ func (d *daemon) resizeLocked() { return } d.cols, d.rows = cols, rows - if d.ptmx != nil { - _ = pty.Setsize(d.ptmx, &pty.Winsize{Cols: cols, Rows: rows}) + if d.con != nil { + _ = d.con.Resize(cols, rows) } } // nudge makes full-screen programs repaint: one row less, then back. func (d *daemon) nudge() { d.mu.Lock() - p, cols, rows := d.ptmx, d.cols, d.rows + p, cols, rows := d.con, d.cols, d.rows d.mu.Unlock() if p == nil || rows < 2 { return } - _ = pty.Setsize(p, &pty.Winsize{Cols: cols, Rows: rows - 1}) + _ = p.Resize(cols, rows-1) time.Sleep(50 * time.Millisecond) d.mu.Lock() - _ = pty.Setsize(p, &pty.Winsize{Cols: d.cols, Rows: d.rows}) + _ = p.Resize(d.cols, d.rows) d.mu.Unlock() } @@ -555,37 +546,28 @@ func (d *daemon) ensureShell(h proto.HelloMsg) error { if d.started { return nil } - shell := os.Getenv("SHELL") - if shell == "" { - shell = "/bin/sh" - } - cmd := exec.Command(shell, "-l") - if home, err := os.UserHomeDir(); err == nil { - cmd.Dir = home - } - cmd.Env = shellEnv(os.Environ(), usableTerm(h.Term), d.session) cols, rows := h.Cols, h.Rows if cols == 0 || rows == 0 { cols, rows = 80, 24 } - ptmx, err := pty.StartWithSize(cmd, &pty.Winsize{Cols: cols, Rows: rows}) + con, wait, err := startShell(d.session, h.Term, cols, rows) if err != nil { return err } - d.ptmx = ptmx + d.con = con d.cols, d.rows = cols, rows d.started = true readDone := make(chan struct{}) - go d.pump(ptmx, readDone) + go d.pump(con, readDone) go func() { - _ = cmd.Wait() + code := wait() // Let the last output drain; a background job still holding the // terminal must not keep the session alive. select { case <-readDone: case <-time.After(500 * time.Millisecond): } - d.finish(exitCode(cmd.ProcessState)) + d.finish(code) }() return nil } @@ -593,11 +575,11 @@ func (d *daemon) ensureShell(h proto.HelloMsg) error { // pump reads the PTY into the ring and wakes every client's writer. It // never waits on one particular client: it pauses only while even the most // caught-up client is more than paceWindow behind (see mustWait). -func (d *daemon) pump(ptmx *os.File, done chan struct{}) { +func (d *daemon) pump(con console, done chan struct{}) { defer close(done) buf := make([]byte, 32<<10) for { - n, err := ptmx.Read(buf) + n, err := con.Read(buf) if n > 0 { d.ring.Append(buf[:n]) d.mu.Lock() @@ -647,8 +629,8 @@ func (d *daemon) finish(code int) { for cl := range d.clients { cl.close() } - if d.ptmx != nil { - d.ptmx.Close() + if d.con != nil { + d.con.Close() } d.mu.Unlock() d.shutdown() @@ -668,60 +650,3 @@ func (d *daemon) shutdown() { close(d.exited) } } - -func exitCode(ps *os.ProcessState) int { - if ps == nil { - return 1 - } - if ws, ok := ps.Sys().(syscall.WaitStatus); ok && ws.Signaled() { - return 128 + int(ws.Signal()) - } - return ps.ExitCode() -} - -func shellEnv(env []string, term, session string) []string { - out := make([]string, 0, len(env)+2) - for _, kv := range env { - if strings.HasPrefix(kv, "TERM=") || strings.HasPrefix(kv, "HQSH_SESSION=") { - continue - } - out = append(out, kv) - } - return append(out, "TERM="+term, "HQSH_SESSION="+session) -} - -// usableTerm keeps the client's TERM when this host has a terminfo entry for -// it (xterm-kitty often is missing), else falls back to xterm-256color. -func usableTerm(term string) string { - const fallback = "xterm-256color" - if term == "" || strings.ContainsAny(term, "/\\\x00") || strings.HasPrefix(term, ".") { - return fallback - } - var dirs []string - if t := os.Getenv("TERMINFO"); t != "" { - dirs = append(dirs, t) - } - if home, err := os.UserHomeDir(); err == nil { - dirs = append(dirs, filepath.Join(home, ".terminfo")) - } - if td := os.Getenv("TERMINFO_DIRS"); td != "" { - dirs = append(dirs, filepath.SplitList(td)...) - } - dirs = append(dirs, "/etc/terminfo", "/lib/terminfo", "/usr/share/terminfo", "/usr/lib/terminfo", "/usr/share/lib/terminfo", "/opt/homebrew/share/terminfo", "/usr/local/share/terminfo") - anyDir := false - for _, dir := range dirs { - if st, err := os.Stat(dir); err != nil || !st.IsDir() { - continue - } - anyDir = true - for _, sub := range []string{term[:1], fmt.Sprintf("%x", term[0])} { - if _, err := os.Stat(filepath.Join(dir, sub, term)); err == nil { - return term - } - } - } - if !anyDir { - return term // no terminfo database to check against; trust the client - } - return fallback -} diff --git a/internal/server/helpers_test.go b/internal/server/helpers_test.go new file mode 100644 index 0000000..077c9b4 --- /dev/null +++ b/internal/server/helpers_test.go @@ -0,0 +1,216 @@ +package server + +// Protocol-level test clients, shared by the Unix and Windows tests. + +import ( + "bytes" + "io" + "strings" + "testing" + "time" + + "github.com/profullstack/hqsh/internal/proto" +) + +// pipeConn is one end of an in-process attach: what ssh would carry. +type pipeConn struct { + r *io.PipeReader + w *io.PipeWriter +} + +func (p *pipeConn) Read(b []byte) (int, error) { return p.r.Read(b) } +func (p *pipeConn) Write(b []byte) (int, error) { return p.w.Write(b) } +func (p *pipeConn) Close() error { + p.w.Close() + p.r.Close() + return nil +} + +// attachPipe runs AttachIO the way `ssh host hqsh server attach` would. +func attachPipe(session string) *pipeConn { + cr, sw := io.Pipe() // server -> client + sr, cw := io.Pipe() // client -> server + go func() { + _ = AttachIO(session, sr, sw) + sw.Close() + sr.Close() + }() + return &pipeConn{r: cr, w: cw} +} + +// peer is a protocol-level client. +type peer struct { + t *testing.T + c *pipeConn + frames chan proto.Frame + out bytes.Buffer + seqs []uint64 + // until never matches before mark (the end of its last match) and has + // already searched up to scanned. + mark, scanned int +} + +func dialPeer(t *testing.T, session string, lastSeq uint64) (*peer, proto.WelcomeMsg) { + t.Helper() + return dialHello(t, proto.HelloMsg{Version: proto.Version, LastSeq: lastSeq, Cols: 80, Rows: 24, Session: session, Term: "xterm-256color"}) +} + +func dialHello(t *testing.T, h proto.HelloMsg) (*peer, proto.WelcomeMsg) { + t.Helper() + session := h.Session + p := &peer{t: t, c: attachPipe(session), frames: make(chan proto.Frame, 256)} + go func() { + defer close(p.frames) + for { + f, err := proto.Read(p.c) + if err != nil { + return + } + p.frames <- f + } + }() + p.send(proto.Hello, h.Encode()) + f := p.next() + if f.Type != proto.Welcome { + t.Fatalf("first frame %d, want WELCOME", f.Type) + } + w, err := proto.DecodeWelcome(f.Payload) + if err != nil { + t.Fatal(err) + } + return p, w +} + +func (p *peer) send(typ proto.Type, payload []byte) { + p.t.Helper() + if err := proto.Write(p.c, proto.Frame{Type: typ, Payload: payload}); err != nil { + p.t.Fatalf("send %d: %v", typ, err) + } +} + +func (p *peer) next() proto.Frame { + p.t.Helper() + select { + case f, ok := <-p.frames: + if !ok { + p.t.Fatal("connection closed") + } + return f + case <-time.After(10 * time.Second): + p.t.Fatalf("timed out; output so far: %q", p.out.String()) + } + return proto.Frame{} +} + +// until reads OUTPUT until the accumulated text contains want (searching +// only what arrived since the last match, so megabytes stay cheap). +func (p *peer) until(want string) { + p.t.Helper() + for { + from := max(p.mark, p.scanned-len(want)+1) + if i := strings.Index(p.out.String()[from:], want); i >= 0 { + p.mark = from + i + len(want) + p.scanned = p.mark + return + } + p.scanned = p.out.Len() + f := p.next() + if f.Type != proto.Output { + continue + } + o, err := proto.DecodeOutput(f.Payload) + if err != nil { + p.t.Fatal(err) + } + p.seqs = append(p.seqs, o.Seq) + p.out.Write(o.Data) + } +} + +// closed waits for the daemon to drop p, returning the frame types it got +// first. +func (p *peer) closed(within time.Duration) []proto.Type { + p.t.Helper() + var types []proto.Type + deadline := time.After(within) + for { + select { + case f, ok := <-p.frames: + if !ok { + return types + } + types = append(types, f.Type) + case <-deadline: + p.t.Fatalf("not dropped within %v", within) + } + } +} + +// detachedAndClosed: a --steal elsewhere sent p DETACHED and dropped it. +func (p *peer) detachedAndClosed() { + p.t.Helper() + types := p.closed(5 * time.Second) + if len(types) == 0 || types[len(types)-1] != proto.Detached { + p.t.Fatalf("frames before the drop %v, want DETACHED last", types) + } +} + +// clients is how many clients `server list` reports for name. +func clients(t *testing.T, name string) int { + t.Helper() + ss, err := List() + if err != nil { + t.Fatal(err) + } + for _, s := range ss { + if s.Name == name { + return s.Clients + } + } + return -1 +} + +func waitClients(t *testing.T, name string, want int) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for { + got := clients(t, name) + if got == want { + return + } + if time.Now().After(deadline) { + t.Fatalf("%d clients listed, want %d", got, want) + } + time.Sleep(20 * time.Millisecond) + } +} + +func (p *peer) last() uint64 { + if len(p.seqs) == 0 { + return 0 + } + return p.seqs[len(p.seqs)-1] +} + +func contiguous(t *testing.T, seqs []uint64, from uint64) { + t.Helper() + for i, s := range seqs { + if s != from+uint64(i) { + t.Fatalf("seqs %v are not contiguous from %d", seqs, from) + } + } +} + +func listed(t *testing.T, name string) (found, attached bool) { + t.Helper() + ss, err := List() + if err != nil { + t.Fatal(err) + } + for _, s := range ss { + if s.Name == name { + return true, s.Attached + } + } + return false, false +} diff --git a/internal/server/integration_test.go b/internal/server/integration_test.go index 76c3ae2..3ad8bda 100644 --- a/internal/server/integration_test.go +++ b/internal/server/integration_test.go @@ -71,209 +71,6 @@ func sandbox(t *testing.T) { t.Setenv("PS1", "$ ") } -// pipeConn is one end of an in-process attach: what ssh would carry. -type pipeConn struct { - r *io.PipeReader - w *io.PipeWriter -} - -func (p *pipeConn) Read(b []byte) (int, error) { return p.r.Read(b) } -func (p *pipeConn) Write(b []byte) (int, error) { return p.w.Write(b) } -func (p *pipeConn) Close() error { - p.w.Close() - p.r.Close() - return nil -} - -// attachPipe runs AttachIO the way `ssh host hqsh server attach` would. -func attachPipe(session string) *pipeConn { - cr, sw := io.Pipe() // server -> client - sr, cw := io.Pipe() // client -> server - go func() { - _ = AttachIO(session, sr, sw) - sw.Close() - sr.Close() - }() - return &pipeConn{r: cr, w: cw} -} - -// peer is a protocol-level client. -type peer struct { - t *testing.T - c *pipeConn - frames chan proto.Frame - out bytes.Buffer - seqs []uint64 - // until never matches before mark (the end of its last match) and has - // already searched up to scanned. - mark, scanned int -} - -func dialPeer(t *testing.T, session string, lastSeq uint64) (*peer, proto.WelcomeMsg) { - t.Helper() - return dialHello(t, proto.HelloMsg{Version: proto.Version, LastSeq: lastSeq, Cols: 80, Rows: 24, Session: session, Term: "xterm-256color"}) -} - -func dialHello(t *testing.T, h proto.HelloMsg) (*peer, proto.WelcomeMsg) { - t.Helper() - session := h.Session - p := &peer{t: t, c: attachPipe(session), frames: make(chan proto.Frame, 256)} - go func() { - defer close(p.frames) - for { - f, err := proto.Read(p.c) - if err != nil { - return - } - p.frames <- f - } - }() - p.send(proto.Hello, h.Encode()) - f := p.next() - if f.Type != proto.Welcome { - t.Fatalf("first frame %d, want WELCOME", f.Type) - } - w, err := proto.DecodeWelcome(f.Payload) - if err != nil { - t.Fatal(err) - } - return p, w -} - -func (p *peer) send(typ proto.Type, payload []byte) { - p.t.Helper() - if err := proto.Write(p.c, proto.Frame{Type: typ, Payload: payload}); err != nil { - p.t.Fatalf("send %d: %v", typ, err) - } -} - -func (p *peer) next() proto.Frame { - p.t.Helper() - select { - case f, ok := <-p.frames: - if !ok { - p.t.Fatal("connection closed") - } - return f - case <-time.After(10 * time.Second): - p.t.Fatalf("timed out; output so far: %q", p.out.String()) - } - return proto.Frame{} -} - -// until reads OUTPUT until the accumulated text contains want (searching -// only what arrived since the last match, so megabytes stay cheap). -func (p *peer) until(want string) { - p.t.Helper() - for { - from := max(p.mark, p.scanned-len(want)+1) - if i := strings.Index(p.out.String()[from:], want); i >= 0 { - p.mark = from + i + len(want) - p.scanned = p.mark - return - } - p.scanned = p.out.Len() - f := p.next() - if f.Type != proto.Output { - continue - } - o, err := proto.DecodeOutput(f.Payload) - if err != nil { - p.t.Fatal(err) - } - p.seqs = append(p.seqs, o.Seq) - p.out.Write(o.Data) - } -} - -// closed waits for the daemon to drop p, returning the frame types it got -// first. -func (p *peer) closed(within time.Duration) []proto.Type { - p.t.Helper() - var types []proto.Type - deadline := time.After(within) - for { - select { - case f, ok := <-p.frames: - if !ok { - return types - } - types = append(types, f.Type) - case <-deadline: - p.t.Fatalf("not dropped within %v", within) - } - } -} - -// detachedAndClosed: a --steal elsewhere sent p DETACHED and dropped it. -func (p *peer) detachedAndClosed() { - p.t.Helper() - types := p.closed(5 * time.Second) - if len(types) == 0 || types[len(types)-1] != proto.Detached { - p.t.Fatalf("frames before the drop %v, want DETACHED last", types) - } -} - -// clients is how many clients `server list` reports for name. -func clients(t *testing.T, name string) int { - t.Helper() - ss, err := List() - if err != nil { - t.Fatal(err) - } - for _, s := range ss { - if s.Name == name { - return s.Clients - } - } - return -1 -} - -func waitClients(t *testing.T, name string, want int) { - t.Helper() - deadline := time.Now().Add(5 * time.Second) - for { - got := clients(t, name) - if got == want { - return - } - if time.Now().After(deadline) { - t.Fatalf("%d clients listed, want %d", got, want) - } - time.Sleep(20 * time.Millisecond) - } -} - -func (p *peer) last() uint64 { - if len(p.seqs) == 0 { - return 0 - } - return p.seqs[len(p.seqs)-1] -} - -func contiguous(t *testing.T, seqs []uint64, from uint64) { - t.Helper() - for i, s := range seqs { - if s != from+uint64(i) { - t.Fatalf("seqs %v are not contiguous from %d", seqs, from) - } - } -} - -func listed(t *testing.T, name string) (found, attached bool) { - t.Helper() - ss, err := List() - if err != nil { - t.Fatal(err) - } - for _, s := range ss { - if s.Name == name { - return true, s.Attached - } - } - return false, false -} - func TestDaemonResumesWithoutLossOrDuplicates(t *testing.T) { sandbox(t) const session = "t1" diff --git a/internal/server/server_windows.go b/internal/server/server_windows.go deleted file mode 100644 index 42d3409..0000000 --- a/internal/server/server_windows.go +++ /dev/null @@ -1,26 +0,0 @@ -//go:build windows - -package server - -import ( - "errors" - "io" -) - -// ErrUnsupported: the server side needs a Unix PTY. -var ErrUnsupported = errors.New("hqsh: the server side runs on Linux and macOS only") - -// Attach is Unix-only. -func Attach(session string) error { return ErrUnsupported } - -// AttachIO is Unix-only. -func AttachIO(session string, in io.Reader, out io.Writer) error { return ErrUnsupported } - -// Daemon is Unix-only. -func Daemon(session string) error { return ErrUnsupported } - -// List is Unix-only. -func List() ([]SessionInfo, error) { return nil, ErrUnsupported } - -// Setup is Unix-only. -func Setup(w io.Writer, check bool) error { return ErrUnsupported } diff --git a/internal/server/supervise_test.go b/internal/server/supervise_test.go index ed4f6b0..a2a870a 100644 --- a/internal/server/supervise_test.go +++ b/internal/server/supervise_test.go @@ -155,10 +155,17 @@ func TestSessionRunsUnderTheRealServiceManager(t *testing.T) { t.Fatalf("unit %s not active: %v %s", UnitName(session), err, out) } case SupervisorLaunchd: - out, err := runCmd("launchctl", "print", launchdDomain()+"/"+LaunchdLabel(session)) + d, ok := launchdUsed.Load(session) + if !ok { + // It fell back to a fork: say why. + cmd, _ := daemonCommand("svc-probe") + t.Fatalf("launchd did not take the session: %v", startLaunchd("svc-probe", cmd)) + } + out, err := runCmd("launchctl", "print", d.(string)+"/"+LaunchdLabel(session)) if err != nil || !strings.Contains(string(out), "state = running") { - t.Fatalf("launchd job not running: %v\n%s", err, out) + t.Fatalf("launchd job not running in %s: %v\n%s", d, err, out) } + t.Logf("session ran as a launchd job in %s", d) } c.send(proto.Input, []byte("echo under-"+r.Supervisor+"\n")) c.until("\nunder-" + r.Supervisor + "\r\n") diff --git a/internal/server/supervise_unix.go b/internal/server/supervise_unix.go index e16e04a..0270849 100644 --- a/internal/server/supervise_unix.go +++ b/internal/server/supervise_unix.go @@ -12,6 +12,7 @@ import ( "path/filepath" "runtime" "strings" + "sync" "syscall" ) @@ -300,12 +301,21 @@ func startSystemd(session string, cmd *exec.Cmd) error { // ---------------------------------------------------------------- launchd --- -func launchdDomain() string { - // Over ssh there is no GUI login, so the per-user domain is the one that - // exists; it survives the ssh connection. - return fmt.Sprintf("user/%d", os.Getuid()) +// launchdDomains are the domains a session job may go in, best first: the +// GUI login's when the user has one (it is where their own agents run), +// else the per-user background domain, which exists for ssh-only users too. +// Both outlive the ssh connection. +func launchdDomains() []string { + gui := fmt.Sprintf("gui/%d", os.Getuid()) + user := fmt.Sprintf("user/%d", os.Getuid()) + if _, err := runCmd("launchctl", "print", gui); err == nil { + return []string{gui, user} + } + return []string{user} } +func launchdDomain() string { return launchdDomains()[0] } + // LaunchdLabel is the launchd job label for a session. func LaunchdLabel(session string) string { return "sh.hqterm.hqsh." + strings.TrimSuffix(strings.TrimPrefix(UnitName(session), "hqsh-"), ".service") @@ -370,11 +380,19 @@ func startLaunchd(session string, cmd *exec.Cmd) error { if err := os.WriteFile(plist, []byte(launchdPlist(label, argv, env)), 0o600); err != nil { return err } - domain := launchdDomain() - // A finished job stays loaded; unload the old one so the new one runs. - _, _ = runCmd("launchctl", "bootout", domain+"/"+label) - if out, err := runCmd("launchctl", "bootstrap", domain, plist); err != nil { - return fmt.Errorf("launchctl bootstrap: %v: %s", err, strings.TrimSpace(string(out))) + var errs []string + for _, domain := range launchdDomains() { + // A finished job stays loaded; unload the old one so the new one runs. + _, _ = runCmd("launchctl", "bootout", domain+"/"+label) + out, err := runCmd("launchctl", "bootstrap", domain, plist) + if err == nil { + launchdUsed.Store(session, domain) + return nil + } + errs = append(errs, fmt.Sprintf("%s: %v: %s", domain, err, strings.TrimSpace(string(out)))) } - return nil + return fmt.Errorf("launchctl bootstrap: %s", strings.Join(errs, "; ")) } + +// launchdUsed remembers the domain each session's job went into. +var launchdUsed sync.Map diff --git a/internal/server/supervise_windows.go b/internal/server/supervise_windows.go new file mode 100644 index 0000000..01fe5b7 --- /dev/null +++ b/internal/server/supervise_windows.go @@ -0,0 +1,102 @@ +//go:build windows + +package server + +import ( + "fmt" + "io" + "os" + "os/exec" + "strings" + "syscall" + + "golang.org/x/sys/windows" +) + +// On Windows the session daemon has to escape the ssh connection's job +// object: Win32-OpenSSH runs each connection in a job that kills everything +// in it when the connection closes. In order of preference: +// +// - breakaway: CREATE_BREAKAWAY_FROM_JOB, when the job allows it; +// - wmi: Win32_Process.Create, whose processes are created by the WMI +// service and so belong to no ssh job; +// - detached: a plain detached process, which lives only as long as the +// connection's job does (still fine for a session you stay attached to). +// +// HQSH_SUPERVISOR=fork forces the plain detached process. +const ( + SupervisorFork = "fork" + SupervisorBreakaway = "breakaway" + SupervisorWMI = "wmi" +) + +const detachFlags = windows.DETACHED_PROCESS | windows.CREATE_NEW_PROCESS_GROUP + +func startDaemon(session string) error { + cmd, err := daemonCommand(session) + if err != nil { + return err + } + if os.Getenv("HQSH_SUPERVISOR") != SupervisorFork { + if err := startDetached(cmd, detachFlags|windows.CREATE_BREAKAWAY_FROM_JOB); err == nil { + return nil + } + if err := startWMI(cmd); err == nil { + return nil + } + // Rebuild: an exec.Cmd cannot be started twice. + if cmd, err = daemonCommand(session); err != nil { + return err + } + } + return startDetached(cmd, detachFlags) +} + +func startDetached(cmd *exec.Cmd, flags uint32) error { + cmd.SysProcAttr = &syscall.SysProcAttr{CreationFlags: flags, HideWindow: true} + if err := cmd.Start(); err != nil { + return fmt.Errorf("hqsh: starting the daemon: %w", err) + } + go func() { _ = cmd.Wait() }() + return nil +} + +// WMICommandLine is the command line handed to Win32_Process.Create. +func WMICommandLine(argv []string) string { + q := make([]string, len(argv)) + for i, a := range argv { + q[i] = windows.EscapeArg(a) + } + return strings.Join(q, " ") +} + +// startWMI creates the daemon through WMI. The new process gets the user's +// default environment rather than ours, which is fine for a real daemon: +// its socket directory comes from the profile, as attach's does. +func startWMI(cmd *exec.Cmd) error { + line := WMICommandLine(append([]string{cmd.Path}, cmd.Args[1:]...)) + ps := exec.Command("powershell.exe", "-NoProfile", "-NonInteractive", "-Command", + "$r = Invoke-CimMethod -ClassName Win32_Process -MethodName Create -Arguments @{CommandLine=$env:HQSH_WMI_CMDLINE}; exit [int]$r.ReturnValue") + ps.Env = append(os.Environ(), "HQSH_WMI_CMDLINE="+line) + ps.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} + out, err := ps.CombinedOutput() + if err != nil { + return fmt.Errorf("Win32_Process.Create: %v: %s", err, strings.TrimSpace(string(out))) + } + return nil +} + +// Setup reports how sessions start on this Windows host. There is nothing to +// enable: no service, port or scheduled task. +func Setup(w io.Writer, check bool) error { + how := "breakaway from the ssh job, else WMI (Win32_Process.Create), else a detached process" + if os.Getenv("HQSH_SUPERVISOR") == SupervisorFork { + how = "a detached process (HQSH_SUPERVISOR=fork)" + } + fmt.Fprintf(w, "supervisor: %s\n", how) + fmt.Fprintf(w, "shell: %s (HQSH_SHELL overrides)\n", windowsShell()) + if dir, err := SocketDir(); err == nil { + fmt.Fprintf(w, "sockets: %s\n", dir) + } + return nil +} diff --git a/internal/server/windows_test.go b/internal/server/windows_test.go new file mode 100644 index 0000000..12387f7 --- /dev/null +++ b/internal/server/windows_test.go @@ -0,0 +1,132 @@ +//go:build windows + +package server + +import ( + "fmt" + "os" + "os/exec" + "strings" + "testing" + + "github.com/profullstack/hqsh/internal/proto" + "golang.org/x/sys/windows" +) + +// As on Unix, attach re-execs this test binary as the daemon. +func TestMain(m *testing.M) { + if os.Getenv("HQSH_TEST_DAEMON") == "1" { + if err := Daemon(os.Getenv("HQSH_TEST_SESSION")); err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + os.Exit(0) + } + os.Setenv("HQSH_SUPERVISOR", SupervisorFork) + daemonCommand = func(session string) (*exec.Cmd, error) { + cmd := exec.Command(os.Args[0]) + cmd.Env = append(os.Environ(), "HQSH_TEST_DAEMON=1", "HQSH_TEST_SESSION="+session) + return cmd, nil + } + os.Exit(m.Run()) +} + +func winSandbox(t *testing.T) { + t.Helper() + dir, err := os.MkdirTemp("", "hqsh-") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { os.RemoveAll(dir) }) + t.Setenv("XDG_RUNTIME_DIR", dir) + t.Setenv("HQSH_SHELL", os.Getenv("ComSpec")) +} + +// A cmd.exe session on a real ConPTY: it starts, computes, resizes, survives +// a disconnect, is listed, and its exit status comes back. +func TestWindowsSessionOnConPTY(t *testing.T) { + winSandbox(t) + const session = "win" + c1, w := dialPeer(t, session, 0) + if w.Gap() { + t.Fatal("fresh session reported a gap") + } + c1.until(">") // the prompt + c1.send(proto.Input, []byte("set /a 6*7\r")) + c1.until("42") + c1.send(proto.Resize, proto.EncodeResize(100, 30)) + if found, attached := listed(t, session); !found || !attached { + t.Fatalf("list while attached: found=%v attached=%v", found, attached) + } + last := c1.last() + c1.c.Close() + + c2, w := dialPeer(t, session, last) + if w.Gap() { + t.Fatalf("resume reported a gap: %+v", w) + } + c2.send(proto.Input, []byte("set /a 7*8\r")) + c2.until("56") + c2.send(proto.Input, []byte("exit 7\r")) + for { + f := c2.next() + if f.Type == proto.Exit { + code, _ := proto.DecodeExit(f.Payload) + if code != 7 { + t.Fatalf("exit status %d, want 7", code) + } + return + } + } +} + +func TestWMICommandLineQuotes(t *testing.T) { + got := WMICommandLine([]string{`C:\Program Files\hqsh\hqsh.exe`, "server", "daemon", "main"}) + if got != `"C:\Program Files\hqsh\hqsh.exe" server daemon main` { + t.Fatalf("got %s", got) + } +} + +func TestEnvBlockIsDoubleNulTerminated(t *testing.T) { + b := envBlock([]string{"A=1", "B=2"}) + if s := windows.UTF16ToString(b); s != "A=1" { + t.Fatalf("first entry %q", s) + } + if b[len(b)-1] != 0 || b[len(b)-2] != 0 { + t.Fatal("not NUL-NUL terminated") + } + if n := strings.Count(string(windowsUTF16(b)), "\x00"); n != 3 { + t.Fatalf("%d NULs, want 3", n) + } +} + +func windowsUTF16(b []uint16) []rune { + r := make([]rune, len(b)) + for i, u := range b { + r[i] = rune(u) + } + return r +} + +// The WMI path (what runs where the ssh job refuses breakaway, as on CI) +// starts a process outside our job. A harmless command: a WMI-started +// process gets the user's default environment, not this test's. +func TestWMIStartsAProcess(t *testing.T) { + cmd := exec.Command(os.Getenv("ComSpec"), "/c", "exit", "0") + if err := startWMI(cmd); err != nil { + t.Fatalf("Win32_Process.Create: %v", err) + } +} + +// Whether this machine's job lets the daemon break away; informational (CI +// runners differ), the start must work either way. +func TestDaemonStartsWithoutTheForkOverride(t *testing.T) { + winSandbox(t) + t.Setenv("HQSH_SUPERVISOR", SupervisorBreakaway) + cmd, _ := daemonCommand("probe") + if err := startDetached(cmd, detachFlags|windows.CREATE_BREAKAWAY_FROM_JOB); err != nil { + t.Logf("breakaway refused here (%v); sessions use WMI", err) + } else { + t.Log("breakaway allowed here") + } +}