From c25243a690ed5823cffa2f815fcd321b04fd250e Mon Sep 17 00:00:00 2001 From: Victor Vargas Date: Tue, 30 Jun 2026 23:53:32 -0700 Subject: [PATCH] feat(tools): add Tool definitions, Registry pattern, and filesystem sandbox with path validation --- pkg/tools/registry.go | 65 +++++++++++++ pkg/tools/registry_test.go | 114 +++++++++++++++++++++++ pkg/tools/sandbox/sandbox.go | 115 +++++++++++++++++++++++ pkg/tools/sandbox/sandbox_test.go | 149 ++++++++++++++++++++++++++++++ pkg/tools/types.go | 79 ++++++++++++++++ 5 files changed, 522 insertions(+) create mode 100644 pkg/tools/registry.go create mode 100644 pkg/tools/registry_test.go create mode 100644 pkg/tools/sandbox/sandbox.go create mode 100644 pkg/tools/sandbox/sandbox_test.go create mode 100644 pkg/tools/types.go diff --git a/pkg/tools/registry.go b/pkg/tools/registry.go new file mode 100644 index 0000000..dc64673 --- /dev/null +++ b/pkg/tools/registry.go @@ -0,0 +1,65 @@ +package tools + +import ( + "sync" +) + +// registry is the default implementation of Registry. +type registry struct { + mu sync.RWMutex + tools map[string]Tool + order []string +} + +// NewRegistry returns a new empty registry. +func NewRegistry() Registry { + return ®istry{ + tools: make(map[string]Tool), + } +} + +func (r *registry) Register(t Tool) error { + r.mu.Lock() + defer r.mu.Unlock() + + if _, exists := r.tools[t.Name]; exists { + return ErrDuplicateTool + } + + r.tools[t.Name] = t + r.order = append(r.order, t.Name) + return nil +} + +func (r *registry) Get(name string) (Tool, bool) { + r.mu.RLock() + defer r.mu.RUnlock() + + t, ok := r.tools[name] + return t, ok +} + +func (r *registry) List() []Tool { + r.mu.RLock() + defer r.mu.RUnlock() + + out := make([]Tool, len(r.order)) + for i, name := range r.order { + out[i] = r.tools[name] + } + return out +} + +func (r *registry) Filter(predicate func(Tool) bool) []Tool { + r.mu.RLock() + defer r.mu.RUnlock() + + var result []Tool + for _, name := range r.order { + t := r.tools[name] + if predicate(t) { + result = append(result, t) + } + } + return result +} diff --git a/pkg/tools/registry_test.go b/pkg/tools/registry_test.go new file mode 100644 index 0000000..31e6c80 --- /dev/null +++ b/pkg/tools/registry_test.go @@ -0,0 +1,114 @@ +package tools + +import ( + "fmt" + "sync" + "testing" +) + +func TestNewRegistry(t *testing.T) { + reg := NewRegistry() + if reg == nil { + t.Fatal("expected non-nil registry") + } + if len(reg.List()) != 0 { + t.Errorf("expected empty registry, got %d tools", len(reg.List())) + } +} + +func TestRegisterAndGet(t *testing.T) { + reg := NewRegistry() + tool := Tool{Name: "test", Description: "test tool"} + if err := reg.Register(tool); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + got, ok := reg.Get("test") + if !ok { + t.Fatal("expected tool to be found") + } + if got.Name != "test" { + t.Errorf("expected name 'test', got %q", got.Name) + } +} + +func TestRegisterDuplicate(t *testing.T) { + reg := NewRegistry() + tool := Tool{Name: "test"} + if err := reg.Register(tool); err != nil { + t.Fatalf("unexpected error on first register: %v", err) + } + if err := reg.Register(tool); err == nil { + t.Error("expected error for duplicate tool") + } +} + +func TestList(t *testing.T) { + reg := NewRegistry() + tools := []Tool{ + {Name: "a"}, + {Name: "b"}, + {Name: "c"}, + } + for _, tw := range tools { + if err := reg.Register(tw); err != nil { + t.Fatalf("unexpected error: %v", err) + } + } + + list := reg.List() + if len(list) != 3 { + t.Errorf("expected 3 tools, got %d", len(list)) + } +} + +func TestFilter(t *testing.T) { + reg := NewRegistry() + tools := []Tool{ + {Name: "read", Permission: Allow}, + {Name: "write", Permission: Ask}, + {Name: "delete", Permission: Deny}, + } + for _, tw := range tools { + if err := reg.Register(tw); err != nil { + t.Fatalf("unexpected error: %v", err) + } + } + + allowed := reg.Filter(func(t Tool) bool { + return t.Permission == Allow + }) + if len(allowed) != 1 { + t.Errorf("expected 1 allowed tool, got %d", len(allowed)) + } + if len(allowed) > 0 && allowed[0].Name != "read" { + t.Errorf("expected 'read' tool, got %q", allowed[0].Name) + } +} + +func TestGetNotFound(t *testing.T) { + reg := NewRegistry() + _, ok := reg.Get("nonexistent") + if ok { + t.Error("expected false for nonexistent tool") + } +} + +func TestConcurrentRegister(t *testing.T) { + reg := NewRegistry() + var wg sync.WaitGroup + for i := 0; i < 100; i++ { + wg.Add(1) + go func(n int) { + defer wg.Done() + name := fmt.Sprintf("tool_%d", n) + _ = reg.Register(Tool{Name: name}) + }(i) + } + wg.Wait() + + list := reg.List() + if len(list) != 100 { + t.Errorf("expected 100 tools, got %d", len(list)) + } +} diff --git a/pkg/tools/sandbox/sandbox.go b/pkg/tools/sandbox/sandbox.go new file mode 100644 index 0000000..588944c --- /dev/null +++ b/pkg/tools/sandbox/sandbox.go @@ -0,0 +1,115 @@ +package sandbox + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" + "strings" + + "github.com/VictorVargas/rony-llm-agent/pkg/llm" + "github.com/VictorVargas/rony-llm-agent/pkg/tools" +) + +// ErrSandboxViolation is returned when a tool call violates sandbox rules. +var ErrSandboxViolation = fmt.Errorf("sandbox violation") + +// Sandbox wraps os.Root to enforce filesystem boundaries for tool calls. +type Sandbox struct { + root *os.Root + path string +} + +// NewSandbox creates a new Sandbox rooted at the given directory. +// The directory is created if it does not exist. +func NewSandbox(rootDir string) (*Sandbox, error) { + if err := os.MkdirAll(rootDir, 0755); err != nil { + return nil, fmt.Errorf("creating sandbox root: %w", err) + } + + absRoot, err := filepath.Abs(rootDir) + if err != nil { + return nil, fmt.Errorf("resolving absolute path: %w", err) + } + + root, err := os.OpenRoot(absRoot) + if err != nil { + return nil, fmt.Errorf("opening sandbox root: %w", err) + } + + return &Sandbox{root: root, path: absRoot}, nil +} + +// ValidateToolCall checks if a tool call's arguments reference paths outside the sandbox. +// It returns nil if the call is allowed, or an error explaining the violation. +func (s *Sandbox) ValidateToolCall(tool tools.Tool, call llm.ToolCall) error { + args := make(map[string]interface{}) + if len(call.Arguments) > 0 { + if err := json.Unmarshal(call.Arguments, &args); err != nil { + return fmt.Errorf("parsing arguments: %w", err) + } + } + + for _, path := range extractPaths(args) { + if err := s.validatePath(path); err != nil { + return fmt.Errorf("%w: %s", ErrSandboxViolation, err) + } + } + + return nil +} + +// validatePath checks that the given path resolves inside the sandbox root. +func (s *Sandbox) validatePath(path string) error { + // Absolute paths are rejected (they escape the sandbox by definition) + if filepath.IsAbs(path) { + return fmt.Errorf("absolute paths not allowed: %s", path) + } + + // Resolve to absolute path relative to sandbox root + abs := filepath.Join(s.path, path) + + // Clean the path to normalize it + abs = filepath.Clean(abs) + + // Check if the path is inside the root + if !strings.HasPrefix(abs, s.path+string(filepath.Separator)) && abs != s.path { + return fmt.Errorf("path escapes sandbox: %s (root: %s)", abs, s.path) + } + + return nil +} + +// extractPaths collects all path-like values from the arguments map. +func extractPaths(args map[string]interface{}) []string { + var paths []string + for _, v := range args { + switch val := v.(type) { + case string: + if isPathLike(val) { + paths = append(paths, val) + } + case []interface{}: + for _, item := range val { + if str, ok := item.(string); ok && isPathLike(str) { + paths = append(paths, str) + } + } + } + } + return paths +} + +// isPathLike checks if a string looks like a filesystem path. +func isPathLike(s string) bool { + // Must start with / or ./ or ../ or contain a file extension + return strings.HasPrefix(s, "/") || + strings.HasPrefix(s, "./") || + strings.HasPrefix(s, "../") || + strings.Contains(s, ".") && strings.Contains(s, "/") +} + +// Open returns the underlying os.Root for testing. +func (s *Sandbox) Open() *os.Root { + return s.root +} diff --git a/pkg/tools/sandbox/sandbox_test.go b/pkg/tools/sandbox/sandbox_test.go new file mode 100644 index 0000000..905e9ca --- /dev/null +++ b/pkg/tools/sandbox/sandbox_test.go @@ -0,0 +1,149 @@ +package sandbox_test + +import ( + "encoding/json" + "os" + "path/filepath" + "testing" + + "github.com/VictorVargas/rony-llm-agent/pkg/llm" + "github.com/VictorVargas/rony-llm-agent/pkg/tools" + "github.com/VictorVargas/rony-llm-agent/pkg/tools/sandbox" +) + +func TestNewSandbox(t *testing.T) { + dir := t.TempDir() + + sb, err := sandbox.NewSandbox(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if sb == nil { + t.Fatal("expected non-nil sandbox") + } +} + +func TestValidatePath_Allowed(t *testing.T) { + dir := t.TempDir() + sb, err := sandbox.NewSandbox(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + // Create a file inside the sandbox + testFile := filepath.Join(dir, "test.txt") + if err := os.WriteFile(testFile, []byte("hello"), 0644); err != nil { + t.Fatalf("creating test file: %v", err) + } + + call := llm.ToolCall{ + Name: "read_file", + Arguments: json.RawMessage(`{"path": "test.txt"}`), + } + + err = sb.ValidateToolCall(tools.Tool{}, call) + if err != nil { + t.Fatalf("expected no error, got: %v", err) + } +} + +func TestValidatePath_EscapesSandbox(t *testing.T) { + dir := t.TempDir() + sb, err := sandbox.NewSandbox(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + call := llm.ToolCall{ + Name: "read_file", + Arguments: json.RawMessage(`{"path": "/etc/passwd"}`), + } + + err = sb.ValidateToolCall(tools.Tool{}, call) + if err == nil { + t.Fatal("expected sandbox violation error") + } +} + +func TestValidatePath_SymlinkEscape(t *testing.T) { + dir := t.TempDir() + + // Create a symlink that escapes the sandbox + escapeDir := t.TempDir() + symlinkPath := filepath.Join(dir, "link") + if err := os.Symlink(escapeDir, symlinkPath); err != nil { + t.Fatalf("creating symlink: %v", err) + } + + sb, err := sandbox.NewSandbox(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + // The sandbox should allow symlinks inside the root + call := llm.ToolCall{ + Name: "read_file", + Arguments: json.RawMessage(`{"path": "link"}`), + } + + err = sb.ValidateToolCall(tools.Tool{}, call) + if err != nil { + t.Fatalf("expected symlink to be allowed (path is inside sandbox), got: %v", err) + } +} + +func TestValidatePath_NonExistentFile(t *testing.T) { + dir := t.TempDir() + sb, err := sandbox.NewSandbox(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + // Non-existent file inside sandbox should be allowed + call := llm.ToolCall{ + Name: "read_file", + Arguments: json.RawMessage(`{"path": "nonexistent.txt"}`), + } + + err = sb.ValidateToolCall(tools.Tool{}, call) + if err != nil { + t.Fatalf("expected no error for non-existent file, got: %v", err) + } +} + +func TestValidatePath_InvalidJSON(t *testing.T) { + dir := t.TempDir() + sb, err := sandbox.NewSandbox(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + call := llm.ToolCall{ + Name: "read_file", + Arguments: json.RawMessage(`not json`), + } + + err = sb.ValidateToolCall(tools.Tool{}, call) + if err == nil { + t.Fatal("expected error for invalid JSON") + } +} + +func TestValidatePath_NoPathsInArgs(t *testing.T) { + dir := t.TempDir() + sb, err := sandbox.NewSandbox(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + // Arguments with no paths should not cause errors + call := llm.ToolCall{ + Name: "math_add", + Arguments: json.RawMessage(`{"a": 5, "b": 10}`), + } + + err = sb.ValidateToolCall(tools.Tool{}, call) + if err != nil { + t.Fatalf("expected no error, got: %v", err) + } +} diff --git a/pkg/tools/types.go b/pkg/tools/types.go new file mode 100644 index 0000000..1db51cd --- /dev/null +++ b/pkg/tools/types.go @@ -0,0 +1,79 @@ +package tools + +import ( + "context" + "encoding/json" + "fmt" +) + +// Permission controls whether a tool can be executed without user approval. +type Permission int + +const ( + Allow Permission = iota + Ask + Deny +) + +func (p Permission) String() string { + switch p { + case Allow: + return "allow" + case Ask: + return "ask" + case Deny: + return "deny" + default: + return "unknown" + } +} + +// Tool is a function callable by the LLM with JSON Schema input. +type Tool struct { + Name string + Description string + InputSchema json.RawMessage // JSON Schema draft-07+ + Required []string + Handler ToolHandler + Permission Permission + Examples []ToolExample +} + +// ToolHandler is the function type that implements a tool. +// The handler receives raw JSON arguments and returns a result. +type ToolHandler func(ctx context.Context, args json.RawMessage) (ToolResult, error) + +// ToolResult is returned by a ToolHandler. +type ToolResult struct { + Content string + IsError bool + Metadata map[string]string + Artifacts []Artifact +} + +// Artifact represents a file or data artifact produced by a tool. +type Artifact struct { + Path string + Content []byte + MIME string +} + +// ToolExample provides few-shot examples for the LLM to improve tool usage. +type ToolExample struct { + Input map[string]interface{} + Output string +} + +// Registry manages tool registration and lookup. +type Registry interface { + Register(tool Tool) error + Get(name string) (Tool, bool) + List() []Tool + Filter(predicate func(Tool) bool) []Tool +} + +// ErrToolNotFound is returned when a tool is not found in the registry. +var ErrToolNotFound = fmt.Errorf("tool not found") + +// ErrDuplicateTool is returned when trying to register a tool with an existing name. +var ErrDuplicateTool = fmt.Errorf("duplicate tool name")