feat(web_search): add load balance and failover for api keys

This commit is contained in:
stark 2026-02-27 11:11:26 +08:00
parent 7cbfa89a96
commit 9e02fb45fd
7 changed files with 211 additions and 95 deletions

View file

@ -94,16 +94,16 @@ func registerSharedTools(
// Web tools // Web tools
if searchTool := tools.NewWebSearchTool(tools.WebSearchToolOptions{ if searchTool := tools.NewWebSearchTool(tools.WebSearchToolOptions{
BraveAPIKey: cfg.Tools.Web.Brave.APIKey, BraveAPIKeys: cfg.Tools.Web.Brave.APIKeys,
BraveMaxResults: cfg.Tools.Web.Brave.MaxResults, BraveMaxResults: cfg.Tools.Web.Brave.MaxResults,
BraveEnabled: cfg.Tools.Web.Brave.Enabled, BraveEnabled: cfg.Tools.Web.Brave.Enabled,
TavilyAPIKey: cfg.Tools.Web.Tavily.APIKey, TavilyAPIKeys: cfg.Tools.Web.Tavily.APIKeys,
TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL, TavilyBaseURL: cfg.Tools.Web.Tavily.BaseURL,
TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults, TavilyMaxResults: cfg.Tools.Web.Tavily.MaxResults,
TavilyEnabled: cfg.Tools.Web.Tavily.Enabled, TavilyEnabled: cfg.Tools.Web.Tavily.Enabled,
DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults, DuckDuckGoMaxResults: cfg.Tools.Web.DuckDuckGo.MaxResults,
DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled, DuckDuckGoEnabled: cfg.Tools.Web.DuckDuckGo.Enabled,
PerplexityAPIKey: cfg.Tools.Web.Perplexity.APIKey, PerplexityAPIKeys: cfg.Tools.Web.Perplexity.APIKeys,
PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults, PerplexityMaxResults: cfg.Tools.Web.Perplexity.MaxResults,
PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled, PerplexityEnabled: cfg.Tools.Web.Perplexity.Enabled,
}); searchTool != nil { }); searchTool != nil {

View file

@ -416,13 +416,13 @@ type GatewayConfig struct {
type BraveConfig struct { type BraveConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_BRAVE_ENABLED"` 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"` MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_BRAVE_MAX_RESULTS"`
} }
type TavilyConfig struct { type TavilyConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_TAVILY_ENABLED"` 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"` BaseURL string `json:"base_url" env:"PICOCLAW_TOOLS_WEB_TAVILY_BASE_URL"`
MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_TAVILY_MAX_RESULTS"` MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_TAVILY_MAX_RESULTS"`
} }
@ -434,7 +434,7 @@ type DuckDuckGoConfig struct {
type PerplexityConfig struct { type PerplexityConfig struct {
Enabled bool `json:"enabled" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_ENABLED"` 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"` MaxResults int `json:"max_results" env:"PICOCLAW_TOOLS_WEB_PERPLEXITY_MAX_RESULTS"`
} }

View file

@ -292,7 +292,7 @@ func TestDefaultConfig_WebTools(t *testing.T) {
if cfg.Tools.Web.Brave.MaxResults != 5 { if cfg.Tools.Web.Brave.MaxResults != 5 {
t.Error("Expected Brave MaxResults 5, got ", cfg.Tools.Web.Brave.MaxResults) 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") t.Error("Brave API key should be empty by default")
} }
if cfg.Tools.Web.DuckDuckGo.MaxResults != 5 { if cfg.Tools.Web.DuckDuckGo.MaxResults != 5 {

View file

@ -279,7 +279,7 @@ func DefaultConfig() *Config {
Web: WebToolsConfig{ Web: WebToolsConfig{
Brave: BraveConfig{ Brave: BraveConfig{
Enabled: false, Enabled: false,
APIKey: "", APIKeys: "",
MaxResults: 5, MaxResults: 5,
}, },
DuckDuckGo: DuckDuckGoConfig{ DuckDuckGo: DuckDuckGoConfig{
@ -288,7 +288,7 @@ func DefaultConfig() *Config {
}, },
Perplexity: PerplexityConfig{ Perplexity: PerplexityConfig{
Enabled: false, Enabled: false,
APIKey: "", APIKeys: "",
MaxResults: 5, MaxResults: 5,
}, },
}, },

View file

@ -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 // Migrate old "search" config to "brave" if api_key is present
if search, ok := getMap(web, "search"); ok { if search, ok := getMap(web, "search"); ok {
if v, ok := getString(search, "api_key"); ok { if v, ok := getString(search, "api_key"); ok {
cfg.Tools.Web.Brave.APIKey = v cfg.Tools.Web.Brave.APIKeys = v
if v != "" { if v != "" {
cfg.Tools.Web.Brave.Enabled = true cfg.Tools.Web.Brave.Enabled = true
} }
@ -292,7 +292,7 @@ func MergeConfig(existing, incoming *config.Config) *config.Config {
existing.Channels.MaixCam = incoming.Channels.MaixCam 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 existing.Tools.Web.Brave = incoming.Tools.Web.Brave
} }

View file

@ -10,6 +10,7 @@ import (
"net/url" "net/url"
"regexp" "regexp"
"strings" "strings"
"sync/atomic"
"time" "time"
) )
@ -21,32 +22,97 @@ type SearchProvider interface {
Search(ctx context.Context, query string, count int) (string, error) 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 { type BraveSearchProvider struct {
apiKey string keyPool *APIKeyPool
} }
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) {
searchURL := fmt.Sprintf("https://api.search.brave.com/res/v1/web/search?q=%s&count=%d", searchURL := fmt.Sprintf("https://api.search.brave.com/res/v1/web/search?q=%s&count=%d",
url.QueryEscape(query), count) url.QueryEscape(query), count)
req, err := http.NewRequestWithContext(ctx, "GET", searchURL, nil) maxTries := p.keyPool.Len()
if err != nil { if maxTries == 0 {
return "", fmt.Errorf("failed to create request: %w", err) return "", fmt.Errorf("no brave api key available")
} }
req.Header.Set("Accept", "application/json") var lastErr error
req.Header.Set("X-Subscription-Token", p.apiKey) var body []byte
client := &http.Client{Timeout: 10 * time.Second} for try := 0; try < maxTries; try++ {
resp, err := client.Do(req) apiKey := p.keyPool.Get()
if err != nil {
return "", fmt.Errorf("request failed: %w", err) 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 lastErr != nil {
if err != nil { return "", fmt.Errorf("all brave api keys failed, last error: %w", lastErr)
return "", fmt.Errorf("failed to read response: %w", err)
} }
var searchResp struct { 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 { 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) 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 { type TavilySearchProvider struct {
apiKey string keyPool *APIKeyPool
baseURL string baseURL string
} }
@ -96,43 +160,69 @@ func (p *TavilySearchProvider) Search(ctx context.Context, query string, count i
searchURL = "https://api.tavily.com/search" searchURL = "https://api.tavily.com/search"
} }
payload := map[string]any{ maxTries := p.keyPool.Len()
"api_key": p.apiKey, if maxTries == 0 {
"query": query, return "", fmt.Errorf("no tavily api key available")
"search_depth": "advanced",
"include_answer": false,
"include_images": false,
"include_raw_content": false,
"max_results": count,
} }
bodyBytes, err := json.Marshal(payload) var lastErr error
if err != nil { var body []byte
return "", fmt.Errorf("failed to marshal payload: %w", err)
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 lastErr != nil {
if err != nil { return "", fmt.Errorf("all tavily api keys failed, last error: %w", lastErr)
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))
} }
var searchResp struct { var searchResp struct {
@ -260,12 +350,17 @@ func stripTags(content string) string {
} }
type PerplexitySearchProvider struct { type PerplexitySearchProvider struct {
apiKey string keyPool *APIKeyPool
} }
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) {
searchURL := "https://api.perplexity.ai/chat/completions" 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{ payload := map[string]any{
"model": "sonar", "model": "sonar",
"messages": []map[string]string{ "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) return "", fmt.Errorf("failed to marshal request: %w", err)
} }
req, err := http.NewRequestWithContext(ctx, "POST", searchURL, strings.NewReader(string(payloadBytes))) var lastErr error
if err != nil { var body []byte
return "", fmt.Errorf("failed to create request: %w", err)
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") if lastErr != nil {
req.Header.Set("Authorization", "Bearer "+p.apiKey) return "", fmt.Errorf("all perplexity api keys failed, last error: %w", lastErr)
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))
} }
var searchResp struct { var searchResp struct {
@ -336,16 +452,16 @@ type WebSearchTool struct {
} }
type WebSearchToolOptions struct { type WebSearchToolOptions struct {
BraveAPIKey string BraveAPIKeys string
BraveMaxResults int BraveMaxResults int
BraveEnabled bool BraveEnabled bool
TavilyAPIKey string TavilyAPIKeys string
TavilyBaseURL string TavilyBaseURL string
TavilyMaxResults int TavilyMaxResults int
TavilyEnabled bool TavilyEnabled bool
DuckDuckGoMaxResults int DuckDuckGoMaxResults int
DuckDuckGoEnabled bool DuckDuckGoEnabled bool
PerplexityAPIKey string PerplexityAPIKeys string
PerplexityMaxResults int PerplexityMaxResults int
PerplexityEnabled bool PerplexityEnabled bool
} }
@ -355,19 +471,19 @@ func NewWebSearchTool(opts WebSearchToolOptions) *WebSearchTool {
maxResults := 5 maxResults := 5
// Priority: Perplexity > Brave > Tavily > DuckDuckGo // Priority: Perplexity > Brave > Tavily > DuckDuckGo
if opts.PerplexityEnabled && opts.PerplexityAPIKey != "" { if opts.PerplexityEnabled && opts.PerplexityAPIKeys != "" {
provider = &PerplexitySearchProvider{apiKey: opts.PerplexityAPIKey} provider = &PerplexitySearchProvider{keyPool: NewAPIKeyPool(opts.PerplexityAPIKeys)}
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.BraveAPIKeys != "" {
provider = &BraveSearchProvider{apiKey: opts.BraveAPIKey} provider = &BraveSearchProvider{keyPool: NewAPIKeyPool(opts.BraveAPIKeys)}
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.TavilyAPIKeys != "" {
provider = &TavilySearchProvider{ provider = &TavilySearchProvider{
apiKey: opts.TavilyAPIKey, keyPool: NewAPIKeyPool(opts.TavilyAPIKeys),
baseURL: opts.TavilyBaseURL, baseURL: opts.TavilyBaseURL,
} }
if opts.TavilyMaxResults > 0 { if opts.TavilyMaxResults > 0 {

View file

@ -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 // 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 := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKeys: ""})
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")
} }
@ -189,7 +189,7 @@ 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 := NewWebSearchTool(WebSearchToolOptions{BraveEnabled: true, BraveAPIKeys: "test-key", BraveMaxResults: 5})
ctx := context.Background() ctx := context.Background()
args := map[string]any{} args := map[string]any{}
@ -377,7 +377,7 @@ func TestWebTool_TavilySearch_Success(t *testing.T) {
tool := NewWebSearchTool(WebSearchToolOptions{ tool := NewWebSearchTool(WebSearchToolOptions{
TavilyEnabled: true, TavilyEnabled: true,
TavilyAPIKey: "test-key", TavilyAPIKeys: "test-key",
TavilyBaseURL: server.URL, TavilyBaseURL: server.URL,
TavilyMaxResults: 5, TavilyMaxResults: 5,
}) })