ollie/cmd/olliesrv/internal/toolclient/spawn.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:
}
}
}