Merge remote-tracking branch 'upstream/main' into feat/agent-sandbox
This commit is contained in:
commit
4b8dd643bf
15 changed files with 308 additions and 155 deletions
25
cmd/picoclaw/internal/onboard/helpers_test.go
Normal file
25
cmd/picoclaw/internal/onboard/helpers_test.go
Normal 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -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),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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())
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
||||||
}
|
|
||||||
|
|
|
||||||
109
pkg/tools/web.go
109
pkg/tools/web.go
|
|
@ -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))
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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{
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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 })
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue