232 lines
6.2 KiB
Go
232 lines
6.2 KiB
Go
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
|
|
}
|