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:
parent
e4d1701dab
commit
b3cf5a09d9
10 changed files with 1017 additions and 44 deletions
|
|
@ -367,6 +367,8 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
var nextResponse *context.NextHookResponse = nil
|
var nextResponse *context.NextHookResponse = nil
|
||||||
|
|
||||||
if ast.HookScript != nil {
|
if ast.HookScript != nil {
|
||||||
|
ctx.Logger.HookStart("Next")
|
||||||
|
|
||||||
// Begin step tracking for hook_next
|
// Begin step tracking for hook_next
|
||||||
ast.BeginStep(ctx, context.StepTypeHookNext, map[string]interface{}{
|
ast.BeginStep(ctx, context.StepTypeHookNext, map[string]interface{}{
|
||||||
"messages": fullMessages,
|
"messages": fullMessages,
|
||||||
|
|
@ -393,6 +395,8 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
"response": nextResponse,
|
"response": nextResponse,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
ctx.Logger.HookComplete("Next")
|
||||||
|
|
||||||
// Process Next hook response
|
// Process Next hook response
|
||||||
finalResponse, err = ast.processNextResponse(&NextProcessContext{
|
finalResponse, err = ast.processNextResponse(&NextProcessContext{
|
||||||
Context: ctx,
|
Context: ctx,
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,7 @@ import (
|
||||||
"path"
|
"path"
|
||||||
|
|
||||||
"github.com/yaoapp/gou/fs"
|
"github.com/yaoapp/gou/fs"
|
||||||
"github.com/yaoapp/yao/agent/content"
|
"github.com/yaoapp/yao/agent/caller"
|
||||||
agentContext "github.com/yaoapp/yao/agent/context"
|
agentContext "github.com/yaoapp/yao/agent/context"
|
||||||
"github.com/yaoapp/yao/agent/i18n"
|
"github.com/yaoapp/yao/agent/i18n"
|
||||||
searchTypes "github.com/yaoapp/yao/agent/search/types"
|
searchTypes "github.com/yaoapp/yao/agent/search/types"
|
||||||
|
|
@ -14,8 +14,8 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
// Initialize AgentGetterFunc to allow content package to call agents
|
// Initialize AgentGetterFunc to allow content and search packages to call agents
|
||||||
content.AgentGetterFunc = func(agentID string) (content.AgentCaller, error) {
|
caller.AgentGetterFunc = func(agentID string) (caller.AgentCaller, error) {
|
||||||
ast, err := Get(agentID)
|
ast, err := Get(agentID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
|
|
|
||||||
17
agent/caller/caller.go
Normal file
17
agent/caller/caller.go
Normal 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)
|
||||||
|
|
@ -9,29 +9,22 @@ import (
|
||||||
jsoniter "github.com/json-iterator/go"
|
jsoniter "github.com/json-iterator/go"
|
||||||
"github.com/yaoapp/gou/mcp"
|
"github.com/yaoapp/gou/mcp"
|
||||||
"github.com/yaoapp/kun/log"
|
"github.com/yaoapp/kun/log"
|
||||||
|
"github.com/yaoapp/yao/agent/caller"
|
||||||
agentContext "github.com/yaoapp/yao/agent/context"
|
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
|
// fileInfoMutex protects concurrent access to files_info list in Space
|
||||||
var fileInfoMutex sync.Mutex
|
var fileInfoMutex sync.Mutex
|
||||||
|
|
||||||
// CallAgent calls an agent to process content (vision, audio, etc.)
|
// CallAgent calls an agent to process content (vision, audio, etc.)
|
||||||
// This is a generic function that can be used by any handler
|
// This is a generic function that can be used by any handler
|
||||||
func CallAgent(ctx *agentContext.Context, agentID string, message agentContext.Message) (string, error) {
|
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")
|
return "", fmt.Errorf("AgentGetterFunc not initialized")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load the agent by ID using the injected function
|
// Load the agent by ID using the injected function
|
||||||
agent, err := AgentGetterFunc(agentID)
|
agent, err := caller.AgentGetterFunc(agentID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("failed to load agent %s: %w", agentID, err)
|
return "", fmt.Errorf("failed to load agent %s: %w", agentID, err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1110,7 +1110,7 @@ Tool format: `"builtin"`, `"<assistant-id>"` (Agent), `"mcp:<server>.<tool>"` (M
|
||||||
|
|
||||||
| Mode | Example | Description |
|
| 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 |
|
| 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 |
|
| 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 |
|
| 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 |
|
| Agent | `"workers.search.web"` | AI-powered search with intent understanding |
|
||||||
| MCP | `"mcp:my-server.web_search"` | External search tool via MCP |
|
| MCP | `"mcp:my-server.web_search"` | External search tool via MCP |
|
||||||
|
|
||||||
|
|
|
||||||
232
agent/search/handlers/web/agent.go
Normal file
232
agent/search/handlers/web/agent.go
Normal 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, ""
|
||||||
|
}
|
||||||
284
agent/search/handlers/web/agent_test.go
Normal file
284
agent/search/handlers/web/agent_test.go
Normal 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
|
||||||
|
}
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
agentContext "github.com/yaoapp/yao/agent/context"
|
||||||
"github.com/yaoapp/yao/agent/search/types"
|
"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
|
// 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) {
|
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 {
|
switch {
|
||||||
case h.usesWeb == "builtin" || h.usesWeb == "":
|
case h.usesWeb == "builtin" || h.usesWeb == "":
|
||||||
return h.builtinSearch(req)
|
return h.builtinSearch(req)
|
||||||
|
|
@ -32,7 +40,17 @@ func (h *Handler) Search(req *types.Request) (*types.Result, error) {
|
||||||
return h.mcpSearch(req)
|
return h.mcpSearch(req)
|
||||||
default:
|
default:
|
||||||
// Agent mode: delegate to assistant for AI-powered search
|
// 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
|
// agentSearch delegates to an assistant for AI-powered search
|
||||||
func (h *Handler) agentSearch(req *types.Request) (*types.Result, error) {
|
func (h *Handler) agentSearch(ctx *agentContext.Context, req *types.Request) (*types.Result, error) {
|
||||||
// TODO: Implement agent mode
|
provider := NewAgentProvider(h.usesWeb)
|
||||||
// 1. Call assistant with search request
|
return provider.Search(ctx, req)
|
||||||
// 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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// mcpSearch calls external MCP tool
|
// mcpSearch calls external MCP tool
|
||||||
func (h *Handler) mcpSearch(req *types.Request) (*types.Result, error) {
|
func (h *Handler) mcpSearch(req *types.Request) (*types.Result, error) {
|
||||||
// TODO: Implement MCP mode
|
|
||||||
// Parse "mcp:server.tool"
|
// Parse "mcp:server.tool"
|
||||||
mcpRef := strings.TrimPrefix(h.usesWeb, "mcp:")
|
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{
|
return &types.Result{
|
||||||
Type: types.SearchTypeWeb,
|
Type: types.SearchTypeWeb,
|
||||||
Query: req.Query,
|
Query: req.Query,
|
||||||
Source: req.Source,
|
Source: req.Source,
|
||||||
Items: []*types.ResultItem{},
|
Items: []*types.ResultItem{},
|
||||||
Total: 0,
|
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
|
}, nil
|
||||||
}
|
}
|
||||||
// serverID, toolName := parts[0], parts[1]
|
|
||||||
// Call MCP tool
|
return provider.Search(req)
|
||||||
return &types.Result{
|
|
||||||
Type: types.SearchTypeWeb,
|
|
||||||
Query: req.Query,
|
|
||||||
Source: req.Source,
|
|
||||||
Items: []*types.ResultItem{},
|
|
||||||
Total: 0,
|
|
||||||
Error: "MCP mode not yet implemented",
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
|
||||||
203
agent/search/handlers/web/mcp.go
Normal file
203
agent/search/handlers/web/mcp.go
Normal 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
|
||||||
|
}
|
||||||
241
agent/search/handlers/web/mcp_test.go
Normal file
241
agent/search/handlers/web/mcp_test.go
Normal 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")
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue