Enhance Assistant and Search Modules with Contextual Logging and Refactorings

- Added logging for hook start and completion in the Assistant's Stream method to improve traceability during execution.
- Refactored the AgentGetterFunc to utilize the caller package, enhancing modularity and reducing circular dependencies.
- Updated the CallAgent function to check for the initialized AgentGetterFunc from the caller package, ensuring proper agent loading.
- Enhanced the Search handler to support an optional context parameter, improving flexibility for agent mode operations.
- Refined the agentSearch function to delegate search requests to a new AgentProvider, streamlining the search process.
- Updated DESIGN.md to reflect changes in search modes and the integration of the caller package, ensuring comprehensive documentation.
This commit is contained in:
Max 2025-12-13 14:05:26 +08:00
parent e4d1701dab
commit b3cf5a09d9
10 changed files with 1017 additions and 44 deletions

View file

@ -367,6 +367,8 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
var nextResponse *context.NextHookResponse = nil
if ast.HookScript != nil {
ctx.Logger.HookStart("Next")
// Begin step tracking for hook_next
ast.BeginStep(ctx, context.StepTypeHookNext, map[string]interface{}{
"messages": fullMessages,
@ -393,6 +395,8 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
"response": nextResponse,
})
ctx.Logger.HookComplete("Next")
// Process Next hook response
finalResponse, err = ast.processNextResponse(&NextProcessContext{
Context: ctx,

View file

@ -5,7 +5,7 @@ import (
"path"
"github.com/yaoapp/gou/fs"
"github.com/yaoapp/yao/agent/content"
"github.com/yaoapp/yao/agent/caller"
agentContext "github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/agent/i18n"
searchTypes "github.com/yaoapp/yao/agent/search/types"
@ -14,8 +14,8 @@ import (
)
func init() {
// Initialize AgentGetterFunc to allow content package to call agents
content.AgentGetterFunc = func(agentID string) (content.AgentCaller, error) {
// Initialize AgentGetterFunc to allow content and search packages to call agents
caller.AgentGetterFunc = func(agentID string) (caller.AgentCaller, error) {
ast, err := Get(agentID)
if err != nil {
return nil, err

17
agent/caller/caller.go Normal file
View file

@ -0,0 +1,17 @@
// Package caller provides a shared interface for calling agents
// This package is used by both content and search packages to avoid circular dependencies
package caller
import (
agentContext "github.com/yaoapp/yao/agent/context"
)
// AgentCaller interface for calling agents (to avoid circular dependency)
// Used by content handlers (vision, audio, etc.) and search handlers (agent mode)
type AgentCaller interface {
Stream(ctx *agentContext.Context, messages []agentContext.Message, options ...*agentContext.Options) (interface{}, error)
}
// AgentGetterFunc is a function type that gets an agent by ID
// This should be set by the assistant package during initialization
var AgentGetterFunc func(agentID string) (AgentCaller, error)

View file

@ -9,29 +9,22 @@ import (
jsoniter "github.com/json-iterator/go"
"github.com/yaoapp/gou/mcp"
"github.com/yaoapp/kun/log"
"github.com/yaoapp/yao/agent/caller"
agentContext "github.com/yaoapp/yao/agent/context"
)
// AgentCaller interface for calling agents (to avoid circular dependency)
type AgentCaller interface {
Stream(ctx *agentContext.Context, messages []agentContext.Message, options ...*agentContext.Options) (interface{}, error)
}
// AgentGetterFunc is a function type that gets an agent by ID
var AgentGetterFunc func(agentID string) (AgentCaller, error)
// fileInfoMutex protects concurrent access to files_info list in Space
var fileInfoMutex sync.Mutex
// CallAgent calls an agent to process content (vision, audio, etc.)
// This is a generic function that can be used by any handler
func CallAgent(ctx *agentContext.Context, agentID string, message agentContext.Message) (string, error) {
if AgentGetterFunc == nil {
if caller.AgentGetterFunc == nil {
return "", fmt.Errorf("AgentGetterFunc not initialized")
}
// Load the agent by ID using the injected function
agent, err := AgentGetterFunc(agentID)
agent, err := caller.AgentGetterFunc(agentID)
if err != nil {
return "", fmt.Errorf("failed to load agent %s: %w", agentID, err)
}

View file

@ -1110,7 +1110,7 @@ Tool format: `"builtin"`, `"<assistant-id>"` (Agent), `"mcp:<server>.<tool>"` (M
| Mode | Example | Description |
| --------- | ---------------------------- | -------------------------------------------------------------------------- |
| `builtin` | `"builtin"` | Use built-in providers (Tavily, Serper) |
| `builtin` | `"builtin"` | Use built-in providers (Tavily, Serper, SerpAPI) |
| Agent | `"workers.search.web"` | AI-powered search: understand intent → optimize query → search → summarize |
| MCP | `"mcp:my-server.web_search"` | External search tool via MCP protocol |
@ -1707,7 +1707,7 @@ Web search supports three modes via `uses.web`:
| Mode | Value | Description |
| ------- | ---------------------------- | ------------------------------------------- |
| Builtin | `"builtin"` | Direct API calls to Tavily/Serper |
| Builtin | `"builtin"` | Direct API calls to Tavily/Serper/SerpAPI |
| Agent | `"workers.search.web"` | AI-powered search with intent understanding |
| MCP | `"mcp:my-server.web_search"` | External search tool via MCP |

View file

@ -0,0 +1,232 @@
package web
import (
"encoding/json"
"fmt"
"time"
"github.com/yaoapp/yao/agent/caller"
agentContext "github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/agent/search/types"
)
// AgentProvider implements web search using another agent (AI Search)
type AgentProvider struct {
agentID string // Agent/Assistant ID (e.g., "workers.search.web")
}
// NewAgentProvider creates a new Agent provider
func NewAgentProvider(agentID string) *AgentProvider {
return &AgentProvider{
agentID: agentID,
}
}
// Search executes web search via agent delegation
// The agent can understand intent, generate optimized queries, and return structured results
func (p *AgentProvider) Search(ctx *agentContext.Context, req *types.Request) (*types.Result, error) {
startTime := time.Now()
// Check if context is provided
if ctx == nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: "Agent mode requires context",
}, nil
}
// Check if AgentGetterFunc is initialized
if caller.AgentGetterFunc == nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: "AgentGetterFunc not initialized",
}, nil
}
// Get the agent
agent, err := caller.AgentGetterFunc(p.agentID)
if err != nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: fmt.Sprintf("Agent '%s' not found: %v", p.agentID, err),
}, nil
}
// Build message for the agent
// Include search parameters in the message content
searchParams := map[string]interface{}{
"query": req.Query,
"type": "web",
"source": string(req.Source),
}
if req.Limit > 0 {
searchParams["limit"] = req.Limit
}
if len(req.Sites) > 0 {
searchParams["sites"] = req.Sites
}
if req.TimeRange != "" {
searchParams["time_range"] = req.TimeRange
}
// Convert to JSON for the message
paramsJSON, err := json.Marshal(searchParams)
if err != nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: fmt.Sprintf("Failed to serialize search params: %v", err),
}, nil
}
// Create message for the agent
message := agentContext.Message{
Role: "user",
Content: string(paramsJSON),
}
// Call the agent with skip options (no history, no output)
opts := &agentContext.Options{
Skip: &agentContext.Skip{
History: true,
Output: true,
},
}
response, err := agent.Stream(ctx, []agentContext.Message{message}, opts)
if err != nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: fmt.Sprintf("Agent call failed: %v", err),
}, nil
}
// Parse the agent response
items, total, parseErr := p.parseAgentResponse(response, req.Source)
if parseErr != "" {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: parseErr,
}, nil
}
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: items,
Total: total,
Duration: time.Since(startTime).Milliseconds(),
}, nil
}
// parseAgentResponse parses the agent response into search result items
// The agent should return a JSON structure with search results
func (p *AgentProvider) parseAgentResponse(response interface{}, source types.SourceType) ([]*types.ResultItem, int, string) {
if response == nil {
return nil, 0, "Agent returned nil response"
}
// Try to extract data from response
var data map[string]interface{}
// Handle different response types
switch v := response.(type) {
case map[string]interface{}:
data = v
case string:
// Try to parse as JSON
if err := json.Unmarshal([]byte(v), &data); err != nil {
return nil, 0, fmt.Sprintf("Failed to parse agent response as JSON: %v", err)
}
default:
// Try to marshal and unmarshal
jsonBytes, err := json.Marshal(response)
if err != nil {
return nil, 0, fmt.Sprintf("Failed to serialize agent response: %v", err)
}
if err := json.Unmarshal(jsonBytes, &data); err != nil {
return nil, 0, fmt.Sprintf("Failed to parse agent response: %v", err)
}
}
// Check for "next" field (custom hook data)
if next, hasNext := data["next"]; hasNext && next != nil {
if nextMap, ok := next.(map[string]interface{}); ok {
data = nextMap
} else if nextStr, ok := next.(string); ok {
// Try to parse as JSON
if err := json.Unmarshal([]byte(nextStr), &data); err != nil {
return nil, 0, fmt.Sprintf("Failed to parse next hook data: %v", err)
}
}
}
// Extract items from data
items := []*types.ResultItem{}
total := 0
if itemsData, ok := data["items"].([]interface{}); ok {
for _, itemData := range itemsData {
if item, ok := itemData.(map[string]interface{}); ok {
resultItem := &types.ResultItem{
Type: types.SearchTypeWeb,
Source: source,
}
if title, ok := item["title"].(string); ok {
resultItem.Title = title
}
if content, ok := item["content"].(string); ok {
resultItem.Content = content
}
if url, ok := item["url"].(string); ok {
resultItem.URL = url
}
if score, ok := item["score"].(float64); ok {
resultItem.Score = score
}
items = append(items, resultItem)
}
}
}
if totalVal, ok := data["total"].(float64); ok {
total = int(totalVal)
} else {
total = len(items)
}
return items, total, ""
}

View file

@ -0,0 +1,284 @@
package web_test
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/yaoapp/yao/agent/assistant"
agentContext "github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/agent/search/handlers/web"
"github.com/yaoapp/yao/agent/search/types"
"github.com/yaoapp/yao/agent/testutils"
oauthTypes "github.com/yaoapp/yao/openapi/oauth/types"
)
// TestAgentProviderWithAssistantConfig tests AgentProvider using web-agent-caller assistant config
func TestAgentProviderWithAssistantConfig(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Load the web-agent-caller test assistant to get its config
ast, err := assistant.LoadPath("/assistants/tests/web-agent-caller")
require.NoError(t, err)
require.NotNil(t, ast)
require.NotNil(t, ast.Uses)
// Verify assistant config
assert.Equal(t, "tests.web-agent-caller", ast.ID)
assert.Equal(t, "tests.web-agent", ast.Uses.Web)
// Create AgentProvider from uses.web
provider := web.NewAgentProvider(ast.Uses.Web)
require.NotNil(t, provider)
// Create a mock context
ctx := createTestContext(t)
// Execute search
req := &types.Request{
Query: "Yao App Engine",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
Limit: 5,
}
result, err := provider.Search(ctx, req)
require.NoError(t, err)
require.NotNil(t, result)
// Verify result structure
assert.Equal(t, types.SearchTypeWeb, result.Type)
assert.Equal(t, "Yao App Engine", result.Query)
assert.Equal(t, types.SourceAuto, result.Source)
// Agent should return mock results from Next hook
if result.Error == "" {
assert.Greater(t, result.Total, 0)
assert.NotEmpty(t, result.Items)
assert.Greater(t, result.Duration, int64(0))
// Verify result item structure
for _, item := range result.Items {
assert.Equal(t, types.SearchTypeWeb, item.Type)
assert.Equal(t, types.SourceAuto, item.Source)
assert.NotEmpty(t, item.Title)
assert.NotEmpty(t, item.URL)
}
t.Logf("Agent search returned %d results in %dms", result.Total, result.Duration)
} else {
t.Logf("Agent search returned error: %s", result.Error)
}
}
// TestAgentProviderWithSiteRestriction tests AgentProvider with domain restriction
func TestAgentProviderWithSiteRestriction(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create AgentProvider
provider := web.NewAgentProvider("tests.web-agent")
// Create a mock context
ctx := createTestContext(t)
// Execute search with site restriction
req := &types.Request{
Query: "documentation",
Type: types.SearchTypeWeb,
Source: types.SourceHook,
Sites: []string{"github.com"},
Limit: 3,
}
result, err := provider.Search(ctx, req)
require.NoError(t, err)
require.NotNil(t, result)
assert.Equal(t, types.SearchTypeWeb, result.Type)
assert.Equal(t, types.SourceHook, result.Source)
if result.Error == "" {
// All results should be from github.com (mock data respects sites)
for _, item := range result.Items {
assert.Contains(t, item.URL, "github.com", "Result URL should be from github.com")
}
t.Logf("Site-restricted agent search returned %d results", result.Total)
} else {
t.Logf("Agent search returned error: %s", result.Error)
}
}
// TestAgentProviderWithTimeRange tests AgentProvider with time range filter
func TestAgentProviderWithTimeRange(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create AgentProvider
provider := web.NewAgentProvider("tests.web-agent")
// Create a mock context
ctx := createTestContext(t)
// Execute search with time range
req := &types.Request{
Query: "artificial intelligence news",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
TimeRange: "week",
Limit: 5,
}
result, err := provider.Search(ctx, req)
require.NoError(t, err)
require.NotNil(t, result)
if result.Error == "" {
t.Logf("Time-ranged agent search (last week) returned %d results in %dms", result.Total, result.Duration)
} else {
t.Logf("Agent search returned error: %s", result.Error)
}
}
// TestAgentProviderNotFound tests AgentProvider when agent is not found
func TestAgentProviderNotFound(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create AgentProvider with non-existent agent
provider := web.NewAgentProvider("nonexistent.agent")
// Create a mock context
ctx := createTestContext(t)
req := &types.Request{
Query: "test query",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
}
result, err := provider.Search(ctx, req)
// Should not return error, but result should have error message
require.NoError(t, err)
require.NotNil(t, result)
assert.NotEmpty(t, result.Error)
assert.Contains(t, result.Error, "not found")
}
// TestAgentProviderWithoutContext tests AgentProvider without context
func TestAgentProviderWithoutContext(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create AgentProvider
provider := web.NewAgentProvider("tests.web-agent")
req := &types.Request{
Query: "test query",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
}
// Call without context (nil)
result, err := provider.Search(nil, req)
// Should still work - agent provider handles nil context
require.NoError(t, err)
require.NotNil(t, result)
// May have error if context is required for agent call
t.Logf("Agent search without context: error=%s, total=%d", result.Error, result.Total)
}
// TestWebHandlerAgentMode tests the web handler in agent mode
func TestWebHandlerAgentMode(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create handler with agent mode
handler := web.NewHandler("tests.web-agent", nil)
require.NotNil(t, handler)
// Verify type
assert.Equal(t, types.SearchTypeWeb, handler.Type())
// Create a mock context
ctx := createTestContext(t)
// Execute search with context
req := &types.Request{
Query: "Yao framework",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
Limit: 5,
}
result, err := handler.SearchWithContext(ctx, req)
require.NoError(t, err)
require.NotNil(t, result)
assert.Equal(t, types.SearchTypeWeb, result.Type)
assert.Equal(t, "Yao framework", result.Query)
if result.Error == "" {
t.Logf("Handler agent mode returned %d results", result.Total)
} else {
t.Logf("Handler agent mode returned error: %s", result.Error)
}
}
// TestWebHandlerAgentModeWithoutContext tests the web handler in agent mode without context
func TestWebHandlerAgentModeWithoutContext(t *testing.T) {
// Create handler with agent mode
handler := web.NewHandler("tests.web-agent", nil)
require.NotNil(t, handler)
req := &types.Request{
Query: "test",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
}
// Call Search() without context (uses SearchWithContext with nil)
result, err := handler.Search(req)
require.NoError(t, err)
require.NotNil(t, result)
assert.NotEmpty(t, result.Error)
assert.Contains(t, result.Error, "requires context")
}
// createTestContext creates a test context for agent calls
func createTestContext(t *testing.T) *agentContext.Context {
authorized := &oauthTypes.AuthorizedInfo{
UserID: "test-user",
TenantID: "test-tenant",
}
ctx := agentContext.New(nil, authorized, "test-chat-id")
ctx.AssistantID = "tests.web-agent-caller"
return ctx
}

View file

@ -4,6 +4,7 @@ import (
"fmt"
"strings"
agentContext "github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/agent/search/types"
)
@ -24,7 +25,14 @@ func (h *Handler) Type() types.SearchType {
}
// Search executes web search based on uses.web mode
// ctx is optional and only required for agent mode
func (h *Handler) Search(req *types.Request) (*types.Result, error) {
return h.SearchWithContext(nil, req)
}
// SearchWithContext executes web search with optional agent context
// ctx is required for agent mode, optional for builtin and MCP modes
func (h *Handler) SearchWithContext(ctx *agentContext.Context, req *types.Request) (*types.Result, error) {
switch {
case h.usesWeb == "builtin" || h.usesWeb == "":
return h.builtinSearch(req)
@ -32,7 +40,17 @@ func (h *Handler) Search(req *types.Request) (*types.Result, error) {
return h.mcpSearch(req)
default:
// Agent mode: delegate to assistant for AI-powered search
return h.agentSearch(req)
if ctx == nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Error: "Agent mode requires context",
}, nil
}
return h.agentSearch(ctx, req)
}
}
@ -66,46 +84,27 @@ func (h *Handler) builtinSearch(req *types.Request) (*types.Result, error) {
}
// agentSearch delegates to an assistant for AI-powered search
func (h *Handler) agentSearch(req *types.Request) (*types.Result, error) {
// TODO: Implement agent mode
// 1. Call assistant with search request
// 2. Assistant understands intent, generates optimized queries
// 3. Assistant executes searches (may call builtin internally)
// 4. Assistant analyzes and returns structured results
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Error: "Agent mode not yet implemented",
}, nil
func (h *Handler) agentSearch(ctx *agentContext.Context, req *types.Request) (*types.Result, error) {
provider := NewAgentProvider(h.usesWeb)
return provider.Search(ctx, req)
}
// mcpSearch calls external MCP tool
func (h *Handler) mcpSearch(req *types.Request) (*types.Result, error) {
// TODO: Implement MCP mode
// Parse "mcp:server.tool"
mcpRef := strings.TrimPrefix(h.usesWeb, "mcp:")
parts := strings.SplitN(mcpRef, ".", 2)
if len(parts) != 2 {
provider, err := NewMCPProvider(mcpRef)
if err != nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Error: fmt.Sprintf("Invalid MCP format, expected 'mcp:server.tool', got '%s'", h.usesWeb),
Error: fmt.Sprintf("Invalid MCP format: %v", err),
}, nil
}
// serverID, toolName := parts[0], parts[1]
// Call MCP tool
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Error: "MCP mode not yet implemented",
}, nil
return provider.Search(req)
}

View file

@ -0,0 +1,203 @@
package web
import (
"context"
"encoding/json"
"fmt"
"strings"
"time"
"github.com/yaoapp/gou/mcp"
gouMCPTypes "github.com/yaoapp/gou/mcp/types"
"github.com/yaoapp/yao/agent/search/types"
)
// MCPProvider implements web search using MCP tool
type MCPProvider struct {
serverID string // MCP server ID (e.g., "search")
toolName string // MCP tool name (e.g., "web_search")
}
// NewMCPProvider creates a new MCP provider from "mcp:server.tool" format
func NewMCPProvider(mcpRef string) (*MCPProvider, error) {
// Parse "server.tool" format
parts := strings.SplitN(mcpRef, ".", 2)
if len(parts) != 2 {
return nil, fmt.Errorf("invalid MCP format, expected 'server.tool', got '%s'", mcpRef)
}
return &MCPProvider{
serverID: parts[0],
toolName: parts[1],
}, nil
}
// Search executes web search via MCP tool
func (p *MCPProvider) Search(req *types.Request) (*types.Result, error) {
startTime := time.Now()
// Select MCP client
client, err := mcp.Select(p.serverID)
if err != nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: fmt.Sprintf("MCP client '%s' not found: %v", p.serverID, err),
}, nil
}
// Build MCP tool arguments
args := map[string]interface{}{
"query": req.Query,
}
if req.Limit > 0 {
args["limit"] = req.Limit
}
if len(req.Sites) > 0 {
args["sites"] = req.Sites
}
if req.TimeRange != "" {
args["time_range"] = req.TimeRange
}
// Call MCP tool
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
result, err := client.CallTool(ctx, p.toolName, args)
if err != nil {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: fmt.Sprintf("MCP tool call failed: %v", err),
}, nil
}
// Parse MCP result
items, total, parseErr := p.parseResult(result, req.Source)
if parseErr != "" {
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: []*types.ResultItem{},
Total: 0,
Duration: time.Since(startTime).Milliseconds(),
Error: parseErr,
}, nil
}
return &types.Result{
Type: types.SearchTypeWeb,
Query: req.Query,
Source: req.Source,
Items: items,
Total: total,
Duration: time.Since(startTime).Milliseconds(),
}, nil
}
// parseResult parses MCP tool result into search result items
func (p *MCPProvider) parseResult(result *gouMCPTypes.CallToolResponse, source types.SourceType) ([]*types.ResultItem, int, string) {
if result == nil {
return nil, 0, "MCP returned nil result"
}
// Check for errors in result
if result.IsError {
errMsg := "MCP tool returned error"
if len(result.Content) > 0 && result.Content[0].Text != "" {
errMsg = result.Content[0].Text
}
return nil, 0, errMsg
}
// Parse content - expect JSON data
if len(result.Content) == 0 {
return []*types.ResultItem{}, 0, ""
}
// Try to extract data from content
var data map[string]interface{}
for _, content := range result.Content {
// Check text content type
if content.Type == gouMCPTypes.ToolContentTypeText && content.Text != "" {
// Try to parse as JSON
if parsed, ok := parseJSON(content.Text); ok {
data = parsed
break
}
}
}
if data == nil {
return []*types.ResultItem{}, 0, ""
}
// Extract items from data
items := []*types.ResultItem{}
total := 0
if itemsData, ok := data["items"].([]interface{}); ok {
for _, itemData := range itemsData {
if item, ok := itemData.(map[string]interface{}); ok {
resultItem := &types.ResultItem{
Type: types.SearchTypeWeb,
Source: source,
}
if title, ok := item["title"].(string); ok {
resultItem.Title = title
}
if content, ok := item["content"].(string); ok {
resultItem.Content = content
}
if url, ok := item["url"].(string); ok {
resultItem.URL = url
}
if score, ok := item["score"].(float64); ok {
resultItem.Score = score
}
items = append(items, resultItem)
}
}
}
if totalVal, ok := data["total"].(float64); ok {
total = int(totalVal)
} else {
total = len(items)
}
return items, total, ""
}
// parseJSON attempts to parse a string as JSON
func parseJSON(s string) (map[string]interface{}, bool) {
// Simple JSON detection - if it starts with { and ends with }
s = strings.TrimSpace(s)
if !strings.HasPrefix(s, "{") || !strings.HasSuffix(s, "}") {
return nil, false
}
// Use encoding/json for parsing
var result map[string]interface{}
if err := json.Unmarshal([]byte(s), &result); err != nil {
return nil, false
}
return result, true
}

View file

@ -0,0 +1,241 @@
package web_test
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/yaoapp/yao/agent/assistant"
"github.com/yaoapp/yao/agent/search/handlers/web"
"github.com/yaoapp/yao/agent/search/types"
"github.com/yaoapp/yao/agent/testutils"
)
// TestMCPProviderWithAssistantConfig tests MCPProvider using web-mcp assistant config
func TestMCPProviderWithAssistantConfig(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Load the web-mcp test assistant to get its config
ast, err := assistant.LoadPath("/assistants/tests/web-mcp")
require.NoError(t, err)
require.NotNil(t, ast)
require.NotNil(t, ast.Uses)
// Verify assistant config
assert.Equal(t, "tests.web-mcp", ast.ID)
assert.Equal(t, "mcp:search.web_search", ast.Uses.Web)
// Create MCPProvider from uses.web
mcpRef := ast.Uses.Web[4:] // Remove "mcp:" prefix
provider, err := web.NewMCPProvider(mcpRef)
require.NoError(t, err)
require.NotNil(t, provider)
// Execute search
req := &types.Request{
Query: "Yao App Engine",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
Limit: 5,
}
result, err := provider.Search(req)
require.NoError(t, err)
require.NotNil(t, result)
// Verify result structure
assert.Equal(t, types.SearchTypeWeb, result.Type)
assert.Equal(t, "Yao App Engine", result.Query)
assert.Equal(t, types.SourceAuto, result.Source)
// MCP should return mock results
if result.Error == "" {
assert.Greater(t, result.Total, 0)
assert.NotEmpty(t, result.Items)
assert.Greater(t, result.Duration, int64(0))
// Verify result item structure
for _, item := range result.Items {
assert.Equal(t, types.SearchTypeWeb, item.Type)
assert.Equal(t, types.SourceAuto, item.Source)
assert.NotEmpty(t, item.Title)
assert.NotEmpty(t, item.URL)
}
t.Logf("MCP search returned %d results in %dms", result.Total, result.Duration)
} else {
t.Logf("MCP search returned error (expected if MCP not loaded): %s", result.Error)
}
}
// TestMCPProviderWithSiteRestriction tests MCPProvider with domain restriction
func TestMCPProviderWithSiteRestriction(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create MCPProvider
provider, err := web.NewMCPProvider("search.web_search")
require.NoError(t, err)
// Execute search with site restriction
req := &types.Request{
Query: "documentation",
Type: types.SearchTypeWeb,
Source: types.SourceHook,
Sites: []string{"github.com"},
Limit: 3,
}
result, err := provider.Search(req)
require.NoError(t, err)
require.NotNil(t, result)
assert.Equal(t, types.SearchTypeWeb, result.Type)
assert.Equal(t, types.SourceHook, result.Source)
if result.Error == "" {
t.Logf("Site-restricted MCP search returned %d results", result.Total)
} else {
t.Logf("MCP search returned error (expected if MCP not loaded): %s", result.Error)
}
}
// TestMCPProviderWithTimeRange tests MCPProvider with time range filter
func TestMCPProviderWithTimeRange(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create MCPProvider
provider, err := web.NewMCPProvider("search.web_search")
require.NoError(t, err)
// Execute search with time range
req := &types.Request{
Query: "artificial intelligence news",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
TimeRange: "week",
Limit: 5,
}
result, err := provider.Search(req)
require.NoError(t, err)
require.NotNil(t, result)
if result.Error == "" {
t.Logf("Time-ranged MCP search (last week) returned %d results in %dms", result.Total, result.Duration)
} else {
t.Logf("MCP search returned error (expected if MCP not loaded): %s", result.Error)
}
}
// TestMCPProviderInvalidFormat tests MCPProvider with invalid format
func TestMCPProviderInvalidFormat(t *testing.T) {
// Test invalid format without dot
_, err := web.NewMCPProvider("invalid")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid MCP format")
// Test empty string
_, err = web.NewMCPProvider("")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid MCP format")
}
// TestMCPProviderNotFound tests MCPProvider when MCP server is not found
func TestMCPProviderNotFound(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create MCPProvider with non-existent server
provider, err := web.NewMCPProvider("nonexistent.web_search")
require.NoError(t, err)
req := &types.Request{
Query: "test query",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
}
result, err := provider.Search(req)
// Should not return error, but result should have error message
require.NoError(t, err)
require.NotNil(t, result)
assert.NotEmpty(t, result.Error)
assert.Contains(t, result.Error, "not found")
}
// TestWebHandlerMCPMode tests the web handler in MCP mode
func TestWebHandlerMCPMode(t *testing.T) {
if testing.Short() {
t.Skip("Skipping integration test")
}
testutils.Prepare(t)
defer testutils.Clean(t)
// Create handler with MCP mode
handler := web.NewHandler("mcp:search.web_search", nil)
require.NotNil(t, handler)
// Verify type
assert.Equal(t, types.SearchTypeWeb, handler.Type())
// Execute search
req := &types.Request{
Query: "Yao framework",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
Limit: 5,
}
result, err := handler.Search(req)
require.NoError(t, err)
require.NotNil(t, result)
assert.Equal(t, types.SearchTypeWeb, result.Type)
assert.Equal(t, "Yao framework", result.Query)
if result.Error == "" {
t.Logf("Handler MCP mode returned %d results", result.Total)
} else {
t.Logf("Handler MCP mode returned error (expected if MCP not loaded): %s", result.Error)
}
}
// TestWebHandlerInvalidMCPFormat tests the web handler with invalid MCP format
func TestWebHandlerInvalidMCPFormat(t *testing.T) {
// Create handler with invalid MCP format
handler := web.NewHandler("mcp:invalid", nil)
require.NotNil(t, handler)
req := &types.Request{
Query: "test",
Type: types.SearchTypeWeb,
Source: types.SourceAuto,
}
result, err := handler.Search(req)
require.NoError(t, err)
require.NotNil(t, result)
assert.NotEmpty(t, result.Error)
assert.Contains(t, result.Error, "Invalid MCP format")
}