feat(tools): refactor + tool optionality
This commit is contained in:
parent
9a6b49f9ba
commit
54c178d5d8
14 changed files with 700 additions and 604 deletions
|
|
@ -48,24 +48,8 @@ func NewAgentInstance(
|
||||||
|
|
||||||
restrict := defaults.RestrictToWorkspace
|
restrict := defaults.RestrictToWorkspace
|
||||||
toolsRegistry := tools.NewToolRegistry()
|
toolsRegistry := tools.NewToolRegistry()
|
||||||
if cfg.Tools.Filesystem.EnableRead {
|
// initialize workspace tools
|
||||||
toolsRegistry.Register(tools.NewReadFileTool(workspace, restrict))
|
tools.SetupWorkspaceTools(toolsRegistry, cfg, workspace, restrict)
|
||||||
}
|
|
||||||
if cfg.Tools.Filesystem.EnableWrite {
|
|
||||||
toolsRegistry.Register(tools.NewWriteFileTool(workspace, restrict))
|
|
||||||
}
|
|
||||||
if cfg.Tools.Filesystem.EnableList {
|
|
||||||
toolsRegistry.Register(tools.NewListDirTool(workspace, restrict))
|
|
||||||
}
|
|
||||||
if cfg.Tools.Filesystem.EnableEdit {
|
|
||||||
toolsRegistry.Register(tools.NewEditFileTool(workspace, restrict))
|
|
||||||
}
|
|
||||||
if cfg.Tools.Filesystem.EnableAppend {
|
|
||||||
toolsRegistry.Register(tools.NewAppendFileTool(workspace, restrict))
|
|
||||||
}
|
|
||||||
if cfg.Tools.Exec.Enabled {
|
|
||||||
toolsRegistry.Register(tools.NewExecToolWithConfig(workspace, restrict, cfg))
|
|
||||||
}
|
|
||||||
|
|
||||||
sessionsDir := filepath.Join(workspace, "sessions")
|
sessionsDir := filepath.Join(workspace, "sessions")
|
||||||
sessionsManager := session.NewSessionManager(sessionsDir)
|
sessionsManager := session.NewSessionManager(sessionsDir)
|
||||||
|
|
|
||||||
|
|
@ -23,7 +23,6 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
"github.com/sipeed/picoclaw/pkg/routing"
|
"github.com/sipeed/picoclaw/pkg/routing"
|
||||||
"github.com/sipeed/picoclaw/pkg/skills"
|
|
||||||
"github.com/sipeed/picoclaw/pkg/state"
|
"github.com/sipeed/picoclaw/pkg/state"
|
||||||
"github.com/sipeed/picoclaw/pkg/tools"
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
|
@ -92,76 +91,25 @@ func registerSharedTools(
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Web tools
|
// specific context for this agent
|
||||||
if searchTool := tools.NewWebSearchTool(tools.WebSearchToolOptions{
|
agentCtx := tools.AgentContext{
|
||||||
BraveAPIKey: cfg.Tools.Web.Brave.APIKey,
|
AgentID: agentID,
|
||||||
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
|
Workspace: agent.Workspace,
|
||||||
BraveEnabled: cfg.Tools.Web.Brave.Enabled,
|
Model: agent.Model,
|
||||||
TavilyAPIKey: cfg.Tools.Web.Tavily.APIKey,
|
MaxTokens: agent.MaxTokens,
|
||||||
TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL,
|
Temperature: agent.Temperature,
|
||||||
TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults,
|
|
||||||
TavilyEnabled: cfg.Tools.Web.Tavily.Enabled,
|
|
||||||
DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults,
|
|
||||||
DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled,
|
|
||||||
PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey,
|
|
||||||
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
|
||||||
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
|
||||||
}); searchTool != nil {
|
|
||||||
agent.Tools.Register(searchTool)
|
|
||||||
}
|
|
||||||
if cfg.Tools.Core.EnableWebFetch {
|
|
||||||
agent.Tools.Register(tools.NewWebFetchTool(50000))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Hardware tools (I2C, SPI) - Linux only, returns error on other platforms
|
// subagent security checker
|
||||||
if cfg.Tools.Hardware.EnableI2C {
|
|
||||||
agent.Tools.Register(tools.NewI2CTool())
|
|
||||||
}
|
|
||||||
if cfg.Tools.Hardware.EnableSPI {
|
|
||||||
agent.Tools.Register(tools.NewSPITool())
|
|
||||||
}
|
|
||||||
|
|
||||||
// Message tool
|
|
||||||
if cfg.Tools.Core.EnableMessage {
|
|
||||||
messageTool := tools.NewMessageTool()
|
|
||||||
messageTool.SetSendCallback(func(channel, chatID, content string) error {
|
|
||||||
msgBus.PublishOutbound(bus.OutboundMessage{
|
|
||||||
Channel: channel,
|
|
||||||
ChatID: chatID,
|
|
||||||
Content: content,
|
|
||||||
})
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
agent.Tools.Register(messageTool)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Skill discovery and installation tools
|
|
||||||
if cfg.Tools.Skills.Enabled {
|
|
||||||
registryMgr := skills.NewRegistryManagerFromConfig(skills.RegistryConfig{
|
|
||||||
MaxConcurrentSearches: cfg.Tools.Skills.MaxConcurrentSearches,
|
|
||||||
ClawHub: skills.ClawHubConfig(cfg.Tools.Skills.Registries.ClawHub),
|
|
||||||
})
|
|
||||||
searchCache := skills.NewSearchCache(
|
|
||||||
cfg.Tools.Skills.SearchCache.MaxSize,
|
|
||||||
time.Duration(cfg.Tools.Skills.SearchCache.TTLSeconds)*time.Second,
|
|
||||||
)
|
|
||||||
agent.Tools.Register(tools.NewFindSkillsTool(registryMgr, searchCache))
|
|
||||||
agent.Tools.Register(tools.NewInstallSkillTool(registryMgr, agent.Workspace))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Spawn tool with allowlist checker
|
|
||||||
if cfg.Tools.Core.EnableSpawn {
|
|
||||||
subagentManager := tools.NewSubagentManager(provider, agent.Model, agent.Workspace, msgBus)
|
|
||||||
subagentManager.SetLLMOptions(agent.MaxTokens, agent.Temperature)
|
|
||||||
spawnTool := tools.NewSpawnTool(subagentManager)
|
|
||||||
currentAgentID := agentID
|
currentAgentID := agentID
|
||||||
spawnTool.SetAllowlistChecker(func(targetAgentID string) bool {
|
canSpawn := func(targetAgentID string) bool {
|
||||||
return registry.CanSpawnSubagent(currentAgentID, targetAgentID)
|
return registry.CanSpawnSubagent(currentAgentID, targetAgentID)
|
||||||
})
|
|
||||||
agent.Tools.Register(spawnTool)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Update context builder with the complete tools registry
|
// initialization
|
||||||
|
tools.SetupSharedTools(agent.Tools, cfg, msgBus, provider, agentCtx, canSpawn)
|
||||||
|
|
||||||
|
// update context builder
|
||||||
agent.ContextBuilder.SetToolsRegistry(agent.Tools)
|
agent.ContextBuilder.SetToolsRegistry(agent.Tools)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -304,8 +304,8 @@ func TestAgentLoop_GetStartupInfo(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Should have default tools registered
|
// Should have default tools registered
|
||||||
if count.(int) == 0 {
|
if count.(int) != 0 {
|
||||||
t.Error("Expected at least some tools to be registered")
|
t.Error("registered tools that have not been enabled")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
93
pkg/tools/append.go
Normal file
93
pkg/tools/append.go
Normal file
|
|
@ -0,0 +1,93 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io/fs"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
type AppendFileTool struct {
|
||||||
|
fs fileSystem
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewAppendFileTool(workspace string, restrict bool) *AppendFileTool {
|
||||||
|
var fs fileSystem
|
||||||
|
if restrict {
|
||||||
|
fs = &sandboxFs{workspace: workspace}
|
||||||
|
} else {
|
||||||
|
fs = &hostFs{}
|
||||||
|
}
|
||||||
|
return &AppendFileTool{fs: fs}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *AppendFileTool) Name() string {
|
||||||
|
return "append_file"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *AppendFileTool) Description() string {
|
||||||
|
return "Append content to the end of a file"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *AppendFileTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"path": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "The file path to append to",
|
||||||
|
},
|
||||||
|
"content": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "The content to append",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"path", "content"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *AppendFileTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
path, ok := args["path"].(string)
|
||||||
|
if !ok {
|
||||||
|
return ErrorResult("path is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
content, ok := args["content"].(string)
|
||||||
|
if !ok {
|
||||||
|
return ErrorResult("content is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := appendFile(t.fs, path, content); err != nil {
|
||||||
|
return ErrorResult(err.Error())
|
||||||
|
}
|
||||||
|
return SilentResult(fmt.Sprintf("Appended to %s", path))
|
||||||
|
}
|
||||||
|
|
||||||
|
// appendFile reads the existing content (if any) via sysFs, appends new content, and writes back.
|
||||||
|
func appendFile(sysFs fileSystem, path, appendContent string) error {
|
||||||
|
content, err := sysFs.ReadFile(path)
|
||||||
|
if err != nil && !errors.Is(err, fs.ErrNotExist) {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
newContent := append(content, []byte(appendContent)...)
|
||||||
|
return sysFs.WriteFile(path, newContent)
|
||||||
|
}
|
||||||
|
|
||||||
|
// replaceEditContent handles the core logic of finding and replacing a single occurrence of oldText.
|
||||||
|
func replaceEditContent(content []byte, oldText, newText string) ([]byte, error) {
|
||||||
|
contentStr := string(content)
|
||||||
|
|
||||||
|
if !strings.Contains(contentStr, oldText) {
|
||||||
|
return nil, fmt.Errorf("old_text not found in file. Make sure it matches exactly")
|
||||||
|
}
|
||||||
|
|
||||||
|
count := strings.Count(contentStr, oldText)
|
||||||
|
if count > 1 {
|
||||||
|
return nil, fmt.Errorf("old_text appears %d times. Please provide more context to make it unique", count)
|
||||||
|
}
|
||||||
|
|
||||||
|
newContent := strings.Replace(contentStr, oldText, newText, 1)
|
||||||
|
return []byte(newContent), nil
|
||||||
|
}
|
||||||
|
|
@ -2,10 +2,7 @@ package tools
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io/fs"
|
|
||||||
"strings"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// EditFileTool edits a file by replacing old_text with new_text.
|
// EditFileTool edits a file by replacing old_text with new_text.
|
||||||
|
|
@ -76,62 +73,6 @@ func (t *EditFileTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
return SilentResult(fmt.Sprintf("File edited: %s", path))
|
return SilentResult(fmt.Sprintf("File edited: %s", path))
|
||||||
}
|
}
|
||||||
|
|
||||||
type AppendFileTool struct {
|
|
||||||
fs fileSystem
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewAppendFileTool(workspace string, restrict bool) *AppendFileTool {
|
|
||||||
var fs fileSystem
|
|
||||||
if restrict {
|
|
||||||
fs = &sandboxFs{workspace: workspace}
|
|
||||||
} else {
|
|
||||||
fs = &hostFs{}
|
|
||||||
}
|
|
||||||
return &AppendFileTool{fs: fs}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *AppendFileTool) Name() string {
|
|
||||||
return "append_file"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *AppendFileTool) Description() string {
|
|
||||||
return "Append content to the end of a file"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *AppendFileTool) Parameters() map[string]any {
|
|
||||||
return map[string]any{
|
|
||||||
"type": "object",
|
|
||||||
"properties": map[string]any{
|
|
||||||
"path": map[string]any{
|
|
||||||
"type": "string",
|
|
||||||
"description": "The file path to append to",
|
|
||||||
},
|
|
||||||
"content": map[string]any{
|
|
||||||
"type": "string",
|
|
||||||
"description": "The content to append",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"required": []string{"path", "content"},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *AppendFileTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
|
||||||
path, ok := args["path"].(string)
|
|
||||||
if !ok {
|
|
||||||
return ErrorResult("path is required")
|
|
||||||
}
|
|
||||||
|
|
||||||
content, ok := args["content"].(string)
|
|
||||||
if !ok {
|
|
||||||
return ErrorResult("content is required")
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := appendFile(t.fs, path, content); err != nil {
|
|
||||||
return ErrorResult(err.Error())
|
|
||||||
}
|
|
||||||
return SilentResult(fmt.Sprintf("Appended to %s", path))
|
|
||||||
}
|
|
||||||
|
|
||||||
// editFile reads the file via sysFs, performs the replacement, and writes back.
|
// editFile reads the file via sysFs, performs the replacement, and writes back.
|
||||||
// It uses a fileSystem interface, allowing the same logic for both restricted and unrestricted modes.
|
// It uses a fileSystem interface, allowing the same logic for both restricted and unrestricted modes.
|
||||||
func editFile(sysFs fileSystem, path, oldText, newText string) error {
|
func editFile(sysFs fileSystem, path, oldText, newText string) error {
|
||||||
|
|
@ -147,31 +88,3 @@ func editFile(sysFs fileSystem, path, oldText, newText string) error {
|
||||||
|
|
||||||
return sysFs.WriteFile(path, newContent)
|
return sysFs.WriteFile(path, newContent)
|
||||||
}
|
}
|
||||||
|
|
||||||
// appendFile reads the existing content (if any) via sysFs, appends new content, and writes back.
|
|
||||||
func appendFile(sysFs fileSystem, path, appendContent string) error {
|
|
||||||
content, err := sysFs.ReadFile(path)
|
|
||||||
if err != nil && !errors.Is(err, fs.ErrNotExist) {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
newContent := append(content, []byte(appendContent)...)
|
|
||||||
return sysFs.WriteFile(path, newContent)
|
|
||||||
}
|
|
||||||
|
|
||||||
// replaceEditContent handles the core logic of finding and replacing a single occurrence of oldText.
|
|
||||||
func replaceEditContent(content []byte, oldText, newText string) ([]byte, error) {
|
|
||||||
contentStr := string(content)
|
|
||||||
|
|
||||||
if !strings.Contains(contentStr, oldText) {
|
|
||||||
return nil, fmt.Errorf("old_text not found in file. Make sure it matches exactly")
|
|
||||||
}
|
|
||||||
|
|
||||||
count := strings.Count(contentStr, oldText)
|
|
||||||
if count > 1 {
|
|
||||||
return nil, fmt.Errorf("old_text appears %d times. Please provide more context to make it unique", count)
|
|
||||||
}
|
|
||||||
|
|
||||||
newContent := strings.Replace(contentStr, oldText, newText, 1)
|
|
||||||
return []byte(newContent), nil
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,6 @@
|
||||||
package tools
|
package tools
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io/fs"
|
"io/fs"
|
||||||
"os"
|
"os"
|
||||||
|
|
@ -81,159 +80,6 @@ func isWithinWorkspace(candidate, workspace string) bool {
|
||||||
return err == nil && filepath.IsLocal(rel)
|
return err == nil && filepath.IsLocal(rel)
|
||||||
}
|
}
|
||||||
|
|
||||||
type ReadFileTool struct {
|
|
||||||
fs fileSystem
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewReadFileTool(workspace string, restrict bool) *ReadFileTool {
|
|
||||||
var fs fileSystem
|
|
||||||
if restrict {
|
|
||||||
fs = &sandboxFs{workspace: workspace}
|
|
||||||
} else {
|
|
||||||
fs = &hostFs{}
|
|
||||||
}
|
|
||||||
return &ReadFileTool{fs: fs}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *ReadFileTool) Name() string {
|
|
||||||
return "read_file"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *ReadFileTool) Description() string {
|
|
||||||
return "Read the contents of a file"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *ReadFileTool) Parameters() map[string]any {
|
|
||||||
return map[string]any{
|
|
||||||
"type": "object",
|
|
||||||
"properties": map[string]any{
|
|
||||||
"path": map[string]any{
|
|
||||||
"type": "string",
|
|
||||||
"description": "Path to the file to read",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"required": []string{"path"},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *ReadFileTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
|
||||||
path, ok := args["path"].(string)
|
|
||||||
if !ok {
|
|
||||||
return ErrorResult("path is required")
|
|
||||||
}
|
|
||||||
|
|
||||||
content, err := t.fs.ReadFile(path)
|
|
||||||
if err != nil {
|
|
||||||
return ErrorResult(err.Error())
|
|
||||||
}
|
|
||||||
return NewToolResult(string(content))
|
|
||||||
}
|
|
||||||
|
|
||||||
type WriteFileTool struct {
|
|
||||||
fs fileSystem
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewWriteFileTool(workspace string, restrict bool) *WriteFileTool {
|
|
||||||
var fs fileSystem
|
|
||||||
if restrict {
|
|
||||||
fs = &sandboxFs{workspace: workspace}
|
|
||||||
} else {
|
|
||||||
fs = &hostFs{}
|
|
||||||
}
|
|
||||||
return &WriteFileTool{fs: fs}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WriteFileTool) Name() string {
|
|
||||||
return "write_file"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WriteFileTool) Description() string {
|
|
||||||
return "Write content to a file"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WriteFileTool) Parameters() map[string]any {
|
|
||||||
return map[string]any{
|
|
||||||
"type": "object",
|
|
||||||
"properties": map[string]any{
|
|
||||||
"path": map[string]any{
|
|
||||||
"type": "string",
|
|
||||||
"description": "Path to the file to write",
|
|
||||||
},
|
|
||||||
"content": map[string]any{
|
|
||||||
"type": "string",
|
|
||||||
"description": "Content to write to the file",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"required": []string{"path", "content"},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WriteFileTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
|
||||||
path, ok := args["path"].(string)
|
|
||||||
if !ok {
|
|
||||||
return ErrorResult("path is required")
|
|
||||||
}
|
|
||||||
|
|
||||||
content, ok := args["content"].(string)
|
|
||||||
if !ok {
|
|
||||||
return ErrorResult("content is required")
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := t.fs.WriteFile(path, []byte(content)); err != nil {
|
|
||||||
return ErrorResult(err.Error())
|
|
||||||
}
|
|
||||||
|
|
||||||
return SilentResult(fmt.Sprintf("File written: %s", path))
|
|
||||||
}
|
|
||||||
|
|
||||||
type ListDirTool struct {
|
|
||||||
fs fileSystem
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewListDirTool(workspace string, restrict bool) *ListDirTool {
|
|
||||||
var fs fileSystem
|
|
||||||
if restrict {
|
|
||||||
fs = &sandboxFs{workspace: workspace}
|
|
||||||
} else {
|
|
||||||
fs = &hostFs{}
|
|
||||||
}
|
|
||||||
return &ListDirTool{fs: fs}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *ListDirTool) Name() string {
|
|
||||||
return "list_dir"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *ListDirTool) Description() string {
|
|
||||||
return "List files and directories in a path"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *ListDirTool) Parameters() map[string]any {
|
|
||||||
return map[string]any{
|
|
||||||
"type": "object",
|
|
||||||
"properties": map[string]any{
|
|
||||||
"path": map[string]any{
|
|
||||||
"type": "string",
|
|
||||||
"description": "Path to list",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"required": []string{"path"},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *ListDirTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
|
||||||
path, ok := args["path"].(string)
|
|
||||||
if !ok {
|
|
||||||
path = "."
|
|
||||||
}
|
|
||||||
|
|
||||||
entries, err := t.fs.ReadDir(path)
|
|
||||||
if err != nil {
|
|
||||||
return ErrorResult(fmt.Sprintf("failed to read directory: %v", err))
|
|
||||||
}
|
|
||||||
return formatDirEntries(entries)
|
|
||||||
}
|
|
||||||
|
|
||||||
func formatDirEntries(entries []os.DirEntry) *ToolResult {
|
func formatDirEntries(entries []os.DirEntry) *ToolResult {
|
||||||
var result strings.Builder
|
var result strings.Builder
|
||||||
for _, entry := range entries {
|
for _, entry := range entries {
|
||||||
|
|
|
||||||
119
pkg/tools/init.go
Normal file
119
pkg/tools/init.go
Normal file
|
|
@ -0,0 +1,119 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/skills"
|
||||||
|
)
|
||||||
|
|
||||||
|
type AgentContext struct {
|
||||||
|
AgentID string
|
||||||
|
Workspace string
|
||||||
|
Model string
|
||||||
|
MaxTokens int
|
||||||
|
Temperature float64
|
||||||
|
}
|
||||||
|
|
||||||
|
type SpawnAllowlistChecker func(targetAgentID string) bool
|
||||||
|
|
||||||
|
func SetupSharedTools(
|
||||||
|
registry *ToolRegistry,
|
||||||
|
cfg *config.Config,
|
||||||
|
msgBus *bus.MessageBus,
|
||||||
|
provider providers.LLMProvider,
|
||||||
|
agentCtx AgentContext,
|
||||||
|
canSpawn SpawnAllowlistChecker,
|
||||||
|
) {
|
||||||
|
// Web tools
|
||||||
|
if searchTool := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
BraveAPIKey: cfg.Tools.Web.Brave.APIKey,
|
||||||
|
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
|
||||||
|
BraveEnabled: cfg.Tools.Web.Brave.Enabled,
|
||||||
|
TavilyAPIKey: cfg.Tools.Web.Tavily.APIKey,
|
||||||
|
TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL,
|
||||||
|
TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults,
|
||||||
|
TavilyEnabled: cfg.Tools.Web.Tavily.Enabled,
|
||||||
|
DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults,
|
||||||
|
DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled,
|
||||||
|
PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey,
|
||||||
|
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
|
||||||
|
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
|
||||||
|
}); searchTool != nil {
|
||||||
|
registry.Register(searchTool)
|
||||||
|
}
|
||||||
|
|
||||||
|
if cfg.Tools.Core.EnableWebFetch {
|
||||||
|
registry.Register(NewWebFetchTool(50000))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hardware tools
|
||||||
|
if cfg.Tools.Hardware.EnableI2C {
|
||||||
|
registry.Register(NewI2CTool())
|
||||||
|
}
|
||||||
|
if cfg.Tools.Hardware.EnableSPI {
|
||||||
|
registry.Register(NewSPITool())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Message tool
|
||||||
|
if cfg.Tools.Core.EnableMessage {
|
||||||
|
messageTool := NewMessageTool()
|
||||||
|
messageTool.SetSendCallback(func(channel, chatID, content string) error {
|
||||||
|
msgBus.PublishOutbound(bus.OutboundMessage{
|
||||||
|
Channel: channel,
|
||||||
|
ChatID: chatID,
|
||||||
|
Content: content,
|
||||||
|
})
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
registry.Register(messageTool)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Skills tools
|
||||||
|
if cfg.Tools.Skills.Enabled {
|
||||||
|
registryMgr := skills.NewRegistryManagerFromConfig(skills.RegistryConfig{
|
||||||
|
MaxConcurrentSearches: cfg.Tools.Skills.MaxConcurrentSearches,
|
||||||
|
ClawHub: skills.ClawHubConfig(cfg.Tools.Skills.Registries.ClawHub),
|
||||||
|
})
|
||||||
|
searchCache := skills.NewSearchCache(
|
||||||
|
cfg.Tools.Skills.SearchCache.MaxSize,
|
||||||
|
time.Duration(cfg.Tools.Skills.SearchCache.TTLSeconds)*time.Second,
|
||||||
|
)
|
||||||
|
registry.Register(NewFindSkillsTool(registryMgr, searchCache))
|
||||||
|
registry.Register(NewInstallSkillTool(registryMgr, agentCtx.Workspace))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Spawn tool
|
||||||
|
if cfg.Tools.Core.EnableSpawn {
|
||||||
|
subagentManager := NewSubagentManager(provider, agentCtx.Model, agentCtx.Workspace, msgBus)
|
||||||
|
subagentManager.SetLLMOptions(agentCtx.MaxTokens, agentCtx.Temperature)
|
||||||
|
spawnTool := NewSpawnTool(subagentManager)
|
||||||
|
spawnTool.SetAllowlistChecker(canSpawn)
|
||||||
|
registry.Register(spawnTool)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetupWorkspaceTools registers tools related to file system and execution
|
||||||
|
// centralizing the logic and decoupling it from the agent.
|
||||||
|
func SetupWorkspaceTools(registry *ToolRegistry, cfg *config.Config, workspace string, restrict bool) {
|
||||||
|
if cfg.Tools.Filesystem.EnableRead {
|
||||||
|
registry.Register(NewReadFileTool(workspace, restrict))
|
||||||
|
}
|
||||||
|
if cfg.Tools.Filesystem.EnableWrite {
|
||||||
|
registry.Register(NewWriteFileTool(workspace, restrict))
|
||||||
|
}
|
||||||
|
if cfg.Tools.Filesystem.EnableList {
|
||||||
|
registry.Register(NewListDirTool(workspace, restrict))
|
||||||
|
}
|
||||||
|
if cfg.Tools.Filesystem.EnableEdit {
|
||||||
|
registry.Register(NewEditFileTool(workspace, restrict))
|
||||||
|
}
|
||||||
|
if cfg.Tools.Filesystem.EnableAppend {
|
||||||
|
registry.Register(NewAppendFileTool(workspace, restrict))
|
||||||
|
}
|
||||||
|
if cfg.Tools.Exec.Enabled {
|
||||||
|
registry.Register(NewExecToolWithConfig(workspace, restrict, cfg))
|
||||||
|
}
|
||||||
|
}
|
||||||
54
pkg/tools/list_dir.go
Normal file
54
pkg/tools/list_dir.go
Normal file
|
|
@ -0,0 +1,54 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ListDirTool struct {
|
||||||
|
fs fileSystem
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewListDirTool(workspace string, restrict bool) *ListDirTool {
|
||||||
|
var fs fileSystem
|
||||||
|
if restrict {
|
||||||
|
fs = &sandboxFs{workspace: workspace}
|
||||||
|
} else {
|
||||||
|
fs = &hostFs{}
|
||||||
|
}
|
||||||
|
return &ListDirTool{fs: fs}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ListDirTool) Name() string {
|
||||||
|
return "list_dir"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ListDirTool) Description() string {
|
||||||
|
return "List files and directories in a path"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ListDirTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"path": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Path to list",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"path"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ListDirTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
path, ok := args["path"].(string)
|
||||||
|
if !ok {
|
||||||
|
path = "."
|
||||||
|
}
|
||||||
|
|
||||||
|
entries, err := t.fs.ReadDir(path)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to read directory: %v", err))
|
||||||
|
}
|
||||||
|
return formatDirEntries(entries)
|
||||||
|
}
|
||||||
53
pkg/tools/read_file.go
Normal file
53
pkg/tools/read_file.go
Normal file
|
|
@ -0,0 +1,53 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ReadFileTool struct {
|
||||||
|
fs fileSystem
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewReadFileTool(workspace string, restrict bool) *ReadFileTool {
|
||||||
|
var fs fileSystem
|
||||||
|
if restrict {
|
||||||
|
fs = &sandboxFs{workspace: workspace}
|
||||||
|
} else {
|
||||||
|
fs = &hostFs{}
|
||||||
|
}
|
||||||
|
return &ReadFileTool{fs: fs}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ReadFileTool) Name() string {
|
||||||
|
return "read_file"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ReadFileTool) Description() string {
|
||||||
|
return "Read the contents of a file"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ReadFileTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"path": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Path to the file to read",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"path"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *ReadFileTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
path, ok := args["path"].(string)
|
||||||
|
if !ok {
|
||||||
|
return ErrorResult("path is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
content, err := t.fs.ReadFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(err.Error())
|
||||||
|
}
|
||||||
|
return NewToolResult(string(content))
|
||||||
|
}
|
||||||
190
pkg/tools/web_fetch.go
Normal file
190
pkg/tools/web_fetch.go
Normal file
|
|
@ -0,0 +1,190 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type WebFetchTool struct {
|
||||||
|
maxChars int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewWebFetchTool(maxChars int) *WebFetchTool {
|
||||||
|
if maxChars <= 0 {
|
||||||
|
maxChars = 50000
|
||||||
|
}
|
||||||
|
return &WebFetchTool{
|
||||||
|
maxChars: maxChars,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WebFetchTool) Name() string {
|
||||||
|
return "web_fetch"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WebFetchTool) Description() string {
|
||||||
|
return "Fetch a URL and extract readable content (HTML to text). Use this to get weather info, news, articles, or any web content."
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WebFetchTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"url": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "URL to fetch",
|
||||||
|
},
|
||||||
|
"maxChars": map[string]any{
|
||||||
|
"type": "integer",
|
||||||
|
"description": "Maximum characters to extract",
|
||||||
|
"minimum": 100.0,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"url"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WebFetchTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
urlStr, ok := args["url"].(string)
|
||||||
|
if !ok {
|
||||||
|
return ErrorResult("url is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
parsedURL, err := url.Parse(urlStr)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("invalid URL: %v", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
if parsedURL.Scheme != "http" && parsedURL.Scheme != "https" {
|
||||||
|
return ErrorResult("only http/https URLs are allowed")
|
||||||
|
}
|
||||||
|
|
||||||
|
if parsedURL.Host == "" {
|
||||||
|
return ErrorResult("missing domain in URL")
|
||||||
|
}
|
||||||
|
|
||||||
|
maxChars := t.maxChars
|
||||||
|
if mc, ok := args["maxChars"].(float64); ok {
|
||||||
|
if int(mc) > 100 {
|
||||||
|
maxChars = int(mc)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "GET", urlStr, nil)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to create request: %v", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("User-Agent", userAgent)
|
||||||
|
|
||||||
|
client := &http.Client{
|
||||||
|
Timeout: 60 * time.Second,
|
||||||
|
Transport: &http.Transport{
|
||||||
|
MaxIdleConns: 10,
|
||||||
|
IdleConnTimeout: 30 * time.Second,
|
||||||
|
DisableCompression: false,
|
||||||
|
TLSHandshakeTimeout: 15 * time.Second,
|
||||||
|
},
|
||||||
|
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||||
|
if len(via) >= 5 {
|
||||||
|
return fmt.Errorf("stopped after 5 redirects")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("request failed: %v", err))
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
return ErrorResult(fmt.Sprintf("failed to read response: %v", err))
|
||||||
|
}
|
||||||
|
|
||||||
|
contentType := resp.Header.Get("Content-Type")
|
||||||
|
|
||||||
|
var text, extractor string
|
||||||
|
|
||||||
|
if strings.Contains(contentType, "application/json") {
|
||||||
|
var jsonData any
|
||||||
|
if err := json.Unmarshal(body, &jsonData); err == nil {
|
||||||
|
formatted, _ := json.MarshalIndent(jsonData, "", " ")
|
||||||
|
text = string(formatted)
|
||||||
|
extractor = "json"
|
||||||
|
} else {
|
||||||
|
text = string(body)
|
||||||
|
extractor = "raw"
|
||||||
|
}
|
||||||
|
} else if strings.Contains(contentType, "text/html") || len(body) > 0 &&
|
||||||
|
(strings.HasPrefix(string(body), "<!DOCTYPE") || strings.HasPrefix(strings.ToLower(string(body)), "<html")) {
|
||||||
|
text = t.extractText(string(body))
|
||||||
|
extractor = "text"
|
||||||
|
} else {
|
||||||
|
text = string(body)
|
||||||
|
extractor = "raw"
|
||||||
|
}
|
||||||
|
|
||||||
|
truncated := len(text) > maxChars
|
||||||
|
if truncated {
|
||||||
|
text = text[:maxChars]
|
||||||
|
}
|
||||||
|
|
||||||
|
result := map[string]any{
|
||||||
|
"url": urlStr,
|
||||||
|
"status": resp.StatusCode,
|
||||||
|
"extractor": extractor,
|
||||||
|
"truncated": truncated,
|
||||||
|
"length": len(text),
|
||||||
|
"text": text,
|
||||||
|
}
|
||||||
|
|
||||||
|
resultJSON, _ := json.MarshalIndent(result, "", " ")
|
||||||
|
|
||||||
|
return &ToolResult{
|
||||||
|
ForLLM: fmt.Sprintf(
|
||||||
|
"Fetched %d bytes from %s (extractor: %s, truncated: %v)",
|
||||||
|
len(text),
|
||||||
|
urlStr,
|
||||||
|
extractor,
|
||||||
|
truncated,
|
||||||
|
),
|
||||||
|
ForUser: string(resultJSON),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WebFetchTool) extractText(htmlContent string) string {
|
||||||
|
re := regexp.MustCompile(`<script[\s\S]*?</script>`)
|
||||||
|
result := re.ReplaceAllLiteralString(htmlContent, "")
|
||||||
|
re = regexp.MustCompile(`<style[\s\S]*?</style>`)
|
||||||
|
result = re.ReplaceAllLiteralString(result, "")
|
||||||
|
re = regexp.MustCompile(`<[^>]+>`)
|
||||||
|
result = re.ReplaceAllLiteralString(result, "")
|
||||||
|
|
||||||
|
result = strings.TrimSpace(result)
|
||||||
|
|
||||||
|
re = regexp.MustCompile(`[^\S\n]+`)
|
||||||
|
result = re.ReplaceAllString(result, " ")
|
||||||
|
re = regexp.MustCompile(`\n{3,}`)
|
||||||
|
result = re.ReplaceAllString(result, "\n\n")
|
||||||
|
|
||||||
|
lines := strings.Split(result, "\n")
|
||||||
|
var cleanLines []string
|
||||||
|
for _, line := range lines {
|
||||||
|
line = strings.TrimSpace(line)
|
||||||
|
if line != "" {
|
||||||
|
cleanLines = append(cleanLines, line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.Join(cleanLines, "\n")
|
||||||
|
}
|
||||||
|
|
@ -173,34 +173,6 @@ func TestWebTool_WebFetch_Truncation(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing
|
|
||||||
func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
|
|
||||||
tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: ""})
|
|
||||||
if tool != nil {
|
|
||||||
t.Errorf("Expected nil tool when Brave API key is empty")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Also nil when nothing is enabled
|
|
||||||
tool = NewWebSearchTool(WebSearchToolOptions{})
|
|
||||||
if tool != nil {
|
|
||||||
t.Errorf("Expected nil tool when no provider is enabled")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWebTool_WebSearch_MissingQuery verifies error handling for missing query
|
|
||||||
func TestWebTool_WebSearch_MissingQuery(t *testing.T) {
|
|
||||||
tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: "test-key", BraveMaxResults: 5})
|
|
||||||
ctx := context.Background()
|
|
||||||
args := map[string]any{}
|
|
||||||
|
|
||||||
result := tool.Execute(ctx, args)
|
|
||||||
|
|
||||||
// Should return error result
|
|
||||||
if !result.IsError {
|
|
||||||
t.Errorf("Expected error when query is missing")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWebTool_WebFetch_HTMLExtraction verifies HTML text extraction
|
// TestWebTool_WebFetch_HTMLExtraction verifies HTML text extraction
|
||||||
func TestWebTool_WebFetch_HTMLExtraction(t *testing.T) {
|
func TestWebTool_WebFetch_HTMLExtraction(t *testing.T) {
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
|
@ -333,75 +305,3 @@ func TestWebTool_WebFetch_MissingDomain(t *testing.T) {
|
||||||
t.Errorf("Expected domain error message, got ForLLM: %s", result.ForLLM)
|
t.Errorf("Expected domain error message, got ForLLM: %s", result.ForLLM)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestWebTool_TavilySearch_Success verifies successful Tavily search
|
|
||||||
func TestWebTool_TavilySearch_Success(t *testing.T) {
|
|
||||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
||||||
if r.Method != "POST" {
|
|
||||||
t.Errorf("Expected POST request, got %s", r.Method)
|
|
||||||
}
|
|
||||||
if r.Header.Get("Content-Type") != "application/json" {
|
|
||||||
t.Errorf("Expected Content-Type application/json, got %s", r.Header.Get("Content-Type"))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify payload
|
|
||||||
var payload map[string]any
|
|
||||||
json.NewDecoder(r.Body).Decode(&payload)
|
|
||||||
if payload["api_key"] != "test-key" {
|
|
||||||
t.Errorf("Expected api_key test-key, got %v", payload["api_key"])
|
|
||||||
}
|
|
||||||
if payload["query"] != "test query" {
|
|
||||||
t.Errorf("Expected query 'test query', got %v", payload["query"])
|
|
||||||
}
|
|
||||||
|
|
||||||
// Return mock response
|
|
||||||
response := map[string]any{
|
|
||||||
"results": []map[string]any{
|
|
||||||
{
|
|
||||||
"title": "Test Result 1",
|
|
||||||
"url": "https://example.com/1",
|
|
||||||
"content": "Content for result 1",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"title": "Test Result 2",
|
|
||||||
"url": "https://example.com/2",
|
|
||||||
"content": "Content for result 2",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
w.Header().Set("Content-Type", "application/json")
|
|
||||||
w.WriteHeader(http.StatusOK)
|
|
||||||
json.NewEncoder(w).Encode(response)
|
|
||||||
}))
|
|
||||||
defer server.Close()
|
|
||||||
|
|
||||||
tool := NewWebSearchTool(WebSearchToolOptions{
|
|
||||||
TavilyEnabled: true,
|
|
||||||
TavilyAPIKey: "test-key",
|
|
||||||
TavilyBaseURL: server.URL,
|
|
||||||
TavilyMaxResults: 5,
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx := context.Background()
|
|
||||||
args := map[string]any{
|
|
||||||
"query": "test query",
|
|
||||||
}
|
|
||||||
|
|
||||||
result := tool.Execute(ctx, args)
|
|
||||||
|
|
||||||
// Success should not be an error
|
|
||||||
if result.IsError {
|
|
||||||
t.Errorf("Expected success, got IsError=true: %s", result.ForLLM)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ForUser should contain result titles and URLs
|
|
||||||
if !strings.Contains(result.ForUser, "Test Result 1") ||
|
|
||||||
!strings.Contains(result.ForUser, "https://example.com/1") {
|
|
||||||
t.Errorf("Expected results in output, got: %s", result.ForUser)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Should mention via Tavily
|
|
||||||
if !strings.Contains(result.ForUser, "via Tavily") {
|
|
||||||
t.Errorf("Expected 'via Tavily' in output, got: %s", result.ForUser)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -438,180 +438,3 @@ func (t *WebSearchTool) Execute(ctx context.Context, args map[string]any) *ToolR
|
||||||
ForUser: result,
|
ForUser: result,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
type WebFetchTool struct {
|
|
||||||
maxChars int
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewWebFetchTool(maxChars int) *WebFetchTool {
|
|
||||||
if maxChars <= 0 {
|
|
||||||
maxChars = 50000
|
|
||||||
}
|
|
||||||
return &WebFetchTool{
|
|
||||||
maxChars: maxChars,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WebFetchTool) Name() string {
|
|
||||||
return "web_fetch"
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WebFetchTool) Description() string {
|
|
||||||
return "Fetch a URL and extract readable content (HTML to text). Use this to get weather info, news, articles, or any web content."
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WebFetchTool) Parameters() map[string]any {
|
|
||||||
return map[string]any{
|
|
||||||
"type": "object",
|
|
||||||
"properties": map[string]any{
|
|
||||||
"url": map[string]any{
|
|
||||||
"type": "string",
|
|
||||||
"description": "URL to fetch",
|
|
||||||
},
|
|
||||||
"maxChars": map[string]any{
|
|
||||||
"type": "integer",
|
|
||||||
"description": "Maximum characters to extract",
|
|
||||||
"minimum": 100.0,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"required": []string{"url"},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WebFetchTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
|
||||||
urlStr, ok := args["url"].(string)
|
|
||||||
if !ok {
|
|
||||||
return ErrorResult("url is required")
|
|
||||||
}
|
|
||||||
|
|
||||||
parsedURL, err := url.Parse(urlStr)
|
|
||||||
if err != nil {
|
|
||||||
return ErrorResult(fmt.Sprintf("invalid URL: %v", err))
|
|
||||||
}
|
|
||||||
|
|
||||||
if parsedURL.Scheme != "http" && parsedURL.Scheme != "https" {
|
|
||||||
return ErrorResult("only http/https URLs are allowed")
|
|
||||||
}
|
|
||||||
|
|
||||||
if parsedURL.Host == "" {
|
|
||||||
return ErrorResult("missing domain in URL")
|
|
||||||
}
|
|
||||||
|
|
||||||
maxChars := t.maxChars
|
|
||||||
if mc, ok := args["maxChars"].(float64); ok {
|
|
||||||
if int(mc) > 100 {
|
|
||||||
maxChars = int(mc)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, "GET", urlStr, nil)
|
|
||||||
if err != nil {
|
|
||||||
return ErrorResult(fmt.Sprintf("failed to create request: %v", err))
|
|
||||||
}
|
|
||||||
|
|
||||||
req.Header.Set("User-Agent", userAgent)
|
|
||||||
|
|
||||||
client := &http.Client{
|
|
||||||
Timeout: 60 * time.Second,
|
|
||||||
Transport: &http.Transport{
|
|
||||||
MaxIdleConns: 10,
|
|
||||||
IdleConnTimeout: 30 * time.Second,
|
|
||||||
DisableCompression: false,
|
|
||||||
TLSHandshakeTimeout: 15 * time.Second,
|
|
||||||
},
|
|
||||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
|
||||||
if len(via) >= 5 {
|
|
||||||
return fmt.Errorf("stopped after 5 redirects")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
resp, err := client.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
return ErrorResult(fmt.Sprintf("request failed: %v", err))
|
|
||||||
}
|
|
||||||
defer resp.Body.Close()
|
|
||||||
|
|
||||||
body, err := io.ReadAll(resp.Body)
|
|
||||||
if err != nil {
|
|
||||||
return ErrorResult(fmt.Sprintf("failed to read response: %v", err))
|
|
||||||
}
|
|
||||||
|
|
||||||
contentType := resp.Header.Get("Content-Type")
|
|
||||||
|
|
||||||
var text, extractor string
|
|
||||||
|
|
||||||
if strings.Contains(contentType, "application/json") {
|
|
||||||
var jsonData any
|
|
||||||
if err := json.Unmarshal(body, &jsonData); err == nil {
|
|
||||||
formatted, _ := json.MarshalIndent(jsonData, "", " ")
|
|
||||||
text = string(formatted)
|
|
||||||
extractor = "json"
|
|
||||||
} else {
|
|
||||||
text = string(body)
|
|
||||||
extractor = "raw"
|
|
||||||
}
|
|
||||||
} else if strings.Contains(contentType, "text/html") || len(body) > 0 &&
|
|
||||||
(strings.HasPrefix(string(body), "<!DOCTYPE") || strings.HasPrefix(strings.ToLower(string(body)), "<html")) {
|
|
||||||
text = t.extractText(string(body))
|
|
||||||
extractor = "text"
|
|
||||||
} else {
|
|
||||||
text = string(body)
|
|
||||||
extractor = "raw"
|
|
||||||
}
|
|
||||||
|
|
||||||
truncated := len(text) > maxChars
|
|
||||||
if truncated {
|
|
||||||
text = text[:maxChars]
|
|
||||||
}
|
|
||||||
|
|
||||||
result := map[string]any{
|
|
||||||
"url": urlStr,
|
|
||||||
"status": resp.StatusCode,
|
|
||||||
"extractor": extractor,
|
|
||||||
"truncated": truncated,
|
|
||||||
"length": len(text),
|
|
||||||
"text": text,
|
|
||||||
}
|
|
||||||
|
|
||||||
resultJSON, _ := json.MarshalIndent(result, "", " ")
|
|
||||||
|
|
||||||
return &ToolResult{
|
|
||||||
ForLLM: fmt.Sprintf(
|
|
||||||
"Fetched %d bytes from %s (extractor: %s, truncated: %v)",
|
|
||||||
len(text),
|
|
||||||
urlStr,
|
|
||||||
extractor,
|
|
||||||
truncated,
|
|
||||||
),
|
|
||||||
ForUser: string(resultJSON),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *WebFetchTool) extractText(htmlContent string) string {
|
|
||||||
re := regexp.MustCompile(`<script[\s\S]*?</script>`)
|
|
||||||
result := re.ReplaceAllLiteralString(htmlContent, "")
|
|
||||||
re = regexp.MustCompile(`<style[\s\S]*?</style>`)
|
|
||||||
result = re.ReplaceAllLiteralString(result, "")
|
|
||||||
re = regexp.MustCompile(`<[^>]+>`)
|
|
||||||
result = re.ReplaceAllLiteralString(result, "")
|
|
||||||
|
|
||||||
result = strings.TrimSpace(result)
|
|
||||||
|
|
||||||
re = regexp.MustCompile(`[^\S\n]+`)
|
|
||||||
result = re.ReplaceAllString(result, " ")
|
|
||||||
re = regexp.MustCompile(`\n{3,}`)
|
|
||||||
result = re.ReplaceAllString(result, "\n\n")
|
|
||||||
|
|
||||||
lines := strings.Split(result, "\n")
|
|
||||||
var cleanLines []string
|
|
||||||
for _, line := range lines {
|
|
||||||
line = strings.TrimSpace(line)
|
|
||||||
if line != "" {
|
|
||||||
cleanLines = append(cleanLines, line)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return strings.Join(cleanLines, "\n")
|
|
||||||
}
|
|
||||||
110
pkg/tools/web_search_test.go
Normal file
110
pkg/tools/web_search_test.go
Normal file
|
|
@ -0,0 +1,110 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing
|
||||||
|
func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
|
||||||
|
tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: ""})
|
||||||
|
if tool != nil {
|
||||||
|
t.Errorf("Expected nil tool when Brave API key is empty")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Also nil when nothing is enabled
|
||||||
|
tool = NewWebSearchTool(WebSearchToolOptions{})
|
||||||
|
if tool != nil {
|
||||||
|
t.Errorf("Expected nil tool when no provider is enabled")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebTool_WebSearch_MissingQuery verifies error handling for missing query
|
||||||
|
func TestWebTool_WebSearch_MissingQuery(t *testing.T) {
|
||||||
|
tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: "test-key", BraveMaxResults: 5})
|
||||||
|
ctx := context.Background()
|
||||||
|
args := map[string]any{}
|
||||||
|
|
||||||
|
result := tool.Execute(ctx, args)
|
||||||
|
|
||||||
|
// Should return error result
|
||||||
|
if !result.IsError {
|
||||||
|
t.Errorf("Expected error when query is missing")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebTool_TavilySearch_Success verifies successful Tavily search
|
||||||
|
func TestWebTool_TavilySearch_Success(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Method != "POST" {
|
||||||
|
t.Errorf("Expected POST request, got %s", r.Method)
|
||||||
|
}
|
||||||
|
if r.Header.Get("Content-Type") != "application/json" {
|
||||||
|
t.Errorf("Expected Content-Type application/json, got %s", r.Header.Get("Content-Type"))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify payload
|
||||||
|
var payload map[string]any
|
||||||
|
json.NewDecoder(r.Body).Decode(&payload)
|
||||||
|
if payload["api_key"] != "test-key" {
|
||||||
|
t.Errorf("Expected api_key test-key, got %v", payload["api_key"])
|
||||||
|
}
|
||||||
|
if payload["query"] != "test query" {
|
||||||
|
t.Errorf("Expected query 'test query', got %v", payload["query"])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Return mock response
|
||||||
|
response := map[string]any{
|
||||||
|
"results": []map[string]any{
|
||||||
|
{
|
||||||
|
"title": "Test Result 1",
|
||||||
|
"url": "https://example.com/1",
|
||||||
|
"content": "Content for result 1",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"title": "Test Result 2",
|
||||||
|
"url": "https://example.com/2",
|
||||||
|
"content": "Content for result 2",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
json.NewEncoder(w).Encode(response)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
tool := NewWebSearchTool(WebSearchToolOptions{
|
||||||
|
TavilyEnabled: true,
|
||||||
|
TavilyAPIKey: "test-key",
|
||||||
|
TavilyBaseURL: server.URL,
|
||||||
|
TavilyMaxResults: 5,
|
||||||
|
})
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
args := map[string]any{
|
||||||
|
"query": "test query",
|
||||||
|
}
|
||||||
|
|
||||||
|
result := tool.Execute(ctx, args)
|
||||||
|
|
||||||
|
// Success should not be an error
|
||||||
|
if result.IsError {
|
||||||
|
t.Errorf("Expected success, got IsError=true: %s", result.ForLLM)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ForUser should contain result titles and URLs
|
||||||
|
if !strings.Contains(result.ForUser, "Test Result 1") ||
|
||||||
|
!strings.Contains(result.ForUser, "https://example.com/1") {
|
||||||
|
t.Errorf("Expected results in output, got: %s", result.ForUser)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should mention via Tavily
|
||||||
|
if !strings.Contains(result.ForUser, "via Tavily") {
|
||||||
|
t.Errorf("Expected 'via Tavily' in output, got: %s", result.ForUser)
|
||||||
|
}
|
||||||
|
}
|
||||||
63
pkg/tools/write_file.go
Normal file
63
pkg/tools/write_file.go
Normal file
|
|
@ -0,0 +1,63 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
type WriteFileTool struct {
|
||||||
|
fs fileSystem
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewWriteFileTool(workspace string, restrict bool) *WriteFileTool {
|
||||||
|
var fs fileSystem
|
||||||
|
if restrict {
|
||||||
|
fs = &sandboxFs{workspace: workspace}
|
||||||
|
} else {
|
||||||
|
fs = &hostFs{}
|
||||||
|
}
|
||||||
|
return &WriteFileTool{fs: fs}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WriteFileTool) Name() string {
|
||||||
|
return "write_file"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WriteFileTool) Description() string {
|
||||||
|
return "Write content to a file"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WriteFileTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"path": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Path to the file to write",
|
||||||
|
},
|
||||||
|
"content": map[string]any{
|
||||||
|
"type": "string",
|
||||||
|
"description": "Content to write to the file",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"path", "content"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *WriteFileTool) Execute(ctx context.Context, args map[string]any) *ToolResult {
|
||||||
|
path, ok := args["path"].(string)
|
||||||
|
if !ok {
|
||||||
|
return ErrorResult("path is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
content, ok := args["content"].(string)
|
||||||
|
if !ok {
|
||||||
|
return ErrorResult("content is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := t.fs.WriteFile(path, []byte(content)); err != nil {
|
||||||
|
return ErrorResult(err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
return SilentResult(fmt.Sprintf("File written: %s", path))
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue