ollie/toolsrv/spawn.go

403 lines
9.7 KiB
Go

// spawn.go - Process spawning and lifecycle management for toolsrv.
package toolsrv
import (
"bufio"
"context"
"crypto/rand"
"encoding/hex"
"fmt"
"io"
"net"
"os"
"os/exec"
"path/filepath"
"strings"
"sync"
"syscall"
"time"
"ollie/paths"
)
// Process represents a running toolsrv process.
type Process struct {
Cmd *exec.Cmd
SocketPath string
Secret string
Info ProcessInfo
cancel context.CancelFunc
}
// 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
socketPath := filepath.Join(os.TempDir(), fmt.Sprintf("toolsrv-%d-%s.sock", os.Getpid(), randomHex(8)))
secret := randomHex(32)
// 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
if err := cmd.Start(); err != nil {
cancel()
return nil, fmt.Errorf("start toolsrv: %w", err)
}
// Wait for socket to be ready
if err := waitForSocket(socketPath, 5*time.Second); err != nil {
cancel()
cmd.Process.Kill()
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()
}
if p.Cmd != nil && p.Cmd.Process != nil {
return p.Cmd.Process.Kill()
}
return nil
}
// 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 p.Cmd.Wait()
}
return nil
}
// 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 ~/.config/ollie directory (including tools/*.meta) to the
// remote host, then starts toolsrv with SSH Unix socket forwarding.
func SpawnRemote(ctx context.Context, cfg RemoteConfig) (*Process, error) {
// Verify local config dir exists
cfgDir := paths.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(paths.RuntimeDir(), "ollie")
os.MkdirAll(sockDir, 0700)
localSock := filepath.Join(sockDir, fmt.Sprintf("remote-%d-%s.sock", os.Getpid(), randomHex(8)))
remoteSock := fmt.Sprintf("/tmp/ollie-toolsrv-%d-%s.sock", os.Getpid(), randomHex(8))
// 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 {
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()
return nil, fmt.Errorf("send tarball: %w", err)
}
// Send bootstrap script
if _, err := stdin.Write([]byte(bootstrap)); err != nil {
cancel()
cmd.Process.Kill()
return nil, fmt.Errorf("send bootstrap: %w", err)
}
stdin.Close()
// Wait for toolsrv to signal ready
readyCh := make(chan struct{})
go func() {
scanner := bufio.NewScanner(stderrPipe)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
// toolsrv prints "toolsrv: listening on <path>"
if strings.Contains(line, "listening on") {
close(readyCh)
return
}
}
}()
select {
case <-readyCh:
case <-time.After(30 * time.Second):
cancel()
cmd.Process.Kill()
return nil, fmt.Errorf("remote toolsrv startup timeout")
}
// Wait for local socket forwarding to be ready
if err := waitForSocket(localSock, 5*time.Second); err != nil {
cancel()
cmd.Process.Kill()
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
}
// 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() (*Conn, error) {
pk.mu.Lock()
defer pk.mu.Unlock()
// Try to connect to existing process
if pk.proc != nil {
conn, err := Dial(pk.proc.SocketPath, pk.proc.Secret)
if err == nil {
// Save the secret if this was first connection
if pk.proc.Secret == "" {
pk.proc.Secret = conn.Secret()
}
return conn, nil
}
// Process may have died, try to respawn
}
if pk.respawn != nil {
proc, err := pk.respawn(pk.ctx)
if err != nil {
return nil, fmt.Errorf("respawn: %w", err)
}
pk.proc = proc
conn, err := Dial(proc.SocketPath, proc.Secret)
if err != nil {
return nil, err
}
return conn, nil
}
return nil, fmt.Errorf("no process available")
}
// Close terminates the managed process.
func (pk *ProcessKeeper) Close() error {
pk.mu.Lock()
defer pk.mu.Unlock()
if pk.proc != nil {
return pk.proc.Kill()
}
return nil
}
// --- Helpers ---
func randomHex(n int) string {
b := make([]byte, n)
rand.Read(b)
return hex.EncodeToString(b)
}
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(path string, timeout time.Duration) error {
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
conn, err := net.DialTimeout("unix", path, 100*time.Millisecond)
if err == nil {
conn.Close()
return nil
}
time.Sleep(50 * time.Millisecond)
}
return fmt.Errorf("timeout waiting for socket %s", path)
}