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:
parent
f5efd90fbe
commit
2b4d577e39
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
107
toolsrv/shell.go
107
toolsrv/shell.go
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
Loading…
Reference in New Issue