fix(tools): close resp.Body on retry cancel and cache http.Client instances

Fix resp.Body leak in DoRequestWithRetry where req.Body (request) was
incorrectly closed instead of resp.Body (response) on context cancel.
Cache http.Client on web search/fetch provider structs and channel
adapters (WeCom, LINE) to avoid per-call allocation overhead.
This commit is contained in:
ItsT0ng 2026-03-01 14:57:16 +11:00
parent 33f67e8275
commit dd9d6e40cf
7 changed files with 118 additions and 67 deletions

View file

@ -99,7 +99,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,
@ -113,10 +113,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,6 +45,7 @@ type replyTokenEntry struct {
type LINEChannel struct { type LINEChannel struct {
*channels.BaseChannel *channels.BaseChannel
config config.LINEConfig config config.LINEConfig
client *http.Client
botUserID string // Bot's user ID botUserID string // Bot's user ID
botBasicID string // Bot's basic ID (e.g. @216ru...) botBasicID string // Bot's basic ID (e.g. @216ru...)
botDisplayName string // Bot's display name for text-based mention detection botDisplayName string // Bot's display name for text-based mention detection
@ -69,6 +70,7 @@ func NewLINEChannel(cfg config.LINEConfig, messageBus *bus.MessageBus) (*LINECha
return &LINEChannel{ return &LINEChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
client: &http.Client{Timeout: 60 * time.Second},
}, nil }, nil
} }
@ -104,8 +106,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.client.Do(req)
resp, err := client.Do(req)
if err != nil { if err != nil {
return err return err
} }
@ -644,8 +645,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.client.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
@ -133,6 +134,7 @@ func NewWeComAppChannel(cfg config.WeComAppConfig, messageBus *bus.MessageBus) (
return &WeComAppChannel{ return &WeComAppChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
client: &http.Client{Timeout: 60 * time.Second},
ctx: ctx, ctx: ctx,
cancel: cancel, cancel: cancel,
processedMsgs: make(map[string]bool), processedMsgs: make(map[string]bool),
@ -306,8 +308,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)
} }
@ -364,8 +365,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)
} }
@ -746,8 +746,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
@ -97,6 +98,7 @@ func NewWeComBotChannel(cfg config.WeComConfig, messageBus *bus.MessageBus) (*We
return &WeComBotChannel{ return &WeComBotChannel{
BaseChannel: base, BaseChannel: base,
config: cfg, config: cfg,
client: &http.Client{Timeout: 60 * time.Second},
ctx: ctx, ctx: ctx,
cancel: cancel, cancel: cancel,
processedMsgs: make(map[string]bool), processedMsgs: make(map[string]bool),
@ -450,8 +452,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

@ -74,6 +74,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 +89,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 +140,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 +172,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)
} }
@ -227,6 +221,7 @@ 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 +234,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 +313,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 +348,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 +403,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, 30*time.Second)
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, 10*time.Second)
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, 10*time.Second)
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, 10*time.Second)
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 +513,46 @@ 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 { if maxChars <= 0 {
maxChars = 50000 maxChars = 50000
} }
// No proxy — createHTTPClient cannot fail with an empty proxy string.
client, _ := createHTTPClient("", 60*time.Second)
client.CheckRedirect = func(req *http.Request, via []*http.Request) error {
if len(via) >= 5 {
return fmt.Errorf("stopped after 5 redirects")
}
return nil
}
return &WebFetchTool{ return &WebFetchTool{
maxChars: maxChars, maxChars: maxChars,
client: client,
} }
} }
func NewWebFetchToolWithProxy(maxChars int, proxy string) *WebFetchTool { func NewWebFetchToolWithProxy(maxChars int, proxy string) (*WebFetchTool, error) {
if maxChars <= 0 { if maxChars <= 0 {
maxChars = 50000 maxChars = 50000
} }
client, err := createHTTPClient(proxy, 60*time.Second)
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) >= 5 {
return fmt.Errorf("stopped after 5 redirects")
}
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 +614,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)
} }
} }