384 lines
10 KiB
Go
384 lines
10 KiB
Go
// Package exec handles tool execution in sandboxed environments.
|
|
package exec
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"os"
|
|
osExec "os/exec"
|
|
"path/filepath"
|
|
"strings"
|
|
"syscall"
|
|
"time"
|
|
|
|
"ollie/cmd/toolsrv/internal/bypass"
|
|
"ollie/cmd/toolsrv/internal/sandbox"
|
|
"ollie/toolsrv/metadata"
|
|
"ollie/toolsrv/protocol"
|
|
"ollie/util"
|
|
)
|
|
|
|
const outputLimit = 10 * 1024 * 1024 // 10 MiB
|
|
|
|
// StartResult reports whether an external command started.
|
|
type StartResult struct {
|
|
Process *os.Process
|
|
Err error
|
|
}
|
|
|
|
// Config contains all context needed for tool execution.
|
|
type Config struct {
|
|
CWD string
|
|
Env map[string]string
|
|
Yolo bool
|
|
Timeout int // 0 = no timeout
|
|
Output io.Writer // if set, stream stdout/stderr here in real-time
|
|
Started chan StartResult // receives exactly one startup result if set
|
|
}
|
|
|
|
// ExecuteTool runs a tool script with the given args and returns the result.
|
|
func ExecuteTool(ctx context.Context, info protocol.ToolInfo, args json.RawMessage, cfg Config) (json.RawMessage, error) {
|
|
// Resolve script path
|
|
toolPath, err := metadata.ResolveTool(info.Name)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("resolve tool %s: %w", info.Name, err)
|
|
}
|
|
|
|
// Extract dispatch-level flags from args
|
|
bypassed := false
|
|
timeout := cfg.Timeout // 0 = no timeout; caller sets default
|
|
|
|
var argMap map[string]interface{}
|
|
if err := json.Unmarshal(args, &argMap); err == nil {
|
|
if e, ok := argMap["bypass"]; ok {
|
|
switch v := e.(type) {
|
|
case bool:
|
|
bypassed = v
|
|
case string:
|
|
bypassed = v == "true" || v == "1"
|
|
}
|
|
delete(argMap, "bypass")
|
|
}
|
|
if t, ok := argMap["timeout"]; ok {
|
|
switch v := t.(type) {
|
|
case float64:
|
|
timeout = int(v)
|
|
case string:
|
|
var n int
|
|
if _, err := fmt.Sscanf(v, "%d", &n); err == nil {
|
|
timeout = n
|
|
}
|
|
}
|
|
delete(argMap, "timeout")
|
|
}
|
|
delete(argMap, "sandbox") // ignored, single config
|
|
args, _ = json.Marshal(argMap)
|
|
}
|
|
|
|
stdinData := string(args)
|
|
cwd := cfg.CWD
|
|
if cwd == "" {
|
|
cwd, _ = os.Getwd()
|
|
}
|
|
|
|
var result string
|
|
if bypassed {
|
|
bypassCode := fmt.Sprintf("cat <<'OLLIE_EOF' | %s\n%s\nOLLIE_EOF", toolPath, stdinData)
|
|
result, err = executeBypassDirect(ctx, bypassCode, cwd, cfg.Env, timeout, cfg.Output, cfg.Started)
|
|
} else {
|
|
result, err = executeSandboxed(ctx, toolPath, stdinData, cwd, cfg.Env, timeout, cfg.Yolo, cfg.Output, cfg.Started)
|
|
}
|
|
|
|
if err != nil {
|
|
return json.Marshal(protocol.ToolResult{
|
|
IsError: true,
|
|
Content: []protocol.ToolResultContent{{Type: "text", Text: result + ": " + err.Error()}},
|
|
})
|
|
}
|
|
return json.Marshal(protocol.ToolResult{
|
|
Content: []protocol.ToolResultContent{{Type: "text", Text: result}},
|
|
})
|
|
}
|
|
|
|
// executeSandboxed runs a tool script inside the sandbox.
|
|
func executeSandboxed(ctx context.Context, toolPath, stdinData, cwd string, envExtra map[string]string, timeout int, yolo bool, streamOut io.Writer, started chan StartResult) (string, error) {
|
|
startedSent := false
|
|
defer func() {
|
|
if started != nil && !startedSent {
|
|
started <- StartResult{Err: fmt.Errorf("execution did not start")}
|
|
}
|
|
}()
|
|
|
|
var sandboxCfg sandbox.Sandbox
|
|
if !yolo {
|
|
cfgPath := filepath.Join(util.CfgDir(), "sandbox.yaml")
|
|
f, err := os.Open(cfgPath)
|
|
if err != nil {
|
|
return "", fmt.Errorf("sandbox config not found: %w", err)
|
|
}
|
|
defer f.Close()
|
|
cfg, loadErr := sandbox.LoadSandbox(f)
|
|
if loadErr != nil {
|
|
return "", loadErr
|
|
}
|
|
sandboxCfg = cfg
|
|
}
|
|
|
|
if timeout > 0 {
|
|
var cancel context.CancelFunc
|
|
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
|
defer cancel()
|
|
}
|
|
|
|
interpreter := []string{"bash", "-c", toolPath}
|
|
|
|
// Build environment
|
|
envMap := make(map[string]string)
|
|
for _, ev := range os.Environ() {
|
|
k, v, _ := strings.Cut(ev, "=")
|
|
envMap[k] = v
|
|
}
|
|
for k, v := range envExtra {
|
|
envMap[k] = v
|
|
}
|
|
// Prepend tools directory to PATH so meta-only cmd fields can reference other tools
|
|
if toolsDir := metadata.ToolsPath(); toolsDir != "" {
|
|
if existing := envMap["PATH"]; existing != "" {
|
|
envMap["PATH"] = toolsDir + ":" + existing
|
|
} else {
|
|
envMap["PATH"] = toolsDir
|
|
}
|
|
}
|
|
// Compute NAMESPACE if not set
|
|
if _, ok := envMap["NAMESPACE"]; !ok {
|
|
if ns := Plan9Namespace(envMap); ns != "" {
|
|
envMap["NAMESPACE"] = ns
|
|
}
|
|
}
|
|
// Neutralize repo-controlled git config (GitSpawn / core.fsmonitor):
|
|
// tools run with CWD inside a repository whose .git/config may name
|
|
// commands (fsmonitor) that execute on any index refresh. GIT_CONFIG_*
|
|
// acts like `-c core.fsmonitor=false` and overrides repo/global config.
|
|
envMap["GIT_CONFIG_COUNT"] = "1"
|
|
envMap["GIT_CONFIG_KEY_0"] = "core.fsmonitor"
|
|
envMap["GIT_CONFIG_VALUE_0"] = "false"
|
|
|
|
var cmd *osExec.Cmd
|
|
if yolo {
|
|
cmd = osExec.CommandContext(ctx, interpreter[0], interpreter[1:]...)
|
|
} else {
|
|
wrapped, wrapErr := sandboxCfg.Command(interpreter, cwd, envMap)
|
|
if wrapErr != nil {
|
|
return "", wrapErr
|
|
}
|
|
cmd = osExec.CommandContext(ctx, wrapped[0], wrapped[1:]...)
|
|
}
|
|
|
|
cmd.Dir = cwd
|
|
if stdinData != "" {
|
|
cmd.Stdin = strings.NewReader(stdinData)
|
|
}
|
|
|
|
// Set environment
|
|
for k, v := range envMap {
|
|
cmd.Env = append(cmd.Env, k+"="+v)
|
|
}
|
|
|
|
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true, Pdeathsig: syscall.SIGKILL}
|
|
cmd.Cancel = func() error {
|
|
if cmd.Process != nil {
|
|
syscall.Kill(-cmd.Process.Pid, syscall.SIGTERM)
|
|
}
|
|
return nil
|
|
}
|
|
cmd.WaitDelay = 5 * time.Second
|
|
|
|
var outputBuf bytes.Buffer
|
|
var w io.Writer = &outputBuf
|
|
if streamOut != nil {
|
|
w = io.MultiWriter(&outputBuf, streamOut)
|
|
}
|
|
lw := &limitedWriter{
|
|
w: w,
|
|
limit: outputLimit,
|
|
}
|
|
cmd.Stdout = lw
|
|
cmd.Stderr = lw
|
|
|
|
if err := cmd.Start(); err != nil {
|
|
if started != nil {
|
|
started <- StartResult{Err: fmt.Errorf("execution failed: %w", err)}
|
|
startedSent = true
|
|
}
|
|
return "", fmt.Errorf("execution failed: %w", err)
|
|
}
|
|
|
|
if started != nil {
|
|
started <- StartResult{Process: cmd.Process}
|
|
startedSent = true
|
|
}
|
|
|
|
err := cmd.Wait()
|
|
output := outputBuf.Bytes()
|
|
if lw.truncated {
|
|
output = append(output, []byte("\n[output truncated at 10MB]")...)
|
|
}
|
|
|
|
if ctx.Err() == context.DeadlineExceeded {
|
|
return string(output), fmt.Errorf("execution timeout after %d seconds", timeout)
|
|
}
|
|
if err != nil {
|
|
errOutput := string(output)
|
|
const maxErrOutput = 8192
|
|
if len(errOutput) > maxErrOutput {
|
|
errOutput = errOutput[:maxErrOutput] + fmt.Sprintf("\n[...truncated %d bytes in error]", len(output)-maxErrOutput)
|
|
}
|
|
return string(output), fmt.Errorf("execution failed: %w\nOutput: %s", err, errOutput)
|
|
}
|
|
return string(output), nil
|
|
}
|
|
|
|
// executeBypassDirect requests bypass approval and executes the command directly.
|
|
// The approval comes from olliesrv via the 9P bypass/pending and bypass/resolve files.
|
|
func executeBypassDirect(ctx context.Context, cmd, cwd string, envExtra map[string]string, timeout int, streamOut io.Writer, started chan StartResult) (string, error) {
|
|
// Build environment map
|
|
envMap := make(map[string]string)
|
|
for _, kv := range os.Environ() {
|
|
if k, v, ok := strings.Cut(kv, "="); ok {
|
|
envMap[k] = v
|
|
}
|
|
}
|
|
for k, v := range envExtra {
|
|
envMap[k] = v
|
|
}
|
|
|
|
// For background procs, signal "started" immediately so the caller unblocks
|
|
// and the agent sees the proc ID. The actual OS process is stored later.
|
|
// We send nil Process here; it will be updated when the command actually starts.
|
|
if started != nil {
|
|
started <- StartResult{Process: nil}
|
|
started = nil // prevent double-send
|
|
}
|
|
|
|
// Submit bypass request and wait for approval
|
|
approved, err := bypass.Submit(ctx, cmd, cwd, envMap)
|
|
if err != nil {
|
|
return "", fmt.Errorf("bypass request failed: %w", err)
|
|
}
|
|
if !approved {
|
|
return "", fmt.Errorf("bypass denied")
|
|
}
|
|
|
|
// Approved - execute directly without sandbox
|
|
if timeout > 0 {
|
|
var cancel context.CancelFunc
|
|
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
|
defer cancel()
|
|
}
|
|
|
|
return executeDirectUnsandboxed(ctx, cmd, cwd, envMap, streamOut, nil)
|
|
}
|
|
|
|
// executeDirectUnsandboxed runs a command without any sandbox.
|
|
func executeDirectUnsandboxed(ctx context.Context, cmd, cwd string, env map[string]string, streamOut io.Writer, started chan StartResult) (string, error) {
|
|
execCmd := osExec.CommandContext(ctx, "bash", "-c", cmd)
|
|
execCmd.Dir = cwd
|
|
if len(env) > 0 {
|
|
execCmd.Env = make([]string, 0, len(env))
|
|
for k, v := range env {
|
|
execCmd.Env = append(execCmd.Env, k+"="+v)
|
|
}
|
|
}
|
|
|
|
// Set process group so we can signal all children
|
|
execCmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true, Pdeathsig: syscall.SIGKILL}
|
|
execCmd.Cancel = func() error {
|
|
if execCmd.Process != nil {
|
|
syscall.Kill(-execCmd.Process.Pid, syscall.SIGTERM)
|
|
}
|
|
return nil
|
|
}
|
|
execCmd.WaitDelay = 5 * time.Second
|
|
|
|
var stdout, stderr bytes.Buffer
|
|
if streamOut != nil {
|
|
// Streaming mode: tee to both the stream and local buffer
|
|
execCmd.Stdout = io.MultiWriter(&stdout, streamOut)
|
|
execCmd.Stderr = io.MultiWriter(&stderr, streamOut)
|
|
} else {
|
|
execCmd.Stdout = &stdout
|
|
execCmd.Stderr = &stderr
|
|
}
|
|
|
|
if err := execCmd.Start(); err != nil {
|
|
if started != nil {
|
|
started <- StartResult{Err: fmt.Errorf("bypass execution failed: %w", err)}
|
|
}
|
|
return "", err
|
|
}
|
|
if started != nil {
|
|
started <- StartResult{Process: execCmd.Process}
|
|
}
|
|
err := execCmd.Wait()
|
|
output := stdout.String()
|
|
if stderr.Len() > 0 {
|
|
if output != "" {
|
|
output += "\n"
|
|
}
|
|
output += stderr.String()
|
|
}
|
|
|
|
if err != nil {
|
|
if exitErr, ok := err.(*osExec.ExitError); ok {
|
|
return output, fmt.Errorf("exit %d", exitErr.ExitCode())
|
|
}
|
|
return output, err
|
|
}
|
|
return output, nil
|
|
}
|
|
|
|
// Plan9Namespace computes the Plan 9 namespace directory.
|
|
func Plan9Namespace(env map[string]string) string {
|
|
if ns := env["NAMESPACE"]; ns != "" {
|
|
return ns
|
|
}
|
|
disp := env["DISPLAY"]
|
|
if disp == "" {
|
|
return "/tmp/ns." + env["USER"] + "." + "unix"
|
|
}
|
|
// Strip leading "localhost" if present
|
|
disp = strings.TrimPrefix(disp, "localhost")
|
|
disp = strings.ReplaceAll(disp, "/", "_")
|
|
return "/tmp/ns." + env["USER"] + "." + disp
|
|
}
|
|
|
|
// limitedWriter wraps a writer with a size limit.
|
|
type limitedWriter struct {
|
|
w io.Writer
|
|
limit int
|
|
written int
|
|
truncated bool
|
|
}
|
|
|
|
func (lw *limitedWriter) Write(p []byte) (int, error) {
|
|
if lw.truncated {
|
|
return len(p), nil // discard but report success
|
|
}
|
|
remaining := lw.limit - lw.written
|
|
if remaining <= 0 {
|
|
lw.truncated = true
|
|
return len(p), nil
|
|
}
|
|
toWrite := p
|
|
if len(p) > remaining {
|
|
toWrite = p[:remaining]
|
|
lw.truncated = true
|
|
}
|
|
n, err := lw.w.Write(toWrite)
|
|
lw.written += n
|
|
return len(p), err
|
|
}
|