refactor: architectural cleanup (5 items)

1. Standardize ctl dispatch: hoist rootCtlHandlers and sessionCtlHandlers
   to package-level vars, matching the agent pattern.

2. Fix error wrapping: %v → %w in toolsrv/shell.go (only 2 instances;
   backend/ was already clean).

3. Canonicalize tool server access: LoadTool now only uses
   session.ToolsConn() — removed fallback to agent.ToolServer().

4. Centralize event handler wiring: wireAgentEvents(sessID, ag, al)
   replaces separate SetOutput + wireAgentStateEvents calls. Documents
   the contract.

5. Split toolsrv/shell.go: validation patterns, ValidateCode, and rate
   limiting moved to shell_validate.go (86 lines). shell.go retains
   execution logic (594 lines).
This commit is contained in:
Ollie Agent 2026-08-09 12:46:36 +02:00
parent f5efd90fbe
commit 2b4d577e39
7 changed files with 165 additions and 156 deletions

View File

@ -22,8 +22,7 @@ func requestAgentNew(ctx HandlerCtx, data []byte) ([]byte, error) {
// Create AgentLog for this agent and wire up the event handler
al := NewAgentLog("")
ctx.Session.SetAgentLog(ag.ID(), al)
ag.SetOutput(NewEventHandler(al))
wireAgentStateEvents(ctx.Session.ID(), ag)
wireAgentEvents(ctx.Session.ID(), ag, al)
return []byte(ag.ID() + "\n"), nil
}

View File

@ -24,8 +24,14 @@ func mergeCtx(parent, child context.Context) (context.Context, context.CancelFun
return merged, func() { stop(); cancel() }
}
// wireAgentStateEvents sets up the agent to publish state change events to the event bus.
func wireAgentStateEvents(sessID string, ag *agent.Agent) {
// wireAgentEvents connects an agent's output and state change events.
// Called once per agent: at creation (handlers_agent.go) and on restore (newroot.go).
//
// Contract:
// - SetOutput → events flow to the AgentLog (chat stream, tool blocks, etc.)
// - SetOnStateChange → state transitions publish to the event bus
func wireAgentEvents(sessID string, ag *agent.Agent, al *AgentLog) {
ag.SetOutput(NewEventHandler(al))
ag.SetOnStateChange(func(agentID, state string) {
session.PublishEvent("session."+sessID+".agent."+agentID+".state", state)
})
@ -95,24 +101,25 @@ func readTools(_ HandlerCtx) ([]byte, error) {
}
func requestRootCtl(ctx HandlerCtx, data []byte) ([]byte, error) {
handlers := map[string]rdwrHandler{
"invalidate": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if ctx.Models != nil {
ctx.Models.Invalidate()
}
if ctx.Invalidate != nil {
ctx.Invalidate()
}
return []byte("ok\n"), nil
},
"kill": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if ctx.Shutdown != nil {
ctx.Shutdown()
}
return []byte("ok\n"), nil
},
}
return rdwrDispatch(handlers)(ctx, data)
return rdwrDispatch(rootCtlHandlers)(ctx, data)
}
var rootCtlHandlers = map[string]rdwrHandler{
"invalidate": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if ctx.Models != nil {
ctx.Models.Invalidate()
}
if ctx.Invalidate != nil {
ctx.Invalidate()
}
return []byte("ok\n"), nil
},
"kill": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if ctx.Shutdown != nil {
ctx.Shutdown()
}
return []byte("ok\n"), nil
},
}
func readEventwait(_ HandlerCtx) ([]byte, error) {

View File

@ -114,39 +114,40 @@ func readSessionConnected(ctx HandlerCtx) ([]byte, error) {
}
func requestSessionCtl(ctx HandlerCtx, data []byte) ([]byte, error) {
handlers := map[string]rdwrHandler{
"kill": func(ctx HandlerCtx, _ []string) ([]byte, error) {
ctx.Remove()
return []byte("ok\n"), nil
},
".": func(ctx HandlerCtx, _ []string) ([]byte, error) {
ctx.Remove()
return []byte("ok\n"), nil
},
"save": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if ctx.Session.Core != nil {
ctx.Session.Core.Save()
}
return []byte("ok\n"), nil
},
"invalidate": func(ctx HandlerCtx, _ []string) ([]byte, error) {
ctx.Session.InvalidateModelsCache()
return []byte("ok\n"), nil
},
"pause": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if err := ctx.Session.Pause(); err != nil {
return nil, err
}
return []byte("ok\n"), nil
},
"resume": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if err := ctx.Session.Resume(); err != nil {
return nil, err
}
return []byte("ok\n"), nil
},
}
return rdwrDispatch(handlers)(ctx, data)
return rdwrDispatch(sessionCtlHandlers)(ctx, data)
}
var sessionCtlHandlers = map[string]rdwrHandler{
"kill": func(ctx HandlerCtx, _ []string) ([]byte, error) {
ctx.Remove()
return []byte("ok\n"), nil
},
".": func(ctx HandlerCtx, _ []string) ([]byte, error) {
ctx.Remove()
return []byte("ok\n"), nil
},
"save": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if ctx.Session.Core != nil {
ctx.Session.Core.Save()
}
return []byte("ok\n"), nil
},
"invalidate": func(ctx HandlerCtx, _ []string) ([]byte, error) {
ctx.Session.InvalidateModelsCache()
return []byte("ok\n"), nil
},
"pause": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if err := ctx.Session.Pause(); err != nil {
return nil, err
}
return []byte("ok\n"), nil
},
"resume": func(ctx HandlerCtx, _ []string) ([]byte, error) {
if err := ctx.Session.Resume(); err != nil {
return nil, err
}
return []byte("ok\n"), nil
},
}
func readSessionName(ctx HandlerCtx) ([]byte, error) {

View File

@ -81,8 +81,7 @@ func NewRoot(cfg Config) *Tree {
replayMessagesToLog(al, ra.Messages)
// Wire up the event handler so agent events flow to the AgentLog
if ag := r.Session.FindAgent(ra.ID); ag != nil {
ag.SetOutput(NewEventHandler(al))
wireAgentStateEvents(r.Session.ID, ag)
wireAgentEvents(r.Session.ID, ag, al)
}
}
}

View File

@ -283,13 +283,7 @@ func (s *Session) LoadTool(name string, ag *agent.Agent) error {
if name == "" {
return nil
}
var conn *toolsrv.Conn
if s.ToolsConn() != nil {
conn = s.ToolsConn()
} else if ag != nil {
conn = ag.ToolServer()
}
return LoadToolOnConn(conn, name)
return LoadToolOnConn(s.ToolsConn(), name)
}
// LoadToolOnConn loads a tool on the given connection.

View File

@ -11,7 +11,6 @@ import (
"os"
"os/exec"
"path/filepath"
"regexp"
"strings"
"sync"
"syscall"
@ -21,51 +20,6 @@ import (
"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 (s *Server) ValidateCode(code, language string) error {
if err := s.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) {
s.recordValidationFailure()
return fmt.Errorf("dangerous pattern detected")
}
}
return nil
}
func loadSandboxConfig(name string) (*sandbox.Config, error) {
if name == "" {
name = "default"
@ -113,6 +67,7 @@ func (lw *limitedWriter) Write(p []byte) (n int, err error) {
}
return len(p), nil
}
// Connects to the broker socket, sends the request with the current env,
// and streams the framed response back.
func (s *Server) executeBypass(ctx context.Context, cmd, dir string, timeout int, doDetach ...bool) (string, error) {
@ -243,12 +198,12 @@ func (s *Server) executeBypassOpts(ctx context.Context, cmd, dir string, timeout
ring := NewRingBuffer(RingBufSize)
pid := int(time.Now().UnixNano() & 0x7FFFFFFF) // synthetic PID
proc := &DetachedProcess{
proc := &DetachedProcess{
PID: pid,
Command: cmdStr,
Started: time.Now(),
Ring: ring,
Done: make(chan struct{}),
Command: cmdStr,
Started: time.Now(),
Ring: ring,
Done: make(chan struct{}),
}
s.detachMu.Lock()
s.detached = append(s.detached, proc)
@ -370,38 +325,6 @@ func (s *Server) executeBypassOpts(ctx context.Context, cmd, dir string, timeout
}
}
var whitespacePattern = regexp.MustCompile(`\s+`)
func (s *Server) checkRateLimit() error {
s.rateLimitMu.Lock()
defer s.rateLimitMu.Unlock()
now := time.Now()
if now.Before(s.blockedUntil) {
remaining := s.blockedUntil.Sub(now).Round(time.Second)
return &RateLimitedError{Remaining: remaining}
}
return nil
}
func (s *Server) recordValidationFailure() {
s.rateLimitMu.Lock()
defer s.rateLimitMu.Unlock()
now := time.Now()
if now.Sub(s.lastFailure) > failureWindow {
s.validationFailures = 0
}
s.validationFailures++
s.lastFailure = now
if s.validationFailures >= maxFailures {
s.blockedUntil = now.Add(blockDuration)
s.validationFailures = 0
}
}
// executeWithStdin runs code in a sandbox and returns combined stdout+stderr.
// For languages where code is itself passed via stdin (ed, expect, bc), stdinData is ignored.
func (s *Server) executeWithStdin(ctx context.Context, code, language string, timeout int, sandboxName string, trusted bool, stdinData string, doDetach ...bool) (string, error) {
@ -435,7 +358,7 @@ func (s *Server) executeWithStdin(ctx context.Context, code, language string, ti
if timeout > 0 {
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
} else {
// timeout=0 means no timeout (run indefinitely, s.g. for daemons)
// 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
@ -542,7 +465,7 @@ func (s *Server) executeWithStdin(ctx context.Context, code, language string, ti
if err = cmd.Start(); err != nil {
cancel()
return "", fmt.Errorf("execution failed: %v", err)
return "", fmt.Errorf("execution failed: %w", err)
}
// If detach requested, immediately signal the detach channel
@ -573,7 +496,7 @@ func (s *Server) executeWithStdin(ctx context.Context, code, language string, ti
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), fmt.Errorf("execution failed: %w\nOutput: %s", err, errOutput)
}
return string(output), nil
@ -607,12 +530,12 @@ func (s *Server) executeWithStdin(ctx context.Context, code, language string, ti
lw.stream = nil
lw.mu.Unlock()
proc := &DetachedProcess{
PID: cmd.Process.Pid,
Command: cmdStr,
Started: time.Now(),
Ring: ring,
Done: make(chan struct{}),
proc := &DetachedProcess{
PID: cmd.Process.Pid,
Command: cmdStr,
Started: time.Now(),
Ring: ring,
Done: make(chan struct{}),
}
s.detachMu.Lock()
s.detached = append(s.detached, proc)

86
toolsrv/shell_validate.go Normal file
View File

@ -0,0 +1,86 @@
package toolsrv
import (
"fmt"
"regexp"
"strings"
"time"
)
// 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,
}
var whitespacePattern = regexp.MustCompile(`\s+`)
// ValidateCode checks code against dangerous patterns.
func (s *Server) ValidateCode(code, language string) error {
if err := s.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) {
s.recordValidationFailure()
return fmt.Errorf("dangerous pattern detected")
}
}
return nil
}
func (s *Server) checkRateLimit() error {
s.rateLimitMu.Lock()
defer s.rateLimitMu.Unlock()
now := time.Now()
if now.Before(s.blockedUntil) {
remaining := s.blockedUntil.Sub(now).Round(time.Second)
return &RateLimitedError{Remaining: remaining}
}
return nil
}
func (s *Server) recordValidationFailure() {
s.rateLimitMu.Lock()
defer s.rateLimitMu.Unlock()
now := time.Now()
if now.Sub(s.lastFailure) > failureWindow {
s.validationFailures = 0
}
s.validationFailures++
s.lastFailure = now
if s.validationFailures >= maxFailures {
s.blockedUntil = now.Add(blockDuration)
s.validationFailures = 0
}
}