Merge remote-tracking branch 'upstream/main' into feat/agent-sandbox

This commit is contained in:
0x5487 2026-03-01 14:29:35 +08:00
commit 4b8dd643bf
15 changed files with 308 additions and 155 deletions

View file

@ -0,0 +1,25 @@
package onboard
import (
"os"
"path/filepath"
"testing"
)
func TestCopyEmbeddedToTargetUsesAgentsMarkdown(t *testing.T) {
targetDir := t.TempDir()
if err := copyEmbeddedToTarget(targetDir); err != nil {
t.Fatalf("copyEmbeddedToTarget() error = %v", err)
}
agentsPath := filepath.Join(targetDir, "AGENTS.md")
if _, err := os.Stat(agentsPath); err != nil {
t.Fatalf("expected %s to exist: %v", agentsPath, err)
}
legacyPath := filepath.Join(targetDir, "AGENT.md")
if _, err := os.Stat(legacyPath); !os.IsNotExist(err) {
t.Fatalf("expected legacy file %s to be absent, got err=%v", legacyPath, err)
}
}

View file

@ -71,7 +71,7 @@ func NewSkillsCommand() *cobra.Command {
newInstallBuiltinCommand(workspaceFn), newInstallBuiltinCommand(workspaceFn),
newListBuiltinCommand(), newListBuiltinCommand(),
newRemoveCommand(installerFn), newRemoveCommand(installerFn),
newSearchCommand(installerFn), newSearchCommand(),
newShowCommand(loaderFn), newShowCommand(loaderFn),
) )

View file

@ -15,6 +15,8 @@ import (
"github.com/sipeed/picoclaw/pkg/utils" "github.com/sipeed/picoclaw/pkg/utils"
) )
const skillsSearchMaxResults = 20
func skillsListCmd(loader *skills.SkillsLoader) { func skillsListCmd(loader *skills.SkillsLoader) {
allSkills := loader.ListSkills() allSkills := loader.ListSkills()
@ -215,34 +217,43 @@ func skillsListBuiltinCmd() {
} }
} }
func skillsSearchCmd(installer *skills.SkillInstaller) { func skillsSearchCmd(query string) {
fmt.Println("Searching for available skills...") fmt.Println("Searching for available skills...")
cfg, err := internal.LoadConfig()
if err != nil {
fmt.Printf("✗ Failed to load config: %v\n", err)
return
}
registryMgr := skills.NewRegistryManagerFromConfig(skills.RegistryConfig{
MaxConcurrentSearches: cfg.Tools.Skills.MaxConcurrentSearches,
ClawHub: skills.ClawHubConfig(cfg.Tools.Skills.Registries.ClawHub),
})
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel() defer cancel()
availableSkills, err := installer.ListAvailableSkills(ctx) results, err := registryMgr.SearchAll(ctx, query, skillsSearchMaxResults)
if err != nil { if err != nil {
fmt.Printf("✗ Failed to fetch skills list: %v\n", err) fmt.Printf("✗ Failed to fetch skills list: %v\n", err)
return return
} }
if len(availableSkills) == 0 { if len(results) == 0 {
fmt.Println("No skills available.") fmt.Println("No skills available.")
return return
} }
fmt.Printf("\nAvailable Skills (%d):\n", len(availableSkills)) fmt.Printf("\nAvailable Skills (%d):\n", len(results))
fmt.Println("--------------------") fmt.Println("--------------------")
for _, skill := range availableSkills { for _, result := range results {
fmt.Printf(" 📦 %s\n", skill.Name) fmt.Printf(" 📦 %s\n", result.DisplayName)
fmt.Printf(" %s\n", skill.Description) fmt.Printf(" %s\n", result.Summary)
fmt.Printf(" Repo: %s\n", skill.Repository) fmt.Printf(" Slug: %s\n", result.Slug)
if skill.Author != "" { fmt.Printf(" Registry: %s\n", result.RegistryName)
fmt.Printf(" Author: %s\n", skill.Author) if result.Version != "" {
} fmt.Printf(" Version: %s\n", result.Version)
if len(skill.Tags) > 0 {
fmt.Printf(" Tags: %v\n", skill.Tags)
} }
fmt.Println() fmt.Println()
} }

View file

@ -2,20 +2,19 @@ package skills
import ( import (
"github.com/spf13/cobra" "github.com/spf13/cobra"
"github.com/sipeed/picoclaw/pkg/skills"
) )
func newSearchCommand(installerFn func() (*skills.SkillInstaller, error)) *cobra.Command { func newSearchCommand() *cobra.Command {
cmd := &cobra.Command{ cmd := &cobra.Command{
Use: "search", Use: "search [query]",
Short: "Search available skills", Short: "Search available skills",
RunE: func(_ *cobra.Command, _ []string) error { Args: cobra.MaximumNArgs(1),
installer, err := installerFn() RunE: func(_ *cobra.Command, args []string) error {
if err != nil { query := ""
return err if len(args) == 1 {
query = args[0]
} }
skillsSearchCmd(installer) skillsSearchCmd(query)
return nil return nil
}, },
} }

View file

@ -8,11 +8,11 @@ import (
) )
func TestNewSearchSubcommand(t *testing.T) { func TestNewSearchSubcommand(t *testing.T) {
cmd := newSearchCommand(nil) cmd := newSearchCommand()
require.NotNil(t, cmd) require.NotNil(t, cmd)
assert.Equal(t, "search", cmd.Use) assert.Equal(t, "search [query]", cmd.Use)
assert.Equal(t, "Search available skills", cmd.Short) assert.Equal(t, "Search available skills", cmd.Short)
assert.Nil(t, cmd.Run) assert.Nil(t, cmd.Run)

View file

@ -100,7 +100,7 @@ func registerSharedTools(
} }
// Web tools // Web tools
if searchTool := tools.NewWebSearchTool(tools.WebSearchToolOptions{ searchTool, err := tools.NewWebSearchTool(tools.WebSearchToolOptions{
BraveAPIKey: cfg.Tools.Web.Brave.APIKey, BraveAPIKey: cfg.Tools.Web.Brave.APIKey,
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults, BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
BraveEnabled: cfg.Tools.Web.Brave.Enabled, BraveEnabled: cfg.Tools.Web.Brave.Enabled,
@ -114,10 +114,18 @@ func registerSharedTools(
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults, PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled, PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
Proxy: cfg.Tools.Web.Proxy, Proxy: cfg.Tools.Web.Proxy,
}); searchTool != nil { })
if err != nil {
logger.ErrorCF("agent", "Failed to create web search tool", map[string]any{"error": err.Error()})
} else if searchTool != nil {
agent.Tools.Register(searchTool) agent.Tools.Register(searchTool)
} }
agent.Tools.Register(tools.NewWebFetchToolWithProxy(50000, cfg.Tools.Web.Proxy)) fetchTool, err := tools.NewWebFetchToolWithProxy(50000, cfg.Tools.Web.Proxy)
if err != nil {
logger.ErrorCF("agent", "Failed to create web fetch tool", map[string]any{"error": err.Error()})
} else {
agent.Tools.Register(fetchTool)
}
// Hardware tools (I2C, SPI) - Linux only, returns error on other platforms // Hardware tools (I2C, SPI) - Linux only, returns error on other platforms
agent.Tools.Register(tools.NewI2CTool()) agent.Tools.Register(tools.NewI2CTool())

View file

@ -45,11 +45,13 @@ type replyTokenEntry struct {
type LINEChannel struct { type LINEChannel struct {
*channels.BaseChannel *channels.BaseChannel
config config.LINEConfig config config.LINEConfig
botUserID string // Bot's user ID infoClient *http.Client // for bot info lookups (short timeout)
botBasicID string // Bot's basic ID (e.g. @216ru...) apiClient *http.Client // for messaging API calls
botDisplayName string // Bot's display name for text-based mention detection botUserID string // Bot's user ID
replyTokens sync.Map // chatID -> replyTokenEntry botBasicID string // Bot's basic ID (e.g. @216ru...)
quoteTokens sync.Map // chatID -> quoteToken (string) botDisplayName string // Bot's display name for text-based mention detection
replyTokens sync.Map // chatID -> replyTokenEntry
quoteTokens sync.Map // chatID -> quoteToken (string)
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
} }
@ -69,6 +71,8 @@ func NewLINEChannel(cfg config.LINEConfig, messageBus *bus.MessageBus) (*LINECha
return &LINEChannel{ return &LINEChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
infoClient: &http.Client{Timeout: 10 * time.Second},
apiClient: &http.Client{Timeout: 30 * time.Second},
}, nil }, nil
} }
@ -104,8 +108,7 @@ func (c *LINEChannel) fetchBotInfo() error {
} }
req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken) req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken)
client := &http.Client{Timeout: 10 * time.Second} resp, err := c.infoClient.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return err return err
} }
@ -644,8 +647,7 @@ func (c *LINEChannel) callAPI(ctx context.Context, endpoint string, payload any)
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken) req.Header.Set("Authorization", "Bearer "+c.config.ChannelAccessToken)
client := &http.Client{Timeout: 30 * time.Second} resp, err := c.apiClient.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return channels.ClassifyNetError(err) return channels.ClassifyNetError(err)
} }

View file

@ -32,6 +32,7 @@ const (
type WeComAppChannel struct { type WeComAppChannel struct {
*channels.BaseChannel *channels.BaseChannel
config config.WeComAppConfig config config.WeComAppConfig
client *http.Client
accessToken string accessToken string
tokenExpiry time.Time tokenExpiry time.Time
tokenMu sync.RWMutex tokenMu sync.RWMutex
@ -129,10 +130,18 @@ func NewWeComAppChannel(cfg config.WeComAppConfig, messageBus *bus.MessageBus) (
channels.WithReasoningChannelID(cfg.ReasoningChannelID), channels.WithReasoningChannelID(cfg.ReasoningChannelID),
) )
// Client timeout must be >= the configured ReplyTimeout so the
// per-request context deadline is always the effective limit.
clientTimeout := 30 * time.Second
if d := time.Duration(cfg.ReplyTimeout) * time.Second; d > clientTimeout {
clientTimeout = d
}
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
return &WeComAppChannel{ return &WeComAppChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
client: &http.Client{Timeout: clientTimeout},
ctx: ctx, ctx: ctx,
cancel: cancel, cancel: cancel,
processedMsgs: make(map[string]bool), processedMsgs: make(map[string]bool),
@ -148,6 +157,10 @@ func (c *WeComAppChannel) Name() string {
func (c *WeComAppChannel) Start(ctx context.Context) error { func (c *WeComAppChannel) Start(ctx context.Context) error {
logger.InfoC("wecom_app", "Starting WeCom App channel...") logger.InfoC("wecom_app", "Starting WeCom App channel...")
// Cancel the context created in the constructor to avoid a resource leak.
if c.cancel != nil {
c.cancel()
}
c.ctx, c.cancel = context.WithCancel(ctx) c.ctx, c.cancel = context.WithCancel(ctx)
// Get initial access token // Get initial access token
@ -302,8 +315,7 @@ func (c *WeComAppChannel) uploadMedia(ctx context.Context, accessToken, mediaTyp
} }
req.Header.Set("Content-Type", writer.FormDataContentType()) req.Header.Set("Content-Type", writer.FormDataContentType())
client := &http.Client{Timeout: 30 * time.Second} resp, err := c.client.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return "", channels.ClassifyNetError(err) return "", channels.ClassifyNetError(err)
} }
@ -360,8 +372,7 @@ func (c *WeComAppChannel) sendImageMessage(ctx context.Context, accessToken, use
} }
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: time.Duration(timeout) * time.Second} resp, err := c.client.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return channels.ClassifyNetError(err) return channels.ClassifyNetError(err)
} }
@ -601,14 +612,14 @@ func (c *WeComAppChannel) processMessage(ctx context.Context, msg WeComXMLMessag
return return
} }
c.processedMsgs[msgID] = true c.processedMsgs[msgID] = true
c.msgMu.Unlock() // Clean up old messages while still holding the lock to avoid a data race
// on len(). Reset the map but re-insert the current msgID so it remains
// Clean up old messages periodically (keep last 1000) // deduplicated.
if len(c.processedMsgs) > 1000 { if len(c.processedMsgs) > 1000 {
c.msgMu.Lock()
c.processedMsgs = make(map[string]bool) c.processedMsgs = make(map[string]bool)
c.msgMu.Unlock() c.processedMsgs[msgID] = true
} }
c.msgMu.Unlock()
senderID := msg.FromUserName senderID := msg.FromUserName
chatID := senderID // WeCom App uses user ID as chat ID for direct messages chatID := senderID // WeCom App uses user ID as chat ID for direct messages
@ -742,8 +753,7 @@ func (c *WeComAppChannel) sendTextMessage(ctx context.Context, accessToken, user
} }
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: time.Duration(timeout) * time.Second} resp, err := c.client.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return channels.ClassifyNetError(err) return channels.ClassifyNetError(err)
} }

View file

@ -25,6 +25,7 @@ import (
type WeComBotChannel struct { type WeComBotChannel struct {
*channels.BaseChannel *channels.BaseChannel
config config.WeComConfig config config.WeComConfig
client *http.Client
ctx context.Context ctx context.Context
cancel context.CancelFunc cancel context.CancelFunc
processedMsgs map[string]bool // Message deduplication: msg_id -> processed processedMsgs map[string]bool // Message deduplication: msg_id -> processed
@ -93,10 +94,18 @@ func NewWeComBotChannel(cfg config.WeComConfig, messageBus *bus.MessageBus) (*We
channels.WithReasoningChannelID(cfg.ReasoningChannelID), channels.WithReasoningChannelID(cfg.ReasoningChannelID),
) )
// Client timeout must be >= the configured ReplyTimeout so the
// per-request context deadline is always the effective limit.
clientTimeout := 30 * time.Second
if d := time.Duration(cfg.ReplyTimeout) * time.Second; d > clientTimeout {
clientTimeout = d
}
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
return &WeComBotChannel{ return &WeComBotChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
client: &http.Client{Timeout: clientTimeout},
ctx: ctx, ctx: ctx,
cancel: cancel, cancel: cancel,
processedMsgs: make(map[string]bool), processedMsgs: make(map[string]bool),
@ -112,6 +121,10 @@ func (c *WeComBotChannel) Name() string {
func (c *WeComBotChannel) Start(ctx context.Context) error { func (c *WeComBotChannel) Start(ctx context.Context) error {
logger.InfoC("wecom", "Starting WeCom Bot channel...") logger.InfoC("wecom", "Starting WeCom Bot channel...")
// Cancel the context created in the constructor to avoid a resource leak.
if c.cancel != nil {
c.cancel()
}
c.ctx, c.cancel = context.WithCancel(ctx) c.ctx, c.cancel = context.WithCancel(ctx)
c.SetRunning(true) c.SetRunning(true)
@ -326,14 +339,14 @@ func (c *WeComBotChannel) processMessage(ctx context.Context, msg WeComBotMessag
return return
} }
c.processedMsgs[msgID] = true c.processedMsgs[msgID] = true
c.msgMu.Unlock() // Clean up old messages while still holding the lock to avoid a data race
// on len(). Reset the map but re-insert the current msgID so it remains
// Clean up old messages periodically (keep last 1000) // deduplicated.
if len(c.processedMsgs) > 1000 { if len(c.processedMsgs) > 1000 {
c.msgMu.Lock()
c.processedMsgs = make(map[string]bool) c.processedMsgs = make(map[string]bool)
c.msgMu.Unlock() c.processedMsgs[msgID] = true
} }
c.msgMu.Unlock()
senderID := msg.From.UserID senderID := msg.From.UserID
@ -446,8 +459,7 @@ func (c *WeComBotChannel) sendWebhookReply(ctx context.Context, userID, content
} }
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: time.Duration(timeout) * time.Second} resp, err := c.client.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return channels.ClassifyNetError(err) return channels.ClassifyNetError(err)
} }

View file

@ -2,7 +2,6 @@ package skills
import ( import (
"context" "context"
"encoding/json"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
@ -18,14 +17,6 @@ type SkillInstaller struct {
workspace string workspace string
} }
type AvailableSkill struct {
Name string `json:"name"`
Repository string `json:"repository"`
Description string `json:"description"`
Author string `json:"author"`
Tags []string `json:"tags"`
}
func NewSkillInstaller(workspace string) *SkillInstaller { func NewSkillInstaller(workspace string) *SkillInstaller {
return &SkillInstaller{ return &SkillInstaller{
workspace: workspace, workspace: workspace,
@ -89,35 +80,3 @@ func (si *SkillInstaller) Uninstall(skillName string) error {
return nil return nil
} }
func (si *SkillInstaller) ListAvailableSkills(ctx context.Context) ([]AvailableSkill, error) {
url := "https://raw.githubusercontent.com/sipeed/picoclaw-skills/main/skills.json"
client := &http.Client{Timeout: 15 * time.Second}
req, err := http.NewRequestWithContext(ctx, "GET", url, nil)
if err != nil {
return nil, fmt.Errorf("failed to create request: %w", err)
}
resp, err := utils.DoRequestWithRetry(client, req)
if err != nil {
return nil, fmt.Errorf("failed to fetch skills list: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
return nil, fmt.Errorf("failed to fetch skills list: HTTP %d", resp.StatusCode)
}
body, err := io.ReadAll(resp.Body)
if err != nil {
return nil, fmt.Errorf("failed to read response: %w", err)
}
var skills []AvailableSkill
if err := json.Unmarshal(body, &skills); err != nil {
return nil, fmt.Errorf("failed to parse skills list: %w", err)
}
return skills, nil
}

View file

@ -15,6 +15,14 @@ import (
const ( const (
userAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36" userAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"
// HTTP client timeouts for web tool providers.
searchTimeout = 10 * time.Second // Brave, Tavily, DuckDuckGo
perplexityTimeout = 30 * time.Second // Perplexity (LLM-based, slower)
fetchTimeout = 60 * time.Second // WebFetchTool
defaultMaxChars = 50000
maxRedirects = 5
) )
// Pre-compiled regexes for HTML text extraction // Pre-compiled regexes for HTML text extraction
@ -74,6 +82,7 @@ type SearchProvider interface {
type BraveSearchProvider struct { type BraveSearchProvider struct {
apiKey string apiKey string
proxy string proxy string
client *http.Client
} }
func (p *BraveSearchProvider) Search(ctx context.Context, query string, count int) (string, error) { func (p *BraveSearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
@ -88,11 +97,7 @@ func (p *BraveSearchProvider) Search(ctx context.Context, query string, count in
req.Header.Set("Accept", "application/json") req.Header.Set("Accept", "application/json")
req.Header.Set("X-Subscription-Token", p.apiKey) req.Header.Set("X-Subscription-Token", p.apiKey)
client, err := createHTTPClient(p.proxy, 10*time.Second) resp, err := p.client.Do(req)
if err != nil {
return "", fmt.Errorf("failed to create HTTP client: %w", err)
}
resp, err := client.Do(req)
if err != nil { if err != nil {
return "", fmt.Errorf("request failed: %w", err) return "", fmt.Errorf("request failed: %w", err)
} }
@ -143,6 +148,7 @@ type TavilySearchProvider struct {
apiKey string apiKey string
baseURL string baseURL string
proxy string proxy string
client *http.Client
} }
func (p *TavilySearchProvider) Search(ctx context.Context, query string, count int) (string, error) { func (p *TavilySearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
@ -174,11 +180,7 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", userAgent) req.Header.Set("User-Agent", userAgent)
client, err := createHTTPClient(p.proxy, 10*time.Second) resp, err := p.client.Do(req)
if err != nil {
return "", fmt.Errorf("failed to create HTTP client: %w", err)
}
resp, err := client.Do(req)
if err != nil { if err != nil {
return "", fmt.Errorf("request failed: %w", err) return "", fmt.Errorf("request failed: %w", err)
} }
@ -226,7 +228,8 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i
} }
type DuckDuckGoSearchProvider struct { type DuckDuckGoSearchProvider struct {
proxy string proxy string
client *http.Client
} }
func (p *DuckDuckGoSearchProvider) Search(ctx context.Context, query string, count int) (string, error) { func (p *DuckDuckGoSearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
@ -239,11 +242,7 @@ func (p *DuckDuckGoSearchProvider) Search(ctx context.Context, query string, cou
req.Header.Set("User-Agent", userAgent) req.Header.Set("User-Agent", userAgent)
client, err := createHTTPClient(p.proxy, 10*time.Second) resp, err := p.client.Do(req)
if err != nil {
return "", fmt.Errorf("failed to create HTTP client: %w", err)
}
resp, err := client.Do(req)
if err != nil { if err != nil {
return "", fmt.Errorf("request failed: %w", err) return "", fmt.Errorf("request failed: %w", err)
} }
@ -322,6 +321,7 @@ func stripTags(content string) string {
type PerplexitySearchProvider struct { type PerplexitySearchProvider struct {
apiKey string apiKey string
proxy string proxy string
client *http.Client
} }
func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) { func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) {
@ -356,11 +356,7 @@ func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, cou
req.Header.Set("Authorization", "Bearer "+p.apiKey) req.Header.Set("Authorization", "Bearer "+p.apiKey)
req.Header.Set("User-Agent", userAgent) req.Header.Set("User-Agent", userAgent)
client, err := createHTTPClient(p.proxy, 30*time.Second) resp, err := p.client.Do(req)
if err != nil {
return "", fmt.Errorf("failed to create HTTP client: %w", err)
}
resp, err := client.Do(req)
if err != nil { if err != nil {
return "", fmt.Errorf("request failed: %w", err) return "", fmt.Errorf("request failed: %w", err)
} }
@ -415,43 +411,60 @@ type WebSearchToolOptions struct {
Proxy string Proxy string
} }
func NewWebSearchTool(opts WebSearchToolOptions) *WebSearchTool { func NewWebSearchTool(opts WebSearchToolOptions) (*WebSearchTool, error) {
var provider SearchProvider var provider SearchProvider
maxResults := 5 maxResults := 5
// Priority: Perplexity > Brave > Tavily > DuckDuckGo // Priority: Perplexity > Brave > Tavily > DuckDuckGo
if opts.PerplexityEnabled && opts.PerplexityAPIKey != "" { if opts.PerplexityEnabled && opts.PerplexityAPIKey != "" {
provider = &PerplexitySearchProvider{apiKey: opts.PerplexityAPIKey, proxy: opts.Proxy} client, err := createHTTPClient(opts.Proxy, perplexityTimeout)
if err != nil {
return nil, fmt.Errorf("failed to create HTTP client for Perplexity: %w", err)
}
provider = &PerplexitySearchProvider{apiKey: opts.PerplexityAPIKey, proxy: opts.Proxy, client: client}
if opts.PerplexityMaxResults > 0 { if opts.PerplexityMaxResults > 0 {
maxResults = opts.PerplexityMaxResults maxResults = opts.PerplexityMaxResults
} }
} else if opts.BraveEnabled && opts.BraveAPIKey != "" { } else if opts.BraveEnabled && opts.BraveAPIKey != "" {
provider = &BraveSearchProvider{apiKey: opts.BraveAPIKey, proxy: opts.Proxy} client, err := createHTTPClient(opts.Proxy, searchTimeout)
if err != nil {
return nil, fmt.Errorf("failed to create HTTP client for Brave: %w", err)
}
provider = &BraveSearchProvider{apiKey: opts.BraveAPIKey, proxy: opts.Proxy, client: client}
if opts.BraveMaxResults > 0 { if opts.BraveMaxResults > 0 {
maxResults = opts.BraveMaxResults maxResults = opts.BraveMaxResults
} }
} else if opts.TavilyEnabled && opts.TavilyAPIKey != "" { } else if opts.TavilyEnabled && opts.TavilyAPIKey != "" {
client, err := createHTTPClient(opts.Proxy, searchTimeout)
if err != nil {
return nil, fmt.Errorf("failed to create HTTP client for Tavily: %w", err)
}
provider = &TavilySearchProvider{ provider = &TavilySearchProvider{
apiKey: opts.TavilyAPIKey, apiKey: opts.TavilyAPIKey,
baseURL: opts.TavilyBaseURL, baseURL: opts.TavilyBaseURL,
proxy: opts.Proxy, proxy: opts.Proxy,
client: client,
} }
if opts.TavilyMaxResults > 0 { if opts.TavilyMaxResults > 0 {
maxResults = opts.TavilyMaxResults maxResults = opts.TavilyMaxResults
} }
} else if opts.DuckDuckGoEnabled { } else if opts.DuckDuckGoEnabled {
provider = &DuckDuckGoSearchProvider{proxy: opts.Proxy} client, err := createHTTPClient(opts.Proxy, searchTimeout)
if err != nil {
return nil, fmt.Errorf("failed to create HTTP client for DuckDuckGo: %w", err)
}
provider = &DuckDuckGoSearchProvider{proxy: opts.Proxy, client: client}
if opts.DuckDuckGoMaxResults > 0 { if opts.DuckDuckGoMaxResults > 0 {
maxResults = opts.DuckDuckGoMaxResults maxResults = opts.DuckDuckGoMaxResults
} }
} else { } else {
return nil return nil, nil
} }
return &WebSearchTool{ return &WebSearchTool{
provider: provider, provider: provider,
maxResults: maxResults, maxResults: maxResults,
} }, nil
} }
func (t *WebSearchTool) Name() string { func (t *WebSearchTool) Name() string {
@ -508,25 +521,34 @@ func (t *WebSearchTool) Execute(ctx context.Context, args map[string]any) *ToolR
type WebFetchTool struct { type WebFetchTool struct {
maxChars int maxChars int
proxy string proxy string
client *http.Client
} }
func NewWebFetchTool(maxChars int) *WebFetchTool { func NewWebFetchTool(maxChars int) *WebFetchTool {
if maxChars <= 0 { // createHTTPClient cannot fail with an empty proxy string.
maxChars = 50000 tool, _ := NewWebFetchToolWithProxy(maxChars, "")
} return tool
return &WebFetchTool{
maxChars: maxChars,
}
} }
func NewWebFetchToolWithProxy(maxChars int, proxy string) *WebFetchTool { func NewWebFetchToolWithProxy(maxChars int, proxy string) (*WebFetchTool, error) {
if maxChars <= 0 { if maxChars <= 0 {
maxChars = 50000 maxChars = defaultMaxChars
}
client, err := createHTTPClient(proxy, fetchTimeout)
if err != nil {
return nil, fmt.Errorf("failed to create HTTP client for web fetch: %w", err)
}
client.CheckRedirect = func(req *http.Request, via []*http.Request) error {
if len(via) >= maxRedirects {
return fmt.Errorf("stopped after %d redirects", maxRedirects)
}
return nil
} }
return &WebFetchTool{ return &WebFetchTool{
maxChars: maxChars, maxChars: maxChars,
proxy: proxy, proxy: proxy,
} client: client,
}, nil
} }
func (t *WebFetchTool) Name() string { func (t *WebFetchTool) Name() string {
@ -588,20 +610,7 @@ func (t *WebFetchTool) Execute(ctx context.Context, args map[string]any) *ToolRe
req.Header.Set("User-Agent", userAgent) req.Header.Set("User-Agent", userAgent)
client, err := createHTTPClient(t.proxy, 60*time.Second) resp, err := t.client.Do(req)
if err != nil {
return ErrorResult(fmt.Sprintf("failed to create HTTP client: %v", err))
}
// Configure redirect handling
client.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 { if err != nil {
return ErrorResult(fmt.Sprintf("request failed: %v", err)) return ErrorResult(fmt.Sprintf("request failed: %v", err))
} }

View file

@ -176,13 +176,19 @@ func TestWebTool_WebFetch_Truncation(t *testing.T) {
// TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing // TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing
func TestWebTool_WebSearch_NoApiKey(t *testing.T) { func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: ""}) tool, err := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: ""})
if err != nil {
t.Fatalf("Unexpected error: %v", err)
}
if tool != nil { if tool != nil {
t.Errorf("Expected nil tool when Brave API key is empty") t.Errorf("Expected nil tool when Brave API key is empty")
} }
// Also nil when nothing is enabled // Also nil when nothing is enabled
tool = NewWebSearchTool(WebSearchToolOptions{}) tool, err = NewWebSearchTool(WebSearchToolOptions{})
if err != nil {
t.Fatalf("Unexpected error: %v", err)
}
if tool != nil { if tool != nil {
t.Errorf("Expected nil tool when no provider is enabled") t.Errorf("Expected nil tool when no provider is enabled")
} }
@ -190,7 +196,10 @@ func TestWebTool_WebSearch_NoApiKey(t *testing.T) {
// TestWebTool_WebSearch_MissingQuery verifies error handling for missing query // TestWebTool_WebSearch_MissingQuery verifies error handling for missing query
func TestWebTool_WebSearch_MissingQuery(t *testing.T) { func TestWebTool_WebSearch_MissingQuery(t *testing.T) {
tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: "test-key", BraveMaxResults: 5}) tool, err := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: "test-key", BraveMaxResults: 5})
if err != nil {
t.Fatalf("Unexpected error: %v", err)
}
ctx := context.Background() ctx := context.Background()
args := map[string]any{} args := map[string]any{}
@ -438,7 +447,10 @@ func TestCreateHTTPClient_ProxyFromEnvironmentWhenConfigEmpty(t *testing.T) {
} }
func TestNewWebFetchToolWithProxy(t *testing.T) { func TestNewWebFetchToolWithProxy(t *testing.T) {
tool := NewWebFetchToolWithProxy(1024, "http://127.0.0.1:7890") tool, err := NewWebFetchToolWithProxy(1024, "http://127.0.0.1:7890")
if err != nil {
t.Fatalf("NewWebFetchToolWithProxy() error: %v", err)
}
if tool.maxChars != 1024 { if tool.maxChars != 1024 {
t.Fatalf("maxChars = %d, want %d", tool.maxChars, 1024) t.Fatalf("maxChars = %d, want %d", tool.maxChars, 1024)
} }
@ -446,7 +458,10 @@ func TestNewWebFetchToolWithProxy(t *testing.T) {
t.Fatalf("proxy = %q, want %q", tool.proxy, "http://127.0.0.1:7890") t.Fatalf("proxy = %q, want %q", tool.proxy, "http://127.0.0.1:7890")
} }
tool = NewWebFetchToolWithProxy(0, "http://127.0.0.1:7890") tool, err = NewWebFetchToolWithProxy(0, "http://127.0.0.1:7890")
if err != nil {
t.Fatalf("NewWebFetchToolWithProxy() error: %v", err)
}
if tool.maxChars != 50000 { if tool.maxChars != 50000 {
t.Fatalf("default maxChars = %d, want %d", tool.maxChars, 50000) t.Fatalf("default maxChars = %d, want %d", tool.maxChars, 50000)
} }
@ -454,12 +469,15 @@ func TestNewWebFetchToolWithProxy(t *testing.T) {
func TestNewWebSearchTool_PropagatesProxy(t *testing.T) { func TestNewWebSearchTool_PropagatesProxy(t *testing.T) {
t.Run("perplexity", func(t *testing.T) { t.Run("perplexity", func(t *testing.T) {
tool := NewWebSearchTool(WebSearchToolOptions{ tool, err := NewWebSearchTool(WebSearchToolOptions{
PerplexityEnabled: true, PerplexityEnabled: true,
PerplexityAPIKey: "k", PerplexityAPIKey: "k",
PerplexityMaxResults: 3, PerplexityMaxResults: 3,
Proxy: "http://127.0.0.1:7890", Proxy: "http://127.0.0.1:7890",
}) })
if err != nil {
t.Fatalf("NewWebSearchTool() error: %v", err)
}
p, ok := tool.provider.(*PerplexitySearchProvider) p, ok := tool.provider.(*PerplexitySearchProvider)
if !ok { if !ok {
t.Fatalf("provider type = %T, want *PerplexitySearchProvider", tool.provider) t.Fatalf("provider type = %T, want *PerplexitySearchProvider", tool.provider)
@ -470,12 +488,15 @@ func TestNewWebSearchTool_PropagatesProxy(t *testing.T) {
}) })
t.Run("brave", func(t *testing.T) { t.Run("brave", func(t *testing.T) {
tool := NewWebSearchTool(WebSearchToolOptions{ tool, err := NewWebSearchTool(WebSearchToolOptions{
BraveEnabled: true, BraveEnabled: true,
BraveAPIKey: "k", BraveAPIKey: "k",
BraveMaxResults: 3, BraveMaxResults: 3,
Proxy: "http://127.0.0.1:7890", Proxy: "http://127.0.0.1:7890",
}) })
if err != nil {
t.Fatalf("NewWebSearchTool() error: %v", err)
}
p, ok := tool.provider.(*BraveSearchProvider) p, ok := tool.provider.(*BraveSearchProvider)
if !ok { if !ok {
t.Fatalf("provider type = %T, want *BraveSearchProvider", tool.provider) t.Fatalf("provider type = %T, want *BraveSearchProvider", tool.provider)
@ -486,11 +507,14 @@ func TestNewWebSearchTool_PropagatesProxy(t *testing.T) {
}) })
t.Run("duckduckgo", func(t *testing.T) { t.Run("duckduckgo", func(t *testing.T) {
tool := NewWebSearchTool(WebSearchToolOptions{ tool, err := NewWebSearchTool(WebSearchToolOptions{
DuckDuckGoEnabled: true, DuckDuckGoEnabled: true,
DuckDuckGoMaxResults: 3, DuckDuckGoMaxResults: 3,
Proxy: "http://127.0.0.1:7890", Proxy: "http://127.0.0.1:7890",
}) })
if err != nil {
t.Fatalf("NewWebSearchTool() error: %v", err)
}
p, ok := tool.provider.(*DuckDuckGoSearchProvider) p, ok := tool.provider.(*DuckDuckGoSearchProvider)
if !ok { if !ok {
t.Fatalf("provider type = %T, want *DuckDuckGoSearchProvider", tool.provider) t.Fatalf("provider type = %T, want *DuckDuckGoSearchProvider", tool.provider)
@ -542,12 +566,15 @@ func TestWebTool_TavilySearch_Success(t *testing.T) {
})) }))
defer server.Close() defer server.Close()
tool := NewWebSearchTool(WebSearchToolOptions{ tool, err := NewWebSearchTool(WebSearchToolOptions{
TavilyEnabled: true, TavilyEnabled: true,
TavilyAPIKey: "test-key", TavilyAPIKey: "test-key",
TavilyBaseURL: server.URL, TavilyBaseURL: server.URL,
TavilyMaxResults: 5, TavilyMaxResults: 5,
}) })
if err != nil {
t.Fatalf("NewWebSearchTool() error: %v", err)
}
ctx := context.Background() ctx := context.Background()
args := map[string]any{ args := map[string]any{

View file

@ -37,6 +37,9 @@ func DoRequestWithRetry(client *http.Client, req *http.Request) (*http.Response,
if i < maxRetries-1 { if i < maxRetries-1 {
if err = sleepWithCtx(req.Context(), retryDelayUnit*time.Duration(i+1)); err != nil { if err = sleepWithCtx(req.Context(), retryDelayUnit*time.Duration(i+1)); err != nil {
if resp != nil {
resp.Body.Close()
}
return nil, fmt.Errorf("failed to sleep: %w", err) return nil, fmt.Errorf("failed to sleep: %w", err)
} }
} }

View file

@ -1,8 +1,11 @@
package utils package utils
import ( import (
"context"
"io"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"strings"
"testing" "testing"
"time" "time"
@ -77,6 +80,91 @@ func TestDoRequestWithRetry(t *testing.T) {
} }
} }
func TestDoRequestWithRetry_ContextCancel(t *testing.T) {
// Use a long retry delay so cancellation always hits during sleepWithCtx.
retryDelayUnit = 10 * time.Second
t.Cleanup(func() { retryDelayUnit = time.Second })
bodyClosed := false
firstRoundTripDone := make(chan struct{}, 1)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
w.Write([]byte("error"))
}))
defer server.Close()
client := server.Client()
client.Timeout = 30 * time.Second
client.Transport = &bodyCloseTracker{
rt: client.Transport,
onClose: func() { bodyClosed = true },
// Signal after the first round-trip response is fully constructed on the client side.
onRoundTrip: func() {
select {
case firstRoundTripDone <- struct{}{}:
default:
}
},
trackURL: server.URL,
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
// Cancel the context after the first round-trip completes on the client side.
// This ensures client.Do has returned a valid resp (with body) and the retry
// loop is about to enter sleepWithCtx, where the cancel will be detected.
go func() {
<-firstRoundTripDone
cancel()
}()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, server.URL, nil)
require.NoError(t, err)
resp, err := DoRequestWithRetry(client, req)
if resp != nil {
resp.Body.Close()
}
require.Error(t, err, "expected error from context cancellation")
assert.Nil(t, resp, "expected nil response when context is canceled")
assert.True(t, bodyClosed, "expected resp.Body to be closed on context cancellation")
}
// bodyCloseTracker wraps an http.RoundTripper and records when response bodies are closed.
type bodyCloseTracker struct {
rt http.RoundTripper
onClose func()
onRoundTrip func() // called after each successful round-trip
trackURL string
}
func (t *bodyCloseTracker) RoundTrip(req *http.Request) (*http.Response, error) {
resp, err := t.rt.RoundTrip(req)
if err != nil {
return resp, err
}
if strings.HasPrefix(req.URL.String(), t.trackURL) {
resp.Body = &closeNotifier{ReadCloser: resp.Body, onClose: t.onClose}
if t.onRoundTrip != nil {
t.onRoundTrip()
}
}
return resp, nil
}
// closeNotifier wraps an io.ReadCloser to detect Close calls.
type closeNotifier struct {
io.ReadCloser
onClose func()
}
func (c *closeNotifier) Close() error {
c.onClose()
return c.ReadCloser.Close()
}
func TestDoRequestWithRetry_Delay(t *testing.T) { func TestDoRequestWithRetry_Delay(t *testing.T) {
retryDelayUnit = time.Millisecond retryDelayUnit = time.Millisecond
t.Cleanup(func() { retryDelayUnit = time.Second }) t.Cleanup(func() { retryDelayUnit = time.Second })