feat(tools): add Tool definitions, Registry pattern, and filesystem sandbox with path validation

This commit is contained in:
Victor Hugo Vargas Servin 2026-06-30 23:53:32 -07:00
parent 9f33d6df93
commit c25243a690
5 changed files with 522 additions and 0 deletions

65
pkg/tools/registry.go Normal file
View file

@ -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 &registry{
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
}

114
pkg/tools/registry_test.go Normal file
View file

@ -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))
}
}

View file

@ -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
}

View file

@ -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)
}
}

79
pkg/tools/types.go Normal file
View file

@ -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")