diff --git a/.gitignore b/.gitignore index b90fb8ad..4a28c084 100644 --- a/.gitignore +++ b/.gitignore @@ -85,3 +85,4 @@ sandbox/v2/*.md POSTGRESQL_COMPAT.md openapi/setting/*.md agent/docs/design/*.md +tools/README.md diff --git a/agent/search/handlers/web/handler.go b/agent/search/handlers/web/handler.go index d45eaf54..be84fb86 100644 --- a/agent/search/handlers/web/handler.go +++ b/agent/search/handlers/web/handler.go @@ -3,9 +3,11 @@ package web import ( "fmt" "strings" + "time" agentContext "github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/agent/search/types" + "github.com/yaoapp/yao/tools/websearch" ) // Handler implements web search @@ -35,7 +37,7 @@ func (h *Handler) Search(req *types.Request) (*types.Result, error) { func (h *Handler) SearchWithContext(ctx *agentContext.Context, req *types.Request) (*types.Result, error) { switch { case h.usesWeb == "builtin" || h.usesWeb == "": - return h.builtinSearch(req) + return h.builtinSearch(ctx, req) case strings.HasPrefix(h.usesWeb, "mcp:"): return h.mcpSearch(req) default: @@ -54,33 +56,43 @@ func (h *Handler) SearchWithContext(ctx *agentContext.Context, req *types.Reques } } -// builtinSearch uses Tavily/Serper/SerpAPI directly -func (h *Handler) builtinSearch(req *types.Request) (*types.Result, error) { - // Determine provider from config - providerName := "tavily" // default - if h.config != nil && h.config.Provider != "" { - providerName = h.config.Provider +// builtinSearch delegates to tools/websearch which reads Settings → ENV config. +func (h *Handler) builtinSearch(ctx *agentContext.Context, req *types.Request) (*types.Result, error) { + startTime := time.Now() + + limit := req.Limit + if limit <= 0 { + limit = 10 } - switch providerName { - case "tavily": - return NewTavilyProvider(h.config).Search(req) - case "serper": - // Serper (serper.dev) - POST request with X-API-KEY header - return NewSerperProvider(h.config).Search(req) - case "serpapi": - // SerpAPI (serpapi.com) - GET request with api_key parameter - return NewSerpAPIProvider(h.config).Search(req) - default: - return &types.Result{ - Type: types.SearchTypeWeb, - Query: req.Query, - Source: req.Source, - Items: []*types.ResultItem{}, - Total: 0, - Error: fmt.Sprintf("Unknown provider: %s (supported: tavily, serper, serpapi)", providerName), - }, nil + var userID, teamID string + if ctx != nil && ctx.Authorized != nil { + userID = ctx.Authorized.UserID + teamID = ctx.Authorized.TeamID } + + results := websearch.Search(req.Query, limit, userID, teamID) + + items := make([]*types.ResultItem, 0, len(results)) + for _, r := range results { + items = append(items, &types.ResultItem{ + Type: types.SearchTypeWeb, + Title: r.Title, + Content: r.Content, + URL: r.URL, + Score: r.Score, + Source: req.Source, + }) + } + + return &types.Result{ + Type: types.SearchTypeWeb, + Query: req.Query, + Source: req.Source, + Items: items, + Total: len(items), + Duration: time.Since(startTime).Milliseconds(), + }, nil } // agentSearch delegates to an assistant for AI-powered search diff --git a/config/decrypt.go b/config/decrypt.go new file mode 100644 index 00000000..d3a0a1a3 --- /dev/null +++ b/config/decrypt.go @@ -0,0 +1,54 @@ +package config + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/sha256" + "encoding/base64" + "strings" +) + +const encPrefix = "enc:" + +// DecryptValue decrypts a value encrypted by cloud settings. +// Returns the original string if not encrypted (no "enc:" prefix) +// or if no AES key is configured. +func DecryptValue(s string) string { + if !strings.HasPrefix(s, encPrefix) { + return s + } + secret := Conf.DB.AESKey + if secret == "" { + return strings.TrimPrefix(s, encPrefix) + } + dec, err := aesGCMDecrypt(strings.TrimPrefix(s, encPrefix), secret) + if err != nil { + return s + } + return dec +} + +func aesGCMDecrypt(encoded, secret string) (string, error) { + key := sha256.Sum256([]byte(secret)) + data, err := base64.StdEncoding.DecodeString(encoded) + if err != nil { + return "", err + } + block, err := aes.NewCipher(key[:]) + if err != nil { + return "", err + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return "", err + } + nonceSize := gcm.NonceSize() + if len(data) < nonceSize { + return "", aes.KeySizeError(len(data)) + } + plaintext, err := gcm.Open(nil, data[:nonceSize], data[nonceSize:], nil) + if err != nil { + return "", err + } + return string(plaintext), nil +} diff --git a/main.go b/main.go index b0e8511f..a6ef86d7 100644 --- a/main.go +++ b/main.go @@ -13,6 +13,7 @@ import ( _ "github.com/yaoapp/yao/rss" _ "github.com/yaoapp/yao/seed" _ "github.com/yaoapp/yao/sitemap" + _ "github.com/yaoapp/yao/tools" _ "github.com/yaoapp/yao/trace/jsapi" _ "github.com/yaoapp/yao/wework" diff --git a/openapi/setting/cloud.go b/openapi/setting/cloud.go index 4f9163e9..91d762d1 100644 --- a/openapi/setting/cloud.go +++ b/openapi/setting/cloud.go @@ -386,6 +386,12 @@ func cloudDecrypt(value string) string { return dec } +// DecryptValue decrypts a value encrypted by cloudEncrypt. +// Delegates to config.DecryptValue for the actual decryption. +func DecryptValue(s string) string { + return config.DecryptValue(s) +} + func cloudMaskKey(key string) string { if key == "" { return "" diff --git a/tools/docs/docs.go b/tools/docs/docs.go new file mode 100644 index 00000000..2d2b5ede --- /dev/null +++ b/tools/docs/docs.go @@ -0,0 +1,54 @@ +package docs + +import ( + _ "embed" + + goudoc "github.com/yaoapp/gou/doc" + "github.com/yaoapp/gou/process" +) + +//go:embed list.json +var ListSchemaJSON []byte + +//go:embed inspect.json +var InspectSchemaJSON []byte + +//go:embed validate.json +var ValidateSchemaJSON []byte + +// ListHandler is the tools.doclist process handler. +// Args[0]: keyword (string, optional — empty lists all) +// Args[1]: limit (int, default 20) +func ListHandler(proc *process.Process) interface{} { + keyword := proc.ArgsString(0) + limit := proc.ArgsInt(1, 20) + + var results []*goudoc.Entry + if keyword != "" { + results = goudoc.List(goudoc.TypeProcess, goudoc.ListOption{Search: keyword}) + } else { + results = goudoc.List(goudoc.TypeProcess) + } + if len(results) > limit { + results = results[:limit] + } + return results +} + +// InspectHandler is the tools.docinspect process handler. +// Args[0]: name (string — process name, e.g. "models.user.Find") +func InspectHandler(proc *process.Process) interface{} { + name := proc.ArgsString(0) + entry, ok := goudoc.Get(goudoc.TypeProcess, name) + if !ok { + return nil + } + return entry +} + +// ValidateHandler is the tools.docvalidate process handler. +// Args[0]: name (string — process name) +func ValidateHandler(proc *process.Process) interface{} { + name := proc.ArgsString(0) + return goudoc.Validate(goudoc.TypeProcess, name) +} diff --git a/tools/docs/docs_test.go b/tools/docs/docs_test.go new file mode 100644 index 00000000..cbf5a37f --- /dev/null +++ b/tools/docs/docs_test.go @@ -0,0 +1,111 @@ +package docs + +import ( + "testing" + + goudoc "github.com/yaoapp/gou/doc" + "github.com/yaoapp/gou/process" +) + +func init() { + goudoc.Register(&goudoc.Entry{ + Name: "models.Find", + Type: goudoc.TypeProcess, + Group: "models", + Desc: "Find records by conditions", + Args: []goudoc.TypeValue{ + {Type: "object", Desc: "query conditions"}, + }, + Return: &goudoc.TypeValue{Type: "array", Desc: "matched records"}, + }) + goudoc.Register(&goudoc.Entry{ + Name: "models.Save", + Type: goudoc.TypeProcess, + Group: "models", + Desc: "Save a record", + Args: []goudoc.TypeValue{ + {Type: "object", Desc: "record data"}, + }, + Return: &goudoc.TypeValue{Type: "number", Desc: "record ID"}, + }) +} + +func TestListHandler_All(t *testing.T) { + proc := process.New("tools.doclist", "", 20) + result := ListHandler(proc) + entries, ok := result.([]*goudoc.Entry) + if !ok { + t.Fatalf("expected []*goudoc.Entry, got %T", result) + } + if len(entries) < 2 { + t.Errorf("expected at least 2 entries, got %d", len(entries)) + } +} + +func TestListHandler_Search(t *testing.T) { + proc := process.New("tools.doclist", "Find", 10) + result := ListHandler(proc) + entries, ok := result.([]*goudoc.Entry) + if !ok { + t.Fatalf("expected []*goudoc.Entry, got %T", result) + } + if len(entries) == 0 { + t.Error("expected at least one result for 'Find'") + } + for _, e := range entries { + t.Logf("found: %s - %s", e.Name, e.Desc) + } +} + +func TestInspectHandler(t *testing.T) { + proc := process.New("tools.docinspect", "models.Find") + result := InspectHandler(proc) + if result == nil { + t.Fatal("expected non-nil result for models.Find") + } + entry, ok := result.(*goudoc.Entry) + if !ok { + t.Fatalf("expected *goudoc.Entry, got %T", result) + } + if entry.Name != "models.Find" { + t.Errorf("expected name 'models.Find', got '%s'", entry.Name) + } +} + +func TestInspectHandler_NotFound(t *testing.T) { + proc := process.New("tools.docinspect", "nonexistent.process") + result := InspectHandler(proc) + if result != nil { + t.Error("expected nil for non-existent process") + } +} + +func TestValidateHandler_Valid(t *testing.T) { + proc := process.New("tools.docvalidate", "models.Find") + result := ValidateHandler(proc) + if result == nil { + t.Fatal("expected non-nil result") + } + vr, ok := result.(*goudoc.ValidationResult) + if !ok { + t.Fatalf("expected *goudoc.ValidationResult, got %T", result) + } + if !vr.Valid { + t.Error("expected valid=true for models.Find") + } +} + +func TestValidateHandler_Invalid(t *testing.T) { + proc := process.New("tools.docvalidate", "nonexistent.process") + result := ValidateHandler(proc) + if result == nil { + t.Fatal("expected non-nil result") + } + vr, ok := result.(*goudoc.ValidationResult) + if !ok { + t.Fatalf("expected *goudoc.ValidationResult, got %T", result) + } + if vr.Valid { + t.Error("expected valid=false for non-existent process") + } +} diff --git a/tools/docs/inspect.json b/tools/docs/inspect.json new file mode 100644 index 00000000..9ac4d66a --- /dev/null +++ b/tools/docs/inspect.json @@ -0,0 +1,16 @@ +{ + "name": "docinspect", + "description": "Get detailed documentation for a specific Yao process, including arguments, return type, and methods.", + "process": "tools.docinspect", + "inputSchema": { + "type": "object", + "properties": { + "name": { + "type": "string", + "description": "Process name (e.g. models.user.Find)" + } + }, + "required": ["name"] + }, + "x-process-args": ["$args.name"] +} diff --git a/tools/docs/list.json b/tools/docs/list.json new file mode 100644 index 00000000..e69f99c6 --- /dev/null +++ b/tools/docs/list.json @@ -0,0 +1,20 @@ +{ + "name": "doclist", + "description": "List or search Yao process documentation. Returns matching entries with name, group, and description.", + "process": "tools.doclist", + "inputSchema": { + "type": "object", + "properties": { + "keyword": { + "type": "string", + "description": "Search keyword (empty to list all)" + }, + "limit": { + "type": "integer", + "description": "Max results (default 20)", + "default": 20 + } + } + }, + "x-process-args": ["$args.keyword", "$args.limit"] +} diff --git a/tools/docs/validate.json b/tools/docs/validate.json new file mode 100644 index 00000000..5c50c97e --- /dev/null +++ b/tools/docs/validate.json @@ -0,0 +1,13 @@ +{ + "name": "docvalidate", + "description": "Check if a Yao process has documentation. Returns validation status and suggestions for similar processes if not found.", + "process": "tools.docvalidate", + "inputSchema": { + "type": "object", + "properties": { + "name": { "type": "string", "description": "Process name to validate" } + }, + "required": ["name"] + }, + "x-process-args": ["$args.name"] +} diff --git a/tools/mcps/doc.json b/tools/mcps/doc.json new file mode 100644 index 00000000..a8eef9ad --- /dev/null +++ b/tools/mcps/doc.json @@ -0,0 +1,10 @@ +{ + "name": "yao-doc", + "transport": "process", + "description": "Yao process documentation tools", + "tools": { + "doclist": "tools.doclist", + "docinspect": "tools.docinspect", + "docvalidate": "tools.docvalidate" + } +} diff --git a/tools/mcps/process.json b/tools/mcps/process.json new file mode 100644 index 00000000..bbca5dbd --- /dev/null +++ b/tools/mcps/process.json @@ -0,0 +1,8 @@ +{ + "name": "yao-process", + "transport": "process", + "description": "Yao process execution tool", + "tools": { + "processcall": "tools.processcall" + } +} diff --git a/tools/mcps/web.json b/tools/mcps/web.json new file mode 100644 index 00000000..79e3ce4e --- /dev/null +++ b/tools/mcps/web.json @@ -0,0 +1,9 @@ +{ + "name": "yao-web", + "transport": "process", + "description": "Web search and fetch tools", + "tools": { + "websearch": "tools.websearch", + "webfetch": "tools.webfetch" + } +} diff --git a/tools/proc/proc.go b/tools/proc/proc.go new file mode 100644 index 00000000..8eb2cf0a --- /dev/null +++ b/tools/proc/proc.go @@ -0,0 +1,84 @@ +package proc + +import ( + _ "embed" + "strings" + + "github.com/yaoapp/gou/process" + "github.com/yaoapp/kun/exception" +) + +//go:embed schema.json +var SchemaJSON []byte + +// Allowed process prefixes — blocks system/internal processes. +var allowedPrefixes = []string{ + "models.", + "schemas.", + "stores.", + "flows.", + "scripts.", + "services.", + "tasks.", + "schedules.", + "widgets.", +} + +// Explicitly blocked prefixes for safety. +var blockedPrefixes = []string{ + "yao.sys.", + "yao.env.", + "utils.", + "tools.", +} + +// Handler is the tools.processcall process handler. +// Args[0]: name (string — process name, e.g. "models.user.Find") +// Args[1]: args ([]interface{} — process arguments, optional) +func Handler(p *process.Process) interface{} { + name := p.ArgsString(0) + + if !isAllowedProcess(name) { + exception.New("process %s is not allowed", 403, name).Throw() + } + + var args []interface{} + if len(p.Args) > 1 { + if arr, ok := p.Args[1].([]interface{}); ok { + args = arr + } + } + + target, err := process.Of(name, args...) + if err != nil { + exception.New("process %s not found: %s", 404, name, err.Error()).Throw() + } + if p.Authorized != nil { + target.WithAuthorized(p.Authorized) + } + target.WithSID(p.Sid) + target.WithContext(p.Context) + if err := target.Execute(); err != nil { + exception.New("process %s execution failed: %s", 500, name, err.Error()).Throw() + } + defer target.Release() + return target.Value() +} + +func isAllowedProcess(name string) bool { + lower := strings.ToLower(name) + + for _, prefix := range blockedPrefixes { + if strings.HasPrefix(lower, prefix) { + return false + } + } + + for _, prefix := range allowedPrefixes { + if strings.HasPrefix(lower, prefix) { + return true + } + } + + return false +} diff --git a/tools/proc/proc_test.go b/tools/proc/proc_test.go new file mode 100644 index 00000000..856bc6bf --- /dev/null +++ b/tools/proc/proc_test.go @@ -0,0 +1,39 @@ +package proc + +import ( + "testing" +) + +func TestIsAllowedProcess(t *testing.T) { + allowed := []string{ + "models.user.Find", + "schemas.user.Setting", + "stores.cache.Set", + "flows.login.Run", + "scripts.helper.Format", + "services.user.Create", + "tasks.send.Run", + "schedules.cleanup.Run", + "widgets.chart.Data", + } + for _, name := range allowed { + if !isAllowedProcess(name) { + t.Errorf("expected %q to be allowed", name) + } + } +} + +func TestIsBlockedProcess(t *testing.T) { + blocked := []string{ + "yao.sys.Exec", + "yao.env.Get", + "utils.str.Join", + "tools.websearch", + "unknown.process", + } + for _, name := range blocked { + if isAllowedProcess(name) { + t.Errorf("expected %q to be blocked", name) + } + } +} diff --git a/tools/proc/schema.json b/tools/proc/schema.json new file mode 100644 index 00000000..7bce8685 --- /dev/null +++ b/tools/proc/schema.json @@ -0,0 +1,21 @@ +{ + "name": "processcall", + "description": "Execute a Yao process by name. Supports models, schemas, stores, flows, and scripts.", + "process": "tools.processcall", + "inputSchema": { + "type": "object", + "properties": { + "name": { + "type": "string", + "description": "Process name (e.g. models.user.Find)" + }, + "args": { + "type": "array", + "description": "Process arguments", + "items": {} + } + }, + "required": ["name"] + }, + "x-process-args": ["$args.name", "$args.args"] +} diff --git a/tools/tools.go b/tools/tools.go new file mode 100644 index 00000000..5538696f --- /dev/null +++ b/tools/tools.go @@ -0,0 +1,61 @@ +package tools + +import ( + _ "embed" + "encoding/json" + + "github.com/yaoapp/gou/mcp" + mcpTypes "github.com/yaoapp/gou/mcp/types" + "github.com/yaoapp/gou/process" + "github.com/yaoapp/kun/log" + "github.com/yaoapp/yao/tools/docs" + "github.com/yaoapp/yao/tools/proc" + "github.com/yaoapp/yao/tools/webfetch" + "github.com/yaoapp/yao/tools/websearch" +) + +//go:embed mcps/web.json +var mcpWebDSL []byte + +//go:embed mcps/process.json +var mcpProcessDSL []byte + +//go:embed mcps/doc.json +var mcpDocDSL []byte + +func init() { + process.RegisterGroup("tools", map[string]process.Handler{ + "websearch": websearch.Handler, + "webfetch": webfetch.Handler, + "processcall": proc.Handler, + "doclist": docs.ListHandler, + "docinspect": docs.InspectHandler, + "docvalidate": docs.ValidateHandler, + }) + + registerMCPServer(mcpWebDSL, "yao-web", + websearch.SchemaJSON, webfetch.SchemaJSON) + registerMCPServer(mcpProcessDSL, "yao-process", + proc.SchemaJSON) + registerMCPServer(mcpDocDSL, "yao-doc", + docs.ListSchemaJSON, docs.InspectSchemaJSON, docs.ValidateSchemaJSON) +} + +func registerMCPServer(dsl []byte, id string, schemas ...[]byte) { + mapping := &mcpTypes.MappingData{ + Tools: map[string]*mcpTypes.ToolSchema{}, + Resources: map[string]*mcpTypes.ResourceSchema{}, + Prompts: map[string]*mcpTypes.PromptSchema{}, + } + for _, raw := range schemas { + var s mcpTypes.ToolSchema + if err := json.Unmarshal(raw, &s); err != nil { + log.Error("[tools] failed to parse schema: %s", err.Error()) + continue + } + mapping.Tools[s.Name] = &s + } + if _, err := mcp.LoadClientSourceWithType(string(dsl), id, "", mapping); err != nil { + log.Error("[tools] failed to register MCP server %s: %s", id, err.Error()) + } +} diff --git a/tools/webfetch/cloud.go b/tools/webfetch/cloud.go new file mode 100644 index 00000000..4f371cb5 --- /dev/null +++ b/tools/webfetch/cloud.go @@ -0,0 +1,88 @@ +package webfetch + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "time" +) + +func cloudFetch(cfg *fetchConfig, url, format string) *FetchResponse { + if cfg.APIURL == "" || cfg.APIKey == "" { + return &FetchResponse{ + URL: url, + Content: "cloud service not configured", + Format: format, + } + } + + if format != "markdown" && format != "html" { + format = "markdown" + } + + payload, _ := json.Marshal(map[string]string{ + "url": url, + }) + + endpoint := cfg.APIURL + "/v1/scrape/" + format + req, err := http.NewRequest("POST", endpoint, bytes.NewReader(payload)) + if err != nil { + return &FetchResponse{ + URL: url, + Content: fmt.Sprintf("cloud request build failed: %s", err.Error()), + Format: format, + } + } + req.Header.Set("Authorization", "Bearer "+cfg.APIKey) + req.Header.Set("Content-Type", "application/json") + + client := &http.Client{Timeout: 60 * time.Second} + resp, err := client.Do(req) + if err != nil { + return &FetchResponse{ + URL: url, + Content: fmt.Sprintf("cloud request failed: %s", err.Error()), + Format: format, + } + } + defer resp.Body.Close() + + body, err := io.ReadAll(io.LimitReader(resp.Body, maxBodySize)) + if err != nil { + return &FetchResponse{ + URL: url, + Content: fmt.Sprintf("cloud read body failed: %s", err.Error()), + Format: format, + } + } + + if resp.StatusCode != http.StatusOK { + return &FetchResponse{ + URL: url, + Content: fmt.Sprintf("cloud HTTP %d: %s", resp.StatusCode, truncate(string(body), 200)), + Format: format, + } + } + + var result struct { + Title string `json:"title"` + Content string `json:"content"` + } + if err := json.Unmarshal(body, &result); err != nil { + return &FetchResponse{ + URL: url, + Title: "", + Content: string(body), + Format: format, + } + } + + return &FetchResponse{ + URL: url, + Title: result.Title, + Content: result.Content, + Format: format, + } +} diff --git a/tools/webfetch/convert.go b/tools/webfetch/convert.go new file mode 100644 index 00000000..3bfbcf37 --- /dev/null +++ b/tools/webfetch/convert.go @@ -0,0 +1,218 @@ +package webfetch + +import ( + "html" + "regexp" + "strings" +) + +const maxContentLen = 15000 + +var ( + reScript = regexp.MustCompile(`(?is)`) + reStyle = regexp.MustCompile(`(?is)`) + reComment = regexp.MustCompile(`(?s)`) + reTitle = regexp.MustCompile(`(?is)
]*>(.*?)
`) + reBody = regexp.MustCompile(`(?is)]*>(.*?)`) + reTags = regexp.MustCompile(`<[^>]+>`) + reSpaces = regexp.MustCompile(`\s+`) + + reHeading = regexp.MustCompile(`(?is)]*>(.*?)`) + reCode = regexp.MustCompile(`(?is)
]*>(.*?)`)
+ reBr = regexp.MustCompile(`(?is)]*>(.*?)`) + reStrong = regexp.MustCompile(`(?is)<(?:strong|b)[^>]*>(.*?)(?:strong|b)>`) + reEm = regexp.MustCompile(`(?is)<(?:em|i)[^>]*>(.*?)(?:em|i)>`) +) + +// ExtractTitle extracts the content of the
Paragraph text
+ Link +func main() {}`
+ md := HtmlToMarkdown(html)
+ if !strings.Contains(md, "```") {
+ t.Error("expected fenced code block")
+ }
+ if !strings.Contains(md, "func main()") {
+ t.Error("expected code content")
+ }
+}
+
+func TestCapStr(t *testing.T) {
+ s := "hello world"
+ if capStr(s, 5) != "hello" {
+ t.Errorf("expected 'hello', got '%s'", capStr(s, 5))
+ }
+ if capStr(s, 100) != s {
+ t.Errorf("expected full string, got '%s'", capStr(s, 100))
+ }
+}
diff --git a/tools/webfetch/fetch.go b/tools/webfetch/fetch.go
new file mode 100644
index 00000000..98690aaf
--- /dev/null
+++ b/tools/webfetch/fetch.go
@@ -0,0 +1,232 @@
+package webfetch
+
+import (
+ "bytes"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "path"
+ "strings"
+ "time"
+)
+
+const (
+ directTimeout = 15 * time.Second
+ brightdataTimeout = 90 * time.Second
+ headTimeout = 5 * time.Second
+ maxBodySize = 10 * 1024 * 1024 // 10 MB
+ minDirectBody = 500
+ minMarkdownBody = 100
+ botUserAgent = "Mozilla/5.0 (compatible; YaoBot/1.0; +https://yao.run)"
+ browserUserAgent = "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36"
+)
+
+type fetchResult struct {
+ Body []byte
+ StatusCode int
+ ContentType string
+}
+
+func directFetch(targetURL string, useBot bool) (*fetchResult, error) {
+ client := &http.Client{Timeout: directTimeout}
+ req, err := http.NewRequest(http.MethodGet, targetURL, nil)
+ if err != nil {
+ return nil, fmt.Errorf("build request: %w", err)
+ }
+
+ if useBot {
+ req.Header.Set("User-Agent", botUserAgent)
+ req.Header.Set("Accept", "text/markdown,text/plain,text/html,*/*;q=0.8")
+ } else {
+ req.Header.Set("User-Agent", browserUserAgent)
+ req.Header.Set("Accept", "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8")
+ }
+
+ resp, err := client.Do(req)
+ if err != nil {
+ return nil, fmt.Errorf("http get: %w", err)
+ }
+ defer resp.Body.Close()
+
+ body, err := io.ReadAll(io.LimitReader(resp.Body, maxBodySize))
+ if err != nil {
+ return nil, fmt.Errorf("read body: %w", err)
+ }
+
+ return &fetchResult{
+ Body: body,
+ StatusCode: resp.StatusCode,
+ ContentType: resp.Header.Get("Content-Type"),
+ }, nil
+}
+
+func brightdataFetch(targetURL, apiKey, zone string) ([]byte, error) {
+ payload, _ := json.Marshal(map[string]string{
+ "zone": zone,
+ "url": targetURL,
+ "format": "raw",
+ })
+
+ client := &http.Client{Timeout: brightdataTimeout}
+ req, err := http.NewRequest(http.MethodPost, "https://api.brightdata.com/request", bytes.NewReader(payload))
+ if err != nil {
+ return nil, fmt.Errorf("build brightdata request: %w", err)
+ }
+ req.Header.Set("Authorization", "Bearer "+apiKey)
+ req.Header.Set("Content-Type", "application/json")
+
+ resp, err := client.Do(req)
+ if err != nil {
+ return nil, fmt.Errorf("brightdata request: %w", err)
+ }
+ defer resp.Body.Close()
+
+ body, err := io.ReadAll(io.LimitReader(resp.Body, maxBodySize))
+ if err != nil {
+ return nil, fmt.Errorf("read brightdata response: %w", err)
+ }
+
+ if resp.StatusCode != http.StatusOK {
+ return nil, fmt.Errorf("brightdata HTTP %d: %s", resp.StatusCode, truncate(string(body), 200))
+ }
+ return body, nil
+}
+
+func headCheck(url string) (int, error) {
+ client := &http.Client{Timeout: headTimeout}
+ req, err := http.NewRequest(http.MethodHead, url, nil)
+ if err != nil {
+ return 0, fmt.Errorf("build head request: %w", err)
+ }
+ req.Header.Set("User-Agent", botUserAgent)
+ resp, err := client.Do(req)
+ if err != nil {
+ return 0, fmt.Errorf("head request: %w", err)
+ }
+ resp.Body.Close()
+ return resp.StatusCode, nil
+}
+
+// fetchHTML tries direct GET first, falls back to Brightdata.
+func fetchHTML(cfg *fetchConfig, targetURL string) *FetchResponse {
+ res, err := directFetch(targetURL, false)
+ if err == nil && res.StatusCode == 200 && len(res.Body) >= minDirectBody {
+ htmlStr := string(res.Body)
+ return &FetchResponse{
+ URL: targetURL,
+ Title: ExtractTitle(htmlStr),
+ Content: ExtractContent(htmlStr),
+ Format: "html",
+ }
+ }
+
+ if cfg.BrightdataKey != "" {
+ body, err := brightdataFetch(targetURL, cfg.BrightdataKey, cfg.BrightdataZone)
+ if err == nil {
+ htmlStr := string(body)
+ return &FetchResponse{
+ URL: targetURL,
+ Title: ExtractTitle(htmlStr),
+ Content: ExtractContent(htmlStr),
+ Format: "html",
+ }
+ }
+ }
+
+ return &FetchResponse{
+ URL: targetURL,
+ Content: fmt.Sprintf("Failed to fetch %s", targetURL),
+ Format: "html",
+ }
+}
+
+// fetchMarkdown tries .md probing, then HTML -> markdown conversion.
+func fetchMarkdown(cfg *fetchConfig, targetURL string) *FetchResponse {
+ lower := strings.ToLower(targetURL)
+ ext := strings.ToLower(path.Ext(strings.TrimSuffix(lower, "/")))
+
+ // Already a .md file
+ if ext == ".md" || ext == ".mdx" {
+ res, err := directFetch(targetURL, true)
+ if err == nil && res.StatusCode == 200 && len(res.Body) >= minMarkdownBody {
+ return &FetchResponse{
+ URL: targetURL,
+ Content: string(res.Body),
+ Format: "markdown",
+ }
+ }
+ }
+
+ // Probe for .md version
+ mdURL := buildMdURL(targetURL)
+ if mdURL != "" {
+ code, err := headCheck(mdURL)
+ if err == nil && code == 200 {
+ res, err := directFetch(mdURL, true)
+ if err == nil && res.StatusCode == 200 && len(res.Body) >= minMarkdownBody {
+ return &FetchResponse{
+ URL: targetURL,
+ Content: string(res.Body),
+ Format: "markdown",
+ }
+ }
+ }
+ }
+
+ // Fallback: fetch HTML and convert
+ htmlRes := fetchRawHTML(cfg, targetURL)
+ if htmlRes == nil {
+ return &FetchResponse{
+ URL: targetURL,
+ Content: fmt.Sprintf("Failed to fetch %s", targetURL),
+ Format: "markdown",
+ }
+ }
+
+ htmlStr := string(htmlRes)
+ md := HtmlToMarkdown(htmlStr)
+ title := ExtractTitle(htmlStr)
+ if title != "" {
+ md = "# " + title + "\n\n" + md
+ }
+
+ return &FetchResponse{
+ URL: targetURL,
+ Title: title,
+ Content: md,
+ Format: "markdown",
+ }
+}
+
+// fetchRawHTML fetches HTML with direct -> brightdata fallback.
+func fetchRawHTML(cfg *fetchConfig, targetURL string) []byte {
+ res, err := directFetch(targetURL, false)
+ if err == nil && res.StatusCode == 200 && len(res.Body) >= minDirectBody {
+ return res.Body
+ }
+
+ if cfg.BrightdataKey != "" {
+ body, err := brightdataFetch(targetURL, cfg.BrightdataKey, cfg.BrightdataZone)
+ if err == nil {
+ return body
+ }
+ }
+ return nil
+}
+
+func buildMdURL(u string) string {
+ ext := strings.ToLower(path.Ext(strings.TrimSuffix(u, "/")))
+ if ext == ".md" || ext == ".mdx" {
+ return ""
+ }
+ trimmed := strings.TrimSuffix(u, "/")
+ return trimmed + ".md"
+}
+
+func truncate(s string, n int) string {
+ if len(s) <= n {
+ return s
+ }
+ return s[:n] + "..."
+}
diff --git a/tools/webfetch/schema.json b/tools/webfetch/schema.json
new file mode 100644
index 00000000..712862ab
--- /dev/null
+++ b/tools/webfetch/schema.json
@@ -0,0 +1,19 @@
+{
+ "name": "webfetch",
+ "description": "Fetch a web page and return its content. Supports markdown and HTML output formats.",
+ "process": "tools.webfetch",
+ "inputSchema": {
+ "type": "object",
+ "properties": {
+ "url": { "type": "string", "description": "URL to fetch" },
+ "format": {
+ "type": "string",
+ "description": "Output format: markdown or html (default markdown)",
+ "enum": ["markdown", "html"],
+ "default": "markdown"
+ }
+ },
+ "required": ["url"]
+ },
+ "x-process-args": ["$args.url", "$args.format"]
+}
diff --git a/tools/webfetch/webfetch.go b/tools/webfetch/webfetch.go
new file mode 100644
index 00000000..e7b1d674
--- /dev/null
+++ b/tools/webfetch/webfetch.go
@@ -0,0 +1,119 @@
+package webfetch
+
+import (
+ _ "embed"
+ "os"
+
+ "github.com/yaoapp/gou/process"
+ "github.com/yaoapp/yao/config"
+ "github.com/yaoapp/yao/setting"
+)
+
+//go:embed schema.json
+var SchemaJSON []byte
+
+// FetchResponse is the return type for the webfetch tool.
+type FetchResponse struct {
+ URL string `json:"url"`
+ Title string `json:"title"`
+ Content string `json:"content"`
+ Format string `json:"format"`
+}
+
+type fetchConfig struct {
+ Provider string // "cloud" / "brightdata" / "" (direct)
+ APIKey string
+ APIURL string // cloud mode endpoint
+ BrightdataKey string
+ BrightdataZone string
+}
+
+// Handler is the tools.webfetch process handler.
+// Args[0]: url (string)
+// Args[1]: format (string, default "markdown")
+func Handler(proc *process.Process) interface{} {
+ url := proc.ArgsString(0)
+ format := proc.ArgsString(1, "markdown")
+ userID, teamID := getAuthInfo(proc)
+ cfg := getConfig(userID, teamID)
+
+ switch cfg.Provider {
+ case "cloud":
+ return cloudFetch(cfg, url, format)
+ default:
+ return localFetch(cfg, url, format)
+ }
+}
+
+func getAuthInfo(proc *process.Process) (userID, teamID string) {
+ if proc.Authorized != nil {
+ userID = proc.Authorized.UserID
+ teamID = proc.Authorized.TeamID
+ }
+ return
+}
+
+func getConfig(userID, teamID string) *fetchConfig {
+ cfg := &fetchConfig{}
+
+ if setting.Global != nil {
+ assignment, _ := setting.Global.GetMerged(userID, teamID, "search.tool_assignment")
+ if v, ok := assignment["web_scrape"].(string); ok && v != "" {
+ cfg.Provider = v
+ }
+ }
+
+ switch cfg.Provider {
+ case "cloud":
+ cfg.APIKey, cfg.APIURL = getCloudConfig(userID, teamID)
+ case "brightdata":
+ cfg.BrightdataKey, cfg.BrightdataZone = getBrightdataConfig(userID, teamID)
+ default:
+ cfg.BrightdataKey, cfg.BrightdataZone = getBrightdataConfig(userID, teamID)
+ }
+ return cfg
+}
+
+func getCloudConfig(userID, teamID string) (apiKey, apiURL string) {
+ if setting.Global == nil {
+ return
+ }
+ saved, _ := setting.Global.GetMerged(userID, teamID, "cloud")
+ if v, ok := saved["api_url"].(string); ok {
+ apiURL = v
+ }
+ if v, ok := saved["api_key"].(string); ok {
+ apiKey = config.DecryptValue(v)
+ }
+ return
+}
+
+func getBrightdataConfig(userID, teamID string) (apiKey, zone string) {
+ if setting.Global != nil {
+ saved, _ := setting.Global.GetMerged(userID, teamID, "search.providers.brightdata")
+ if fv, ok := saved["field_values"].(map[string]interface{}); ok {
+ if v, ok := fv["api_key"].(string); ok {
+ apiKey = config.DecryptValue(v)
+ }
+ if v, ok := fv["zone"].(string); ok {
+ zone = v
+ }
+ }
+ }
+ if apiKey == "" {
+ apiKey = os.Getenv("BRIGHTDATA_API_KEY")
+ }
+ if zone == "" {
+ zone = os.Getenv("BRIGHTDATA_ZONE")
+ }
+ return
+}
+
+func localFetch(cfg *fetchConfig, url, format string) *FetchResponse {
+ switch format {
+ case "html":
+ return fetchHTML(cfg, url)
+ default:
+ return fetchMarkdown(cfg, url)
+ }
+}
diff --git a/tools/webfetch/webfetch_test.go b/tools/webfetch/webfetch_test.go
new file mode 100644
index 00000000..88815968
--- /dev/null
+++ b/tools/webfetch/webfetch_test.go
@@ -0,0 +1,74 @@
+package webfetch
+
+import (
+ "testing"
+)
+
+func TestDirectFetch_Success(t *testing.T) {
+ res, err := directFetch("https://example.com", false)
+ if err != nil {
+ t.Fatalf("directFetch failed: %v", err)
+ }
+ if res.StatusCode != 200 {
+ t.Errorf("expected 200, got %d", res.StatusCode)
+ }
+ if len(res.Body) < 100 {
+ t.Error("expected body with at least 100 bytes")
+ }
+}
+
+func TestDirectFetch_Bot(t *testing.T) {
+ res, err := directFetch("https://example.com", true)
+ if err != nil {
+ t.Fatalf("directFetch with bot failed: %v", err)
+ }
+ if res.StatusCode != 200 {
+ t.Errorf("expected 200, got %d", res.StatusCode)
+ }
+}
+
+func TestFetchHTML_Local(t *testing.T) {
+ cfg := &fetchConfig{}
+ resp := fetchHTML(cfg, "https://example.com")
+ if resp == nil {
+ t.Fatal("expected non-nil response")
+ }
+ if resp.Format != "html" {
+ t.Errorf("expected format 'html', got '%s'", resp.Format)
+ }
+ if resp.Content == "" {
+ t.Error("expected non-empty content")
+ }
+}
+
+func TestFetchMarkdown_Local(t *testing.T) {
+ cfg := &fetchConfig{}
+ resp := fetchMarkdown(cfg, "https://example.com")
+ if resp == nil {
+ t.Fatal("expected non-nil response")
+ }
+ if resp.Format != "markdown" {
+ t.Errorf("expected format 'markdown', got '%s'", resp.Format)
+ }
+ if resp.Content == "" {
+ t.Error("expected non-empty content")
+ }
+}
+
+func TestBuildMdURL(t *testing.T) {
+ tests := []struct {
+ input string
+ expected string
+ }{
+ {"https://example.com/docs/page", "https://example.com/docs/page.md"},
+ {"https://example.com/docs/page/", "https://example.com/docs/page.md"},
+ {"https://example.com/docs/page.md", ""},
+ {"https://example.com/docs/page.MDX", ""},
+ }
+ for _, tt := range tests {
+ got := buildMdURL(tt.input)
+ if got != tt.expected {
+ t.Errorf("buildMdURL(%q) = %q, want %q", tt.input, got, tt.expected)
+ }
+ }
+}
diff --git a/tools/websearch/cloud.go b/tools/websearch/cloud.go
new file mode 100644
index 00000000..8ea42cc0
--- /dev/null
+++ b/tools/websearch/cloud.go
@@ -0,0 +1,77 @@
+package websearch
+
+import (
+ "bytes"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "time"
+)
+
+func cloudSearch(cfg *searchConfig, query string, limit int) []SearchResult {
+ if cfg.APIURL == "" || cfg.APIKey == "" {
+ return nil
+ }
+
+ payload, _ := json.Marshal(map[string]interface{}{
+ "query": query,
+ "max_results": limit,
+ })
+
+ tool := cfg.CloudTool
+ if tool == "" {
+ tool = "serper-search"
+ }
+ url := cfg.APIURL + "/v1/search/" + tool
+ req, err := http.NewRequest("POST", url, bytes.NewReader(payload))
+ if err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("cloud request build failed: %s", err.Error())}}
+ }
+ req.Header.Set("Authorization", "Bearer "+cfg.APIKey)
+ req.Header.Set("Content-Type", "application/json")
+
+ client := &http.Client{Timeout: 30 * time.Second}
+ resp, err := client.Do(req)
+ if err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("cloud request failed: %s", err.Error())}}
+ }
+ defer resp.Body.Close()
+
+ body, err := io.ReadAll(resp.Body)
+ if err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("cloud read body failed: %s", err.Error())}}
+ }
+
+ if resp.StatusCode != http.StatusOK {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("cloud HTTP %d: %s", resp.StatusCode, string(body))}}
+ }
+
+ var result struct {
+ Results []struct {
+ Title string `json:"title"`
+ URL string `json:"url"`
+ Snippet string `json:"snippet"`
+ Content string `json:"content"`
+ Score float64 `json:"score"`
+ } `json:"results"`
+ }
+ if err := json.Unmarshal(body, &result); err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("cloud parse failed: %s", err.Error())}}
+ }
+
+ out := make([]SearchResult, 0, len(result.Results))
+ for _, r := range result.Results {
+ text := r.Snippet
+ if text == "" {
+ text = r.Content
+ }
+ out = append(out, SearchResult{
+ Title: r.Title,
+ URL: r.URL,
+ Content: text,
+ Score: r.Score,
+ })
+ }
+ return out
+}
diff --git a/tools/websearch/schema.json b/tools/websearch/schema.json
new file mode 100644
index 00000000..7c2e5d33
--- /dev/null
+++ b/tools/websearch/schema.json
@@ -0,0 +1,18 @@
+{
+ "name": "websearch",
+ "description": "Search the web for real-time information. Returns structured results with title, URL, and content snippet.",
+ "process": "tools.websearch",
+ "inputSchema": {
+ "type": "object",
+ "properties": {
+ "query": { "type": "string", "description": "Search query" },
+ "limit": {
+ "type": "integer",
+ "description": "Max results (default 10)",
+ "default": 10
+ }
+ },
+ "required": ["query"]
+ },
+ "x-process-args": ["$args.query", "$args.limit"]
+}
diff --git a/tools/websearch/serper.go b/tools/websearch/serper.go
new file mode 100644
index 00000000..f3c42ad3
--- /dev/null
+++ b/tools/websearch/serper.go
@@ -0,0 +1,65 @@
+package websearch
+
+import (
+ "bytes"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "time"
+)
+
+func serperSearch(apiKey, query string, limit int) []SearchResult {
+ if apiKey == "" {
+ return nil
+ }
+
+ payload, _ := json.Marshal(map[string]interface{}{
+ "q": query,
+ "num": limit,
+ })
+
+ req, err := http.NewRequest("POST", "https://google.serper.dev/search", bytes.NewReader(payload))
+ if err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("serper request build failed: %s", err.Error())}}
+ }
+ req.Header.Set("X-API-KEY", apiKey)
+ req.Header.Set("Content-Type", "application/json")
+
+ client := &http.Client{Timeout: 30 * time.Second}
+ resp, err := client.Do(req)
+ if err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("serper request failed: %s", err.Error())}}
+ }
+ defer resp.Body.Close()
+
+ body, err := io.ReadAll(resp.Body)
+ if err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("serper read body failed: %s", err.Error())}}
+ }
+
+ if resp.StatusCode != http.StatusOK {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("serper HTTP %d: %s", resp.StatusCode, string(body))}}
+ }
+
+ var result struct {
+ Organic []struct {
+ Title string `json:"title"`
+ Link string `json:"link"`
+ Snippet string `json:"snippet"`
+ } `json:"organic"`
+ }
+ if err := json.Unmarshal(body, &result); err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("serper parse failed: %s", err.Error())}}
+ }
+
+ out := make([]SearchResult, 0, len(result.Organic))
+ for _, r := range result.Organic {
+ out = append(out, SearchResult{
+ Title: r.Title,
+ URL: r.Link,
+ Content: r.Snippet,
+ })
+ }
+ return out
+}
diff --git a/tools/websearch/tavily.go b/tools/websearch/tavily.go
new file mode 100644
index 00000000..33122be1
--- /dev/null
+++ b/tools/websearch/tavily.go
@@ -0,0 +1,62 @@
+package websearch
+
+import (
+ "bytes"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net/http"
+ "time"
+)
+
+func tavilySearch(apiKey, query string, limit int) []SearchResult {
+ if apiKey == "" {
+ return nil
+ }
+
+ payload, _ := json.Marshal(map[string]interface{}{
+ "api_key": apiKey,
+ "query": query,
+ "max_results": limit,
+ "include_raw_content": false,
+ })
+
+ client := &http.Client{Timeout: 30 * time.Second}
+ resp, err := client.Post("https://api.tavily.com/search", "application/json", bytes.NewReader(payload))
+ if err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("tavily request failed: %s", err.Error())}}
+ }
+ defer resp.Body.Close()
+
+ body, err := io.ReadAll(resp.Body)
+ if err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("tavily read body failed: %s", err.Error())}}
+ }
+
+ if resp.StatusCode != http.StatusOK {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("tavily HTTP %d: %s", resp.StatusCode, string(body))}}
+ }
+
+ var result struct {
+ Results []struct {
+ Title string `json:"title"`
+ URL string `json:"url"`
+ Content string `json:"content"`
+ Score float64 `json:"score"`
+ } `json:"results"`
+ }
+ if err := json.Unmarshal(body, &result); err != nil {
+ return []SearchResult{{Title: "Error", Content: fmt.Sprintf("tavily parse failed: %s", err.Error())}}
+ }
+
+ out := make([]SearchResult, 0, len(result.Results))
+ for _, r := range result.Results {
+ out = append(out, SearchResult{
+ Title: r.Title,
+ URL: r.URL,
+ Content: r.Content,
+ Score: r.Score,
+ })
+ }
+ return out
+}
diff --git a/tools/websearch/websearch.go b/tools/websearch/websearch.go
new file mode 100644
index 00000000..62dfcefd
--- /dev/null
+++ b/tools/websearch/websearch.go
@@ -0,0 +1,117 @@
+package websearch
+
+import (
+ _ "embed"
+ "os"
+
+ "github.com/yaoapp/gou/process"
+ "github.com/yaoapp/yao/config"
+ "github.com/yaoapp/yao/setting"
+)
+
+//go:embed schema.json
+var SchemaJSON []byte
+
+// SearchResult is the unified return type for all search providers.
+type SearchResult struct {
+ Title string `json:"title"`
+ URL string `json:"url"`
+ Content string `json:"content"`
+ Score float64 `json:"score,omitempty"`
+}
+
+type searchConfig struct {
+ Provider string // "tavily" / "serper" / "cloud"
+ APIKey string
+ APIURL string // cloud mode endpoint
+ CloudTool string // cloud search tool name, e.g. "serper-search", "tavily-search"
+}
+
+// Handler is the tools.websearch process handler.
+// Args[0]: query (string)
+// Args[1]: limit (int, default 10)
+func Handler(proc *process.Process) interface{} {
+ query := proc.ArgsString(0)
+ limit := proc.ArgsInt(1, 10)
+ userID, teamID := getAuthInfo(proc)
+ return Search(query, limit, userID, teamID)
+}
+
+// Search executes a web search using the configured provider.
+// Reads provider/key from Settings (with ENV fallback).
+func Search(query string, limit int, userID, teamID string) []SearchResult {
+ cfg := getConfig(userID, teamID)
+ switch cfg.Provider {
+ case "cloud":
+ return cloudSearch(cfg, query, limit)
+ case "serper":
+ return serperSearch(cfg.APIKey, query, limit)
+ default:
+ return tavilySearch(cfg.APIKey, query, limit)
+ }
+}
+
+func getAuthInfo(proc *process.Process) (userID, teamID string) {
+ if proc.Authorized != nil {
+ userID = proc.Authorized.UserID
+ teamID = proc.Authorized.TeamID
+ }
+ return
+}
+
+func getConfig(userID, teamID string) *searchConfig {
+ cfg := &searchConfig{Provider: "tavily"}
+
+ if setting.Global != nil {
+ assignment, _ := setting.Global.GetMerged(userID, teamID, "search.tool_assignment")
+ if v, ok := assignment["web_search"].(string); ok && v != "" {
+ cfg.Provider = v
+ }
+ }
+
+ switch cfg.Provider {
+ case "cloud":
+ cfg.APIKey, cfg.APIURL, cfg.CloudTool = getCloudConfig(userID, teamID)
+ case "tavily":
+ cfg.APIKey = getProviderKey(userID, teamID, "tavily")
+ if cfg.APIKey == "" {
+ cfg.APIKey = os.Getenv("TAVILY_API_KEY")
+ }
+ case "serper":
+ cfg.APIKey = getProviderKey(userID, teamID, "serper")
+ if cfg.APIKey == "" {
+ cfg.APIKey = os.Getenv("SERPER_API_KEY")
+ }
+ }
+ return cfg
+}
+
+func getCloudConfig(userID, teamID string) (apiKey, apiURL, cloudTool string) {
+ if setting.Global == nil {
+ return
+ }
+ saved, _ := setting.Global.GetMerged(userID, teamID, "cloud")
+ if v, ok := saved["api_url"].(string); ok {
+ apiURL = v
+ }
+ if v, ok := saved["api_key"].(string); ok {
+ apiKey = config.DecryptValue(v)
+ }
+ if v, ok := saved["search_tool"].(string); ok && v != "" {
+ cloudTool = v
+ }
+ return
+}
+
+func getProviderKey(userID, teamID, presetKey string) string {
+ if setting.Global == nil {
+ return ""
+ }
+ saved, _ := setting.Global.GetMerged(userID, teamID, "search.providers."+presetKey)
+ if fv, ok := saved["field_values"].(map[string]interface{}); ok {
+ if v, ok := fv["api_key"].(string); ok {
+ return config.DecryptValue(v)
+ }
+ }
+ return ""
+}
diff --git a/tools/websearch/websearch_test.go b/tools/websearch/websearch_test.go
new file mode 100644
index 00000000..431de135
--- /dev/null
+++ b/tools/websearch/websearch_test.go
@@ -0,0 +1,69 @@
+package websearch
+
+import (
+ "os"
+ "testing"
+)
+
+func TestTavilySearch(t *testing.T) {
+ key := os.Getenv("TAVILY_API_KEY")
+ if key == "" {
+ t.Skip("TAVILY_API_KEY not set")
+ }
+
+ results := tavilySearch(key, "Yao application engine", 3)
+ if len(results) == 0 {
+ t.Fatal("expected at least one result from tavily")
+ }
+ for _, r := range results {
+ if r.Title == "Error" {
+ t.Fatalf("tavily returned error: %s", r.Content)
+ }
+ if r.URL == "" {
+ t.Error("expected non-empty URL")
+ }
+ }
+ t.Logf("got %d results", len(results))
+}
+
+func TestSerperSearch(t *testing.T) {
+ key := os.Getenv("SERPER_API_KEY")
+ if key == "" {
+ t.Skip("SERPER_API_KEY not set")
+ }
+
+ results := serperSearch(key, "Yao application engine", 3)
+ if len(results) == 0 {
+ t.Fatal("expected at least one result from serper")
+ }
+ for _, r := range results {
+ if r.Title == "Error" {
+ t.Fatalf("serper returned error: %s", r.Content)
+ }
+ if r.URL == "" {
+ t.Error("expected non-empty URL")
+ }
+ }
+ t.Logf("got %d results", len(results))
+}
+
+func TestTavilySearch_NoKey(t *testing.T) {
+ results := tavilySearch("", "test", 5)
+ if results != nil {
+ t.Error("expected nil results with empty key")
+ }
+}
+
+func TestSerperSearch_NoKey(t *testing.T) {
+ results := serperSearch("", "test", 5)
+ if results != nil {
+ t.Error("expected nil results with empty key")
+ }
+}
+
+func TestGetConfig_Defaults(t *testing.T) {
+ cfg := getConfig("", "")
+ if cfg.Provider != "tavily" {
+ t.Errorf("expected default provider 'tavily', got '%s'", cfg.Provider)
+ }
+}