396 lines
9.9 KiB
Go
396 lines
9.9 KiB
Go
// Package exec handles tool execution in sandboxed environments.
|
|
package exec
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/binary"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"os"
|
|
"os/exec"
|
|
"path/filepath"
|
|
"strings"
|
|
"syscall"
|
|
"time"
|
|
|
|
"ollie/cmd/toolsrv/internal/sandbox"
|
|
"ollie/paths"
|
|
"ollie/toolsrv"
|
|
)
|
|
|
|
// Config contains all context needed for tool execution.
|
|
type Config struct {
|
|
CWD string
|
|
Env map[string]string
|
|
Yolo bool
|
|
Timeout int
|
|
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 toolsrv.ToolInfo, args json.RawMessage, cfg Config) (json.RawMessage, error) {
|
|
// Resolve script path
|
|
toolPath, err := toolsrv.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
|
|
timeoutExplicit := timeout > 0
|
|
if !timeoutExplicit {
|
|
timeout = 30
|
|
}
|
|
sandboxName := "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 {
|
|
timeoutExplicit = true
|
|
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")
|
|
}
|
|
if s, ok := argMap["sandbox"]; ok {
|
|
if v, ok := s.(string); ok && v != "" {
|
|
sandboxName = v
|
|
}
|
|
delete(argMap, "sandbox")
|
|
}
|
|
args, _ = json.Marshal(argMap)
|
|
}
|
|
|
|
// Check if tool requires sudo
|
|
needsSudo := false
|
|
if m, err := toolsrv.LoadMetaFile(info.Name); err == nil && m != nil {
|
|
resolved := m.Resolve()
|
|
if resolved != nil && resolved.Sudo {
|
|
needsSudo = true
|
|
bypassed = true
|
|
}
|
|
}
|
|
|
|
stdinData := string(args)
|
|
cwd := cfg.CWD
|
|
if cwd == "" {
|
|
cwd, _ = os.Getwd()
|
|
}
|
|
|
|
var result string
|
|
if needsSudo {
|
|
sudoCode := fmt.Sprintf("cat <<'OLLIE_EOF' | %s\n%s\nOLLIE_EOF", toolPath, stdinData)
|
|
result, err = executeBypassDirect(ctx, sudoCode, cwd, cfg.Env, timeout, true)
|
|
} else if bypassed {
|
|
bypassCode := fmt.Sprintf("cat <<'OLLIE_EOF' | %s\n%s\nOLLIE_EOF", toolPath, stdinData)
|
|
result, err = executeBypassDirect(ctx, bypassCode, cwd, cfg.Env, timeout, false)
|
|
} else {
|
|
result, err = executeSandboxed(ctx, toolPath, stdinData, cwd, cfg.Env, timeout, sandboxName, cfg.Yolo, cfg.Output, cfg.Started)
|
|
}
|
|
|
|
if err != nil {
|
|
return json.Marshal(toolsrv.ToolResult{
|
|
IsError: true,
|
|
Content: []toolsrv.ToolResultContent{{Type: "text", Text: result + ": " + err.Error()}},
|
|
})
|
|
}
|
|
return json.Marshal(toolsrv.ToolResult{
|
|
Content: []toolsrv.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, sandboxName string, yolo bool, streamOut io.Writer, started chan *os.Process) (string, error) {
|
|
var sandboxCfg *sandbox.Config
|
|
if !yolo {
|
|
cfgPath := filepath.Join(paths.CfgDir(), "sandbox", sandboxName+".yaml")
|
|
f, err := os.Open(cfgPath)
|
|
if err != nil {
|
|
return "", fmt.Errorf("sandbox %q not found: %w", sandboxName, 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
|
|
}
|
|
// 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 *exec.Cmd
|
|
if yolo {
|
|
cmd = exec.CommandContext(ctx, interpreter[0], interpreter[1:]...)
|
|
} else {
|
|
wrapped, wrapErr := sandbox.WrapCommand(sandboxCfg, interpreter, cwd, getenv)
|
|
if wrapErr != nil {
|
|
return "", wrapErr
|
|
}
|
|
cmd = exec.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}
|
|
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 runs a command via the bypass broker.
|
|
func executeBypassDirect(ctx context.Context, cmd, cwd string, envExtra map[string]string, timeout int, sudo bool) (string, error) {
|
|
xdg := os.Getenv("XDG_RUNTIME_DIR")
|
|
if xdg == "" {
|
|
return "", fmt.Errorf("bypass not available: no XDG_RUNTIME_DIR")
|
|
}
|
|
sockPath := filepath.Join(xdg, "ollie", "bypass.sock")
|
|
|
|
if timeout > 0 {
|
|
var cancel context.CancelFunc
|
|
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
|
defer cancel()
|
|
}
|
|
|
|
// 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
|
|
}
|
|
|
|
// Connect and send request
|
|
conn, err := net.DialTimeout("unix", sockPath, 5*time.Second)
|
|
if err != nil {
|
|
return "", fmt.Errorf("bypass not available: %w", err)
|
|
}
|
|
defer conn.Close()
|
|
|
|
reqJSON, _ := json.Marshal(struct {
|
|
Cmd string `json:"cmd"`
|
|
Cwd string `json:"cwd"`
|
|
Env map[string]string `json:"env"`
|
|
Sudo bool `json:"sudo,omitempty"`
|
|
}{Cmd: cmd, Cwd: cwd, Env: envMap, Sudo: sudo})
|
|
reqJSON = append(reqJSON, '\n')
|
|
|
|
if _, err := conn.Write(reqJSON); err != nil {
|
|
return "", fmt.Errorf("bypass write failed: %w", err)
|
|
}
|
|
|
|
// Read response frames
|
|
var outputBuf bytes.Buffer
|
|
exitCode := readBypassFrames(ctx, conn, &outputBuf, StreamFunc(ctx))
|
|
|
|
output := outputBuf.String()
|
|
if exitCode != 0 {
|
|
return output, fmt.Errorf("bypass execution failed (exit %d)", exitCode)
|
|
}
|
|
return output, nil
|
|
}
|
|
|
|
// readBypassFrames reads framed output from bypass broker.
|
|
func readBypassFrames(ctx context.Context, conn net.Conn, w *bytes.Buffer, stream func(string)) int {
|
|
header := make([]byte, 5)
|
|
for {
|
|
conn.SetReadDeadline(time.Now().Add(1 * time.Second))
|
|
_, err := io.ReadFull(conn, header)
|
|
if err != nil {
|
|
if ne, ok := err.(net.Error); ok && ne.Timeout() {
|
|
select {
|
|
case <-ctx.Done():
|
|
return -1
|
|
default:
|
|
continue
|
|
}
|
|
}
|
|
return -1
|
|
}
|
|
|
|
frameType := header[0]
|
|
length := binary.BigEndian.Uint32(header[1:5])
|
|
|
|
payload := make([]byte, length)
|
|
if length > 0 {
|
|
conn.SetReadDeadline(time.Now().Add(30 * time.Second))
|
|
if _, err := io.ReadFull(conn, payload); err != nil {
|
|
return -1
|
|
}
|
|
}
|
|
|
|
switch frameType {
|
|
case 'd':
|
|
w.Write(payload)
|
|
if stream != nil {
|
|
stream(string(payload))
|
|
}
|
|
case 'x':
|
|
var code int
|
|
fmt.Sscanf(string(payload), "%d", &code)
|
|
return code
|
|
}
|
|
}
|
|
}
|
|
|
|
// 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
|
|
}
|