ollie/cmd/toolsrv/internal/exec/exec.go

364 lines
9.3 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"
)
// 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 *os.Process // if set, receives the os.Process after Start()
}
// StreamFunc returns the streaming output function from context, if any.
func StreamFunc(ctx context.Context) func(string) {
if fn, ok := ctx.Value(streamKey{}).(func(string)); ok {
return fn
}
return nil
}
type streamKey struct{}
// WithStreamFunc attaches a streaming output function to context.
func WithStreamFunc(ctx context.Context, fn func(string)) context.Context {
return context.WithValue(ctx, streamKey{}, fn)
}
// 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)
} 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 *os.Process) (string, error) {
var sandboxCfg *sandbox.Config
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()
var loadErr error
sandboxCfg, loadErr = sandbox.LoadSandbox(f)
if loadErr != nil {
return "", loadErr
}
}
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
}
}
getenv := func(key string) string { return envMap[key] }
var cmd *osExec.Cmd
if yolo {
cmd = osExec.CommandContext(ctx, interpreter[0], interpreter[1:]...)
} else {
wrapped, wrapErr := sandbox.WrapCommand(sandboxCfg, interpreter, cwd, getenv)
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: 10 * 1024 * 1024,
stream: StreamFunc(ctx),
}
cmd.Stdout = lw
cmd.Stderr = lw
if err := cmd.Start(); err != nil {
if started != nil {
close(started)
}
return "", fmt.Errorf("execution failed: %w", err)
}
if started != nil {
started <- cmd.Process
}
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) (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
}
// 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)
}
// executeDirectUnsandboxed runs a command without any sandbox.
func executeDirectUnsandboxed(ctx context.Context, cmd, cwd string, env map[string]string) (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)
}
}
var stdout, stderr bytes.Buffer
execCmd.Stdout = &stdout
execCmd.Stderr = &stderr
// Stream output if context has a stream function
if stream := StreamFunc(ctx); stream != nil {
execCmd.Stdout = io.MultiWriter(&stdout, &streamWriter{stream})
execCmd.Stderr = io.MultiWriter(&stderr, &streamWriter{stream})
}
err := execCmd.Run()
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
}
// streamWriter adapts a stream function to io.Writer.
type streamWriter struct {
fn func(string)
}
func (w *streamWriter) Write(p []byte) (int, error) {
w.fn(string(p))
return len(p), 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 and optional streaming.
type limitedWriter struct {
w io.Writer
limit int
written int
truncated bool
stream func(string)
}
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
if lw.stream != nil && n > 0 {
lw.stream(string(toWrite[:n]))
}
return len(p), err
}