424 lines
11 KiB
Go
424 lines
11 KiB
Go
// Package remote provides the SSH bootstrap and JSON-RPC client for
|
|
// split-brain remote execution. It connects to a remote host over SSH,
|
|
// ensures ollie-remote is deployed, and returns a Server that
|
|
// forwards execution calls over the RPC channel.
|
|
package toolsrv
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"compress/gzip"
|
|
"context"
|
|
"crypto/sha256"
|
|
_ "embed"
|
|
"encoding/base64"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"os"
|
|
"os/exec"
|
|
"path/filepath"
|
|
"strings"
|
|
"sync"
|
|
"sync/atomic"
|
|
"syscall"
|
|
"time"
|
|
|
|
)
|
|
|
|
//go:embed bootstrap.sh
|
|
var bootstrapTemplate string
|
|
|
|
// HostInfo holds environment details from the remote host.
|
|
type HostInfo struct {
|
|
Platform string `json:"platform"`
|
|
Arch string `json:"arch"`
|
|
IsGitRepo bool `json:"is_git_repo"`
|
|
}
|
|
|
|
// Server implements Server by forwarding calls to a remote
|
|
// ollie-remote process over SSH.
|
|
type RemoteServer struct {
|
|
mu sync.Mutex
|
|
stdin io.WriteCloser
|
|
stdout io.ReadCloser
|
|
cmd *exec.Cmd
|
|
enc *json.Encoder
|
|
dec *json.Decoder
|
|
nextID atomic.Int64
|
|
Info HostInfo
|
|
}
|
|
|
|
// Config holds the parameters for connecting to a remote host.
|
|
type RemoteConfig struct {
|
|
// SSHTarget is the SSH destination (e.g., "user@host" or an SSH config alias).
|
|
SSHTarget string
|
|
// CWD is the working directory on the remote host.
|
|
CWD string
|
|
// ToolsPath overrides the remote tool scripts path.
|
|
ToolsPath string
|
|
// Yolo disables sandbox enforcement on the remote.
|
|
Yolo bool
|
|
}
|
|
|
|
// Dial opens an SSH connection to the remote host, bootstraps ollie-remote
|
|
// (deploying the binary if needed), and returns a Server ready for tool calls.
|
|
func RemoteDial(ctx context.Context, cfg RemoteConfig) (*RemoteServer, error) {
|
|
// Find the local ollie-remote binary to compute hash and transfer if needed.
|
|
localBin, err := findLocalBinary()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("find ollie-remote binary: %w", err)
|
|
}
|
|
binData, err := os.ReadFile(localBin)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("read ollie-remote binary: %w", err)
|
|
}
|
|
hash := sha256.Sum256(binData)
|
|
hashStr := hex.EncodeToString(hash[:])
|
|
|
|
// Generate bootstrap script with hash and CWD baked in.
|
|
bootstrap := bootstrapTemplate
|
|
bootstrap = strings.ReplaceAll(bootstrap, "@@HASH@@", hashStr)
|
|
cwdEscaped := shellEscape(cfg.CWD)
|
|
extraArgs := ""
|
|
if cfg.ToolsPath != "" {
|
|
extraArgs += " --tools " + shellEscape(cfg.ToolsPath)
|
|
}
|
|
if cfg.Yolo {
|
|
extraArgs += " --yolo"
|
|
}
|
|
bootstrap = strings.ReplaceAll(bootstrap, "@@CWD@@", cwdEscaped+extraArgs)
|
|
|
|
// Parse optional port from SSHTarget (user@host:port or host:port)
|
|
sshArgs := []string{}
|
|
sshTarget := cfg.SSHTarget
|
|
if host, port, ok := strings.Cut(sshTarget, ":"); ok && port != "" {
|
|
sshTarget = host
|
|
sshArgs = append(sshArgs, "-p", port)
|
|
}
|
|
// Run bash on the remote; we'll send the bootstrap script over stdin.
|
|
sshArgs = append(sshArgs, sshTarget, "bash -s")
|
|
|
|
cmd := exec.CommandContext(ctx, "ssh", sshArgs...)
|
|
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
|
|
|
stdin, err := cmd.StdinPipe()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("ssh stdin pipe: %w", err)
|
|
}
|
|
stdout, err := cmd.StdoutPipe()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("ssh stdout pipe: %w", err)
|
|
}
|
|
stderrPipe, err := cmd.StderrPipe()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("ssh stderr pipe: %w", err)
|
|
}
|
|
|
|
if err := cmd.Start(); err != nil {
|
|
return nil, fmt.Errorf("ssh start: %w", err)
|
|
}
|
|
|
|
// Send bootstrap script over stdin.
|
|
if _, err := stdin.Write([]byte(bootstrap + "\n")); err != nil {
|
|
cmd.Process.Kill()
|
|
return nil, fmt.Errorf("send bootstrap: %w", err)
|
|
}
|
|
|
|
// Monitor stderr for bootstrap messages.
|
|
type loaderMsg struct {
|
|
NeedDownload bool `json:"need_download"`
|
|
}
|
|
|
|
stderrScanner := bufio.NewScanner(stderrPipe)
|
|
loaderCh := make(chan loaderMsg, 1)
|
|
readyCh := make(chan struct{}, 1)
|
|
go func() {
|
|
for stderrScanner.Scan() {
|
|
line := strings.TrimSpace(stderrScanner.Text())
|
|
if after, ok := strings.CutPrefix(line, "OLLIE_LOADER_START "); ok {
|
|
var msg loaderMsg
|
|
json.Unmarshal([]byte(after), &msg)
|
|
loaderCh <- msg
|
|
if !msg.NeedDownload {
|
|
close(readyCh)
|
|
}
|
|
} else if line == "OLLIE_LOADER_READY" {
|
|
close(readyCh)
|
|
}
|
|
}
|
|
}()
|
|
|
|
// Wait for loader start message.
|
|
select {
|
|
case msg := <-loaderCh:
|
|
if msg.NeedDownload {
|
|
if err := transferBinary(stdin, binData); err != nil {
|
|
cmd.Process.Kill()
|
|
return nil, fmt.Errorf("transfer binary: %w", err)
|
|
}
|
|
// Wait for ready signal.
|
|
select {
|
|
case <-readyCh:
|
|
case <-time.After(30 * time.Second):
|
|
cmd.Process.Kill()
|
|
return nil, fmt.Errorf("bootstrap timeout waiting for ready")
|
|
}
|
|
}
|
|
case <-time.After(30 * time.Second):
|
|
cmd.Process.Kill()
|
|
return nil, fmt.Errorf("bootstrap timeout waiting for loader start")
|
|
}
|
|
|
|
// Bootstrap complete — stdin/stdout are now JSON-RPC.
|
|
s := &RemoteServer{
|
|
stdin: stdin,
|
|
stdout: stdout,
|
|
cmd: cmd,
|
|
enc: json.NewEncoder(stdin),
|
|
dec: json.NewDecoder(bufio.NewReader(stdout)),
|
|
}
|
|
|
|
// Verify connectivity with a ping.
|
|
if err := s.ping(ctx); err != nil {
|
|
s.Close()
|
|
return nil, fmt.Errorf("remote ping failed: %w", err)
|
|
}
|
|
|
|
// Fetch remote host info for prompt resolution.
|
|
// FIXME: This duplicates information that prime scripts also try to detect
|
|
// locally (platform, is_git_repo). Eventually unify so prompt resolution
|
|
// always uses these env vars rather than running local checks against a
|
|
// non-existent remote CWD.
|
|
if info, err := s.fetchHostInfo(ctx); err == nil {
|
|
s.Info = info
|
|
}
|
|
|
|
return s, nil
|
|
}
|
|
|
|
// findLocalBinary locates the ollie-remote binary on the local system.
|
|
func findLocalBinary() (string, error) {
|
|
// Check common locations in priority order.
|
|
home, _ := os.UserHomeDir()
|
|
candidates := []string{
|
|
filepath.Join(home, "bin", "ollie-remote"),
|
|
filepath.Join(home, ".local", "bin", "ollie-remote"),
|
|
}
|
|
for _, p := range candidates {
|
|
if _, err := os.Stat(p); err == nil {
|
|
return p, nil
|
|
}
|
|
}
|
|
// Try PATH as last resort.
|
|
if p, err := exec.LookPath("ollie-remote"); err == nil {
|
|
return p, nil
|
|
}
|
|
return "", fmt.Errorf("ollie-remote not found (checked %s)", strings.Join(candidates, ", "))
|
|
}
|
|
|
|
// transferBinary sends the ollie-remote binary gzipped+base64 over stdin.
|
|
func transferBinary(stdin io.Writer, binData []byte) error {
|
|
var compressed bytes.Buffer
|
|
gzw := gzip.NewWriter(&compressed)
|
|
if _, err := gzw.Write(binData); err != nil {
|
|
return err
|
|
}
|
|
if err := gzw.Close(); err != nil {
|
|
return err
|
|
}
|
|
encoded := base64.StdEncoding.EncodeToString(compressed.Bytes())
|
|
if _, err := fmt.Fprintf(stdin, "OLLIE_DOWNLOAD %d\n", len(encoded)); err != nil {
|
|
return err
|
|
}
|
|
if _, err := stdin.Write([]byte(encoded)); err != nil {
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Close shuts down the SSH connection.
|
|
func (s *RemoteServer) Close() error {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
if s.stdin != nil {
|
|
s.stdin.Close()
|
|
}
|
|
if s.cmd != nil && s.cmd.Process != nil {
|
|
syscall.Kill(-s.cmd.Process.Pid, syscall.SIGTERM)
|
|
s.cmd.Wait()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// --- Server interface ---
|
|
|
|
func (s *RemoteServer) ListTools() ([]ToolInfo, error) {
|
|
id := s.nextID.Add(1)
|
|
req := rpcRequest{
|
|
JSONRPC: "2.0",
|
|
ID: id,
|
|
Method: "list_tools",
|
|
}
|
|
|
|
s.mu.Lock()
|
|
if err := s.enc.Encode(req); err != nil {
|
|
s.mu.Unlock()
|
|
return nil, fmt.Errorf("remote list_tools write: %w", err)
|
|
}
|
|
var resp rpcResponse
|
|
if err := s.dec.Decode(&resp); err != nil {
|
|
s.mu.Unlock()
|
|
return nil, fmt.Errorf("remote list_tools read: %w", err)
|
|
}
|
|
s.mu.Unlock()
|
|
|
|
if resp.Error != nil {
|
|
return nil, fmt.Errorf("remote list_tools: %s", resp.Error.Message)
|
|
}
|
|
|
|
var infos []ToolInfo
|
|
if err := json.Unmarshal(resp.Result, &infos); err != nil {
|
|
return nil, fmt.Errorf("remote list_tools unmarshal: %w", err)
|
|
}
|
|
return infos, nil
|
|
}
|
|
|
|
func (s *RemoteServer) CallTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) {
|
|
// Forward the call with the tool name as the RPC method.
|
|
id := s.nextID.Add(1)
|
|
|
|
req := rpcRequest{
|
|
JSONRPC: "2.0",
|
|
ID: id,
|
|
Method: tool,
|
|
Params: args,
|
|
}
|
|
|
|
s.mu.Lock()
|
|
if err := s.enc.Encode(req); err != nil {
|
|
s.mu.Unlock()
|
|
return nil, fmt.Errorf("remote rpc write: %w", err)
|
|
}
|
|
|
|
// Read responses, handling interleaved streaming notifications.
|
|
for {
|
|
var resp rpcResponse
|
|
if err := s.dec.Decode(&resp); err != nil {
|
|
s.mu.Unlock()
|
|
return nil, fmt.Errorf("remote rpc read: %w", err)
|
|
}
|
|
|
|
// Notification: id == 0 and method is "output".
|
|
// (JSON-RPC notifications have no id; our struct decodes missing id as 0.)
|
|
if resp.ID == 0 && resp.Result != nil {
|
|
// Streaming output notification — emit to context callback.
|
|
var notif outputNotification
|
|
if json.Unmarshal(resp.Result, ¬if) == nil && notif.Data != "" {
|
|
StreamOutput(ctx, notif.Data)
|
|
}
|
|
continue
|
|
}
|
|
|
|
// Actual response for our request.
|
|
s.mu.Unlock()
|
|
if resp.Error != nil {
|
|
return nil, fmt.Errorf("remote error: %s", resp.Error.Message)
|
|
}
|
|
return resp.Result, nil
|
|
}
|
|
}
|
|
|
|
type outputNotification struct {
|
|
Data string `json:"data"`
|
|
}
|
|
|
|
// fetchHostInfo retrieves environment details from the remote host.
|
|
func (s *RemoteServer) fetchHostInfo(ctx context.Context) (HostInfo, error) {
|
|
id := s.nextID.Add(1)
|
|
req := rpcRequest{
|
|
JSONRPC: "2.0",
|
|
ID: id,
|
|
Method: "host_info",
|
|
}
|
|
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
|
|
if err := s.enc.Encode(req); err != nil {
|
|
return HostInfo{}, err
|
|
}
|
|
var resp rpcResponse
|
|
if err := s.dec.Decode(&resp); err != nil {
|
|
return HostInfo{}, err
|
|
}
|
|
if resp.Error != nil {
|
|
return HostInfo{}, fmt.Errorf("%s", resp.Error.Message)
|
|
}
|
|
var info HostInfo
|
|
if err := json.Unmarshal(resp.Result, &info); err != nil {
|
|
return HostInfo{}, err
|
|
}
|
|
return info, nil
|
|
}
|
|
|
|
// ping verifies the remote server is responsive.
|
|
func (s *RemoteServer) ping(ctx context.Context) error {
|
|
id := s.nextID.Add(1)
|
|
req := rpcRequest{
|
|
JSONRPC: "2.0",
|
|
ID: id,
|
|
Method: "ping",
|
|
}
|
|
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
|
|
if err := s.enc.Encode(req); err != nil {
|
|
return err
|
|
}
|
|
var resp rpcResponse
|
|
if err := s.dec.Decode(&resp); err != nil {
|
|
return err
|
|
}
|
|
if resp.Error != nil {
|
|
return fmt.Errorf("%s", resp.Error.Message)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// --- JSON-RPC types ---
|
|
|
|
type rpcRequest struct {
|
|
JSONRPC string `json:"jsonrpc"`
|
|
ID int64 `json:"id"`
|
|
Method string `json:"method"`
|
|
Params json.RawMessage `json:"params,omitempty"`
|
|
}
|
|
|
|
type rpcResponse struct {
|
|
JSONRPC string `json:"jsonrpc"`
|
|
ID int64 `json:"id"`
|
|
Result json.RawMessage `json:"result,omitempty"`
|
|
Error *rpcError `json:"error,omitempty"`
|
|
}
|
|
|
|
type rpcError struct {
|
|
Code int `json:"code"`
|
|
Message string `json:"message"`
|
|
}
|
|
|
|
// --- helpers ---
|
|
|
|
func shellEscape(s string) string {
|
|
if s == "" {
|
|
return "''"
|
|
}
|
|
if !strings.ContainsAny(s, " \t\n\r'\"\\$`!#&|;(){}") {
|
|
return s
|
|
}
|
|
return "'" + strings.ReplaceAll(s, "'", "'\\''") + "'"
|
|
}
|
|
|