517 lines
12 KiB
Go
517 lines
12 KiB
Go
// spawn.go - Process spawning and lifecycle management for toolsrvclient.
|
|
package toolclient
|
|
|
|
import (
|
|
"bufio"
|
|
"context"
|
|
"crypto/rand"
|
|
"encoding/hex"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"os"
|
|
"os/exec"
|
|
"path/filepath"
|
|
"strings"
|
|
"sync"
|
|
"syscall"
|
|
"time"
|
|
|
|
"ollie/util"
|
|
)
|
|
|
|
// Process represents a running toolsrv process.
|
|
type Process struct {
|
|
Cmd *exec.Cmd
|
|
SocketPath string
|
|
Secret string
|
|
Info ProcessInfo
|
|
cancel context.CancelFunc
|
|
waitOnce sync.Once
|
|
waitErr error
|
|
}
|
|
|
|
// ProcessInfo contains metadata about the toolsrv process.
|
|
type ProcessInfo struct {
|
|
Platform string
|
|
IsGitRepo bool
|
|
}
|
|
|
|
// Option configures spawning behavior.
|
|
type Option func(*spawnConfig)
|
|
|
|
type spawnConfig struct {
|
|
yolo bool
|
|
sessionID string
|
|
}
|
|
|
|
// WithYolo disables sandbox enforcement.
|
|
func WithYolo() Option {
|
|
return func(c *spawnConfig) { c.yolo = true }
|
|
}
|
|
|
|
// WithSessionID sets the session ID for tool registry.
|
|
func WithSessionID(id string) Option {
|
|
return func(c *spawnConfig) { c.sessionID = id }
|
|
}
|
|
|
|
// Spawn starts a local toolsrv process.
|
|
func Spawn(ctx context.Context, cwd string, opts ...Option) (*Process, error) {
|
|
cfg := &spawnConfig{}
|
|
for _, opt := range opts {
|
|
opt(cfg)
|
|
}
|
|
|
|
// Generate socket path and secret
|
|
socketID, err := randomHex(8)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("generate socket path: %w", err)
|
|
}
|
|
socketPath := filepath.Join(os.TempDir(), fmt.Sprintf("toolsrv-%d-%s.sock", os.Getpid(), socketID))
|
|
secret, err := randomHex(32)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("generate toolsrv secret: %w", err)
|
|
}
|
|
|
|
// Find toolsrv binary
|
|
toolsrvPath, err := findToolsrv()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("find toolsrv: %w", err)
|
|
}
|
|
|
|
// Build command
|
|
args := []string{"serve", "--cwd", cwd, "--listen", socketPath}
|
|
if cfg.yolo {
|
|
args = append(args, "--yolo")
|
|
}
|
|
if cfg.sessionID != "" {
|
|
args = append(args, "--session-id", cfg.sessionID)
|
|
}
|
|
|
|
procCtx, cancel := context.WithCancel(ctx)
|
|
cmd := exec.CommandContext(procCtx, toolsrvPath, args...)
|
|
// Note: secret is NOT passed via env var. It's established on first
|
|
// connection via 9P Tauth (socket permissions + SSH handle security).
|
|
cmd.Stderr = os.Stderr // Let server errors go to stderr for debugging
|
|
cmd.SysProcAttr = &syscall.SysProcAttr{Pdeathsig: syscall.SIGTERM}
|
|
|
|
if err := cmd.Start(); err != nil {
|
|
cancel()
|
|
_ = os.Remove(socketPath)
|
|
return nil, fmt.Errorf("start toolsrv: %w", err)
|
|
}
|
|
|
|
// Wait for socket to be ready
|
|
if err := waitForSocket(ctx, socketPath, 5*time.Second); err != nil {
|
|
cancel()
|
|
_ = cmd.Process.Kill()
|
|
_ = cmd.Wait()
|
|
_ = os.Remove(socketPath)
|
|
return nil, fmt.Errorf("wait for socket: %w", err)
|
|
}
|
|
|
|
return &Process{
|
|
Cmd: cmd,
|
|
SocketPath: socketPath,
|
|
Secret: secret,
|
|
Info: ProcessInfo{
|
|
Platform: "linux",
|
|
IsGitRepo: false, // Caller can override
|
|
},
|
|
cancel: cancel,
|
|
}, nil
|
|
}
|
|
|
|
// Kill terminates the process.
|
|
func (p *Process) Kill() error {
|
|
if p.cancel != nil {
|
|
p.cancel()
|
|
}
|
|
var err error
|
|
if p.Cmd != nil && p.Cmd.Process != nil {
|
|
err = p.Cmd.Process.Kill()
|
|
}
|
|
waitErr := p.Wait()
|
|
if err != nil && !errors.Is(err, os.ErrProcessDone) {
|
|
return err
|
|
}
|
|
return waitErr
|
|
}
|
|
|
|
// Close terminates the process (alias for Kill).
|
|
func (p *Process) Close() error {
|
|
return p.Kill()
|
|
}
|
|
|
|
// Wait waits for the process to exit.
|
|
func (p *Process) Wait() error {
|
|
if p.Cmd == nil {
|
|
return nil
|
|
}
|
|
p.waitOnce.Do(func() {
|
|
p.waitErr = p.Cmd.Wait()
|
|
if p.SocketPath != "" {
|
|
_ = os.Remove(p.SocketPath)
|
|
}
|
|
})
|
|
return p.waitErr
|
|
}
|
|
|
|
// RemoteConfig configures remote toolsrv spawning.
|
|
type RemoteConfig struct {
|
|
SSHTarget string // user@host or host
|
|
CWD string
|
|
SessionID string // session ID for tool registry
|
|
Yolo bool // skip sandbox enforcement
|
|
}
|
|
|
|
// SpawnRemote starts a remote toolsrv process via SSH.
|
|
// Deploys the local configuration directory (including the toolsrv binary and tools/*.meta)
|
|
// to the remote host, then starts toolsrv with SSH Unix socket forwarding. Sandboxing
|
|
// is implemented by the transferred toolsrv binary; no external sandbox executable is copied.
|
|
func SpawnRemote(ctx context.Context, cfg RemoteConfig) (*Process, error) {
|
|
// Verify local config dir exists
|
|
cfgDir := util.CfgDir()
|
|
if _, err := os.Stat(cfgDir); err != nil {
|
|
return nil, fmt.Errorf("local config dir not found at %s", cfgDir)
|
|
}
|
|
|
|
// Generate socket paths
|
|
sockDir := filepath.Join(util.RuntimeDir(), "ollie")
|
|
if err := os.MkdirAll(sockDir, 0700); err != nil {
|
|
return nil, fmt.Errorf("create socket directory: %w", err)
|
|
}
|
|
localID, err := randomHex(8)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("generate local socket path: %w", err)
|
|
}
|
|
remoteID, err := randomHex(8)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("generate remote socket path: %w", err)
|
|
}
|
|
localSock := filepath.Join(sockDir, fmt.Sprintf("remote-%d-%s.sock", os.Getpid(), localID))
|
|
remoteSock := fmt.Sprintf("/tmp/ollie-toolsrv-%d-%s.sock", os.Getpid(), remoteID)
|
|
|
|
// Build remote bootstrap script
|
|
// Extracts tarball to ~/.cache/ollie, sets XDG_CONFIG_HOME so tools/*.meta are found
|
|
var args []string
|
|
args = append(args, "serve", "--cwd", shellEscape(cfg.CWD), "--listen", shellEscape(remoteSock))
|
|
if cfg.SessionID != "" {
|
|
args = append(args, "--session-id", shellEscape(cfg.SessionID))
|
|
}
|
|
if cfg.Yolo {
|
|
args = append(args, "--yolo")
|
|
}
|
|
|
|
bootstrap := fmt.Sprintf(`#!/bin/sh
|
|
set -e
|
|
CACHE_DIR="${XDG_CACHE_HOME:-$HOME/.cache}/ollie"
|
|
mkdir -p "$CACHE_DIR"
|
|
tar xzf - -C "$CACHE_DIR"
|
|
export XDG_CONFIG_HOME="$(dirname "$CACHE_DIR")"
|
|
export PATH="$CACHE_DIR/bin:$PATH"
|
|
exec "$CACHE_DIR/bin/toolsrv" %s
|
|
`, strings.Join(args, " "))
|
|
|
|
// Build SSH command with socket forwarding
|
|
sshArgs := []string{
|
|
"-o", "BatchMode=yes",
|
|
"-o", "StrictHostKeyChecking=accept-new",
|
|
"-o", "StreamLocalBindUnlink=yes",
|
|
"-L", localSock + ":" + remoteSock,
|
|
}
|
|
sshTarget := cfg.SSHTarget
|
|
if host, port, ok := strings.Cut(sshTarget, ":"); ok && port != "" {
|
|
sshTarget = host
|
|
sshArgs = append(sshArgs, "-p", port)
|
|
}
|
|
sshArgs = append(sshArgs, sshTarget, "bash -s")
|
|
|
|
procCtx, cancel := context.WithCancel(ctx)
|
|
cmd := exec.CommandContext(procCtx, "ssh", sshArgs...)
|
|
cmd.SysProcAttr = &syscall.SysProcAttr{
|
|
Setpgid: true,
|
|
Pdeathsig: syscall.SIGTERM,
|
|
}
|
|
|
|
stdin, err := cmd.StdinPipe()
|
|
if err != nil {
|
|
cancel()
|
|
return nil, fmt.Errorf("ssh stdin pipe: %w", err)
|
|
}
|
|
stderrPipe, err := cmd.StderrPipe()
|
|
if err != nil {
|
|
cancel()
|
|
return nil, fmt.Errorf("ssh stderr pipe: %w", err)
|
|
}
|
|
|
|
if err := cmd.Start(); err != nil {
|
|
_ = stdin.Close()
|
|
_ = stderrPipe.Close()
|
|
cancel()
|
|
return nil, fmt.Errorf("ssh start: %w", err)
|
|
}
|
|
|
|
// Send tarball of config dir over stdin
|
|
if err := writeTarball(stdin, cfgDir); err != nil {
|
|
cancel()
|
|
_ = cmd.Process.Kill()
|
|
_ = cmd.Wait()
|
|
_ = os.Remove(localSock)
|
|
return nil, fmt.Errorf("send tarball: %w", err)
|
|
}
|
|
// Send bootstrap script
|
|
if _, err := stdin.Write([]byte(bootstrap)); err != nil {
|
|
cancel()
|
|
_ = cmd.Process.Kill()
|
|
_ = cmd.Wait()
|
|
_ = os.Remove(localSock)
|
|
return nil, fmt.Errorf("send bootstrap: %w", err)
|
|
}
|
|
stdin.Close()
|
|
|
|
// Wait for toolsrv to signal ready while continuing to drain stderr for the
|
|
// lifetime of the SSH process. Leaving stderr undrained can block SSH when
|
|
// the remote server writes enough diagnostics.
|
|
readyCh := make(chan bool, 1)
|
|
stderrDone := make(chan struct{})
|
|
go func() {
|
|
defer close(stderrDone)
|
|
scanner := bufio.NewScanner(stderrPipe)
|
|
ready := false
|
|
for scanner.Scan() {
|
|
line := strings.TrimSpace(scanner.Text())
|
|
if !ready && strings.Contains(line, "listening on") {
|
|
readyCh <- true
|
|
ready = true
|
|
}
|
|
}
|
|
if !ready {
|
|
readyCh <- false
|
|
}
|
|
}()
|
|
|
|
select {
|
|
case ready := <-readyCh:
|
|
if !ready {
|
|
_ = cmd.Wait()
|
|
<-stderrDone
|
|
_ = os.Remove(localSock)
|
|
return nil, fmt.Errorf("remote toolsrv exited before startup")
|
|
}
|
|
case <-time.After(30 * time.Second):
|
|
cancel()
|
|
_ = cmd.Process.Kill()
|
|
_ = cmd.Wait()
|
|
<-stderrDone
|
|
_ = os.Remove(localSock)
|
|
return nil, fmt.Errorf("remote toolsrv startup timeout")
|
|
}
|
|
|
|
// Wait for local socket forwarding to be ready
|
|
if err := waitForSocket(ctx, localSock, 5*time.Second); err != nil {
|
|
cancel()
|
|
_ = cmd.Process.Kill()
|
|
_ = cmd.Wait()
|
|
_ = os.Remove(localSock)
|
|
return nil, fmt.Errorf("wait for forwarded socket: %w", err)
|
|
}
|
|
|
|
return &Process{
|
|
Cmd: cmd,
|
|
SocketPath: localSock,
|
|
Secret: "", // Will be established on first Tauth
|
|
Info: ProcessInfo{
|
|
Platform: "linux", // TODO: detect from remote
|
|
IsGitRepo: false,
|
|
},
|
|
cancel: cancel,
|
|
}, nil
|
|
}
|
|
|
|
// writeTarball creates a gzipped tarball of the directory and writes it to w.
|
|
func writeTarball(w io.Writer, dir string) error {
|
|
cmd := exec.Command("tar", "czf", "-", "-C", dir, ".")
|
|
cmd.Stdout = w
|
|
cmd.Stderr = os.Stderr
|
|
return cmd.Run()
|
|
}
|
|
|
|
// shellEscape escapes a string for safe use in a shell command.
|
|
func shellEscape(s string) string {
|
|
if s == "" {
|
|
return "''"
|
|
}
|
|
if !strings.ContainsAny(s, " \t\n\r'\"\\$`!#&|;(){}") {
|
|
return s
|
|
}
|
|
return "'" + strings.ReplaceAll(s, "'", "'\\''") + "'"
|
|
}
|
|
|
|
// ProcessKeeper manages process lifecycle with automatic respawning.
|
|
type ProcessKeeper struct {
|
|
ctx context.Context
|
|
proc *Process
|
|
respawn func(context.Context) (*Process, error)
|
|
mu sync.Mutex
|
|
dialMu sync.Mutex
|
|
closed bool
|
|
}
|
|
|
|
// NewProcessKeeper creates a new process keeper.
|
|
func NewProcessKeeper(ctx context.Context, proc *Process, respawn func(context.Context) (*Process, error)) *ProcessKeeper {
|
|
return &ProcessKeeper{
|
|
ctx: ctx,
|
|
proc: proc,
|
|
respawn: respawn,
|
|
}
|
|
}
|
|
|
|
// SetContext updates the context used for respawning.
|
|
func (pk *ProcessKeeper) SetContext(ctx context.Context) {
|
|
pk.mu.Lock()
|
|
defer pk.mu.Unlock()
|
|
pk.ctx = ctx
|
|
}
|
|
|
|
// Dial connects to the managed process, respawning if necessary.
|
|
func (pk *ProcessKeeper) Dial() (*ToolsrvConn, error) {
|
|
pk.dialMu.Lock()
|
|
defer pk.dialMu.Unlock()
|
|
|
|
pk.mu.Lock()
|
|
if pk.closed {
|
|
pk.mu.Unlock()
|
|
return nil, fmt.Errorf("process keeper is closed")
|
|
}
|
|
proc := pk.proc
|
|
ctx := pk.ctx
|
|
respawn := pk.respawn
|
|
pk.mu.Unlock()
|
|
|
|
// Try to connect to the existing process without holding the keeper lock.
|
|
if proc != nil {
|
|
conn, err := DialToolsrv(proc.SocketPath, proc.Secret)
|
|
if err == nil {
|
|
pk.mu.Lock()
|
|
closed := pk.closed
|
|
if !closed && proc.Secret == "" {
|
|
proc.Secret = conn.Secret()
|
|
}
|
|
pk.mu.Unlock()
|
|
if closed {
|
|
conn.Close()
|
|
return nil, fmt.Errorf("process keeper is closed")
|
|
}
|
|
return conn, nil
|
|
}
|
|
}
|
|
|
|
if respawn == nil {
|
|
return nil, fmt.Errorf("no process available")
|
|
}
|
|
|
|
// Kill the old process to prevent orphans, outside the keeper lock.
|
|
if proc != nil {
|
|
_ = proc.Kill()
|
|
}
|
|
newProc, err := respawn(ctx)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("respawn: %w", err)
|
|
}
|
|
|
|
pk.mu.Lock()
|
|
if pk.closed {
|
|
pk.mu.Unlock()
|
|
_ = newProc.Kill()
|
|
return nil, fmt.Errorf("process keeper is closed")
|
|
}
|
|
pk.proc = newProc
|
|
pk.mu.Unlock()
|
|
|
|
conn, err := DialToolsrv(newProc.SocketPath, newProc.Secret)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
pk.mu.Lock()
|
|
closed := pk.closed
|
|
pk.mu.Unlock()
|
|
if closed {
|
|
conn.Close()
|
|
return nil, fmt.Errorf("process keeper is closed")
|
|
}
|
|
return conn, nil
|
|
}
|
|
|
|
// Close terminates the managed process.
|
|
func (pk *ProcessKeeper) Close() error {
|
|
pk.mu.Lock()
|
|
if pk.closed {
|
|
pk.mu.Unlock()
|
|
return nil
|
|
}
|
|
pk.closed = true
|
|
proc := pk.proc
|
|
pk.mu.Unlock()
|
|
|
|
if proc != nil {
|
|
return proc.Kill()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// --- Helpers ---
|
|
|
|
func randomHex(n int) (string, error) {
|
|
b := make([]byte, n)
|
|
if _, err := rand.Read(b); err != nil {
|
|
return "", err
|
|
}
|
|
return hex.EncodeToString(b), nil
|
|
}
|
|
|
|
func findToolsrv() (string, error) {
|
|
// Check if toolsrv is in PATH
|
|
if path, err := exec.LookPath("toolsrv"); err == nil {
|
|
return path, nil
|
|
}
|
|
|
|
// Check common locations
|
|
home, _ := os.UserHomeDir()
|
|
candidates := []string{
|
|
filepath.Join(home, ".config", "ollie", "bin", "toolsrv"),
|
|
filepath.Join(home, "go", "bin", "toolsrv"),
|
|
"/usr/local/bin/toolsrv",
|
|
}
|
|
|
|
for _, p := range candidates {
|
|
if _, err := os.Stat(p); err == nil {
|
|
return p, nil
|
|
}
|
|
}
|
|
|
|
return "", fmt.Errorf("toolsrv binary not found in PATH or common locations")
|
|
}
|
|
|
|
func waitForSocket(ctx context.Context, path string, timeout time.Duration) error {
|
|
deadline := time.NewTimer(timeout)
|
|
defer deadline.Stop()
|
|
ticker := time.NewTicker(50 * time.Millisecond)
|
|
defer ticker.Stop()
|
|
for {
|
|
conn, err := net.DialTimeout("unix", path, 100*time.Millisecond)
|
|
if err == nil {
|
|
conn.Close()
|
|
return nil
|
|
}
|
|
select {
|
|
case <-ctx.Done():
|
|
return ctx.Err()
|
|
case <-deadline.C:
|
|
return fmt.Errorf("timeout waiting for socket %s", path)
|
|
case <-ticker.C:
|
|
}
|
|
}
|
|
}
|