Files
goclaw/internal/tools/web_search.go
T
mozaa cfb7f46324 feat(web_search): add provider arg to force a specific engine (#1108)
Currently the web_search tool resolves a tenant-configured provider
chain (e.g. tavily → exa → brave) and stops at the first success.
That's the right default for cost/latency, but it makes cross-engine
corroboration impossible at the agent level: the caller has no way to
ask "search the same query on Exa specifically" once Tavily has
already returned a hit.

Add an optional `provider` argument. When set, the chain is narrowed
to the named provider (case-insensitive); other providers are not
queried. The first-success-wins fallback is preserved when the arg
is omitted.

Use case (concrete): a research agent verifying source freshness
calls web_search twice for the same query — once with provider="tavily"
and once with provider="exa" — and compares the URLs returned. URLs
that appear in both engines are high-confidence original primary
sources; URLs only in one signal a republish or low-circulation outlet
worth verifying further.

Other changes:

  - Cache key now includes the requested provider so per-engine
    results don't collide. Without this, the second call would just
    replay the first engine's cached result, defeating the purpose.
  - Unknown provider returns a clear error listing what IS configured
    for the tenant — better DX than a silent no-op.

Tests: 5 new cases covering narrowing, case-insensitivity, unknown
provider, default first-success behaviour preservation, and cache
isolation across providers.
2026-06-25 16:59:48 +07:00

292 lines
8.5 KiB
Go

package tools
import (
"context"
"fmt"
"log/slog"
"regexp"
"strings"
"time"
"github.com/google/uuid"
"github.com/nextlevelbuilder/goclaw/internal/bus"
"github.com/nextlevelbuilder/goclaw/internal/store"
"github.com/nextlevelbuilder/goclaw/pkg/protocol"
)
// Matching TS src/agents/tools/web-search.ts constants.
const (
defaultSearchCount = 5
maxSearchCount = 10
searchTimeoutSeconds = 30
braveSearchEndpoint = "https://api.search.brave.com/res/v1/web/search"
exaSearchEndpoint = "https://api.exa.ai/search"
tavilySearchEndpoint = "https://api.tavily.com/search"
webSearchUserAgent = "Mozilla/5.0 (Macintosh; Intel Mac OS X 14_7_2) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"
)
const (
searchProviderExa = "exa"
searchProviderTavily = "tavily"
searchProviderBrave = "brave"
searchProviderDuckDuckGo = "duckduckgo"
)
var defaultSearchProviderOrder = []string{
searchProviderExa,
searchProviderTavily,
searchProviderBrave,
searchProviderDuckDuckGo,
}
// SearchProvider abstracts a web search backend.
type SearchProvider interface {
Search(ctx context.Context, params searchParams) ([]searchResult, error)
Name() string
}
type searchParams struct {
Query string
Count int
Country string
SearchLang string
UILang string
Freshness string
}
type searchResult struct {
Title string `json:"title"`
URL string `json:"url"`
Description string `json:"description"`
}
// --- Freshness validation (matching TS) ---
var (
freshnessShortcuts = map[string]bool{"pd": true, "pw": true, "pm": true, "py": true}
freshnessRangeRe = regexp.MustCompile(`^(\d{4}-\d{2}-\d{2})to(\d{4}-\d{2}-\d{2})$`)
)
func normalizeFreshness(value string) string {
v := strings.ToLower(strings.TrimSpace(value))
if v == "" {
return ""
}
if freshnessShortcuts[v] {
return v
}
if m := freshnessRangeRe.FindStringSubmatch(v); len(m) == 3 {
start, errS := time.Parse("2006-01-02", m[1])
end, errE := time.Parse("2006-01-02", m[2])
if errS == nil && errE == nil && !start.After(end) {
return v
}
}
return ""
}
// --- WebSearchTool ---
// WebSearchTool implements the web_search tool with per-tenant provider chain
// resolution. Providers are resolved per-request from config_secrets and
// builtin_tool_tenant_configs.settings, cached per tenant for 60 seconds.
type WebSearchTool struct {
secrets store.ConfigSecretsStore
cache *webCache
chainCache *tenantChainCache
}
// NewWebSearchTool constructs a WebSearchTool. msgBus may be nil (e.g. desktop
// edition) — cache invalidation then relies on TTL alone.
func NewWebSearchTool(secrets store.ConfigSecretsStore, msgBus *bus.MessageBus) *WebSearchTool {
t := &WebSearchTool{
secrets: secrets,
cache: newWebCache(defaultCacheMaxEntries, defaultCacheTTL),
chainCache: newTenantChainCache(),
}
if msgBus != nil {
msgBus.Subscribe("web_search:cache_invalidate", func(event bus.Event) {
if event.Name != protocol.EventCacheInvalidate {
return
}
payload, ok := event.Payload.(bus.CacheInvalidatePayload)
if !ok || payload.Kind != bus.CacheKindBuiltinTools || payload.Key != "web_search" {
return
}
if payload.TenantID == uuid.Nil {
// Master admin write — wipe all tenants.
t.chainCache.InvalidateAll()
} else {
t.chainCache.Invalidate(payload.TenantID)
}
})
}
return t
}
func (t *WebSearchTool) Name() string { return "web_search" }
func (t *WebSearchTool) Description() string {
return "Search the web for current information. Returns titles, URLs, and snippets from search results."
}
func (t *WebSearchTool) Parameters() map[string]any {
return map[string]any{
"type": "object",
"properties": map[string]any{
"query": map[string]any{
"type": "string",
"description": "Search query string.",
},
"count": map[string]any{
"type": "number",
"description": "Number of results to return (1-10).",
"minimum": 1.0,
"maximum": float64(maxSearchCount),
},
"country": map[string]any{
"type": "string",
"description": "2-letter country code for region-specific results (e.g., 'DE', 'US', 'ALL'). Default: 'US'.",
},
"search_lang": map[string]any{
"type": "string",
"description": "ISO language code for search results (e.g., 'de', 'en', 'fr').",
},
"ui_lang": map[string]any{
"type": "string",
"description": "ISO language code for UI elements.",
},
"freshness": map[string]any{
"type": "string",
"description": "Filter results by discovery time. Supports 'pd' (past day), 'pw' (past week), 'pm' (past month), 'py' (past year), and date range 'YYYY-MM-DDtoYYYY-MM-DD'.",
},
"provider": map[string]any{
"type": "string",
"description": "Optional: force a specific provider (e.g., 'tavily', 'exa', 'brave', 'duckduckgo'). When omitted, the tenant's configured provider chain is used (first-success-wins). Use this to force cross-engine corroboration — call once with each provider and compare results.",
},
},
"required": []string{"query"},
}
}
func (t *WebSearchTool) Execute(ctx context.Context, args map[string]any) *Result {
query, _ := args["query"].(string)
if query == "" {
return ErrorResult("query is required")
}
count := defaultSearchCount
if c, ok := args["count"].(float64); ok && int(c) >= 1 && int(c) <= maxSearchCount {
count = int(c)
}
country, _ := args["country"].(string)
searchLang, _ := args["search_lang"].(string)
uiLang, _ := args["ui_lang"].(string)
freshness, _ := args["freshness"].(string)
requestedProvider, _ := args["provider"].(string)
params := searchParams{
Query: query,
Count: count,
Country: country,
SearchLang: searchLang,
UILang: uiLang,
Freshness: freshness,
}
// Check cache (scoped per channel + provider to prevent cross-engine cache mixing)
channel := ToolChannelFromCtx(ctx)
cacheKey := fmt.Sprintf("%s:%s:%s", channel, requestedProvider, buildSearchCacheKey(params))
if cached, ok := t.cache.get(cacheKey); ok {
slog.Debug("web_search cache hit", "query", query, "provider", requestedProvider)
return NewResult(cached)
}
// Resolve per-request provider chain from tenant config_secrets + settings overlay.
chain := t.resolveChain(ctx)
// If caller explicitly named a provider, narrow the chain to just that one.
// This unlocks cross-engine corroboration (caller invokes once per engine
// and compares results) — without a provider param, the first-success-wins
// chain hides everything after the first hit.
if requestedProvider != "" {
filtered := make([]SearchProvider, 0, 1)
for _, p := range chain {
if strings.EqualFold(p.Name(), requestedProvider) {
filtered = append(filtered, p)
break
}
}
if len(filtered) == 0 {
available := make([]string, 0, len(chain))
for _, p := range chain {
available = append(available, p.Name())
}
return ErrorResult(fmt.Sprintf("provider %q not configured for this tenant; available: %v", requestedProvider, available))
}
chain = filtered
}
// Try providers in order (first success wins, unless narrowed above)
var lastErr error
for _, provider := range chain {
results, err := provider.Search(ctx, params)
if err != nil {
slog.Warn("web_search provider failed", "provider", provider.Name(), "error", err)
lastErr = err
continue
}
formatted := formatSearchResults(query, results, provider.Name())
wrapped := wrapExternalContent(formatted, "Web Search", false)
t.cache.set(cacheKey, wrapped)
return NewResult(wrapped)
}
if lastErr != nil {
return ErrorResult(fmt.Sprintf("all search providers failed: %v", lastErr))
}
return ErrorResult("no search providers configured")
}
func buildSearchCacheKey(p searchParams) string {
parts := []string{
p.Query,
fmt.Sprintf("%d", p.Count),
orDefault(p.Country, "default"),
orDefault(p.SearchLang, "default"),
orDefault(p.UILang, "default"),
orDefault(p.Freshness, "default"),
}
return strings.Join(parts, ":")
}
func orDefault(s, def string) string {
if s == "" {
return def
}
return s
}
func formatSearchResults(query string, results []searchResult, provider string) string {
if len(results) == 0 {
return fmt.Sprintf("No results found for: %s", query)
}
var sb strings.Builder
sb.WriteString(fmt.Sprintf("Search results for: %s (via %s)\n\n", query, provider))
for i, r := range results {
sb.WriteString(fmt.Sprintf("%d. %s\n %s\n", i+1, r.Title, r.URL))
if r.Description != "" {
sb.WriteString(fmt.Sprintf(" %s\n", r.Description))
}
sb.WriteByte('\n')
}
return sb.String()
}