This repository has been archived on 2026-08-16. You can view files and clone it, but cannot push or open issues or pull requests.
ollie-core/tools/code.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
}