ollie/cmd/toolsrv/internal/exec/exec.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
}