From 9e02fb45fddd6645cc822c9ca07def1b7bcdcfc2 Mon Sep 17 00:00:00 2001 From: stark Date: Fri, 27 Feb 2026 11:11:26 +0800 Subject: [PATCH] feat(web_search): add load balance and failover for api keys --- pkg/agent/loop.go | 6 +- pkg/config/config.go | 6 +- pkg/config/config_test.go | 2 +- pkg/config/defaults.go | 4 +- pkg/migrate/config.go | 4 +- pkg/tools/web.go | 278 +++++++++++++++++++++++++++----------- pkg/tools/web_test.go | 6 +- 7 files changed, 211 insertions(+), 95 deletions(-) diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index bf229ad74..1e3eb1a01 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -94,16 +94,16 @@ func registerSharedTools( // Web tools if searchTool := tools.NewWebSearchTool(tools.WebSearchToolOptions{ - BraveAPIKey: cfg.Tools.Web.Brave.APIKey, + BraveAPIKeys: cfg.Tools.Web.Brave.APIKeys, BraveMaxResults: cfg.Tools.Web.Brave.MaxResults, BraveEnabled: cfg.Tools.Web.Brave.Enabled, - TavilyAPIKey: cfg.Tools.Web.Tavily.APIKey, + TavilyAPIKeys: cfg.Tools.Web.Tavily.APIKeys, TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL, TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults, TavilyEnabled: cfg.Tools.Web.Tavily.Enabled, DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults, DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled, - PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey, + PerplexityAPIKeys: cfg.Tools.Web.Perplexity.APIKeys, PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults, PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled, }); searchTool != nil { diff --git a/pkg/config/config.go b/pkg/config/config.go index 2595398c7..e0f225243 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -416,13 +416,13 @@ type GatewayConfig struct { type BraveConfig struct { Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"` - APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEY"` + APIKeys string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_BRAVE_API_KEYS"` MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_BRAVE_MAX_RESULTS"` } type TavilyConfig struct { Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_TAVILY_ENABLED"` - APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_TAVILY_API_KEY"` + APIKeys string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_TAVILY_API_KEYS"` BaseURL string `json:"base_url" env:"PICOCLAW_TOOLS_WEB_TAVILY_BASE_URL"` MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_TAVILY_MAX_RESULTS"` } @@ -434,7 +434,7 @@ type DuckDuckGoConfig struct { type PerplexityConfig struct { Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_ENABLED"` - APIKey string `json:"api_key" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEY"` + APIKeys string `json:"api_keys" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_API_KEYS"` MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"` } diff --git a/pkg/config/config_test.go b/pkg/config/config_test.go index f88c0269c..bf6a09a2d 100644 --- a/pkg/config/config_test.go +++ b/pkg/config/config_test.go @@ -292,7 +292,7 @@ func TestDefaultConfig_WebTools(t *testing.T) { if cfg.Tools.Web.Brave.MaxResults != 5 { t.Error("Expected Brave MaxResults 5, got ", cfg.Tools.Web.Brave.MaxResults) } - if cfg.Tools.Web.Brave.APIKey != "" { + if cfg.Tools.Web.Brave.APIKeys != "" { t.Error("Brave API key should be empty by default") } if cfg.Tools.Web.DuckDuckGo.MaxResults != 5 { diff --git a/pkg/config/defaults.go b/pkg/config/defaults.go index b96ee4d89..e7be176d8 100644 --- a/pkg/config/defaults.go +++ b/pkg/config/defaults.go @@ -279,7 +279,7 @@ func DefaultConfig() *Config { Web: WebToolsConfig{ Brave: BraveConfig{ Enabled: false, - APIKey: "", + APIKeys: "", MaxResults: 5, }, DuckDuckGo: DuckDuckGoConfig{ @@ -288,7 +288,7 @@ func DefaultConfig() *Config { }, Perplexity: PerplexityConfig{ Enabled: false, - APIKey: "", + APIKeys: "", MaxResults: 5, }, }, diff --git a/pkg/migrate/config.go b/pkg/migrate/config.go index 24ce33e94..334e83c33 100644 --- a/pkg/migrate/config.go +++ b/pkg/migrate/config.go @@ -222,7 +222,7 @@ func ConvertConfig(data map[string]any) (*config.Config, []string, error) { // Migrate old "search" config to "brave" if api_key is present if search, ok := getMap(web, "search"); ok { if v, ok := getString(search, "api_key"); ok { - cfg.Tools.Web.Brave.APIKey = v + cfg.Tools.Web.Brave.APIKeys = v if v != "" { cfg.Tools.Web.Brave.Enabled = true } @@ -292,7 +292,7 @@ func MergeConfig(existing, incoming *config.Config) *config.Config { existing.Channels.MaixCam = incoming.Channels.MaixCam } - if existing.Tools.Web.Brave.APIKey == "" { + if existing.Tools.Web.Brave.APIKeys == "" { existing.Tools.Web.Brave = incoming.Tools.Web.Brave } diff --git a/pkg/tools/web.go b/pkg/tools/web.go index 452e95e0f..61fd467f6 100644 --- a/pkg/tools/web.go +++ b/pkg/tools/web.go @@ -10,6 +10,7 @@ import ( "net/url" "regexp" "strings" + "sync/atomic" "time" ) @@ -21,32 +22,97 @@ type SearchProvider interface { Search(ctx context.Context, query string, count int) (string, error) } +type APIKeyPool struct { + keys []string + current uint32 +} + +func NewAPIKeyPool(keysStr string) *APIKeyPool { + var keys []string + for _, k := range strings.Split(keysStr, ",") { + if trimmed := strings.TrimSpace(k); trimmed != "" { + keys = append(keys, trimmed) + } + } + return &APIKeyPool{ + keys: keys, + } +} + +func (p *APIKeyPool) Get() string { + if len(p.keys) == 0 { + return "" + } + if len(p.keys) == 1 { + return p.keys[0] + } + idx := atomic.AddUint32(&p.current, 1) - 1 + if idx >= uint32(len(p.keys))-1 { + atomic.CompareAndSwapUint32(&p.current, idx+1, 0) + } + return p.keys[idx%uint32(len(p.keys))] +} + +func (p *APIKeyPool) Len() int { + return len(p.keys) +} + type BraveSearchProvider struct { - apiKey string + keyPool *APIKeyPool } func (p *BraveSearchProvider) Search(ctx context.Context, query string, count int) (string, error) { searchURL := fmt.Sprintf("https://api.search.brave.com/res/v1/web/search?q=%s&count=%d", url.QueryEscape(query), count) - req, err := http.NewRequestWithContext(ctx, "GET", searchURL, nil) - if err != nil { - return "", fmt.Errorf("failed to create request: %w", err) + maxTries := p.keyPool.Len() + if maxTries == 0 { + return "", fmt.Errorf("no brave api key available") } - req.Header.Set("Accept", "application/json") - req.Header.Set("X-Subscription-Token", p.apiKey) + var lastErr error + var body []byte - client := &http.Client{Timeout: 10 * time.Second} - resp, err := client.Do(req) - if err != nil { - return "", fmt.Errorf("request failed: %w", err) + for try := 0; try < maxTries; try++ { + apiKey := p.keyPool.Get() + + req, err := http.NewRequestWithContext(ctx, "GET", searchURL, nil) + if err != nil { + return "", fmt.Errorf("failed to create request: %w", err) + } + + req.Header.Set("Accept", "application/json") + req.Header.Set("X-Subscription-Token", apiKey) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + if err != nil { + lastErr = fmt.Errorf("request failed: %w", err) + continue + } + + body, err = io.ReadAll(resp.Body) + resp.Body.Close() + + if err != nil { + lastErr = fmt.Errorf("failed to read response: %w", err) + continue + } + + if resp.StatusCode == http.StatusOK { + lastErr = nil + break + } + lastErr = fmt.Errorf("brave api error (status %d): %s", resp.StatusCode, string(body)) + if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden { + continue + } else { + break + } } - defer resp.Body.Close() - body, err := io.ReadAll(resp.Body) - if err != nil { - return "", fmt.Errorf("failed to read response: %w", err) + if lastErr != nil { + return "", fmt.Errorf("all brave api keys failed, last error: %w", lastErr) } var searchResp struct { @@ -60,8 +126,6 @@ func (p *BraveSearchProvider) Search(ctx context.Context, query string, count in } if err := json.Unmarshal(body, &searchResp); err != nil { - // Log error body for debugging - fmt.Printf("Brave API Error Body: %s\n", string(body)) return "", fmt.Errorf("failed to parse response: %w", err) } @@ -86,7 +150,7 @@ func (p *BraveSearchProvider) Search(ctx context.Context, query string, count in } type TavilySearchProvider struct { - apiKey string + keyPool *APIKeyPool baseURL string } @@ -96,43 +160,69 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i searchURL = "https://api.tavily.com/search" } - payload := map[string]any{ - "api_key": p.apiKey, - "query": query, - "search_depth": "advanced", - "include_answer": false, - "include_images": false, - "include_raw_content": false, - "max_results": count, + maxTries := p.keyPool.Len() + if maxTries == 0 { + return "", fmt.Errorf("no tavily api key available") } - bodyBytes, err := json.Marshal(payload) - if err != nil { - return "", fmt.Errorf("failed to marshal payload: %w", err) + var lastErr error + var body []byte + + for try := 0; try < maxTries; try++ { + apiKey := p.keyPool.Get() + payload := map[string]any{ + "api_key": apiKey, + "query": query, + "search_depth": "advanced", + "include_answer": false, + "include_images": false, + "include_raw_content": false, + "max_results": count, + } + + bodyBytes, err := json.Marshal(payload) + if err != nil { + return "", fmt.Errorf("failed to marshal payload: %w", err) + } + + req, err := http.NewRequestWithContext(ctx, "POST", searchURL, bytes.NewBuffer(bodyBytes)) + if err != nil { + return "", fmt.Errorf("failed to create request: %w", err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("User-Agent", userAgent) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + if err != nil { + lastErr = fmt.Errorf("request failed: %w", err) + continue + } + + body, err = io.ReadAll(resp.Body) + resp.Body.Close() + + if err != nil { + lastErr = fmt.Errorf("failed to read response: %w", err) + continue + } + + if resp.StatusCode == http.StatusOK { + lastErr = nil + break + } + + lastErr = fmt.Errorf("tavily api error (status %d): %s", resp.StatusCode, string(body)) + if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden { + continue + } else { + break + } } - req, err := http.NewRequestWithContext(ctx, "POST", searchURL, bytes.NewBuffer(bodyBytes)) - if err != nil { - return "", fmt.Errorf("failed to create request: %w", err) - } - - req.Header.Set("Content-Type", "application/json") - req.Header.Set("User-Agent", userAgent) - - client := &http.Client{Timeout: 10 * time.Second} - resp, err := client.Do(req) - if err != nil { - return "", fmt.Errorf("request failed: %w", err) - } - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return "", fmt.Errorf("failed to read response: %w", err) - } - - if resp.StatusCode != http.StatusOK { - return "", fmt.Errorf("tavily api error (status %d): %s", resp.StatusCode, string(body)) + if lastErr != nil { + return "", fmt.Errorf("all tavily api keys failed, last error: %w", lastErr) } var searchResp struct { @@ -260,12 +350,17 @@ func stripTags(content string) string { } type PerplexitySearchProvider struct { - apiKey string + keyPool *APIKeyPool } func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, count int) (string, error) { searchURL := "https://api.perplexity.ai/chat/completions" + maxTries := p.keyPool.Len() + if maxTries == 0 { + return "", fmt.Errorf("no perplexity api key available") + } + payload := map[string]any{ "model": "sonar", "messages": []map[string]string{ @@ -286,29 +381,50 @@ func (p *PerplexitySearchProvider) Search(ctx context.Context, query string, cou return "", fmt.Errorf("failed to marshal request: %w", err) } - req, err := http.NewRequestWithContext(ctx, "POST", searchURL, strings.NewReader(string(payloadBytes))) - if err != nil { - return "", fmt.Errorf("failed to create request: %w", err) + var lastErr error + var body []byte + + for try := 0; try < maxTries; try++ { + apiKey := p.keyPool.Get() + req, err := http.NewRequestWithContext(ctx, "POST", searchURL, bytes.NewReader(payloadBytes)) + if err != nil { + return "", fmt.Errorf("failed to create request: %w", err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+apiKey) + req.Header.Set("User-Agent", userAgent) + + client := &http.Client{Timeout: 30 * time.Second} + resp, err := client.Do(req) + if err != nil { + lastErr = fmt.Errorf("request failed: %w", err) + continue + } + + body, err = io.ReadAll(resp.Body) + resp.Body.Close() + + if err != nil { + lastErr = fmt.Errorf("failed to read response: %w", err) + continue + } + + if resp.StatusCode == http.StatusOK { + lastErr = nil + break + } + + lastErr = fmt.Errorf("Perplexity API error: %s", string(body)) + if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden { + continue + } else { + break + } } - req.Header.Set("Content-Type", "application/json") - req.Header.Set("Authorization", "Bearer "+p.apiKey) - req.Header.Set("User-Agent", userAgent) - - client := &http.Client{Timeout: 30 * time.Second} - resp, err := client.Do(req) - if err != nil { - return "", fmt.Errorf("request failed: %w", err) - } - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - if err != nil { - return "", fmt.Errorf("failed to read response: %w", err) - } - - if resp.StatusCode != http.StatusOK { - return "", fmt.Errorf("Perplexity API error: %s", string(body)) + if lastErr != nil { + return "", fmt.Errorf("all perplexity api keys failed, last error: %w", lastErr) } var searchResp struct { @@ -336,16 +452,16 @@ type WebSearchTool struct { } type WebSearchToolOptions struct { - BraveAPIKey string + BraveAPIKeys string BraveMaxResults int BraveEnabled bool - TavilyAPIKey string + TavilyAPIKeys string TavilyBaseURL string TavilyMaxResults int TavilyEnabled bool DuckDuckGoMaxResults int DuckDuckGoEnabled bool - PerplexityAPIKey string + PerplexityAPIKeys string PerplexityMaxResults int PerplexityEnabled bool } @@ -355,19 +471,19 @@ func NewWebSearchTool(opts WebSearchToolOptions) *WebSearchTool { maxResults := 5 // Priority: Perplexity > Brave > Tavily > DuckDuckGo - if opts.PerplexityEnabled && opts.PerplexityAPIKey != "" { - provider = &PerplexitySearchProvider{apiKey: opts.PerplexityAPIKey} + if opts.PerplexityEnabled && opts.PerplexityAPIKeys != "" { + provider = &PerplexitySearchProvider{keyPool: NewAPIKeyPool(opts.PerplexityAPIKeys)} if opts.PerplexityMaxResults > 0 { maxResults = opts.PerplexityMaxResults } - } else if opts.BraveEnabled && opts.BraveAPIKey != "" { - provider = &BraveSearchProvider{apiKey: opts.BraveAPIKey} + } else if opts.BraveEnabled && opts.BraveAPIKeys != "" { + provider = &BraveSearchProvider{keyPool: NewAPIKeyPool(opts.BraveAPIKeys)} if opts.BraveMaxResults > 0 { maxResults = opts.BraveMaxResults } - } else if opts.TavilyEnabled && opts.TavilyAPIKey != "" { + } else if opts.TavilyEnabled && opts.TavilyAPIKeys != "" { provider = &TavilySearchProvider{ - apiKey: opts.TavilyAPIKey, + keyPool: NewAPIKeyPool(opts.TavilyAPIKeys), baseURL: opts.TavilyBaseURL, } if opts.TavilyMaxResults > 0 { diff --git a/pkg/tools/web_test.go b/pkg/tools/web_test.go index 75e0d8d16..49a19b543 100644 --- a/pkg/tools/web_test.go +++ b/pkg/tools/web_test.go @@ -175,7 +175,7 @@ func TestWebTool_WebFetch_Truncation(t *testing.T) { // TestWebTool_WebSearch_NoApiKey verifies that no tool is created when API key is missing func TestWebTool_WebSearch_NoApiKey(t *testing.T) { - tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: ""}) + tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKeys: ""}) if tool != nil { t.Errorf("Expected nil tool when Brave API key is empty") } @@ -189,7 +189,7 @@ func TestWebTool_WebSearch_NoApiKey(t *testing.T) { // TestWebTool_WebSearch_MissingQuery verifies error handling for missing query func TestWebTool_WebSearch_MissingQuery(t *testing.T) { - tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKey: "test-key", BraveMaxResults: 5}) + tool := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKeys: "test-key", BraveMaxResults: 5}) ctx := context.Background() args := map[string]any{} @@ -377,7 +377,7 @@ func TestWebTool_TavilySearch_Success(t *testing.T) { tool := NewWebSearchTool(WebSearchToolOptions{ TavilyEnabled: true, - TavilyAPIKey: "test-key", + TavilyAPIKeys: "test-key", TavilyBaseURL: server.URL, TavilyMaxResults: 5, })