Enhance Search API with Parallel Search Methods
- Refactored the SearchAPI interface to replace the Parallel method with All, Any, and Race methods, inspired by JavaScript Promise patterns. - Updated the Searcher struct to implement these new parallel search methods, improving flexibility and performance in executing multiple searches. - Revised the JSAPI implementation to support the new parallel search methods, ensuring consistency across the API. - Enhanced documentation in DESIGN.md to detail the new parallel search functionalities and provide usage examples, clarifying their behavior and expected outcomes.
This commit is contained in:
parent
534f4d6ed5
commit
fc9e00c917
14 changed files with 2188 additions and 62 deletions
|
|
@ -16,9 +16,13 @@ type SearchAPI interface {
|
||||||
// Returns *types.Result or error information
|
// Returns *types.Result or error information
|
||||||
DB(query string, opts map[string]interface{}) interface{}
|
DB(query string, opts map[string]interface{}) interface{}
|
||||||
|
|
||||||
// Parallel executes multiple searches in parallel
|
// Parallel search methods - inspired by JavaScript Promise
|
||||||
// Returns []*types.Result
|
// All waits for all searches to complete (like Promise.all)
|
||||||
Parallel(requests []interface{}) []interface{}
|
All(requests []interface{}) []interface{}
|
||||||
|
// Any returns when any search succeeds with results (like Promise.any)
|
||||||
|
Any(requests []interface{}) []interface{}
|
||||||
|
// Race returns when any search completes (like Promise.race)
|
||||||
|
Race(requests []interface{}) []interface{}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SearchAPIFactory is a function type that creates a SearchAPI for a context
|
// SearchAPIFactory is a function type that creates a SearchAPI for a context
|
||||||
|
|
|
||||||
|
|
@ -313,25 +313,33 @@ func (s *Searcher) Search(ctx *context.Context, req *types.Request) (*types.Resu
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SearchMultiple executes multiple searches in parallel
|
// ParallelMode defines how parallel search should behave (inspired by JavaScript Promise)
|
||||||
func (s *Searcher) SearchMultiple(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
type ParallelMode string
|
||||||
results := make([]*types.Result, len(reqs))
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
var mu sync.Mutex
|
|
||||||
|
|
||||||
for i, req := range reqs {
|
// ParallelMode constants (similar to Promise.all, Promise.any, Promise.race)
|
||||||
wg.Add(1)
|
const (
|
||||||
go func(idx int, r *types.Request) {
|
// ModeAll waits for all searches to complete, returns all results (like Promise.all)
|
||||||
defer wg.Done()
|
ModeAll ParallelMode = "all"
|
||||||
result, _ := s.Search(ctx, r)
|
// ModeAny returns as soon as any search succeeds (has results), others continue but are discarded (like Promise.any)
|
||||||
mu.Lock()
|
ModeAny ParallelMode = "any"
|
||||||
results[idx] = result
|
// ModeRace returns as soon as any search completes (success or empty), others continue but are discarded (like Promise.race)
|
||||||
mu.Unlock()
|
ModeRace ParallelMode = "race"
|
||||||
}(i, req)
|
)
|
||||||
}
|
|
||||||
|
|
||||||
wg.Wait()
|
// ParallelOptions configures parallel search behavior
|
||||||
return results, nil
|
// All executes all searches and waits for all to complete (like Promise.all)
|
||||||
|
func (s *Searcher) All(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
|
return s.parallelAll(ctx, reqs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Any returns as soon as any search succeeds with results (like Promise.any)
|
||||||
|
func (s *Searcher) Any(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
|
return s.parallelAny(ctx, reqs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Race returns as soon as any search completes (like Promise.race)
|
||||||
|
func (s *Searcher) Race(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
|
return s.parallelRace(ctx, reqs)
|
||||||
}
|
}
|
||||||
|
|
||||||
// BuildReferences converts search results to unified Reference format
|
// BuildReferences converts search results to unified Reference format
|
||||||
|
|
@ -424,17 +432,26 @@ import (
|
||||||
// Searcher is the main interface exposed to external callers
|
// Searcher is the main interface exposed to external callers
|
||||||
type Searcher interface {
|
type Searcher interface {
|
||||||
// Search executes a single search request
|
// Search executes a single search request
|
||||||
Search(req *types.Request) (*types.Result, error)
|
Search(ctx *context.Context, req *types.Request) (*types.Result, error)
|
||||||
|
|
||||||
// SearchMultiple executes multiple searches (potentially in parallel)
|
// Parallel search methods - inspired by JavaScript Promise
|
||||||
SearchMultiple(reqs []*types.Request) ([]*types.Result, error)
|
// All waits for all searches to complete (like Promise.all)
|
||||||
|
All(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error)
|
||||||
|
// Any returns when any search succeeds with results (like Promise.any)
|
||||||
|
Any(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error)
|
||||||
|
// Race returns when any search completes (like Promise.race)
|
||||||
|
Race(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error)
|
||||||
|
|
||||||
// BuildReferences converts search results to unified Reference format for LLM
|
// BuildReferences converts search results to unified Reference format for LLM
|
||||||
BuildReferences(results []*types.Result) []*types.Reference
|
BuildReferences(results []*types.Result) []*types.Reference
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
> **Note**: The actual `Searcher` struct in `search.go` has `Search(ctx, req)` and `SearchMultiple(ctx, reqs)` signatures that include context for reranking support. The interface is kept minimal for flexibility.
|
> **Note**: Parallel search methods follow JavaScript Promise naming:
|
||||||
|
>
|
||||||
|
> - `All()`: Wait for all searches to complete (like `Promise.all`)
|
||||||
|
> - `Any()`: Return when any search succeeds with results (like `Promise.any`)
|
||||||
|
> - `Race()`: Return when any search completes (like `Promise.race`)
|
||||||
|
|
||||||
### NLP Interfaces (`interfaces/nlp.go`)
|
### NLP Interfaces (`interfaces/nlp.go`)
|
||||||
|
|
||||||
|
|
@ -901,17 +918,15 @@ agent/context/jsapi_search.go agent/search/jsapi.go
|
||||||
```typescript
|
```typescript
|
||||||
// In hook scripts (index.ts)
|
// In hook scripts (index.ts)
|
||||||
|
|
||||||
// Web search
|
// Single search methods
|
||||||
ctx.search.Web(query: string, options?: WebOptions): Result
|
ctx.search.Web(query: string, options?: WebOptions): Result
|
||||||
|
|
||||||
// Knowledge base search
|
|
||||||
ctx.search.KB(query: string, options?: KBOptions): Result
|
ctx.search.KB(query: string, options?: KBOptions): Result
|
||||||
|
|
||||||
// Database search (Yao Model/QueryDSL)
|
|
||||||
ctx.search.DB(query: string, options?: DBOptions): Result
|
ctx.search.DB(query: string, options?: DBOptions): Result
|
||||||
|
|
||||||
// Parallel search (multiple types)
|
// Parallel search methods - inspired by JavaScript Promise
|
||||||
ctx.search.Parallel(requests: Request[]): Result[]
|
ctx.search.All(requests: Request[]): Result[] // Like Promise.all - wait for all
|
||||||
|
ctx.search.Any(requests: Request[]): Result[] // Like Promise.any - first success
|
||||||
|
ctx.search.Race(requests: Request[]): Result[] // Like Promise.race - first complete
|
||||||
```
|
```
|
||||||
|
|
||||||
### Options Types
|
### Options Types
|
||||||
|
|
@ -1056,14 +1071,14 @@ function Create(ctx, messages, options) {
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
#### Example 4: Parallel Web + KB + DB Search
|
#### Example 4: Parallel Search with ctx.search.All()
|
||||||
|
|
||||||
```typescript
|
```typescript
|
||||||
function Create(ctx, messages, options) {
|
function Create(ctx, messages, options) {
|
||||||
const query = messages[messages.length - 1].content;
|
const query = messages[messages.length - 1].content;
|
||||||
|
|
||||||
// Execute web, KB, and DB search in parallel
|
// Execute web, KB, and DB search in parallel (wait for all) - like Promise.all
|
||||||
const [webResult, kbResult, dbResult] = ctx.search.Parallel([
|
const [webResult, kbResult, dbResult] = ctx.search.All([
|
||||||
{ type: "web", query: query, limit: 5 },
|
{ type: "web", query: query, limit: 5 },
|
||||||
{ type: "kb", query: query, collections: ["docs"], limit: 10 },
|
{ type: "kb", query: query, collections: ["docs"], limit: 10 },
|
||||||
{ type: "db", query: query, models: ["product"], limit: 10 },
|
{ type: "db", query: query, models: ["product"], limit: 10 },
|
||||||
|
|
@ -1084,6 +1099,56 @@ function Create(ctx, messages, options) {
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### Example 4b: Parallel Search with ctx.search.Any()
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
function Create(ctx, messages, options) {
|
||||||
|
const query = messages[messages.length - 1].content;
|
||||||
|
|
||||||
|
// Return as soon as any search succeeds (has results) - like Promise.any
|
||||||
|
const results = ctx.search.Any([
|
||||||
|
{ type: "web", query: query, limit: 5 },
|
||||||
|
{ type: "kb", query: query, collections: ["docs"], limit: 10 },
|
||||||
|
]);
|
||||||
|
|
||||||
|
// Use the first successful result
|
||||||
|
const successResult = results.find((r) => r && r.items?.length > 0);
|
||||||
|
if (successResult) {
|
||||||
|
return {
|
||||||
|
messages: [{ role: "system", content: formatContext(successResult) }],
|
||||||
|
uses: { search: "disabled" },
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
return { messages: [] };
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### Example 4c: Parallel Search with ctx.search.Race()
|
||||||
|
|
||||||
|
```typescript
|
||||||
|
function Create(ctx, messages, options) {
|
||||||
|
const query = messages[messages.length - 1].content;
|
||||||
|
|
||||||
|
// Return as soon as any search completes (success or not) - like Promise.race
|
||||||
|
const results = ctx.search.Race([
|
||||||
|
{ type: "web", query: query, limit: 5 },
|
||||||
|
{ type: "kb", query: query, collections: ["docs"], limit: 10 },
|
||||||
|
]);
|
||||||
|
|
||||||
|
// Use the first completed result
|
||||||
|
const firstResult = results.find((r) => r != null);
|
||||||
|
if (firstResult && firstResult.items?.length > 0) {
|
||||||
|
return {
|
||||||
|
messages: [{ role: "system", content: formatContext(firstResult) }],
|
||||||
|
uses: { search: "disabled" },
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
return { messages: [] };
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
#### Example 5: Custom Citation Format
|
#### Example 5: Custom Citation Format
|
||||||
|
|
||||||
```typescript
|
```typescript
|
||||||
|
|
|
||||||
88
agent/search/citation_test.go
Normal file
88
agent/search/citation_test.go
Normal file
|
|
@ -0,0 +1,88 @@
|
||||||
|
package search
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCitationGenerator_Next(t *testing.T) {
|
||||||
|
gen := NewCitationGenerator()
|
||||||
|
|
||||||
|
// First ID should be ref_001
|
||||||
|
id1 := gen.Next()
|
||||||
|
assert.Equal(t, "ref_001", id1)
|
||||||
|
|
||||||
|
// Second ID should be ref_002
|
||||||
|
id2 := gen.Next()
|
||||||
|
assert.Equal(t, "ref_002", id2)
|
||||||
|
|
||||||
|
// Third ID should be ref_003
|
||||||
|
id3 := gen.Next()
|
||||||
|
assert.Equal(t, "ref_003", id3)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCitationGenerator_Reset(t *testing.T) {
|
||||||
|
gen := NewCitationGenerator()
|
||||||
|
|
||||||
|
// Generate some IDs
|
||||||
|
gen.Next()
|
||||||
|
gen.Next()
|
||||||
|
gen.Next()
|
||||||
|
|
||||||
|
// Reset
|
||||||
|
gen.Reset()
|
||||||
|
|
||||||
|
// Next ID should be ref_001 again
|
||||||
|
id := gen.Next()
|
||||||
|
assert.Equal(t, "ref_001", id)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCitationGenerator_Format(t *testing.T) {
|
||||||
|
gen := NewCitationGenerator()
|
||||||
|
|
||||||
|
// Generate 999 IDs to test padding
|
||||||
|
for i := 0; i < 999; i++ {
|
||||||
|
gen.Next()
|
||||||
|
}
|
||||||
|
|
||||||
|
// 1000th ID should be ref_1000 (no padding limit)
|
||||||
|
id := gen.Next()
|
||||||
|
assert.Equal(t, "ref_1000", id)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCitationGenerator_Concurrent(t *testing.T) {
|
||||||
|
gen := NewCitationGenerator()
|
||||||
|
|
||||||
|
// Run 100 goroutines, each generating 10 IDs
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
ids := make(chan string, 1000)
|
||||||
|
|
||||||
|
for i := 0; i < 100; i++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
for j := 0; j < 10; j++ {
|
||||||
|
ids <- gen.Next()
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
close(ids)
|
||||||
|
|
||||||
|
// Collect all IDs
|
||||||
|
idSet := make(map[string]bool)
|
||||||
|
for id := range ids {
|
||||||
|
idSet[id] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
// All 1000 IDs should be unique
|
||||||
|
assert.Equal(t, 1000, len(idSet))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewCitationGenerator(t *testing.T) {
|
||||||
|
gen := NewCitationGenerator()
|
||||||
|
assert.NotNil(t, gen)
|
||||||
|
}
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
package db
|
package db
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/yaoapp/yao/agent/search/types"
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -21,14 +23,71 @@ func (h *Handler) Type() types.SearchType {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Search converts NL to QueryDSL and executes
|
// Search converts NL to QueryDSL and executes
|
||||||
// TODO: Implement actual search logic
|
// TODO: Implement actual QueryDSL generation and model query logic
|
||||||
func (h *Handler) Search(req *types.Request) (*types.Result, error) {
|
func (h *Handler) Search(req *types.Request) (*types.Result, error) {
|
||||||
// Skeleton implementation - returns empty result
|
start := time.Now()
|
||||||
|
|
||||||
|
// Validate request
|
||||||
|
if req.Query == "" {
|
||||||
return &types.Result{
|
return &types.Result{
|
||||||
Type: types.SearchTypeDB,
|
Type: types.SearchTypeDB,
|
||||||
Query: req.Query,
|
Query: req.Query,
|
||||||
Source: req.Source,
|
Source: req.Source,
|
||||||
Items: []*types.ResultItem{},
|
Items: []*types.ResultItem{},
|
||||||
Total: 0,
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
Error: "query is required",
|
||||||
}, nil
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get models from request or config
|
||||||
|
models := req.Models
|
||||||
|
if len(models) == 0 && h.config != nil {
|
||||||
|
models = h.config.Models
|
||||||
|
}
|
||||||
|
|
||||||
|
// If no models specified, return empty result
|
||||||
|
if len(models) == 0 {
|
||||||
|
return &types.Result{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: req.Query,
|
||||||
|
Source: req.Source,
|
||||||
|
Items: []*types.ResultItem{},
|
||||||
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get max results
|
||||||
|
maxResults := req.Limit
|
||||||
|
if maxResults == 0 && h.config != nil && h.config.MaxResults > 0 {
|
||||||
|
maxResults = h.config.MaxResults
|
||||||
|
}
|
||||||
|
if maxResults == 0 {
|
||||||
|
maxResults = 20 // default
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: Implement actual DB search
|
||||||
|
// 1. Get model schemas for specified models
|
||||||
|
// 2. Generate QueryDSL from natural language query using uses.querydsl mode:
|
||||||
|
// - "builtin": template-based generation
|
||||||
|
// - "<assistant-id>": delegate to LLM assistant
|
||||||
|
// - "mcp:<server>.<tool>": call external MCP tool
|
||||||
|
// 3. Execute QueryDSL on each model
|
||||||
|
// 4. Format results and return
|
||||||
|
|
||||||
|
// For now, return empty result (skeleton)
|
||||||
|
result := &types.Result{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: req.Query,
|
||||||
|
Source: req.Source,
|
||||||
|
Items: []*types.ResultItem{},
|
||||||
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store maxResults for later use
|
||||||
|
_ = maxResults
|
||||||
|
|
||||||
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
215
agent/search/handlers/db/handler_test.go
Normal file
215
agent/search/handlers/db/handler_test.go
Normal file
|
|
@ -0,0 +1,215 @@
|
||||||
|
package db
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewHandler(t *testing.T) {
|
||||||
|
t.Run("with nil config", func(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", nil)
|
||||||
|
assert.NotNil(t, h)
|
||||||
|
assert.Equal(t, "builtin", h.usesQueryDSL)
|
||||||
|
assert.Nil(t, h.config)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("with config", func(t *testing.T) {
|
||||||
|
cfg := &types.DBConfig{
|
||||||
|
Models: []string{"product", "order"},
|
||||||
|
MaxResults: 50,
|
||||||
|
}
|
||||||
|
h := NewHandler("workers.nlp.querydsl", cfg)
|
||||||
|
assert.NotNil(t, h)
|
||||||
|
assert.Equal(t, "workers.nlp.querydsl", h.usesQueryDSL)
|
||||||
|
assert.Equal(t, cfg, h.config)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("with mcp mode", func(t *testing.T) {
|
||||||
|
h := NewHandler("mcp:nlp.generate_querydsl", nil)
|
||||||
|
assert.NotNil(t, h)
|
||||||
|
assert.Equal(t, "mcp:nlp.generate_querydsl", h.usesQueryDSL)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_Type(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", nil)
|
||||||
|
assert.Equal(t, types.SearchTypeDB, h.Type())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_Search(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
usesQueryDSL string
|
||||||
|
config *types.DBConfig
|
||||||
|
req *types.Request
|
||||||
|
expectError string
|
||||||
|
expectItems int
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "empty query",
|
||||||
|
usesQueryDSL: "builtin",
|
||||||
|
config: nil,
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "",
|
||||||
|
},
|
||||||
|
expectError: "query is required",
|
||||||
|
expectItems: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no models in request or config",
|
||||||
|
usesQueryDSL: "builtin",
|
||||||
|
config: nil,
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "find products under $100",
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "models from config",
|
||||||
|
usesQueryDSL: "builtin",
|
||||||
|
config: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
MaxResults: 20,
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "find products under $100",
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "models from request",
|
||||||
|
usesQueryDSL: "builtin",
|
||||||
|
config: nil,
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "find products under $100",
|
||||||
|
Models: []string{"product", "order"},
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "with limit",
|
||||||
|
usesQueryDSL: "builtin",
|
||||||
|
config: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "find products",
|
||||||
|
Models: []string{"product"},
|
||||||
|
Limit: 5,
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "with wheres",
|
||||||
|
usesQueryDSL: "builtin",
|
||||||
|
config: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "find products",
|
||||||
|
Models: []string{"product"},
|
||||||
|
// Wheres would be set here in real usage
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "agent mode",
|
||||||
|
usesQueryDSL: "workers.nlp.querydsl",
|
||||||
|
config: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "find products",
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "mcp mode",
|
||||||
|
usesQueryDSL: "mcp:nlp.generate_querydsl",
|
||||||
|
config: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "find products",
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
h := NewHandler(tt.usesQueryDSL, tt.config)
|
||||||
|
result, err := h.Search(tt.req)
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, types.SearchTypeDB, result.Type)
|
||||||
|
assert.Equal(t, tt.req.Query, result.Query)
|
||||||
|
assert.Equal(t, tt.expectItems, len(result.Items))
|
||||||
|
|
||||||
|
if tt.expectError != "" {
|
||||||
|
assert.Equal(t, tt.expectError, result.Error)
|
||||||
|
} else {
|
||||||
|
assert.Empty(t, result.Error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Duration should be set
|
||||||
|
assert.GreaterOrEqual(t, result.Duration, int64(0))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_Search_SourcePreserved(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", &types.DBConfig{Models: []string{"product"}})
|
||||||
|
|
||||||
|
sources := []types.SourceType{types.SourceUser, types.SourceHook, types.SourceAuto}
|
||||||
|
for _, source := range sources {
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "test",
|
||||||
|
Source: source,
|
||||||
|
Models: []string{"product"},
|
||||||
|
}
|
||||||
|
result, err := h.Search(req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, source, result.Source)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_Search_MaxResultsFromConfig(t *testing.T) {
|
||||||
|
cfg := &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
MaxResults: 50,
|
||||||
|
}
|
||||||
|
h := NewHandler("builtin", cfg)
|
||||||
|
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "test",
|
||||||
|
Models: []string{"product"},
|
||||||
|
// No limit in request, should use config's MaxResults
|
||||||
|
}
|
||||||
|
result, err := h.Search(req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
// Skeleton doesn't actually use maxResults yet, but the test ensures the handler runs
|
||||||
|
}
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
package kb
|
package kb
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/yaoapp/yao/agent/search/types"
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -20,14 +22,74 @@ func (h *Handler) Type() types.SearchType {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Search executes vector search and optional graph association
|
// Search executes vector search and optional graph association
|
||||||
// TODO: Implement actual search logic
|
// TODO: Implement actual vector search and graph association logic
|
||||||
func (h *Handler) Search(req *types.Request) (*types.Result, error) {
|
func (h *Handler) Search(req *types.Request) (*types.Result, error) {
|
||||||
// Skeleton implementation - returns empty result
|
start := time.Now()
|
||||||
|
|
||||||
|
// Validate request
|
||||||
|
if req.Query == "" {
|
||||||
return &types.Result{
|
return &types.Result{
|
||||||
Type: types.SearchTypeKB,
|
Type: types.SearchTypeKB,
|
||||||
Query: req.Query,
|
Query: req.Query,
|
||||||
Source: req.Source,
|
Source: req.Source,
|
||||||
Items: []*types.ResultItem{},
|
Items: []*types.ResultItem{},
|
||||||
Total: 0,
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
Error: "query is required",
|
||||||
}, nil
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get collections from request or config
|
||||||
|
collections := req.Collections
|
||||||
|
if len(collections) == 0 && h.config != nil {
|
||||||
|
collections = h.config.Collections
|
||||||
|
}
|
||||||
|
|
||||||
|
// If no collections specified, return empty result
|
||||||
|
if len(collections) == 0 {
|
||||||
|
return &types.Result{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: req.Query,
|
||||||
|
Source: req.Source,
|
||||||
|
Items: []*types.ResultItem{},
|
||||||
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get threshold from request or config
|
||||||
|
threshold := req.Threshold
|
||||||
|
if threshold == 0 && h.config != nil && h.config.Threshold > 0 {
|
||||||
|
threshold = h.config.Threshold
|
||||||
|
}
|
||||||
|
if threshold == 0 {
|
||||||
|
threshold = 0.7 // default
|
||||||
|
}
|
||||||
|
|
||||||
|
// Get limit
|
||||||
|
limit := req.Limit
|
||||||
|
if limit == 0 {
|
||||||
|
limit = 10 // default
|
||||||
|
}
|
||||||
|
|
||||||
|
// TODO: Implement actual vector search
|
||||||
|
// 1. Generate embedding for query using collection's embedding config
|
||||||
|
// 2. Search each collection with vector similarity
|
||||||
|
// 3. If req.Graph is true, perform graph association
|
||||||
|
// 4. Merge and return results
|
||||||
|
|
||||||
|
// For now, return empty result (skeleton)
|
||||||
|
result := &types.Result{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: req.Query,
|
||||||
|
Source: req.Source,
|
||||||
|
Items: []*types.ResultItem{},
|
||||||
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
}
|
||||||
|
|
||||||
|
// Store threshold in result metadata for debugging
|
||||||
|
_ = threshold
|
||||||
|
|
||||||
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
170
agent/search/handlers/kb/handler_test.go
Normal file
170
agent/search/handlers/kb/handler_test.go
Normal file
|
|
@ -0,0 +1,170 @@
|
||||||
|
package kb
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewHandler(t *testing.T) {
|
||||||
|
t.Run("with nil config", func(t *testing.T) {
|
||||||
|
h := NewHandler(nil)
|
||||||
|
assert.NotNil(t, h)
|
||||||
|
assert.Nil(t, h.config)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("with config", func(t *testing.T) {
|
||||||
|
cfg := &types.KBConfig{
|
||||||
|
Collections: []string{"docs", "faq"},
|
||||||
|
Threshold: 0.8,
|
||||||
|
Graph: true,
|
||||||
|
}
|
||||||
|
h := NewHandler(cfg)
|
||||||
|
assert.NotNil(t, h)
|
||||||
|
assert.Equal(t, cfg, h.config)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_Type(t *testing.T) {
|
||||||
|
h := NewHandler(nil)
|
||||||
|
assert.Equal(t, types.SearchTypeKB, h.Type())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_Search(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
config *types.KBConfig
|
||||||
|
req *types.Request
|
||||||
|
expectError string
|
||||||
|
expectItems int
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "empty query",
|
||||||
|
config: nil,
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "",
|
||||||
|
},
|
||||||
|
expectError: "query is required",
|
||||||
|
expectItems: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no collections in request or config",
|
||||||
|
config: nil,
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test query",
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "collections from config",
|
||||||
|
config: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
Threshold: 0.7,
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test query",
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "collections from request",
|
||||||
|
config: nil,
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test query",
|
||||||
|
Collections: []string{"docs", "faq"},
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "with threshold from request",
|
||||||
|
config: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
Threshold: 0.7,
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test query",
|
||||||
|
Threshold: 0.9,
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "with graph enabled",
|
||||||
|
config: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
Graph: true,
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test query",
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
Graph: true,
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "with limit",
|
||||||
|
config: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
req: &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test query",
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
Limit: 5,
|
||||||
|
},
|
||||||
|
expectError: "",
|
||||||
|
expectItems: 0, // skeleton returns empty
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
h := NewHandler(tt.config)
|
||||||
|
result, err := h.Search(tt.req)
|
||||||
|
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, types.SearchTypeKB, result.Type)
|
||||||
|
assert.Equal(t, tt.req.Query, result.Query)
|
||||||
|
assert.Equal(t, tt.expectItems, len(result.Items))
|
||||||
|
|
||||||
|
if tt.expectError != "" {
|
||||||
|
assert.Equal(t, tt.expectError, result.Error)
|
||||||
|
} else {
|
||||||
|
assert.Empty(t, result.Error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Duration should be set
|
||||||
|
assert.GreaterOrEqual(t, result.Duration, int64(0))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_Search_SourcePreserved(t *testing.T) {
|
||||||
|
h := NewHandler(&types.KBConfig{Collections: []string{"docs"}})
|
||||||
|
|
||||||
|
sources := []types.SourceType{types.SourceUser, types.SourceHook, types.SourceAuto}
|
||||||
|
for _, source := range sources {
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test",
|
||||||
|
Source: source,
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
}
|
||||||
|
result, err := h.Search(req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, source, result.Source)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,16 +1,22 @@
|
||||||
package interfaces
|
package interfaces
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"github.com/yaoapp/yao/agent/context"
|
||||||
"github.com/yaoapp/yao/agent/search/types"
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Searcher is the main interface exposed to external callers
|
// Searcher is the main interface exposed to external callers
|
||||||
type Searcher interface {
|
type Searcher interface {
|
||||||
// Search executes a single search request
|
// Search executes a single search request
|
||||||
Search(req *types.Request) (*types.Result, error)
|
Search(ctx *context.Context, req *types.Request) (*types.Result, error)
|
||||||
|
|
||||||
// SearchMultiple executes multiple searches (potentially in parallel)
|
// Parallel search methods - inspired by JavaScript Promise
|
||||||
SearchMultiple(reqs []*types.Request) ([]*types.Result, error)
|
// All waits for all searches to complete (like Promise.all)
|
||||||
|
All(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error)
|
||||||
|
// Any returns when any search succeeds with results (like Promise.any)
|
||||||
|
Any(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error)
|
||||||
|
// Race returns when any search completes (like Promise.race)
|
||||||
|
Race(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error)
|
||||||
|
|
||||||
// BuildReferences converts search results to unified Reference format for LLM
|
// BuildReferences converts search results to unified Reference format for LLM
|
||||||
BuildReferences(results []*types.Result) []*types.Reference
|
BuildReferences(results []*types.Result) []*types.Reference
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,7 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
// JSAPI implements context.SearchAPI interface
|
// JSAPI implements context.SearchAPI interface
|
||||||
// Provides ctx.search.Web(), ctx.search.KB(), ctx.search.DB(), ctx.search.Parallel()
|
// Provides ctx.search.Web(), ctx.search.KB(), ctx.search.DB(), ctx.search.All(), ctx.search.Any(), ctx.search.Race()
|
||||||
type JSAPI struct {
|
type JSAPI struct {
|
||||||
ctx *context.Context
|
ctx *context.Context
|
||||||
config *types.Config
|
config *types.Config
|
||||||
|
|
@ -79,15 +79,53 @@ func (api *JSAPI) DB(query string, opts map[string]interface{}) interface{} {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Parallel executes multiple searches in parallel
|
// All executes all searches and waits for all to complete (like Promise.all)
|
||||||
// Each request should have:
|
// Each request should have:
|
||||||
// - type: string - "web", "kb", or "db"
|
// - type: string - "web", "kb", or "db"
|
||||||
// - query: string - search query
|
// - query: string - search query
|
||||||
// - ... other type-specific options
|
// - ... other type-specific options
|
||||||
func (api *JSAPI) Parallel(requests []interface{}) []interface{} {
|
func (api *JSAPI) All(requests []interface{}) []interface{} {
|
||||||
// TODO: Implement parallel search
|
// TODO: Implement All search
|
||||||
// 1. Parse requests into []Request
|
// 1. Parse requests into []Request
|
||||||
// 2. Call SearchMultiple
|
// 2. Call Searcher.All()
|
||||||
|
// 3. Return []Result
|
||||||
|
results := make([]interface{}, len(requests))
|
||||||
|
for i := range requests {
|
||||||
|
results[i] = &types.Result{
|
||||||
|
Error: "not implemented",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return results
|
||||||
|
}
|
||||||
|
|
||||||
|
// Any returns as soon as any search succeeds with results (like Promise.any)
|
||||||
|
// Each request should have:
|
||||||
|
// - type: string - "web", "kb", or "db"
|
||||||
|
// - query: string - search query
|
||||||
|
// - ... other type-specific options
|
||||||
|
func (api *JSAPI) Any(requests []interface{}) []interface{} {
|
||||||
|
// TODO: Implement Any search
|
||||||
|
// 1. Parse requests into []Request
|
||||||
|
// 2. Call Searcher.Any()
|
||||||
|
// 3. Return []Result
|
||||||
|
results := make([]interface{}, len(requests))
|
||||||
|
for i := range requests {
|
||||||
|
results[i] = &types.Result{
|
||||||
|
Error: "not implemented",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return results
|
||||||
|
}
|
||||||
|
|
||||||
|
// Race returns as soon as any search completes (like Promise.race)
|
||||||
|
// Each request should have:
|
||||||
|
// - type: string - "web", "kb", or "db"
|
||||||
|
// - query: string - search query
|
||||||
|
// - ... other type-specific options
|
||||||
|
func (api *JSAPI) Race(requests []interface{}) []interface{} {
|
||||||
|
// TODO: Implement Race search
|
||||||
|
// 1. Parse requests into []Request
|
||||||
|
// 2. Call Searcher.Race()
|
||||||
// 3. Return []Result
|
// 3. Return []Result
|
||||||
results := make([]interface{}, len(requests))
|
results := make([]interface{}, len(requests))
|
||||||
for i := range requests {
|
for i := range requests {
|
||||||
|
|
|
||||||
484
agent/search/reference_test.go
Normal file
484
agent/search/reference_test.go
Normal file
|
|
@ -0,0 +1,484 @@
|
||||||
|
package search
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestBuildReferences(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
results []*types.Result
|
||||||
|
expected int
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "nil results",
|
||||||
|
results: nil,
|
||||||
|
expected: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty results",
|
||||||
|
results: []*types.Result{},
|
||||||
|
expected: 0,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "single result with items",
|
||||||
|
results: []*types.Result{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Query: "test query",
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{
|
||||||
|
CitationID: "ref_001",
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Weight: 0.6,
|
||||||
|
Score: 0.9,
|
||||||
|
Title: "Test Title",
|
||||||
|
Content: "Test content",
|
||||||
|
URL: "https://example.com",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
CitationID: "ref_002",
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Weight: 0.6,
|
||||||
|
Score: 0.8,
|
||||||
|
Title: "Test Title 2",
|
||||||
|
Content: "Test content 2",
|
||||||
|
URL: "https://example2.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
expected: 2,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "multiple results",
|
||||||
|
results: []*types.Result{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{CitationID: "ref_001", Type: types.SearchTypeWeb, Content: "Web content"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{CitationID: "ref_002", Type: types.SearchTypeKB, Content: "KB content"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{CitationID: "ref_003", Type: types.SearchTypeDB, Content: "DB content"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
expected: 3,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "result with nil items",
|
||||||
|
results: []*types.Result{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{CitationID: "ref_001", Content: "Content 1"},
|
||||||
|
nil,
|
||||||
|
{CitationID: "ref_002", Content: "Content 2"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
expected: 2,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nil result in slice",
|
||||||
|
results: []*types.Result{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{CitationID: "ref_001", Content: "Content"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
nil,
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{CitationID: "ref_002", Content: "Content 2"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
expected: 2,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
refs := BuildReferences(tt.results)
|
||||||
|
assert.Equal(t, tt.expected, len(refs))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildReferences_FieldMapping(t *testing.T) {
|
||||||
|
item := &types.ResultItem{
|
||||||
|
CitationID: "ref_001",
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Source: types.SourceHook,
|
||||||
|
Weight: 0.8,
|
||||||
|
Score: 0.95,
|
||||||
|
Title: "Test Title",
|
||||||
|
Content: "Test Content",
|
||||||
|
URL: "https://example.com",
|
||||||
|
}
|
||||||
|
|
||||||
|
results := []*types.Result{
|
||||||
|
{Items: []*types.ResultItem{item}},
|
||||||
|
}
|
||||||
|
|
||||||
|
refs := BuildReferences(results)
|
||||||
|
assert.Equal(t, 1, len(refs))
|
||||||
|
|
||||||
|
ref := refs[0]
|
||||||
|
assert.Equal(t, "ref_001", ref.ID)
|
||||||
|
assert.Equal(t, types.SearchTypeWeb, ref.Type)
|
||||||
|
assert.Equal(t, types.SourceHook, ref.Source)
|
||||||
|
assert.Equal(t, 0.8, ref.Weight)
|
||||||
|
assert.Equal(t, 0.95, ref.Score)
|
||||||
|
assert.Equal(t, "Test Title", ref.Title)
|
||||||
|
assert.Equal(t, "Test Content", ref.Content)
|
||||||
|
assert.Equal(t, "https://example.com", ref.URL)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatReferencesXML(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
refs []*types.Reference
|
||||||
|
contains []string
|
||||||
|
excludes []string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "nil refs",
|
||||||
|
refs: nil,
|
||||||
|
contains: []string{},
|
||||||
|
excludes: []string{"<references>"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty refs",
|
||||||
|
refs: []*types.Reference{},
|
||||||
|
contains: []string{},
|
||||||
|
excludes: []string{"<references>"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "single ref with all fields",
|
||||||
|
refs: []*types.Reference{
|
||||||
|
{
|
||||||
|
ID: "ref_001",
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Source: types.SourceUser,
|
||||||
|
Weight: 1.0,
|
||||||
|
Score: 0.9,
|
||||||
|
Title: "Test Title",
|
||||||
|
Content: "Test Content",
|
||||||
|
URL: "https://example.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
contains: []string{
|
||||||
|
"<references>",
|
||||||
|
"</references>",
|
||||||
|
`<ref id="ref_001" type="web" weight="1.0" source="user">`,
|
||||||
|
"</ref>",
|
||||||
|
"Test Title",
|
||||||
|
"Test Content",
|
||||||
|
"URL: https://example.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ref without title",
|
||||||
|
refs: []*types.Reference{
|
||||||
|
{
|
||||||
|
ID: "ref_001",
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Source: types.SourceHook,
|
||||||
|
Weight: 0.8,
|
||||||
|
Content: "Content without title",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
contains: []string{
|
||||||
|
`<ref id="ref_001" type="kb" weight="0.8" source="hook">`,
|
||||||
|
"Content without title",
|
||||||
|
},
|
||||||
|
excludes: []string{
|
||||||
|
"URL:",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ref without URL",
|
||||||
|
refs: []*types.Reference{
|
||||||
|
{
|
||||||
|
ID: "ref_001",
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Weight: 0.6,
|
||||||
|
Title: "DB Record",
|
||||||
|
Content: "Database content",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
contains: []string{
|
||||||
|
`<ref id="ref_001" type="db" weight="0.6" source="auto">`,
|
||||||
|
"DB Record",
|
||||||
|
"Database content",
|
||||||
|
},
|
||||||
|
excludes: []string{
|
||||||
|
"URL:",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "multiple refs",
|
||||||
|
refs: []*types.Reference{
|
||||||
|
{ID: "ref_001", Type: types.SearchTypeWeb, Source: types.SourceUser, Weight: 1.0, Content: "Content 1"},
|
||||||
|
{ID: "ref_002", Type: types.SearchTypeKB, Source: types.SourceHook, Weight: 0.8, Content: "Content 2"},
|
||||||
|
{ID: "ref_003", Type: types.SearchTypeDB, Source: types.SourceAuto, Weight: 0.6, Content: "Content 3"},
|
||||||
|
},
|
||||||
|
contains: []string{
|
||||||
|
"<references>",
|
||||||
|
"</references>",
|
||||||
|
`id="ref_001"`,
|
||||||
|
`id="ref_002"`,
|
||||||
|
`id="ref_003"`,
|
||||||
|
"Content 1",
|
||||||
|
"Content 2",
|
||||||
|
"Content 3",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nil ref in slice",
|
||||||
|
refs: []*types.Reference{
|
||||||
|
{ID: "ref_001", Type: types.SearchTypeWeb, Weight: 1.0, Content: "Content 1"},
|
||||||
|
nil,
|
||||||
|
{ID: "ref_002", Type: types.SearchTypeKB, Weight: 0.8, Content: "Content 2"},
|
||||||
|
},
|
||||||
|
contains: []string{
|
||||||
|
`id="ref_001"`,
|
||||||
|
`id="ref_002"`,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
xml := FormatReferencesXML(tt.refs)
|
||||||
|
|
||||||
|
for _, s := range tt.contains {
|
||||||
|
assert.Contains(t, xml, s, "expected XML to contain: %s", s)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, s := range tt.excludes {
|
||||||
|
assert.NotContains(t, xml, s, "expected XML to not contain: %s", s)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFormatReferencesXML_Structure(t *testing.T) {
|
||||||
|
refs := []*types.Reference{
|
||||||
|
{
|
||||||
|
ID: "ref_001",
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Source: types.SourceUser,
|
||||||
|
Weight: 1.0,
|
||||||
|
Title: "Title",
|
||||||
|
Content: "Content",
|
||||||
|
URL: "https://example.com",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
xml := FormatReferencesXML(refs)
|
||||||
|
|
||||||
|
// Check structure
|
||||||
|
assert.True(t, strings.HasPrefix(xml, "<references>\n"))
|
||||||
|
assert.True(t, strings.HasSuffix(xml, "</references>"))
|
||||||
|
assert.Contains(t, xml, "</ref>\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGetCitationPrompt(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
cfg *types.CitationConfig
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "nil config",
|
||||||
|
cfg: nil,
|
||||||
|
expected: DefaultCitationPrompt,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty config",
|
||||||
|
cfg: &types.CitationConfig{},
|
||||||
|
expected: DefaultCitationPrompt,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "config with custom prompt",
|
||||||
|
cfg: &types.CitationConfig{
|
||||||
|
CustomPrompt: "Custom citation instructions",
|
||||||
|
},
|
||||||
|
expected: "Custom citation instructions",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "config with empty custom prompt",
|
||||||
|
cfg: &types.CitationConfig{
|
||||||
|
CustomPrompt: "",
|
||||||
|
},
|
||||||
|
expected: DefaultCitationPrompt,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "config with format but no custom prompt",
|
||||||
|
cfg: &types.CitationConfig{
|
||||||
|
Format: "[{id}]",
|
||||||
|
},
|
||||||
|
expected: DefaultCitationPrompt,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
prompt := GetCitationPrompt(tt.cfg)
|
||||||
|
assert.Equal(t, tt.expected, prompt)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDefaultCitationPrompt(t *testing.T) {
|
||||||
|
// Verify default prompt contains key instructions
|
||||||
|
assert.Contains(t, DefaultCitationPrompt, "<references>")
|
||||||
|
assert.Contains(t, DefaultCitationPrompt, "id: Citation identifier")
|
||||||
|
assert.Contains(t, DefaultCitationPrompt, "type: Data type")
|
||||||
|
assert.Contains(t, DefaultCitationPrompt, "weight: Relevance weight")
|
||||||
|
assert.Contains(t, DefaultCitationPrompt, "source: Origin")
|
||||||
|
assert.Contains(t, DefaultCitationPrompt, `<a class="ref"`)
|
||||||
|
assert.Contains(t, DefaultCitationPrompt, "data-ref-id")
|
||||||
|
assert.Contains(t, DefaultCitationPrompt, "data-ref-type")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildReferenceContext(t *testing.T) {
|
||||||
|
results := []*types.Result{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{
|
||||||
|
CitationID: "ref_001",
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Weight: 0.6,
|
||||||
|
Title: "Test",
|
||||||
|
Content: "Content",
|
||||||
|
URL: "https://example.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("with nil config", func(t *testing.T) {
|
||||||
|
ctx := BuildReferenceContext(results, nil)
|
||||||
|
|
||||||
|
assert.NotNil(t, ctx)
|
||||||
|
assert.Equal(t, 1, len(ctx.References))
|
||||||
|
assert.Contains(t, ctx.XML, "<references>")
|
||||||
|
assert.Contains(t, ctx.XML, "ref_001")
|
||||||
|
assert.Equal(t, DefaultCitationPrompt, ctx.Prompt)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("with custom prompt config", func(t *testing.T) {
|
||||||
|
cfg := &types.CitationConfig{
|
||||||
|
CustomPrompt: "Custom prompt",
|
||||||
|
}
|
||||||
|
ctx := BuildReferenceContext(results, cfg)
|
||||||
|
|
||||||
|
assert.NotNil(t, ctx)
|
||||||
|
assert.Equal(t, "Custom prompt", ctx.Prompt)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("with empty results", func(t *testing.T) {
|
||||||
|
ctx := BuildReferenceContext([]*types.Result{}, nil)
|
||||||
|
|
||||||
|
assert.NotNil(t, ctx)
|
||||||
|
assert.Equal(t, 0, len(ctx.References))
|
||||||
|
assert.Equal(t, "", ctx.XML)
|
||||||
|
assert.Equal(t, DefaultCitationPrompt, ctx.Prompt)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildReferenceContext_Integration(t *testing.T) {
|
||||||
|
// Simulate a real-world scenario with multiple search types
|
||||||
|
results := []*types.Result{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Query: "AI developments",
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{
|
||||||
|
CitationID: "ref_001",
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Weight: 0.6,
|
||||||
|
Score: 0.95,
|
||||||
|
Title: "OpenAI Announces GPT-5",
|
||||||
|
Content: "OpenAI has announced the development of GPT-5...",
|
||||||
|
URL: "https://news.example.com/gpt5",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "AI developments",
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{
|
||||||
|
CitationID: "ref_002",
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Source: types.SourceHook,
|
||||||
|
Weight: 0.8,
|
||||||
|
Score: 0.88,
|
||||||
|
Title: "Internal AI Research Notes",
|
||||||
|
Content: "Our internal research on AI capabilities...",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "AI developments",
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{
|
||||||
|
CitationID: "ref_003",
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Source: types.SourceUser,
|
||||||
|
Weight: 1.0,
|
||||||
|
Score: 0.92,
|
||||||
|
Title: "Product: AI Assistant",
|
||||||
|
Content: "Name: AI Assistant\nPrice: $99\nCategory: Software",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := BuildReferenceContext(results, nil)
|
||||||
|
|
||||||
|
// Verify all references are included
|
||||||
|
assert.Equal(t, 3, len(ctx.References))
|
||||||
|
|
||||||
|
// Verify XML contains all references
|
||||||
|
assert.Contains(t, ctx.XML, "ref_001")
|
||||||
|
assert.Contains(t, ctx.XML, "ref_002")
|
||||||
|
assert.Contains(t, ctx.XML, "ref_003")
|
||||||
|
|
||||||
|
// Verify different source types are represented
|
||||||
|
assert.Contains(t, ctx.XML, `source="auto"`)
|
||||||
|
assert.Contains(t, ctx.XML, `source="hook"`)
|
||||||
|
assert.Contains(t, ctx.XML, `source="user"`)
|
||||||
|
|
||||||
|
// Verify different search types are represented
|
||||||
|
assert.Contains(t, ctx.XML, `type="web"`)
|
||||||
|
assert.Contains(t, ctx.XML, `type="kb"`)
|
||||||
|
assert.Contains(t, ctx.XML, `type="db"`)
|
||||||
|
}
|
||||||
77
agent/search/registry_test.go
Normal file
77
agent/search/registry_test.go
Normal file
|
|
@ -0,0 +1,77 @@
|
||||||
|
package search
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/agent/search/handlers/db"
|
||||||
|
"github.com/yaoapp/yao/agent/search/handlers/kb"
|
||||||
|
"github.com/yaoapp/yao/agent/search/handlers/web"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewRegistry(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
assert.NotNil(t, r)
|
||||||
|
assert.NotNil(t, r.handlers)
|
||||||
|
assert.Equal(t, 0, len(r.handlers))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistry_Register(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
|
||||||
|
// Register web handler
|
||||||
|
webHandler := web.NewHandler("builtin", nil)
|
||||||
|
r.Register(webHandler)
|
||||||
|
|
||||||
|
h, ok := r.Get(types.SearchTypeWeb)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Equal(t, types.SearchTypeWeb, h.Type())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistry_RegisterMultiple(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
|
||||||
|
// Register all handlers
|
||||||
|
r.Register(web.NewHandler("builtin", nil))
|
||||||
|
r.Register(kb.NewHandler(nil))
|
||||||
|
r.Register(db.NewHandler("builtin", nil))
|
||||||
|
|
||||||
|
// Verify all are registered
|
||||||
|
webH, ok := r.Get(types.SearchTypeWeb)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Equal(t, types.SearchTypeWeb, webH.Type())
|
||||||
|
|
||||||
|
kbH, ok := r.Get(types.SearchTypeKB)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Equal(t, types.SearchTypeKB, kbH.Type())
|
||||||
|
|
||||||
|
dbH, ok := r.Get(types.SearchTypeDB)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Equal(t, types.SearchTypeDB, dbH.Type())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistry_Get_NotFound(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
|
||||||
|
h, ok := r.Get(types.SearchTypeWeb)
|
||||||
|
assert.False(t, ok)
|
||||||
|
assert.Nil(t, h)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegistry_RegisterOverwrite(t *testing.T) {
|
||||||
|
r := NewRegistry()
|
||||||
|
|
||||||
|
// Register first handler
|
||||||
|
h1 := web.NewHandler("builtin", nil)
|
||||||
|
r.Register(h1)
|
||||||
|
|
||||||
|
// Register second handler (same type)
|
||||||
|
h2 := web.NewHandler("agent", nil)
|
||||||
|
r.Register(h2)
|
||||||
|
|
||||||
|
// Should get the second handler
|
||||||
|
h, ok := r.Get(types.SearchTypeWeb)
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.NotNil(t, h)
|
||||||
|
}
|
||||||
|
|
@ -84,8 +84,32 @@ func (s *Searcher) Search(ctx *context.Context, req *types.Request) (*types.Resu
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// SearchMultiple executes multiple searches in parallel
|
// All executes all searches and waits for all to complete (like Promise.all)
|
||||||
func (s *Searcher) SearchMultiple(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
func (s *Searcher) All(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
|
if len(reqs) == 0 {
|
||||||
|
return []*types.Result{}, nil
|
||||||
|
}
|
||||||
|
return s.parallelAll(ctx, reqs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Any returns as soon as any search succeeds with results (like Promise.any)
|
||||||
|
func (s *Searcher) Any(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
|
if len(reqs) == 0 {
|
||||||
|
return []*types.Result{}, nil
|
||||||
|
}
|
||||||
|
return s.parallelAny(ctx, reqs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Race returns as soon as any search completes (like Promise.race)
|
||||||
|
func (s *Searcher) Race(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
|
if len(reqs) == 0 {
|
||||||
|
return []*types.Result{}, nil
|
||||||
|
}
|
||||||
|
return s.parallelRace(ctx, reqs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// parallelAll executes all searches and waits for all to complete (like Promise.all)
|
||||||
|
func (s *Searcher) parallelAll(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
results := make([]*types.Result, len(reqs))
|
results := make([]*types.Result, len(reqs))
|
||||||
var wg sync.WaitGroup
|
var wg sync.WaitGroup
|
||||||
var mu sync.Mutex
|
var mu sync.Mutex
|
||||||
|
|
@ -105,6 +129,103 @@ func (s *Searcher) SearchMultiple(ctx *context.Context, reqs []*types.Request) (
|
||||||
return results, nil
|
return results, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// parallelAny returns as soon as any search succeeds (has results) (like Promise.any)
|
||||||
|
// Other searches continue in background but results are discarded
|
||||||
|
func (s *Searcher) parallelAny(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
|
results := make([]*types.Result, len(reqs))
|
||||||
|
resultChan := make(chan struct {
|
||||||
|
idx int
|
||||||
|
result *types.Result
|
||||||
|
}, len(reqs))
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
done := make(chan struct{})
|
||||||
|
|
||||||
|
for i, req := range reqs {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(idx int, r *types.Request) {
|
||||||
|
defer wg.Done()
|
||||||
|
result, _ := s.Search(ctx, r)
|
||||||
|
select {
|
||||||
|
case resultChan <- struct {
|
||||||
|
idx int
|
||||||
|
result *types.Result
|
||||||
|
}{idx, result}:
|
||||||
|
case <-done:
|
||||||
|
// Already found a successful result, discard this one
|
||||||
|
}
|
||||||
|
}(i, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close channel when all goroutines complete
|
||||||
|
go func() {
|
||||||
|
wg.Wait()
|
||||||
|
close(resultChan)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Collect results until we find one with items (success)
|
||||||
|
var mu sync.Mutex
|
||||||
|
for res := range resultChan {
|
||||||
|
mu.Lock()
|
||||||
|
results[res.idx] = res.result
|
||||||
|
// Check if this result has items (success = has results and no error)
|
||||||
|
if res.result != nil && len(res.result.Items) > 0 && res.result.Error == "" {
|
||||||
|
mu.Unlock()
|
||||||
|
close(done) // Signal other goroutines to stop sending
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// No successful result found, return all results
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// parallelRace returns as soon as any search completes (like Promise.race)
|
||||||
|
// Returns immediately when first result arrives, regardless of success/failure
|
||||||
|
func (s *Searcher) parallelRace(ctx *context.Context, reqs []*types.Request) ([]*types.Result, error) {
|
||||||
|
results := make([]*types.Result, len(reqs))
|
||||||
|
resultChan := make(chan struct {
|
||||||
|
idx int
|
||||||
|
result *types.Result
|
||||||
|
}, len(reqs))
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
done := make(chan struct{})
|
||||||
|
|
||||||
|
for i, req := range reqs {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(idx int, r *types.Request) {
|
||||||
|
defer wg.Done()
|
||||||
|
result, _ := s.Search(ctx, r)
|
||||||
|
select {
|
||||||
|
case resultChan <- struct {
|
||||||
|
idx int
|
||||||
|
result *types.Result
|
||||||
|
}{idx, result}:
|
||||||
|
case <-done:
|
||||||
|
// Already got first result, discard this one
|
||||||
|
}
|
||||||
|
}(i, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close channel when all goroutines complete
|
||||||
|
go func() {
|
||||||
|
wg.Wait()
|
||||||
|
close(resultChan)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Return immediately when first result arrives
|
||||||
|
if res, ok := <-resultChan; ok {
|
||||||
|
results[res.idx] = res.result
|
||||||
|
close(done) // Signal other goroutines to stop sending
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// No results (shouldn't happen with valid requests)
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
// BuildReferences converts search results to unified Reference format
|
// BuildReferences converts search results to unified Reference format
|
||||||
func (s *Searcher) BuildReferences(results []*types.Result) []*types.Reference {
|
func (s *Searcher) BuildReferences(results []*types.Result) []*types.Reference {
|
||||||
return BuildReferences(results)
|
return BuildReferences(results)
|
||||||
|
|
|
||||||
402
agent/search/search_test.go
Normal file
402
agent/search/search_test.go
Normal file
|
|
@ -0,0 +1,402 @@
|
||||||
|
package search
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNew(t *testing.T) {
|
||||||
|
t.Run("with nil config and uses", func(t *testing.T) {
|
||||||
|
s := New(nil, nil)
|
||||||
|
assert.NotNil(t, s)
|
||||||
|
assert.NotNil(t, s.config)
|
||||||
|
assert.NotNil(t, s.handlers)
|
||||||
|
assert.NotNil(t, s.citation)
|
||||||
|
assert.Equal(t, 3, len(s.handlers)) // web, kb, db
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("with config", func(t *testing.T) {
|
||||||
|
cfg := &types.Config{
|
||||||
|
Web: &types.WebConfig{
|
||||||
|
Provider: "tavily",
|
||||||
|
MaxResults: 10,
|
||||||
|
},
|
||||||
|
KB: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
Threshold: 0.8,
|
||||||
|
},
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
MaxResults: 20,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, nil)
|
||||||
|
assert.NotNil(t, s)
|
||||||
|
assert.Equal(t, cfg, s.config)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("with uses", func(t *testing.T) {
|
||||||
|
uses := &Uses{
|
||||||
|
Search: "builtin",
|
||||||
|
Web: "builtin",
|
||||||
|
Keyword: "builtin",
|
||||||
|
QueryDSL: "builtin",
|
||||||
|
Rerank: "builtin",
|
||||||
|
}
|
||||||
|
s := New(nil, uses)
|
||||||
|
assert.NotNil(t, s)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Search_UnsupportedType(t *testing.T) {
|
||||||
|
s := New(nil, nil)
|
||||||
|
|
||||||
|
req := &types.Request{
|
||||||
|
Type: "unsupported",
|
||||||
|
Query: "test",
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, "unsupported search type", result.Error)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Search_Web(t *testing.T) {
|
||||||
|
// Note: This test uses skeleton handlers that return empty results
|
||||||
|
// Real tests with actual API calls are in handlers/web/*_test.go
|
||||||
|
cfg := &types.Config{
|
||||||
|
Web: &types.WebConfig{
|
||||||
|
Provider: "tavily",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, &Uses{Web: "builtin"})
|
||||||
|
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Query: "test query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, types.SearchTypeWeb, result.Type)
|
||||||
|
assert.Equal(t, "test query", result.Query)
|
||||||
|
// Note: actual result depends on API key availability
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Search_KB(t *testing.T) {
|
||||||
|
cfg := &types.Config{
|
||||||
|
KB: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
Threshold: 0.7,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, nil)
|
||||||
|
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test query",
|
||||||
|
Source: types.SourceHook,
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, types.SearchTypeKB, result.Type)
|
||||||
|
assert.Equal(t, "test query", result.Query)
|
||||||
|
assert.Equal(t, types.SourceHook, result.Source)
|
||||||
|
// Skeleton returns empty items
|
||||||
|
assert.Equal(t, 0, len(result.Items))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Search_DB(t *testing.T) {
|
||||||
|
cfg := &types.Config{
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
MaxResults: 20,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, &Uses{QueryDSL: "builtin"})
|
||||||
|
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "find products under $100",
|
||||||
|
Source: types.SourceUser,
|
||||||
|
Models: []string{"product"},
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, types.SearchTypeDB, result.Type)
|
||||||
|
assert.Equal(t, "find products under $100", result.Query)
|
||||||
|
assert.Equal(t, types.SourceUser, result.Source)
|
||||||
|
// Skeleton returns empty items
|
||||||
|
assert.Equal(t, 0, len(result.Items))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Search_WeightAssignment(t *testing.T) {
|
||||||
|
cfg := &types.Config{
|
||||||
|
KB: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
Weights: &types.WeightsConfig{
|
||||||
|
User: 1.0,
|
||||||
|
Hook: 0.8,
|
||||||
|
Auto: 0.6,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, nil)
|
||||||
|
|
||||||
|
// Test with different sources
|
||||||
|
sources := []struct {
|
||||||
|
source types.SourceType
|
||||||
|
weight float64
|
||||||
|
}{
|
||||||
|
{types.SourceUser, 1.0},
|
||||||
|
{types.SourceHook, 0.8},
|
||||||
|
{types.SourceAuto, 0.6},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range sources {
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test",
|
||||||
|
Source: tc.source,
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
}
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
// Items are empty in skeleton, so weight assignment can't be verified here
|
||||||
|
// This test ensures the code path works without error
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_All(t *testing.T) {
|
||||||
|
cfg := &types.Config{
|
||||||
|
KB: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, nil)
|
||||||
|
|
||||||
|
reqs := []*types.Request{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "KB query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "DB query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test All() - waits for all searches to complete (like Promise.all)
|
||||||
|
results, err := s.All(nil, reqs)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, 2, len(results))
|
||||||
|
|
||||||
|
// Verify each result corresponds to its request
|
||||||
|
assert.Equal(t, types.SearchTypeKB, results[0].Type)
|
||||||
|
assert.Equal(t, "KB query", results[0].Query)
|
||||||
|
|
||||||
|
assert.Equal(t, types.SearchTypeDB, results[1].Type)
|
||||||
|
assert.Equal(t, "DB query", results[1].Query)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Any(t *testing.T) {
|
||||||
|
cfg := &types.Config{
|
||||||
|
KB: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, nil)
|
||||||
|
|
||||||
|
reqs := []*types.Request{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "KB query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "DB query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test Any() - returns when first search has results (like Promise.any)
|
||||||
|
// Note: With skeleton handlers returning empty results, this will wait for all
|
||||||
|
results, err := s.Any(nil, reqs)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, 2, len(results))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Race(t *testing.T) {
|
||||||
|
cfg := &types.Config{
|
||||||
|
KB: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, nil)
|
||||||
|
|
||||||
|
reqs := []*types.Request{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "KB query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "DB query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Models: []string{"product"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test Race() - returns when first search completes (like Promise.race)
|
||||||
|
results, err := s.Race(nil, reqs)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
// At least one result should be set
|
||||||
|
hasResult := false
|
||||||
|
for _, r := range results {
|
||||||
|
if r != nil {
|
||||||
|
hasResult = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
assert.True(t, hasResult)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_All_Empty(t *testing.T) {
|
||||||
|
s := New(nil, nil)
|
||||||
|
|
||||||
|
results, err := s.All(nil, []*types.Request{})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, 0, len(results))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Any_Empty(t *testing.T) {
|
||||||
|
s := New(nil, nil)
|
||||||
|
|
||||||
|
results, err := s.Any(nil, []*types.Request{})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, 0, len(results))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_Race_Empty(t *testing.T) {
|
||||||
|
s := New(nil, nil)
|
||||||
|
|
||||||
|
results, err := s.Race(nil, []*types.Request{})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, 0, len(results))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_All_ManyRequests(t *testing.T) {
|
||||||
|
cfg := &types.Config{
|
||||||
|
KB: &types.KBConfig{
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
s := New(cfg, nil)
|
||||||
|
|
||||||
|
// Create multiple requests to test parallel execution
|
||||||
|
reqs := make([]*types.Request, 10)
|
||||||
|
for i := 0; i < 10; i++ {
|
||||||
|
reqs[i] = &types.Request{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Query: "test query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Collections: []string{"docs"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
results, err := s.All(nil, reqs)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Equal(t, 10, len(results))
|
||||||
|
|
||||||
|
// All results should be valid
|
||||||
|
for _, result := range results {
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
assert.Equal(t, types.SearchTypeKB, result.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_BuildReferences(t *testing.T) {
|
||||||
|
s := New(nil, nil)
|
||||||
|
|
||||||
|
results := []*types.Result{
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{
|
||||||
|
CitationID: "ref_001",
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Weight: 0.6,
|
||||||
|
Title: "Web Result",
|
||||||
|
Content: "Web content",
|
||||||
|
URL: "https://example.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Items: []*types.ResultItem{
|
||||||
|
{
|
||||||
|
CitationID: "ref_002",
|
||||||
|
Type: types.SearchTypeKB,
|
||||||
|
Source: types.SourceHook,
|
||||||
|
Weight: 0.8,
|
||||||
|
Title: "KB Result",
|
||||||
|
Content: "KB content",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
refs := s.BuildReferences(results)
|
||||||
|
assert.Equal(t, 2, len(refs))
|
||||||
|
assert.Equal(t, "ref_001", refs[0].ID)
|
||||||
|
assert.Equal(t, "ref_002", refs[1].ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSearcher_CitationGeneration(t *testing.T) {
|
||||||
|
s := New(nil, nil)
|
||||||
|
|
||||||
|
// Reset citation generator for predictable IDs
|
||||||
|
s.citation.Reset()
|
||||||
|
|
||||||
|
// Note: This test would need actual results with items to verify citation generation
|
||||||
|
// The skeleton handlers return empty items, so we test the citation generator directly
|
||||||
|
|
||||||
|
id1 := s.citation.Next()
|
||||||
|
id2 := s.citation.Next()
|
||||||
|
id3 := s.citation.Next()
|
||||||
|
|
||||||
|
assert.Equal(t, "ref_001", id1)
|
||||||
|
assert.Equal(t, "ref_002", id2)
|
||||||
|
assert.Equal(t, "ref_003", id3)
|
||||||
|
}
|
||||||
335
agent/search/search_web_test.go
Normal file
335
agent/search/search_web_test.go
Normal file
|
|
@ -0,0 +1,335 @@
|
||||||
|
package search_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"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
|
"github.com/yaoapp/yao/agent/testutils"
|
||||||
|
)
|
||||||
|
|
||||||
|
// =============================================================================
|
||||||
|
// Web Search Integration Tests - Single Search
|
||||||
|
// =============================================================================
|
||||||
|
|
||||||
|
// TestWebSearch_Tavily tests web search using Tavily provider via assistant config
|
||||||
|
func TestWebSearch_Tavily(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Load the web-tavily test assistant
|
||||||
|
ast, err := assistant.LoadPath("/assistants/tests/web-tavily")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, ast.Search)
|
||||||
|
require.NotNil(t, ast.Search.Web)
|
||||||
|
|
||||||
|
// Verify assistant config
|
||||||
|
assert.Equal(t, "tavily", ast.Search.Web.Provider)
|
||||||
|
|
||||||
|
// Create Searcher with assistant's config
|
||||||
|
uses := &search.Uses{Web: "builtin"}
|
||||||
|
s := search.New(ast.Search, uses)
|
||||||
|
|
||||||
|
// Execute search
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Query: "Yao App Engine low-code platform",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Limit: 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
require.Empty(t, result.Error, "Search should succeed, got error: %s", result.Error)
|
||||||
|
|
||||||
|
// Verify results
|
||||||
|
assert.NotEmpty(t, result.Items, "Should have search results")
|
||||||
|
for _, item := range result.Items {
|
||||||
|
assert.NotEmpty(t, item.CitationID, "Each item should have citation ID")
|
||||||
|
assert.NotEmpty(t, item.Content, "Each item should have content")
|
||||||
|
t.Logf(" [%s] %s - %s", item.CitationID, item.Title, item.URL)
|
||||||
|
}
|
||||||
|
t.Logf("Tavily search returned %d results", len(result.Items))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebSearch_Serper tests web search using Serper provider via assistant config
|
||||||
|
func TestWebSearch_Serper(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Load the web-serper test assistant
|
||||||
|
ast, err := assistant.LoadPath("/assistants/tests/web-serper")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, ast.Search)
|
||||||
|
require.NotNil(t, ast.Search.Web)
|
||||||
|
|
||||||
|
// Verify assistant config
|
||||||
|
assert.Equal(t, "serper", ast.Search.Web.Provider)
|
||||||
|
|
||||||
|
// Create Searcher with assistant's config
|
||||||
|
uses := &search.Uses{Web: "builtin"}
|
||||||
|
s := search.New(ast.Search, uses)
|
||||||
|
|
||||||
|
// Execute search
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Query: "Go programming language concurrency",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Limit: 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
require.Empty(t, result.Error, "Search should succeed, got error: %s", result.Error)
|
||||||
|
|
||||||
|
// Verify results
|
||||||
|
assert.NotEmpty(t, result.Items, "Should have search results")
|
||||||
|
for _, item := range result.Items {
|
||||||
|
assert.NotEmpty(t, item.CitationID, "Each item should have citation ID")
|
||||||
|
t.Logf(" [%s] %s - %s", item.CitationID, item.Title, item.URL)
|
||||||
|
}
|
||||||
|
t.Logf("Serper search returned %d results", len(result.Items))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebSearch_SerpAPI tests web search using SerpAPI provider via assistant config
|
||||||
|
func TestWebSearch_SerpAPI(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Load the web-serpapi test assistant
|
||||||
|
ast, err := assistant.LoadPath("/assistants/tests/web-serpapi")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, ast.Search)
|
||||||
|
require.NotNil(t, ast.Search.Web)
|
||||||
|
|
||||||
|
// Verify assistant config
|
||||||
|
assert.Equal(t, "serpapi", ast.Search.Web.Provider)
|
||||||
|
|
||||||
|
// Create Searcher with assistant's config
|
||||||
|
uses := &search.Uses{Web: "builtin"}
|
||||||
|
s := search.New(ast.Search, uses)
|
||||||
|
|
||||||
|
// Execute search
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Query: "Kubernetes container orchestration",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Limit: 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
require.Empty(t, result.Error, "Search should succeed, got error: %s", result.Error)
|
||||||
|
|
||||||
|
// Verify results
|
||||||
|
assert.NotEmpty(t, result.Items, "Should have search results")
|
||||||
|
for _, item := range result.Items {
|
||||||
|
assert.NotEmpty(t, item.CitationID, "Each item should have citation ID")
|
||||||
|
t.Logf(" [%s] %s - %s", item.CitationID, item.Title, item.URL)
|
||||||
|
}
|
||||||
|
t.Logf("SerpAPI search returned %d results", len(result.Items))
|
||||||
|
}
|
||||||
|
|
||||||
|
// =============================================================================
|
||||||
|
// Web Search Integration Tests - Parallel Search
|
||||||
|
// =============================================================================
|
||||||
|
|
||||||
|
// TestWebSearch_All tests parallel web search with All() - like Promise.all
|
||||||
|
func TestWebSearch_All(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Load the web-tavily test assistant
|
||||||
|
ast, err := assistant.LoadPath("/assistants/tests/web-tavily")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, ast.Search)
|
||||||
|
|
||||||
|
// Create Searcher
|
||||||
|
uses := &search.Uses{Web: "builtin"}
|
||||||
|
s := search.New(ast.Search, uses)
|
||||||
|
|
||||||
|
// Multiple queries
|
||||||
|
reqs := []*types.Request{
|
||||||
|
{Type: types.SearchTypeWeb, Query: "artificial intelligence", Source: types.SourceAuto, Limit: 3},
|
||||||
|
{Type: types.SearchTypeWeb, Query: "machine learning", Source: types.SourceAuto, Limit: 3},
|
||||||
|
{Type: types.SearchTypeWeb, Query: "deep learning", Source: types.SourceAuto, Limit: 3},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute parallel search with All() - waits for all searches to complete
|
||||||
|
results, err := s.All(nil, reqs)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, results, 3, "Should have 3 results")
|
||||||
|
|
||||||
|
// Verify all results
|
||||||
|
for i, result := range results {
|
||||||
|
require.NotNil(t, result, "Result %d should not be nil", i)
|
||||||
|
if result.Error == "" {
|
||||||
|
assert.NotEmpty(t, result.Items, "Result %d should have items", i)
|
||||||
|
t.Logf("Query '%s': %d results", reqs[i].Query, len(result.Items))
|
||||||
|
} else {
|
||||||
|
t.Logf("Query '%s': error - %s", reqs[i].Query, result.Error)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebSearch_Any tests parallel web search with Any() - like Promise.any
|
||||||
|
func TestWebSearch_Any(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Load the web-serper test assistant
|
||||||
|
ast, err := assistant.LoadPath("/assistants/tests/web-serper")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, ast.Search)
|
||||||
|
|
||||||
|
// Create Searcher
|
||||||
|
uses := &search.Uses{Web: "builtin"}
|
||||||
|
s := search.New(ast.Search, uses)
|
||||||
|
|
||||||
|
// Multiple queries
|
||||||
|
reqs := []*types.Request{
|
||||||
|
{Type: types.SearchTypeWeb, Query: "golang channels", Source: types.SourceAuto, Limit: 3},
|
||||||
|
{Type: types.SearchTypeWeb, Query: "rust ownership", Source: types.SourceAuto, Limit: 3},
|
||||||
|
{Type: types.SearchTypeWeb, Query: "python asyncio", Source: types.SourceAuto, Limit: 3},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute parallel search with Any() - returns when first search succeeds
|
||||||
|
results, err := s.Any(nil, reqs)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Any() returns as soon as any search succeeds
|
||||||
|
hasSuccess := false
|
||||||
|
for _, result := range results {
|
||||||
|
if result != nil && len(result.Items) > 0 && result.Error == "" {
|
||||||
|
hasSuccess = true
|
||||||
|
t.Logf("First success: '%s' with %d results", result.Query, len(result.Items))
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
assert.True(t, hasSuccess, "At least one search should succeed")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestWebSearch_Race tests parallel web search with Race() - like Promise.race
|
||||||
|
func TestWebSearch_Race(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Load the web-tavily test assistant
|
||||||
|
ast, err := assistant.LoadPath("/assistants/tests/web-tavily")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, ast.Search)
|
||||||
|
|
||||||
|
// Create Searcher
|
||||||
|
uses := &search.Uses{Web: "builtin"}
|
||||||
|
s := search.New(ast.Search, uses)
|
||||||
|
|
||||||
|
// Multiple queries
|
||||||
|
reqs := []*types.Request{
|
||||||
|
{Type: types.SearchTypeWeb, Query: "docker containers", Source: types.SourceAuto, Limit: 3},
|
||||||
|
{Type: types.SearchTypeWeb, Query: "kubernetes pods", Source: types.SourceAuto, Limit: 3},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute parallel search with Race() - returns when first search completes
|
||||||
|
results, err := s.Race(nil, reqs)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
// Race() returns immediately when first result arrives
|
||||||
|
hasResult := false
|
||||||
|
for _, result := range results {
|
||||||
|
if result != nil {
|
||||||
|
hasResult = true
|
||||||
|
t.Logf("First to complete: '%s'", result.Query)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
assert.True(t, hasResult, "Should have at least one result")
|
||||||
|
}
|
||||||
|
|
||||||
|
// =============================================================================
|
||||||
|
// Web Search - Citation and Reference Tests
|
||||||
|
// =============================================================================
|
||||||
|
|
||||||
|
// TestWebSearch_BuildReferences tests building references from web search results
|
||||||
|
func TestWebSearch_BuildReferences(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Load the web-tavily test assistant
|
||||||
|
ast, err := assistant.LoadPath("/assistants/tests/web-tavily")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, ast.Search)
|
||||||
|
|
||||||
|
// Create Searcher with weights config
|
||||||
|
uses := &search.Uses{Web: "builtin"}
|
||||||
|
s := search.New(ast.Search, uses)
|
||||||
|
|
||||||
|
// Execute search
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Query: "OpenAI GPT-4",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Limit: 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
require.Empty(t, result.Error, "Search should succeed")
|
||||||
|
require.NotEmpty(t, result.Items, "Should have results")
|
||||||
|
|
||||||
|
// Build references
|
||||||
|
refs := s.BuildReferences([]*types.Result{result})
|
||||||
|
assert.NotEmpty(t, refs, "Should have references")
|
||||||
|
|
||||||
|
for _, ref := range refs {
|
||||||
|
assert.NotEmpty(t, ref.ID, "Reference should have ID")
|
||||||
|
assert.Equal(t, types.SearchTypeWeb, ref.Type, "Reference type should be web")
|
||||||
|
assert.Equal(t, types.SourceAuto, ref.Source, "Reference source should be auto")
|
||||||
|
t.Logf(" Ref: %s - %s (weight: %.2f)", ref.ID, ref.Title, ref.Weight)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// =============================================================================
|
||||||
|
// Web Search - Error Handling Tests
|
||||||
|
// =============================================================================
|
||||||
|
|
||||||
|
// TestWebSearch_SiteRestriction tests web search with site restriction
|
||||||
|
func TestWebSearch_SiteRestriction(t *testing.T) {
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Load the web-serper test assistant
|
||||||
|
ast, err := assistant.LoadPath("/assistants/tests/web-serper")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, ast.Search)
|
||||||
|
|
||||||
|
// Create Searcher
|
||||||
|
uses := &search.Uses{Web: "builtin"}
|
||||||
|
s := search.New(ast.Search, uses)
|
||||||
|
|
||||||
|
// Execute search with site restriction
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeWeb,
|
||||||
|
Query: "yao-app-engine",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Sites: []string{"github.com"},
|
||||||
|
Limit: 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := s.Search(nil, req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
if result.Error == "" && len(result.Items) > 0 {
|
||||||
|
// Log results
|
||||||
|
for _, item := range result.Items {
|
||||||
|
t.Logf(" %s - %s", item.Title, item.URL)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue