364 lines
9.3 KiB
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
|
|
}
|