tools: decompose code.go into server.go + shell.go
server.go: Server struct, options, lifecycle, dispatch routing, tool_* handlers shell.go: execution engine (sandbox, elevation, detach), validation, limitedWriter
This commit is contained in:
parent
10e0731a0a
commit
873f83464f
231
tools/code.go
231
tools/code.go
|
|
@ -1,231 +0,0 @@
|
|||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
"ollie/sandbox"
|
||||
"ollie/paths"
|
||||
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// universalPatterns apply to all code.
|
||||
var universalPatterns = []*regexp.Regexp{
|
||||
regexp.MustCompile(`\bmkfs\b`),
|
||||
regexp.MustCompile(`\bdd\b.*\bif=/dev/`),
|
||||
regexp.MustCompile(`\b(sudo|su)\s`),
|
||||
regexp.MustCompile(`/etc/(shadow|sudoers)`),
|
||||
}
|
||||
|
||||
// bashPatterns apply to bash (flag syntax, redirects, shell-specific constructs).
|
||||
var bashPatterns = []*regexp.Regexp{
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s*/(home|var|usr|etc|boot|root|bin|sbin|lib|opt|srv)?`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*f[a-z]*\s+)*-[a-z]*r[a-z]*\s*/(home|var|usr|etc|boot|root|bin|sbin|lib|opt|srv)?`),
|
||||
regexp.MustCompile(`rm\s+.*--recursive.*--force`),
|
||||
regexp.MustCompile(`rm\s+.*--force.*--recursive`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+\.\.(/|\s|$)`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+~`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+\*`),
|
||||
regexp.MustCompile(`:\s*\(\s*\)\s*\{\s*:\s*\|\s*:\s*&`), // fork bomb
|
||||
regexp.MustCompile(`>\s*/dev/sd`),
|
||||
regexp.MustCompile(`\beval\s+".*\$`),
|
||||
}
|
||||
|
||||
var languagePatterns = map[string][]*regexp.Regexp{
|
||||
"bash": bashPatterns,
|
||||
"": bashPatterns,
|
||||
}
|
||||
|
||||
// Dispatch routes tool calls to the appropriate handler.
|
||||
func (e *Server) Dispatch(ctx context.Context, name string, args json.RawMessage) (string, error) {
|
||||
if e.OnPreDispatch != nil {
|
||||
e.OnPreDispatch()
|
||||
}
|
||||
|
||||
switch name {
|
||||
case "shell":
|
||||
return dispatchShell(ctx, e, args)
|
||||
case "tool_list":
|
||||
return dispatchToolList(ctx, e, args)
|
||||
case "tool_load":
|
||||
return dispatchToolLoad(ctx, e, args)
|
||||
case "tool_active":
|
||||
return dispatchToolActive(ctx, e, args)
|
||||
case "skill_list":
|
||||
return dispatchSkillList(ctx, e, args)
|
||||
case "skill_load":
|
||||
return dispatchSkillLoad(ctx, e, args)
|
||||
case "skill_active":
|
||||
return dispatchSkillActive(ctx, e, args)
|
||||
default:
|
||||
return "", fmt.Errorf("unknown execute tool: %s", name)
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateCode checks code against dangerous patterns.
|
||||
func (e *Server) ValidateCode(code, language string) error {
|
||||
if err := e.checkRateLimit(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
normalized := strings.ToLower(code)
|
||||
normalized = whitespacePattern.ReplaceAllString(normalized, " ")
|
||||
|
||||
patterns := append(universalPatterns, languagePatterns[language]...)
|
||||
for _, pattern := range patterns {
|
||||
if pattern.MatchString(normalized) {
|
||||
e.recordValidationFailure()
|
||||
return fmt.Errorf("dangerous pattern detected")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadSandboxConfig(name string) (*sandbox.Config, error) {
|
||||
if name == "" {
|
||||
name = "default"
|
||||
}
|
||||
path := filepath.Join(paths.CfgDir(), "sandbox", name+".yaml")
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("sandbox %q not found: %w", name, err)
|
||||
}
|
||||
defer f.Close()
|
||||
return sandbox.LoadSandbox(f)
|
||||
}
|
||||
|
||||
type limitedWriter struct {
|
||||
mu sync.Mutex
|
||||
w io.Writer
|
||||
written int
|
||||
limit int
|
||||
truncated bool
|
||||
stream func(string) // if non-nil, called with each chunk of output
|
||||
}
|
||||
|
||||
func (lw *limitedWriter) Write(p []byte) (n int, err error) {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
if lw.written >= lw.limit {
|
||||
lw.truncated = true
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
remaining := lw.limit - lw.written
|
||||
toWrite := p
|
||||
if len(p) > remaining {
|
||||
toWrite = p[:remaining]
|
||||
lw.truncated = true
|
||||
}
|
||||
|
||||
written, err := lw.w.Write(toWrite)
|
||||
lw.written += written
|
||||
if lw.stream != nil {
|
||||
lw.stream(string(toWrite[:written]))
|
||||
}
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// dispatchShell handles the shell tool: a single bash command.
|
||||
func dispatchShell(ctx context.Context, e *Server, args json.RawMessage) (string, error) {
|
||||
var a struct {
|
||||
Cmd string `json:"cmd"`
|
||||
Timeout int `json:"timeout"`
|
||||
Sandbox string `json:"sandbox"`
|
||||
Elevated bool `json:"elevated"`
|
||||
Detach bool `json:"detach"`
|
||||
}
|
||||
if err := json.Unmarshal(args, &a); err != nil {
|
||||
return "", fmt.Errorf("shell: bad args: %w", err)
|
||||
}
|
||||
if a.Cmd == "" {
|
||||
return "", fmt.Errorf("shell: cmd is required")
|
||||
}
|
||||
timeout := a.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 30
|
||||
}
|
||||
if a.Elevated {
|
||||
e.wdMu.RLock()
|
||||
workDir := e.cwd
|
||||
e.wdMu.RUnlock()
|
||||
return e.executeElevated(ctx, a.Cmd, workDir, timeout, a.Detach)
|
||||
}
|
||||
sandboxName := a.Sandbox
|
||||
if sandboxName == "" {
|
||||
sandboxName = "default"
|
||||
}
|
||||
return e.executeWithStdin(ctx, a.Cmd, "bash", timeout, sandboxName, false, "", a.Detach)
|
||||
}
|
||||
|
||||
// dispatchToolList lists all available tools from the global registry.
|
||||
func dispatchToolList(ctx context.Context, e *Server, args json.RawMessage) (string, error) {
|
||||
if e.toolRegistry == nil {
|
||||
return "", fmt.Errorf("tool_list: no registry available")
|
||||
}
|
||||
summaries := e.toolRegistry.Summaries()
|
||||
if len(summaries) == 0 {
|
||||
return "(no tools found)", nil
|
||||
}
|
||||
var out strings.Builder
|
||||
for _, s := range summaries {
|
||||
out.WriteString(s.Name)
|
||||
if s.Description != "" {
|
||||
out.WriteString(" — ")
|
||||
out.WriteString(s.Description)
|
||||
}
|
||||
out.WriteString("\n")
|
||||
}
|
||||
return strings.TrimRight(out.String(), "\n"), nil
|
||||
}
|
||||
|
||||
// dispatchToolLoad loads a tool into the current session.
|
||||
func dispatchToolLoad(ctx context.Context, e *Server, args json.RawMessage) (string, error) {
|
||||
var a struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
if err := json.Unmarshal(args, &a); err != nil {
|
||||
return "", fmt.Errorf("tool_load: bad args: %w", err)
|
||||
}
|
||||
if a.Name == "" {
|
||||
return "", fmt.Errorf("tool_load: name is required")
|
||||
}
|
||||
if e.toolRegistry == nil || e.sessionID == "" {
|
||||
return "", fmt.Errorf("tool_load: no session registry")
|
||||
}
|
||||
if err := e.toolRegistry.Load(e.sessionID, a.Name); err != nil {
|
||||
return "", fmt.Errorf("tool_load: %w", err)
|
||||
}
|
||||
return fmt.Sprintf("loaded: %s", a.Name), nil
|
||||
}
|
||||
|
||||
// dispatchToolActive lists tools currently loaded (promoted) in this session.
|
||||
func dispatchToolActive(ctx context.Context, e *Server, args json.RawMessage) (string, error) {
|
||||
if e.toolRegistry == nil || e.sessionID == "" {
|
||||
return "(no tools loaded)", nil
|
||||
}
|
||||
loaded := e.toolRegistry.Loaded(e.sessionID)
|
||||
if len(loaded) == 0 {
|
||||
return "(no tools loaded)", nil
|
||||
}
|
||||
var out strings.Builder
|
||||
for _, t := range loaded {
|
||||
out.WriteString(t.Name)
|
||||
if t.Description != "" {
|
||||
out.WriteString(" — ")
|
||||
out.WriteString(t.Description)
|
||||
}
|
||||
out.WriteString("\n")
|
||||
}
|
||||
return strings.TrimRight(out.String(), "\n"), nil
|
||||
}
|
||||
624
tools/server.go
624
tools/server.go
|
|
@ -2,25 +2,17 @@ package tools
|
|||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"ollie/sandbox"
|
||||
"ollie/detach"
|
||||
"ollie/paths"
|
||||
"ollie/skills"
|
||||
"ollie/detach"
|
||||
)
|
||||
|
||||
const (
|
||||
|
|
@ -122,7 +114,6 @@ func (e *Server) AllowTools() []string {
|
|||
return out
|
||||
}
|
||||
|
||||
|
||||
// WithToolRegistry attaches a tool registry and session ID to the Server.
|
||||
func WithToolRegistry(r *Registry, sessionID string) Option {
|
||||
return func(s *Server) {
|
||||
|
|
@ -340,9 +331,6 @@ func (e *Server) callPromotedTool(ctx context.Context, tool string, args json.Ra
|
|||
})
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
// Close is called when the session ends. Calls OnClose hook if registered.
|
||||
func (e *Server) Close() {
|
||||
e.cleanupDetached()
|
||||
|
|
@ -351,571 +339,92 @@ func (e *Server) Close() {
|
|||
}
|
||||
}
|
||||
|
||||
// executeElevated runs cmd outside the sandbox via the integrated elevation broker.
|
||||
// Connects to the broker socket, sends the request with the current env,
|
||||
// and streams the framed response back.
|
||||
func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout int, doDetach ...bool) (string, error) {
|
||||
sockPath := os.Getenv("OLLIE_ELEVATE_SOCKET")
|
||||
if sockPath == "" {
|
||||
xdg := os.Getenv("XDG_RUNTIME_DIR")
|
||||
if xdg == "" {
|
||||
return "", fmt.Errorf("elevation not available: no XDG_RUNTIME_DIR")
|
||||
}
|
||||
sockPath = filepath.Join(xdg, "ollie", "elevate.sock")
|
||||
// Dispatch routes tool calls to the appropriate handler.
|
||||
func (e *Server) Dispatch(ctx context.Context, name string, args json.RawMessage) (string, error) {
|
||||
if e.OnPreDispatch != nil {
|
||||
e.OnPreDispatch()
|
||||
}
|
||||
|
||||
wantDetach := len(doDetach) > 0 && doDetach[0]
|
||||
|
||||
var cancel context.CancelFunc
|
||||
if wantDetach || timeout <= 0 {
|
||||
ctx, cancel = context.WithCancel(ctx)
|
||||
} else {
|
||||
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
||||
}
|
||||
|
||||
// Connect to broker
|
||||
conn, err := net.DialTimeout("unix", sockPath, 5*time.Second)
|
||||
if err != nil {
|
||||
cancel()
|
||||
return "", fmt.Errorf("elevation not available: %w", err)
|
||||
}
|
||||
|
||||
// Send request
|
||||
envMap := make(map[string]string, len(os.Environ()))
|
||||
for _, kv := range os.Environ() {
|
||||
if k, v, ok := strings.Cut(kv, "="); ok {
|
||||
envMap[k] = v
|
||||
}
|
||||
}
|
||||
reqJSON, _ := json.Marshal(struct {
|
||||
Cmd string `json:"cmd"`
|
||||
Cwd string `json:"cwd"`
|
||||
Env map[string]string `json:"env"`
|
||||
Session string `json:"session,omitempty"`
|
||||
}{Cmd: cmd, Cwd: dir, Env: envMap, Session: os.Getenv("OLLIE_SESSION_ID")})
|
||||
reqJSON = append(reqJSON, '\n')
|
||||
if _, err := conn.Write(reqJSON); err != nil {
|
||||
conn.Close()
|
||||
cancel()
|
||||
return "", fmt.Errorf("elevated execution failed: write: %w", err)
|
||||
}
|
||||
|
||||
// Set up detach channel
|
||||
detachCh := make(chan struct{}, 1)
|
||||
e.detachMu.Lock()
|
||||
e.detachCh = detachCh
|
||||
e.detachMu.Unlock()
|
||||
defer func() {
|
||||
e.detachMu.Lock()
|
||||
e.detachCh = nil
|
||||
e.detachMu.Unlock()
|
||||
}()
|
||||
|
||||
if wantDetach {
|
||||
close(detachCh)
|
||||
}
|
||||
|
||||
// readFrames reads from the connection, writing to w and streaming.
|
||||
// Returns exit code when 'x' frame arrives or -1 on error.
|
||||
readFrames := func(w io.Writer, stream func(string)) int {
|
||||
header := make([]byte, 5)
|
||||
for {
|
||||
conn.SetReadDeadline(time.Now().Add(1 * time.Second)) //nolint:errcheck
|
||||
_, err := io.ReadFull(conn, header)
|
||||
if err != nil {
|
||||
if ne, ok := err.(net.Error); ok && ne.Timeout() {
|
||||
// Check for context cancellation
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return -1
|
||||
switch name {
|
||||
case "shell":
|
||||
return dispatchShell(ctx, e, args)
|
||||
case "tool_list":
|
||||
return dispatchToolList(ctx, e, args)
|
||||
case "tool_load":
|
||||
return dispatchToolLoad(ctx, e, args)
|
||||
case "tool_active":
|
||||
return dispatchToolActive(ctx, e, args)
|
||||
case "skill_list":
|
||||
return dispatchSkillList(ctx, e, args)
|
||||
case "skill_load":
|
||||
return dispatchSkillLoad(ctx, e, args)
|
||||
case "skill_active":
|
||||
return dispatchSkillActive(ctx, e, args)
|
||||
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)) //nolint:errcheck
|
||||
if _, err := io.ReadFull(conn, payload); err != nil {
|
||||
return -1
|
||||
return "", fmt.Errorf("unknown execute tool: %s", name)
|
||||
}
|
||||
}
|
||||
|
||||
switch frameType {
|
||||
case 'd':
|
||||
w.Write(payload) //nolint:errcheck
|
||||
if stream != nil {
|
||||
stream(string(payload))
|
||||
// dispatchToolList lists all available tools from the global registry.
|
||||
func dispatchToolList(ctx context.Context, e *Server, args json.RawMessage) (string, error) {
|
||||
if e.toolRegistry == nil {
|
||||
return "", fmt.Errorf("tool_list: no registry available")
|
||||
}
|
||||
case 'x':
|
||||
var code int
|
||||
fmt.Sscanf(string(payload), "%d", &code)
|
||||
return code
|
||||
summaries := e.toolRegistry.Summaries()
|
||||
if len(summaries) == 0 {
|
||||
return "(no tools found)", nil
|
||||
}
|
||||
var out strings.Builder
|
||||
for _, s := range summaries {
|
||||
out.WriteString(s.Name)
|
||||
if s.Description != "" {
|
||||
out.WriteString(" — ")
|
||||
out.WriteString(s.Description)
|
||||
}
|
||||
out.WriteString("\n")
|
||||
}
|
||||
return strings.TrimRight(out.String(), "\n"), nil
|
||||
}
|
||||
|
||||
// Check if detach was requested
|
||||
select {
|
||||
case <-detachCh:
|
||||
// Detach: background the socket read into a goroutine
|
||||
cmdStr := "elevated: " + cmd
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
// dispatchToolLoad loads a tool into the current session.
|
||||
func dispatchToolLoad(ctx context.Context, e *Server, args json.RawMessage) (string, error) {
|
||||
var a struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
ring := detach.NewRingBuffer(detach.RingBufSize)
|
||||
pid := int(time.Now().UnixNano() & 0x7FFFFFFF) // synthetic PID
|
||||
|
||||
proc := &detach.Process{
|
||||
PID: pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Done: make(chan struct{}),
|
||||
if err := json.Unmarshal(args, &a); err != nil {
|
||||
return "", fmt.Errorf("tool_load: bad args: %w", err)
|
||||
}
|
||||
e.detachMu.Lock()
|
||||
e.detached = append(e.detached, proc)
|
||||
e.detachMu.Unlock()
|
||||
|
||||
go func() {
|
||||
defer conn.Close()
|
||||
defer cancel()
|
||||
exitCode := readFrames(ring, nil)
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
proc.ExitCode = exitCode
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if e.OnExit != nil {
|
||||
e.OnExit(proc.PID, proc.ExitCode)
|
||||
if a.Name == "" {
|
||||
return "", fmt.Errorf("tool_load: name is required")
|
||||
}
|
||||
}()
|
||||
|
||||
if e.OnDetach != nil {
|
||||
e.OnDetach(proc.PID, cmdStr)
|
||||
if e.toolRegistry == nil || e.sessionID == "" {
|
||||
return "", fmt.Errorf("tool_load: no session registry")
|
||||
}
|
||||
return fmt.Sprintf("[detached: pid %d]", pid), nil
|
||||
|
||||
default:
|
||||
// Normal (foreground) execution — must also handle manual detach.
|
||||
var outputBuf bytes.Buffer
|
||||
streamFn := StreamFunc(ctx)
|
||||
lw := &limitedWriter{
|
||||
w: &outputBuf,
|
||||
limit: 10 * 1024 * 1024,
|
||||
stream: streamFn,
|
||||
if err := e.toolRegistry.Load(e.sessionID, a.Name); err != nil {
|
||||
return "", fmt.Errorf("tool_load: %w", err)
|
||||
}
|
||||
return fmt.Sprintf("loaded: %s", a.Name), nil
|
||||
}
|
||||
|
||||
// Run readFrames in a goroutine so we can select on detach.
|
||||
type frameResult struct{ exitCode int }
|
||||
frameCh := make(chan frameResult, 1)
|
||||
go func() {
|
||||
code := readFrames(lw, nil)
|
||||
frameCh <- frameResult{code}
|
||||
}()
|
||||
|
||||
select {
|
||||
case fr := <-frameCh:
|
||||
// Normal completion
|
||||
conn.Close()
|
||||
cancel()
|
||||
combined := outputBuf.String()
|
||||
if fr.exitCode != 0 {
|
||||
if combined == "" {
|
||||
return "", fmt.Errorf("elevated execution failed (exit %d)", fr.exitCode)
|
||||
// dispatchToolActive lists tools currently loaded (promoted) in this session.
|
||||
func dispatchToolActive(ctx context.Context, e *Server, args json.RawMessage) (string, error) {
|
||||
if e.toolRegistry == nil || e.sessionID == "" {
|
||||
return "(no tools loaded)", nil
|
||||
}
|
||||
errOutput := combined
|
||||
const maxErrOutput = 8192
|
||||
if len(errOutput) > maxErrOutput {
|
||||
errOutput = errOutput[:maxErrOutput] + fmt.Sprintf("\n[...truncated %d bytes in error]", len(combined)-maxErrOutput)
|
||||
loaded := e.toolRegistry.Loaded(e.sessionID)
|
||||
if len(loaded) == 0 {
|
||||
return "(no tools loaded)", nil
|
||||
}
|
||||
return combined, fmt.Errorf("elevated execution failed (exit %d)\nOutput: %s", fr.exitCode, errOutput)
|
||||
var out strings.Builder
|
||||
for _, t := range loaded {
|
||||
out.WriteString(t.Name)
|
||||
if t.Description != "" {
|
||||
out.WriteString(" — ")
|
||||
out.WriteString(t.Description)
|
||||
}
|
||||
return combined, nil
|
||||
|
||||
case <-ctx.Done():
|
||||
// Timeout or external cancellation
|
||||
conn.Close()
|
||||
cancel()
|
||||
combined := outputBuf.String()
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return combined, fmt.Errorf("elevated execution timeout after %d seconds", timeout)
|
||||
out.WriteString("\n")
|
||||
}
|
||||
return combined, fmt.Errorf("elevated execution interrupted")
|
||||
|
||||
case <-detachCh:
|
||||
// Manual detach: background the connection read
|
||||
cmdStr := "elevated: " + cmd
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
}
|
||||
ring := detach.NewRingBuffer(detach.RingBufSize)
|
||||
pid := int(time.Now().UnixNano() & 0x7FFFFFFF)
|
||||
|
||||
// Splice: future output goes to ring buffer, stop streaming
|
||||
lw.mu.Lock()
|
||||
lw.w = ring
|
||||
lw.stream = nil
|
||||
lw.mu.Unlock()
|
||||
|
||||
proc := &detach.Process{
|
||||
PID: pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Done: make(chan struct{}),
|
||||
}
|
||||
e.detachMu.Lock()
|
||||
e.detached = append(e.detached, proc)
|
||||
e.detachMu.Unlock()
|
||||
|
||||
go func() {
|
||||
fr := <-frameCh
|
||||
conn.Close()
|
||||
cancel()
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
proc.ExitCode = fr.exitCode
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if e.OnExit != nil {
|
||||
e.OnExit(proc.PID, proc.ExitCode)
|
||||
}
|
||||
}()
|
||||
|
||||
if e.OnDetach != nil {
|
||||
e.OnDetach(proc.PID, cmdStr)
|
||||
}
|
||||
|
||||
partial := outputBuf.String()
|
||||
return partial + fmt.Sprintf("\n[detached: pid %d]", pid), nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var whitespacePattern = regexp.MustCompile(`\s+`)
|
||||
|
||||
func (e *Server) checkRateLimit() error {
|
||||
e.rateLimitMu.Lock()
|
||||
defer e.rateLimitMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if now.Before(e.blockedUntil) {
|
||||
remaining := e.blockedUntil.Sub(now).Round(time.Second)
|
||||
return fmt.Errorf("rate limited: too many validation failures, blocked for %v", remaining)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Server) recordValidationFailure() {
|
||||
e.rateLimitMu.Lock()
|
||||
defer e.rateLimitMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if now.Sub(e.lastFailure) > failureWindow {
|
||||
e.validationFailures = 0
|
||||
}
|
||||
|
||||
e.validationFailures++
|
||||
e.lastFailure = now
|
||||
|
||||
if e.validationFailures >= maxFailures {
|
||||
e.blockedUntil = now.Add(blockDuration)
|
||||
e.validationFailures = 0
|
||||
}
|
||||
}
|
||||
|
||||
// Execute runs code in a sandbox and returns combined stdout+stderr.
|
||||
func (e *Server) Execute(ctx context.Context, code, language string, timeout int, sandboxName string, trusted bool) (string, error) {
|
||||
return e.executeWithStdin(ctx, code, language, timeout, sandboxName, trusted, "")
|
||||
}
|
||||
|
||||
// executeWithStdin is like Execute but feeds stdinData to the command's stdin.
|
||||
// For languages where code is itself passed via stdin (ed, expect, bc), stdinData is ignored.
|
||||
func (e *Server) executeWithStdin(ctx context.Context, code, language string, timeout int, sandboxName string, trusted bool, stdinData string, doDetach ...bool) (string, error) {
|
||||
if timeout < 0 {
|
||||
timeout = 30
|
||||
}
|
||||
|
||||
if !trusted {
|
||||
if err := e.ValidateCode(code, language); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
var cfg *sandbox.Config
|
||||
var err error
|
||||
if !e.Yolo {
|
||||
cfg, err = loadSandboxConfig(sandboxName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
e.wdMu.RLock()
|
||||
workDir := e.cwd
|
||||
e.wdMu.RUnlock()
|
||||
if workDir == "" {
|
||||
workDir, _ = os.Getwd()
|
||||
}
|
||||
|
||||
var cancel context.CancelFunc
|
||||
if timeout > 0 {
|
||||
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
||||
} else {
|
||||
// timeout=0 means no timeout (run indefinitely, e.g. for daemons)
|
||||
ctx, cancel = context.WithCancel(ctx)
|
||||
}
|
||||
// NOTE: do NOT defer cancel() here — the detach path must avoid cancelling
|
||||
// the context (which would kill the detached process via cmd.Cancel).
|
||||
// Each exit path calls cancel() explicitly except detach.
|
||||
|
||||
var cmd *exec.Cmd
|
||||
var interpreter []string
|
||||
// codeStdin: non-empty means the code itself is fed via stdin (ed, expect, bc).
|
||||
// In these cases stdinData cannot be used simultaneously.
|
||||
var codeStdin string
|
||||
switch language {
|
||||
case "bash", "":
|
||||
interpreter = []string{"bash", "-c", code}
|
||||
default:
|
||||
cancel()
|
||||
return "", fmt.Errorf("unsupported language: %s (only bash is supported)", language)
|
||||
}
|
||||
e.envMu.RLock()
|
||||
envMap := make(map[string]string, len(e.envExtra))
|
||||
for k, v := range e.envExtra {
|
||||
envMap[k] = v
|
||||
}
|
||||
e.envMu.RUnlock()
|
||||
for _, ev := range os.Environ() {
|
||||
k, v, _ := strings.Cut(ev, "=")
|
||||
if _, ok := envMap[k]; !ok {
|
||||
envMap[k] = v
|
||||
}
|
||||
}
|
||||
if _, ok := envMap["NAMESPACE"]; !ok {
|
||||
if ns := plan9Namespace(envMap); ns != "" {
|
||||
envMap["NAMESPACE"] = ns
|
||||
}
|
||||
}
|
||||
getenv := func(key string) string { return envMap[key] }
|
||||
|
||||
if e.Yolo {
|
||||
cmd = exec.CommandContext(ctx, interpreter[0], interpreter[1:]...)
|
||||
} else {
|
||||
wrapped, wrapErr := sandbox.WrapCommand(cfg, interpreter, workDir, getenv)
|
||||
if wrapErr != nil {
|
||||
cancel()
|
||||
return "", wrapErr
|
||||
}
|
||||
cmd = exec.CommandContext(ctx, wrapped[0], wrapped[1:]...)
|
||||
}
|
||||
cmd.Dir = workDir
|
||||
switch {
|
||||
case codeStdin != "":
|
||||
cmd.Stdin = strings.NewReader(codeStdin)
|
||||
case stdinData != "":
|
||||
cmd.Stdin = strings.NewReader(stdinData)
|
||||
}
|
||||
|
||||
e.envMu.RLock()
|
||||
baseEnv := os.Environ()
|
||||
// Remove keys that envExtra overrides so duplicates don't leak through.
|
||||
filtered := baseEnv[:0]
|
||||
for _, ev := range baseEnv {
|
||||
k, _, _ := strings.Cut(ev, "=")
|
||||
if _, overridden := e.envExtra[k]; !overridden {
|
||||
filtered = append(filtered, ev)
|
||||
}
|
||||
}
|
||||
cmd.Env = prependOlliePath(filtered, paths.CfgDir())
|
||||
cmd.Env = append(cmd.Env, "OLLIE_TOOLS_PATH="+ToolsPath())
|
||||
for k, v := range e.envExtra {
|
||||
cmd.Env = append(cmd.Env, k+"="+v)
|
||||
}
|
||||
e.envMu.RUnlock()
|
||||
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||
cmd.Cancel = func() error {
|
||||
if cmd.Process != nil {
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
cmd.WaitDelay = time.Second
|
||||
|
||||
var outputBuf bytes.Buffer
|
||||
lw := &limitedWriter{
|
||||
w: &outputBuf,
|
||||
limit: 10 * 1024 * 1024,
|
||||
stream: StreamFunc(ctx),
|
||||
}
|
||||
cmd.Stdout = lw
|
||||
cmd.Stderr = lw
|
||||
|
||||
// Set up detach channel for this execution
|
||||
detachCh := make(chan struct{}, 1)
|
||||
e.detachMu.Lock()
|
||||
e.detachCh = detachCh
|
||||
e.detachMu.Unlock()
|
||||
defer func() {
|
||||
e.detachMu.Lock()
|
||||
e.detachCh = nil
|
||||
e.detachMu.Unlock()
|
||||
}()
|
||||
|
||||
if err = cmd.Start(); err != nil {
|
||||
cancel()
|
||||
return "", fmt.Errorf("execution failed: %v", err)
|
||||
}
|
||||
|
||||
// If detach requested, immediately signal the detach channel
|
||||
if len(doDetach) > 0 && doDetach[0] {
|
||||
close(detachCh)
|
||||
}
|
||||
|
||||
// Wait for completion, context cancellation, or detach signal
|
||||
waitCh := make(chan error, 1)
|
||||
go func() { waitCh <- cmd.Wait() }()
|
||||
|
||||
select {
|
||||
case err = <-waitCh:
|
||||
// Normal completion
|
||||
cancel()
|
||||
output := outputBuf.Bytes()
|
||||
if lw.truncated {
|
||||
output = append(output, []byte("\n[output truncated at 10MB]")...)
|
||||
}
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return "", fmt.Errorf("execution timeout after %d seconds", timeout)
|
||||
}
|
||||
if err != nil {
|
||||
// Cap output in error message to avoid flooding the agent context.
|
||||
// The full output is still returned as the first return value.
|
||||
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: %v\nOutput: %s", err, errOutput)
|
||||
}
|
||||
return string(output), nil
|
||||
|
||||
case <-ctx.Done():
|
||||
// Context cancelled (interrupt or timeout)
|
||||
cancel()
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
<-waitCh // reap
|
||||
output := outputBuf.Bytes()
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return string(output), fmt.Errorf("execution timeout after %d seconds", timeout)
|
||||
}
|
||||
return string(output), fmt.Errorf("execution interrupted")
|
||||
|
||||
case <-detachCh:
|
||||
// Detach: move process to registry, continue running
|
||||
// Neutralize cmd.Cancel and WaitDelay so context cancellation
|
||||
// won't kill the process after the timeout fires.
|
||||
cmd.Cancel = nil
|
||||
cmd.WaitDelay = 0
|
||||
cancel()
|
||||
|
||||
cmdStr := strings.Join(interpreter, " ")
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
}
|
||||
ring := detach.NewRingBuffer(detach.RingBufSize)
|
||||
// Splice: future output goes to ring buffer instead of outputBuf.
|
||||
lw.mu.Lock()
|
||||
lw.w = ring
|
||||
lw.stream = nil
|
||||
lw.mu.Unlock()
|
||||
|
||||
proc := &detach.Process{
|
||||
PID: cmd.Process.Pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Cmd: cmd.Process,
|
||||
Done: make(chan struct{}),
|
||||
}
|
||||
e.detachMu.Lock()
|
||||
e.detached = append(e.detached, proc)
|
||||
e.detachMu.Unlock()
|
||||
|
||||
// Monitor for exit in background
|
||||
go func() {
|
||||
waitErr := <-waitCh
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
if waitErr != nil {
|
||||
if exitErr, ok := waitErr.(*exec.ExitError); ok {
|
||||
proc.ExitCode = exitErr.ExitCode()
|
||||
} else {
|
||||
proc.ExitCode = -1
|
||||
}
|
||||
}
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if e.OnExit != nil {
|
||||
e.OnExit(proc.PID, proc.ExitCode)
|
||||
}
|
||||
}()
|
||||
|
||||
if e.OnDetach != nil {
|
||||
e.OnDetach(proc.PID, cmdStr)
|
||||
}
|
||||
|
||||
partial := outputBuf.String()
|
||||
return partial + fmt.Sprintf("\n[detached: pid %d]", proc.PID), nil
|
||||
}
|
||||
}
|
||||
|
||||
// prependOlliePath returns env with $OLLIE_CFG_PATH/scripts/x/ prepended to PATH so wrapper
|
||||
// scripts placed there shadow system binaries.
|
||||
func prependOlliePath(env []string, cfgDir string) []string {
|
||||
if cfgDir == "" {
|
||||
return env
|
||||
}
|
||||
xdir := filepath.Join(cfgDir, "scripts", "x")
|
||||
result := make([]string, 0, len(env))
|
||||
for _, e := range env {
|
||||
if strings.HasPrefix(e, "PATH=") {
|
||||
e = "PATH=" + xdir + string(filepath.ListSeparator) + e[5:]
|
||||
}
|
||||
result = append(result, e)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// plan9Namespace computes the plan9port namespace directory from the environment.
|
||||
// Convention: $NAMESPACE if set, else /tmp/ns.<user>.<display> where display is
|
||||
// $DISPLAY with trailing .0 stripped and / replaced by _.
|
||||
func plan9Namespace(env map[string]string) string {
|
||||
if ns := env["NAMESPACE"]; ns != "" {
|
||||
return ns
|
||||
}
|
||||
user := env["USER"]
|
||||
if user == "" {
|
||||
return ""
|
||||
}
|
||||
disp := env["DISPLAY"]
|
||||
if disp == "" {
|
||||
return ""
|
||||
}
|
||||
// Canonicalize: strip trailing .0
|
||||
if strings.HasSuffix(disp, ".0") {
|
||||
disp = disp[:len(disp)-2]
|
||||
}
|
||||
// Replace / with _
|
||||
disp = strings.ReplaceAll(disp, "/", "_")
|
||||
return "/tmp/ns." + user + "." + disp
|
||||
return strings.TrimRight(out.String(), "\n"), nil
|
||||
}
|
||||
|
||||
// Detach signals the currently running process to be detached from the agent.
|
||||
|
|
@ -1024,4 +533,3 @@ func (e *Server) cleanupDetached() {
|
|||
p.Mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,714 @@
|
|||
package tools
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"ollie/detach"
|
||||
"ollie/paths"
|
||||
"ollie/sandbox"
|
||||
)
|
||||
|
||||
// universalPatterns apply to all code.
|
||||
var universalPatterns = []*regexp.Regexp{
|
||||
regexp.MustCompile(`\bmkfs\b`),
|
||||
regexp.MustCompile(`\bdd\b.*\bif=/dev/`),
|
||||
regexp.MustCompile(`\b(sudo|su)\s`),
|
||||
regexp.MustCompile(`/etc/(shadow|sudoers)`),
|
||||
}
|
||||
|
||||
// bashPatterns apply to bash (flag syntax, redirects, shell-specific constructs).
|
||||
var bashPatterns = []*regexp.Regexp{
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s*/(home|var|usr|etc|boot|root|bin|sbin|lib|opt|srv)?`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*f[a-z]*\s+)*-[a-z]*r[a-z]*\s*/(home|var|usr|etc|boot|root|bin|sbin|lib|opt|srv)?`),
|
||||
regexp.MustCompile(`rm\s+.*--recursive.*--force`),
|
||||
regexp.MustCompile(`rm\s+.*--force.*--recursive`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+\.\.(/|\s|$)`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+~`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+\*`),
|
||||
regexp.MustCompile(`:\s*\(\s*\)\s*\{\s*:\s*\|\s*:\s*&`), // fork bomb
|
||||
regexp.MustCompile(`>\s*/dev/sd`),
|
||||
regexp.MustCompile(`\beval\s+".*\$`),
|
||||
}
|
||||
|
||||
var languagePatterns = map[string][]*regexp.Regexp{
|
||||
"bash": bashPatterns,
|
||||
"": bashPatterns,
|
||||
}
|
||||
// ValidateCode checks code against dangerous patterns.
|
||||
func (e *Server) ValidateCode(code, language string) error {
|
||||
if err := e.checkRateLimit(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
normalized := strings.ToLower(code)
|
||||
normalized = whitespacePattern.ReplaceAllString(normalized, " ")
|
||||
|
||||
patterns := append(universalPatterns, languagePatterns[language]...)
|
||||
for _, pattern := range patterns {
|
||||
if pattern.MatchString(normalized) {
|
||||
e.recordValidationFailure()
|
||||
return fmt.Errorf("dangerous pattern detected")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadSandboxConfig(name string) (*sandbox.Config, error) {
|
||||
if name == "" {
|
||||
name = "default"
|
||||
}
|
||||
path := filepath.Join(paths.CfgDir(), "sandbox", name+".yaml")
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("sandbox %q not found: %w", name, err)
|
||||
}
|
||||
defer f.Close()
|
||||
return sandbox.LoadSandbox(f)
|
||||
}
|
||||
|
||||
type limitedWriter struct {
|
||||
mu sync.Mutex
|
||||
w io.Writer
|
||||
written int
|
||||
limit int
|
||||
truncated bool
|
||||
stream func(string) // if non-nil, called with each chunk of output
|
||||
}
|
||||
|
||||
func (lw *limitedWriter) Write(p []byte) (n int, err error) {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
if lw.written >= lw.limit {
|
||||
lw.truncated = true
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
remaining := lw.limit - lw.written
|
||||
toWrite := p
|
||||
if len(p) > remaining {
|
||||
toWrite = p[:remaining]
|
||||
lw.truncated = true
|
||||
}
|
||||
|
||||
written, err := lw.w.Write(toWrite)
|
||||
lw.written += written
|
||||
if lw.stream != nil {
|
||||
lw.stream(string(toWrite[:written]))
|
||||
}
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// dispatchShell handles the shell tool: a single bash command.
|
||||
func dispatchShell(ctx context.Context, e *Server, args json.RawMessage) (string, error) {
|
||||
var a struct {
|
||||
Cmd string `json:"cmd"`
|
||||
Timeout int `json:"timeout"`
|
||||
Sandbox string `json:"sandbox"`
|
||||
Elevated bool `json:"elevated"`
|
||||
Detach bool `json:"detach"`
|
||||
}
|
||||
if err := json.Unmarshal(args, &a); err != nil {
|
||||
return "", fmt.Errorf("shell: bad args: %w", err)
|
||||
}
|
||||
if a.Cmd == "" {
|
||||
return "", fmt.Errorf("shell: cmd is required")
|
||||
}
|
||||
timeout := a.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 30
|
||||
}
|
||||
if a.Elevated {
|
||||
e.wdMu.RLock()
|
||||
workDir := e.cwd
|
||||
e.wdMu.RUnlock()
|
||||
return e.executeElevated(ctx, a.Cmd, workDir, timeout, a.Detach)
|
||||
}
|
||||
sandboxName := a.Sandbox
|
||||
if sandboxName == "" {
|
||||
sandboxName = "default"
|
||||
}
|
||||
return e.executeWithStdin(ctx, a.Cmd, "bash", timeout, sandboxName, false, "", a.Detach)
|
||||
}
|
||||
// executeElevated runs cmd outside the sandbox via the integrated elevation broker.
|
||||
// Connects to the broker socket, sends the request with the current env,
|
||||
// and streams the framed response back.
|
||||
func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout int, doDetach ...bool) (string, error) {
|
||||
sockPath := os.Getenv("OLLIE_ELEVATE_SOCKET")
|
||||
if sockPath == "" {
|
||||
xdg := os.Getenv("XDG_RUNTIME_DIR")
|
||||
if xdg == "" {
|
||||
return "", fmt.Errorf("elevation not available: no XDG_RUNTIME_DIR")
|
||||
}
|
||||
sockPath = filepath.Join(xdg, "ollie", "elevate.sock")
|
||||
}
|
||||
|
||||
wantDetach := len(doDetach) > 0 && doDetach[0]
|
||||
|
||||
var cancel context.CancelFunc
|
||||
if wantDetach || timeout <= 0 {
|
||||
ctx, cancel = context.WithCancel(ctx)
|
||||
} else {
|
||||
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
||||
}
|
||||
|
||||
// Connect to broker
|
||||
conn, err := net.DialTimeout("unix", sockPath, 5*time.Second)
|
||||
if err != nil {
|
||||
cancel()
|
||||
return "", fmt.Errorf("elevation not available: %w", err)
|
||||
}
|
||||
|
||||
// Send request
|
||||
envMap := make(map[string]string, len(os.Environ()))
|
||||
for _, kv := range os.Environ() {
|
||||
if k, v, ok := strings.Cut(kv, "="); ok {
|
||||
envMap[k] = v
|
||||
}
|
||||
}
|
||||
reqJSON, _ := json.Marshal(struct {
|
||||
Cmd string `json:"cmd"`
|
||||
Cwd string `json:"cwd"`
|
||||
Env map[string]string `json:"env"`
|
||||
Session string `json:"session,omitempty"`
|
||||
}{Cmd: cmd, Cwd: dir, Env: envMap, Session: os.Getenv("OLLIE_SESSION_ID")})
|
||||
reqJSON = append(reqJSON, '\n')
|
||||
if _, err := conn.Write(reqJSON); err != nil {
|
||||
conn.Close()
|
||||
cancel()
|
||||
return "", fmt.Errorf("elevated execution failed: write: %w", err)
|
||||
}
|
||||
|
||||
// Set up detach channel
|
||||
detachCh := make(chan struct{}, 1)
|
||||
e.detachMu.Lock()
|
||||
e.detachCh = detachCh
|
||||
e.detachMu.Unlock()
|
||||
defer func() {
|
||||
e.detachMu.Lock()
|
||||
e.detachCh = nil
|
||||
e.detachMu.Unlock()
|
||||
}()
|
||||
|
||||
if wantDetach {
|
||||
close(detachCh)
|
||||
}
|
||||
|
||||
// readFrames reads from the connection, writing to w and streaming.
|
||||
// Returns exit code when 'x' frame arrives or -1 on error.
|
||||
readFrames := func(w io.Writer, stream func(string)) int {
|
||||
header := make([]byte, 5)
|
||||
for {
|
||||
conn.SetReadDeadline(time.Now().Add(1 * time.Second)) //nolint:errcheck
|
||||
_, err := io.ReadFull(conn, header)
|
||||
if err != nil {
|
||||
if ne, ok := err.(net.Error); ok && ne.Timeout() {
|
||||
// Check for context cancellation
|
||||
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)) //nolint:errcheck
|
||||
if _, err := io.ReadFull(conn, payload); err != nil {
|
||||
return -1
|
||||
}
|
||||
}
|
||||
|
||||
switch frameType {
|
||||
case 'd':
|
||||
w.Write(payload) //nolint:errcheck
|
||||
if stream != nil {
|
||||
stream(string(payload))
|
||||
}
|
||||
case 'x':
|
||||
var code int
|
||||
fmt.Sscanf(string(payload), "%d", &code)
|
||||
return code
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check if detach was requested
|
||||
select {
|
||||
case <-detachCh:
|
||||
// Detach: background the socket read into a goroutine
|
||||
cmdStr := "elevated: " + cmd
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
}
|
||||
ring := detach.NewRingBuffer(detach.RingBufSize)
|
||||
pid := int(time.Now().UnixNano() & 0x7FFFFFFF) // synthetic PID
|
||||
|
||||
proc := &detach.Process{
|
||||
PID: pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Done: make(chan struct{}),
|
||||
}
|
||||
e.detachMu.Lock()
|
||||
e.detached = append(e.detached, proc)
|
||||
e.detachMu.Unlock()
|
||||
|
||||
go func() {
|
||||
defer conn.Close()
|
||||
defer cancel()
|
||||
exitCode := readFrames(ring, nil)
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
proc.ExitCode = exitCode
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if e.OnExit != nil {
|
||||
e.OnExit(proc.PID, proc.ExitCode)
|
||||
}
|
||||
}()
|
||||
|
||||
if e.OnDetach != nil {
|
||||
e.OnDetach(proc.PID, cmdStr)
|
||||
}
|
||||
return fmt.Sprintf("[detached: pid %d]", pid), nil
|
||||
|
||||
default:
|
||||
// Normal (foreground) execution — must also handle manual detach.
|
||||
var outputBuf bytes.Buffer
|
||||
streamFn := StreamFunc(ctx)
|
||||
lw := &limitedWriter{
|
||||
w: &outputBuf,
|
||||
limit: 10 * 1024 * 1024,
|
||||
stream: streamFn,
|
||||
}
|
||||
|
||||
// Run readFrames in a goroutine so we can select on detach.
|
||||
type frameResult struct{ exitCode int }
|
||||
frameCh := make(chan frameResult, 1)
|
||||
go func() {
|
||||
code := readFrames(lw, nil)
|
||||
frameCh <- frameResult{code}
|
||||
}()
|
||||
|
||||
select {
|
||||
case fr := <-frameCh:
|
||||
// Normal completion
|
||||
conn.Close()
|
||||
cancel()
|
||||
combined := outputBuf.String()
|
||||
if fr.exitCode != 0 {
|
||||
if combined == "" {
|
||||
return "", fmt.Errorf("elevated execution failed (exit %d)", fr.exitCode)
|
||||
}
|
||||
errOutput := combined
|
||||
const maxErrOutput = 8192
|
||||
if len(errOutput) > maxErrOutput {
|
||||
errOutput = errOutput[:maxErrOutput] + fmt.Sprintf("\n[...truncated %d bytes in error]", len(combined)-maxErrOutput)
|
||||
}
|
||||
return combined, fmt.Errorf("elevated execution failed (exit %d)\nOutput: %s", fr.exitCode, errOutput)
|
||||
}
|
||||
return combined, nil
|
||||
|
||||
case <-ctx.Done():
|
||||
// Timeout or external cancellation
|
||||
conn.Close()
|
||||
cancel()
|
||||
combined := outputBuf.String()
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return combined, fmt.Errorf("elevated execution timeout after %d seconds", timeout)
|
||||
}
|
||||
return combined, fmt.Errorf("elevated execution interrupted")
|
||||
|
||||
case <-detachCh:
|
||||
// Manual detach: background the connection read
|
||||
cmdStr := "elevated: " + cmd
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
}
|
||||
ring := detach.NewRingBuffer(detach.RingBufSize)
|
||||
pid := int(time.Now().UnixNano() & 0x7FFFFFFF)
|
||||
|
||||
// Splice: future output goes to ring buffer, stop streaming
|
||||
lw.mu.Lock()
|
||||
lw.w = ring
|
||||
lw.stream = nil
|
||||
lw.mu.Unlock()
|
||||
|
||||
proc := &detach.Process{
|
||||
PID: pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Done: make(chan struct{}),
|
||||
}
|
||||
e.detachMu.Lock()
|
||||
e.detached = append(e.detached, proc)
|
||||
e.detachMu.Unlock()
|
||||
|
||||
go func() {
|
||||
fr := <-frameCh
|
||||
conn.Close()
|
||||
cancel()
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
proc.ExitCode = fr.exitCode
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if e.OnExit != nil {
|
||||
e.OnExit(proc.PID, proc.ExitCode)
|
||||
}
|
||||
}()
|
||||
|
||||
if e.OnDetach != nil {
|
||||
e.OnDetach(proc.PID, cmdStr)
|
||||
}
|
||||
|
||||
partial := outputBuf.String()
|
||||
return partial + fmt.Sprintf("\n[detached: pid %d]", pid), nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var whitespacePattern = regexp.MustCompile(`\s+`)
|
||||
|
||||
func (e *Server) checkRateLimit() error {
|
||||
e.rateLimitMu.Lock()
|
||||
defer e.rateLimitMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if now.Before(e.blockedUntil) {
|
||||
remaining := e.blockedUntil.Sub(now).Round(time.Second)
|
||||
return fmt.Errorf("rate limited: too many validation failures, blocked for %v", remaining)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Server) recordValidationFailure() {
|
||||
e.rateLimitMu.Lock()
|
||||
defer e.rateLimitMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if now.Sub(e.lastFailure) > failureWindow {
|
||||
e.validationFailures = 0
|
||||
}
|
||||
|
||||
e.validationFailures++
|
||||
e.lastFailure = now
|
||||
|
||||
if e.validationFailures >= maxFailures {
|
||||
e.blockedUntil = now.Add(blockDuration)
|
||||
e.validationFailures = 0
|
||||
}
|
||||
}
|
||||
|
||||
// Execute runs code in a sandbox and returns combined stdout+stderr.
|
||||
func (e *Server) Execute(ctx context.Context, code, language string, timeout int, sandboxName string, trusted bool) (string, error) {
|
||||
return e.executeWithStdin(ctx, code, language, timeout, sandboxName, trusted, "")
|
||||
}
|
||||
|
||||
// executeWithStdin is like Execute but feeds stdinData to the command's stdin.
|
||||
// For languages where code is itself passed via stdin (ed, expect, bc), stdinData is ignored.
|
||||
func (e *Server) executeWithStdin(ctx context.Context, code, language string, timeout int, sandboxName string, trusted bool, stdinData string, doDetach ...bool) (string, error) {
|
||||
if timeout < 0 {
|
||||
timeout = 30
|
||||
}
|
||||
|
||||
if !trusted {
|
||||
if err := e.ValidateCode(code, language); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
var cfg *sandbox.Config
|
||||
var err error
|
||||
if !e.Yolo {
|
||||
cfg, err = loadSandboxConfig(sandboxName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
e.wdMu.RLock()
|
||||
workDir := e.cwd
|
||||
e.wdMu.RUnlock()
|
||||
if workDir == "" {
|
||||
workDir, _ = os.Getwd()
|
||||
}
|
||||
|
||||
var cancel context.CancelFunc
|
||||
if timeout > 0 {
|
||||
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
||||
} else {
|
||||
// timeout=0 means no timeout (run indefinitely, e.g. for daemons)
|
||||
ctx, cancel = context.WithCancel(ctx)
|
||||
}
|
||||
// NOTE: do NOT defer cancel() here — the detach path must avoid cancelling
|
||||
// the context (which would kill the detached process via cmd.Cancel).
|
||||
// Each exit path calls cancel() explicitly except detach.
|
||||
|
||||
var cmd *exec.Cmd
|
||||
var interpreter []string
|
||||
// codeStdin: non-empty means the code itself is fed via stdin (ed, expect, bc).
|
||||
// In these cases stdinData cannot be used simultaneously.
|
||||
var codeStdin string
|
||||
switch language {
|
||||
case "bash", "":
|
||||
interpreter = []string{"bash", "-c", code}
|
||||
default:
|
||||
cancel()
|
||||
return "", fmt.Errorf("unsupported language: %s (only bash is supported)", language)
|
||||
}
|
||||
e.envMu.RLock()
|
||||
envMap := make(map[string]string, len(e.envExtra))
|
||||
for k, v := range e.envExtra {
|
||||
envMap[k] = v
|
||||
}
|
||||
e.envMu.RUnlock()
|
||||
for _, ev := range os.Environ() {
|
||||
k, v, _ := strings.Cut(ev, "=")
|
||||
if _, ok := envMap[k]; !ok {
|
||||
envMap[k] = v
|
||||
}
|
||||
}
|
||||
if _, ok := envMap["NAMESPACE"]; !ok {
|
||||
if ns := plan9Namespace(envMap); ns != "" {
|
||||
envMap["NAMESPACE"] = ns
|
||||
}
|
||||
}
|
||||
getenv := func(key string) string { return envMap[key] }
|
||||
|
||||
if e.Yolo {
|
||||
cmd = exec.CommandContext(ctx, interpreter[0], interpreter[1:]...)
|
||||
} else {
|
||||
wrapped, wrapErr := sandbox.WrapCommand(cfg, interpreter, workDir, getenv)
|
||||
if wrapErr != nil {
|
||||
cancel()
|
||||
return "", wrapErr
|
||||
}
|
||||
cmd = exec.CommandContext(ctx, wrapped[0], wrapped[1:]...)
|
||||
}
|
||||
cmd.Dir = workDir
|
||||
switch {
|
||||
case codeStdin != "":
|
||||
cmd.Stdin = strings.NewReader(codeStdin)
|
||||
case stdinData != "":
|
||||
cmd.Stdin = strings.NewReader(stdinData)
|
||||
}
|
||||
|
||||
e.envMu.RLock()
|
||||
baseEnv := os.Environ()
|
||||
// Remove keys that envExtra overrides so duplicates don't leak through.
|
||||
filtered := baseEnv[:0]
|
||||
for _, ev := range baseEnv {
|
||||
k, _, _ := strings.Cut(ev, "=")
|
||||
if _, overridden := e.envExtra[k]; !overridden {
|
||||
filtered = append(filtered, ev)
|
||||
}
|
||||
}
|
||||
cmd.Env = prependOlliePath(filtered, paths.CfgDir())
|
||||
cmd.Env = append(cmd.Env, "OLLIE_TOOLS_PATH="+ToolsPath())
|
||||
for k, v := range e.envExtra {
|
||||
cmd.Env = append(cmd.Env, k+"="+v)
|
||||
}
|
||||
e.envMu.RUnlock()
|
||||
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||
cmd.Cancel = func() error {
|
||||
if cmd.Process != nil {
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
cmd.WaitDelay = time.Second
|
||||
|
||||
var outputBuf bytes.Buffer
|
||||
lw := &limitedWriter{
|
||||
w: &outputBuf,
|
||||
limit: 10 * 1024 * 1024,
|
||||
stream: StreamFunc(ctx),
|
||||
}
|
||||
cmd.Stdout = lw
|
||||
cmd.Stderr = lw
|
||||
|
||||
// Set up detach channel for this execution
|
||||
detachCh := make(chan struct{}, 1)
|
||||
e.detachMu.Lock()
|
||||
e.detachCh = detachCh
|
||||
e.detachMu.Unlock()
|
||||
defer func() {
|
||||
e.detachMu.Lock()
|
||||
e.detachCh = nil
|
||||
e.detachMu.Unlock()
|
||||
}()
|
||||
|
||||
if err = cmd.Start(); err != nil {
|
||||
cancel()
|
||||
return "", fmt.Errorf("execution failed: %v", err)
|
||||
}
|
||||
|
||||
// If detach requested, immediately signal the detach channel
|
||||
if len(doDetach) > 0 && doDetach[0] {
|
||||
close(detachCh)
|
||||
}
|
||||
|
||||
// Wait for completion, context cancellation, or detach signal
|
||||
waitCh := make(chan error, 1)
|
||||
go func() { waitCh <- cmd.Wait() }()
|
||||
|
||||
select {
|
||||
case err = <-waitCh:
|
||||
// Normal completion
|
||||
cancel()
|
||||
output := outputBuf.Bytes()
|
||||
if lw.truncated {
|
||||
output = append(output, []byte("\n[output truncated at 10MB]")...)
|
||||
}
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return "", fmt.Errorf("execution timeout after %d seconds", timeout)
|
||||
}
|
||||
if err != nil {
|
||||
// Cap output in error message to avoid flooding the agent context.
|
||||
// The full output is still returned as the first return value.
|
||||
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: %v\nOutput: %s", err, errOutput)
|
||||
}
|
||||
return string(output), nil
|
||||
|
||||
case <-ctx.Done():
|
||||
// Context cancelled (interrupt or timeout)
|
||||
cancel()
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
<-waitCh // reap
|
||||
output := outputBuf.Bytes()
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return string(output), fmt.Errorf("execution timeout after %d seconds", timeout)
|
||||
}
|
||||
return string(output), fmt.Errorf("execution interrupted")
|
||||
|
||||
case <-detachCh:
|
||||
// Detach: move process to registry, continue running
|
||||
// Neutralize cmd.Cancel and WaitDelay so context cancellation
|
||||
// won't kill the process after the timeout fires.
|
||||
cmd.Cancel = nil
|
||||
cmd.WaitDelay = 0
|
||||
cancel()
|
||||
|
||||
cmdStr := strings.Join(interpreter, " ")
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
}
|
||||
ring := detach.NewRingBuffer(detach.RingBufSize)
|
||||
// Splice: future output goes to ring buffer instead of outputBuf.
|
||||
lw.mu.Lock()
|
||||
lw.w = ring
|
||||
lw.stream = nil
|
||||
lw.mu.Unlock()
|
||||
|
||||
proc := &detach.Process{
|
||||
PID: cmd.Process.Pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Cmd: cmd.Process,
|
||||
Done: make(chan struct{}),
|
||||
}
|
||||
e.detachMu.Lock()
|
||||
e.detached = append(e.detached, proc)
|
||||
e.detachMu.Unlock()
|
||||
|
||||
// Monitor for exit in background
|
||||
go func() {
|
||||
waitErr := <-waitCh
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
if waitErr != nil {
|
||||
if exitErr, ok := waitErr.(*exec.ExitError); ok {
|
||||
proc.ExitCode = exitErr.ExitCode()
|
||||
} else {
|
||||
proc.ExitCode = -1
|
||||
}
|
||||
}
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if e.OnExit != nil {
|
||||
e.OnExit(proc.PID, proc.ExitCode)
|
||||
}
|
||||
}()
|
||||
|
||||
if e.OnDetach != nil {
|
||||
e.OnDetach(proc.PID, cmdStr)
|
||||
}
|
||||
|
||||
partial := outputBuf.String()
|
||||
return partial + fmt.Sprintf("\n[detached: pid %d]", proc.PID), nil
|
||||
}
|
||||
}
|
||||
|
||||
// prependOlliePath returns env with $OLLIE_CFG_PATH/scripts/x/ prepended to PATH so wrapper
|
||||
// scripts placed there shadow system binaries.
|
||||
func prependOlliePath(env []string, cfgDir string) []string {
|
||||
if cfgDir == "" {
|
||||
return env
|
||||
}
|
||||
xdir := filepath.Join(cfgDir, "scripts", "x")
|
||||
result := make([]string, 0, len(env))
|
||||
for _, e := range env {
|
||||
if strings.HasPrefix(e, "PATH=") {
|
||||
e = "PATH=" + xdir + string(filepath.ListSeparator) + e[5:]
|
||||
}
|
||||
result = append(result, e)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// plan9Namespace computes the plan9port namespace directory from the environment.
|
||||
// Convention: $NAMESPACE if set, else /tmp/ns.<user>.<display> where display is
|
||||
// $DISPLAY with trailing .0 stripped and / replaced by _.
|
||||
func plan9Namespace(env map[string]string) string {
|
||||
if ns := env["NAMESPACE"]; ns != "" {
|
||||
return ns
|
||||
}
|
||||
user := env["USER"]
|
||||
if user == "" {
|
||||
return ""
|
||||
}
|
||||
disp := env["DISPLAY"]
|
||||
if disp == "" {
|
||||
return ""
|
||||
}
|
||||
// Canonicalize: strip trailing .0
|
||||
if strings.HasSuffix(disp, ".0") {
|
||||
disp = disp[:len(disp)-2]
|
||||
}
|
||||
// Replace / with _
|
||||
disp = strings.ReplaceAll(disp, "/", "_")
|
||||
return "/tmp/ns." + user + "." + disp
|
||||
}
|
||||
Reference in New Issue