403 lines
9.7 KiB
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)
|
|
}
|