This repository has been archived on 2026-08-16. You can view files and clone it, but cannot push or open issues or pull requests.
ollie-core/execute/remote.go

460 lines
12 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 tools.Server that
// forwards execution calls over the RPC channel.
package execute
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"
"ollie/tools"
)
//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 tools.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
}
// --- tools.Server interface ---
func (s *RemoteServer) ListTools() ([]tools.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 []tools.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, &notif) == nil && notif.Data != "" {
tools.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, "'", "'\\''") + "'"
}
// Decl returns a factory function compatible with tools.NewDispatcherFunc.
// It dials the remote on first call and returns the Server.
func RemoteDecl(cfg RemoteConfig) func() tools.Server {
var (
once sync.Once
server *RemoteServer
err error
)
return func() tools.Server {
once.Do(func() {
server, err = RemoteDial(context.Background(), cfg)
if err != nil {
// Return a stub that errors on every call
server = nil
}
})
if server == nil {
return &errServer{err: err}
}
return server
}
}
// errServer is a tools.Server that returns an error for every call.
type errServer struct {
err error
}
func (e *errServer) ListTools() ([]tools.ToolInfo, error) {
return nil, e.err
}
func (e *errServer) CallTool(_ context.Context, _ string, _ json.RawMessage) (json.RawMessage, error) {
return nil, e.err
}