Implement Context-Aware Search Functionality and Enhance QueryDSL Generation
- Refactored the Search method to support context-aware execution, allowing handlers to utilize context when performing searches. - Introduced a new SearchWithContext method in the handler interface to facilitate context-based search operations. - Updated the DB handler to implement the context-aware search, ensuring proper QueryDSL generation and execution. - Enhanced test cases to validate the new context requirements and scenarios for database searches, improving error handling and robustness. - Added scenario type support for QueryDSL generation, allowing for more complex query handling. - Updated documentation to reflect the new context handling and scenario features in search operations.
This commit is contained in:
parent
4abc7e41ef
commit
2f1e4b9f33
12 changed files with 1363 additions and 353 deletions
|
|
@ -1,8 +1,15 @@
|
||||||
package db
|
package db
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/yaoapp/gou/model"
|
||||||
|
"github.com/yaoapp/gou/query"
|
||||||
|
"github.com/yaoapp/gou/query/gou"
|
||||||
|
agentContext "github.com/yaoapp/yao/agent/context"
|
||||||
|
"github.com/yaoapp/yao/agent/search/nlp/querydsl"
|
||||||
"github.com/yaoapp/yao/agent/search/types"
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -23,8 +30,13 @@ func (h *Handler) Type() types.SearchType {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Search converts NL to QueryDSL and executes
|
// Search converts NL to QueryDSL and executes
|
||||||
// TODO: Implement actual QueryDSL generation and model query logic
|
// Note: This method doesn't have context, use SearchWithContext for full functionality
|
||||||
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 DB search with context (required for QueryDSL generation)
|
||||||
|
func (h *Handler) SearchWithContext(ctx *agentContext.Context, req *types.Request) (*types.Result, error) {
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
|
|
||||||
// Validate request
|
// Validate request
|
||||||
|
|
@ -41,13 +53,13 @@ func (h *Handler) Search(req *types.Request) (*types.Result, error) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get models from request or config
|
// Get models from request or config
|
||||||
models := req.Models
|
modelIDs := req.Models
|
||||||
if len(models) == 0 && h.config != nil {
|
if len(modelIDs) == 0 && h.config != nil {
|
||||||
models = h.config.Models
|
modelIDs = h.config.Models
|
||||||
}
|
}
|
||||||
|
|
||||||
// If no models specified, return empty result
|
// If no models specified, return empty result
|
||||||
if len(models) == 0 {
|
if len(modelIDs) == 0 {
|
||||||
return &types.Result{
|
return &types.Result{
|
||||||
Type: types.SearchTypeDB,
|
Type: types.SearchTypeDB,
|
||||||
Query: req.Query,
|
Query: req.Query,
|
||||||
|
|
@ -55,6 +67,7 @@ func (h *Handler) Search(req *types.Request) (*types.Result, error) {
|
||||||
Items: []*types.ResultItem{},
|
Items: []*types.ResultItem{},
|
||||||
Total: 0,
|
Total: 0,
|
||||||
Duration: time.Since(start).Milliseconds(),
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
Error: "no models specified",
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -67,27 +80,300 @@ func (h *Handler) Search(req *types.Request) (*types.Result, error) {
|
||||||
maxResults = 20 // default
|
maxResults = 20 // default
|
||||||
}
|
}
|
||||||
|
|
||||||
// TODO: Implement actual DB search
|
// Context is required for QueryDSL generation
|
||||||
// 1. Get model schemas for specified models
|
if ctx == nil {
|
||||||
// 2. Generate QueryDSL from natural language query using uses.querydsl mode:
|
return &types.Result{
|
||||||
// - "builtin": template-based generation
|
Type: types.SearchTypeDB,
|
||||||
// - "<assistant-id>": delegate to LLM assistant
|
Query: req.Query,
|
||||||
// - "mcp:<server>.<tool>": call external MCP tool
|
Source: req.Source,
|
||||||
// 3. Execute QueryDSL on each model
|
Items: []*types.ResultItem{},
|
||||||
// 4. Format results and return
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
Error: "context is required for DB search",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
// For now, return empty result (skeleton)
|
// 1. Load all models and build combined schema
|
||||||
result := &types.Result{
|
models := make(map[string]*model.Model)
|
||||||
|
schemas := make([]map[string]interface{}, 0, len(modelIDs))
|
||||||
|
|
||||||
|
for _, modelID := range modelIDs {
|
||||||
|
mod := model.Select(modelID)
|
||||||
|
if mod == nil {
|
||||||
|
continue // Skip non-existent models
|
||||||
|
}
|
||||||
|
models[modelID] = mod
|
||||||
|
schemas = append(schemas, h.buildModelSchema(mod))
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(schemas) == 0 {
|
||||||
|
return &types.Result{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: req.Query,
|
||||||
|
Source: req.Source,
|
||||||
|
Items: []*types.ResultItem{},
|
||||||
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
Error: "no valid models found",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. Generate QueryDSL with all schemas
|
||||||
|
generator := querydsl.NewGenerator(h.usesQueryDSL, nil)
|
||||||
|
input := &querydsl.Input{
|
||||||
|
Query: req.Query,
|
||||||
|
ModelIDs: modelIDs,
|
||||||
|
Scenario: req.Scenario, // Pass scenario: filter, aggregation, join, complex
|
||||||
|
Limit: maxResults,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build schema input: single schema or array of schemas
|
||||||
|
var schemaInput interface{}
|
||||||
|
if len(schemas) == 1 {
|
||||||
|
schemaInput = schemas[0]
|
||||||
|
} else {
|
||||||
|
schemaInput = schemas
|
||||||
|
}
|
||||||
|
|
||||||
|
input.ExtraParams = map[string]interface{}{
|
||||||
|
"schema": schemaInput,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := generator.Generate(ctx, input)
|
||||||
|
if err != nil {
|
||||||
|
return &types.Result{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: req.Query,
|
||||||
|
Source: req.Source,
|
||||||
|
Items: []*types.ResultItem{},
|
||||||
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
Error: fmt.Sprintf("QueryDSL generation failed: %v", err),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if result == nil || result.DSL == nil {
|
||||||
|
return &types.Result{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: req.Query,
|
||||||
|
Source: req.Source,
|
||||||
|
Items: []*types.ResultItem{},
|
||||||
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
Error: "no QueryDSL generated",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. Merge preset conditions into generated DSL
|
||||||
|
h.mergeDSLConditions(result.DSL, req)
|
||||||
|
|
||||||
|
// 4. Execute QueryDSL using gou query engine
|
||||||
|
records, err := h.executeDSL(result.DSL)
|
||||||
|
if err != nil {
|
||||||
|
return &types.Result{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: req.Query,
|
||||||
|
Source: req.Source,
|
||||||
|
Items: []*types.ResultItem{},
|
||||||
|
Total: 0,
|
||||||
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
Error: fmt.Sprintf("query execution failed: %v", err),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 5. Determine the primary model for result formatting
|
||||||
|
// Use the "from" table from DSL, or first model
|
||||||
|
primaryModelID := modelIDs[0]
|
||||||
|
if result.DSL.From != nil && result.DSL.From.Name != "" {
|
||||||
|
// Find model by table name
|
||||||
|
for id, mod := range models {
|
||||||
|
if mod.MetaData.Table.Name == result.DSL.From.Name {
|
||||||
|
primaryModelID = id
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
primaryModel := models[primaryModelID]
|
||||||
|
if primaryModel == nil {
|
||||||
|
primaryModel = model.Select(primaryModelID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 6. Convert records to ResultItems
|
||||||
|
items := h.convertToResultItems(records, primaryModelID, primaryModel, req.Source)
|
||||||
|
|
||||||
|
// Apply limit
|
||||||
|
if len(items) > maxResults {
|
||||||
|
items = items[:maxResults]
|
||||||
|
}
|
||||||
|
|
||||||
|
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: items,
|
||||||
Total: 0,
|
Total: len(items),
|
||||||
Duration: time.Since(start).Milliseconds(),
|
Duration: time.Since(start).Milliseconds(),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// mergeDSLConditions merges preset conditions from request into generated DSL
|
||||||
|
func (h *Handler) mergeDSLConditions(dsl *gou.QueryDSL, req *types.Request) {
|
||||||
|
if dsl == nil {
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Store maxResults for later use
|
// Merge preset Wheres (prepend to ensure they take priority)
|
||||||
_ = maxResults
|
if len(req.Wheres) > 0 {
|
||||||
|
dsl.Wheres = append(req.Wheres, dsl.Wheres...)
|
||||||
|
}
|
||||||
|
|
||||||
return result, nil
|
// Merge preset Orders (prepend to ensure they take priority)
|
||||||
|
if len(req.Orders) > 0 {
|
||||||
|
dsl.Orders = append(req.Orders, dsl.Orders...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Merge preset Select fields
|
||||||
|
if len(req.Select) > 0 {
|
||||||
|
// Convert string fields to Expression
|
||||||
|
selectExprs := make([]gou.Expression, 0, len(req.Select))
|
||||||
|
for _, field := range req.Select {
|
||||||
|
selectExprs = append(selectExprs, gou.Expression{Field: field})
|
||||||
|
}
|
||||||
|
// If DSL has no select, use preset; otherwise merge
|
||||||
|
if len(dsl.Select) == 0 {
|
||||||
|
dsl.Select = selectExprs
|
||||||
|
} else {
|
||||||
|
// Prepend preset fields
|
||||||
|
dsl.Select = append(selectExprs, dsl.Select...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ensure limit is set
|
||||||
|
if dsl.Limit == 0 && req.Limit > 0 {
|
||||||
|
dsl.Limit = req.Limit
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildModelSchema builds a simplified schema for QueryDSL generator
|
||||||
|
func (h *Handler) buildModelSchema(mod *model.Model) map[string]interface{} {
|
||||||
|
columns := make([]map[string]interface{}, 0, len(mod.Columns))
|
||||||
|
for _, col := range mod.Columns {
|
||||||
|
colInfo := map[string]interface{}{
|
||||||
|
"name": col.Name,
|
||||||
|
"type": col.Type,
|
||||||
|
}
|
||||||
|
if col.Label != "" {
|
||||||
|
colInfo["label"] = col.Label
|
||||||
|
}
|
||||||
|
if col.Description != "" {
|
||||||
|
colInfo["description"] = col.Description
|
||||||
|
}
|
||||||
|
columns = append(columns, colInfo)
|
||||||
|
}
|
||||||
|
|
||||||
|
return map[string]interface{}{
|
||||||
|
"name": mod.MetaData.Table.Name,
|
||||||
|
"columns": columns,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// executeDSL executes the QueryDSL and returns records
|
||||||
|
func (h *Handler) executeDSL(dsl interface{}) ([]map[string]interface{}, error) {
|
||||||
|
// Get the default query engine
|
||||||
|
engine, err := query.Select("default")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("query engine not found: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Marshal DSL to JSON
|
||||||
|
dslJSON, err := json.Marshal(dsl)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal DSL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Load and execute the query
|
||||||
|
q, err := engine.Load(json.RawMessage(dslJSON))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to load DSL: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute query
|
||||||
|
rawRecords := q.Get(nil)
|
||||||
|
|
||||||
|
// Convert to map[string]interface{}
|
||||||
|
records := make([]map[string]interface{}, 0, len(rawRecords))
|
||||||
|
for _, rec := range rawRecords {
|
||||||
|
records = append(records, map[string]interface{}(rec))
|
||||||
|
}
|
||||||
|
|
||||||
|
return records, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// convertToResultItems converts query results to ResultItems
|
||||||
|
func (h *Handler) convertToResultItems(records []map[string]interface{}, modelID string, mod *model.Model, source types.SourceType) []*types.ResultItem {
|
||||||
|
items := make([]*types.ResultItem, 0, len(records))
|
||||||
|
|
||||||
|
primaryKey := "id"
|
||||||
|
if mod != nil && mod.PrimaryKey != "" {
|
||||||
|
primaryKey = mod.PrimaryKey
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, rec := range records {
|
||||||
|
item := &types.ResultItem{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Source: source,
|
||||||
|
Model: modelID,
|
||||||
|
Data: rec,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to extract title from common fields
|
||||||
|
item.Title = h.extractTitle(rec, mod)
|
||||||
|
|
||||||
|
// Try to extract content/description
|
||||||
|
item.Content = h.extractContent(rec, mod)
|
||||||
|
|
||||||
|
// Try to extract record ID
|
||||||
|
if id, ok := rec[primaryKey]; ok {
|
||||||
|
item.RecordID = id
|
||||||
|
}
|
||||||
|
|
||||||
|
items = append(items, item)
|
||||||
|
}
|
||||||
|
|
||||||
|
return items
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractTitle tries to extract a title from the record
|
||||||
|
func (h *Handler) extractTitle(rec map[string]interface{}, mod *model.Model) string {
|
||||||
|
// Common title fields
|
||||||
|
titleFields := []string{"title", "name", "subject", "label"}
|
||||||
|
for _, field := range titleFields {
|
||||||
|
if val, ok := rec[field]; ok {
|
||||||
|
if str, ok := val.(string); ok && str != "" {
|
||||||
|
return str
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractContent tries to extract content from the record
|
||||||
|
func (h *Handler) extractContent(rec map[string]interface{}, mod *model.Model) string {
|
||||||
|
// Common content fields
|
||||||
|
contentFields := []string{"content", "description", "summary", "text", "body"}
|
||||||
|
for _, field := range contentFields {
|
||||||
|
if val, ok := rec[field]; ok {
|
||||||
|
if str, ok := val.(string); ok && str != "" {
|
||||||
|
return str
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: serialize first few fields as content
|
||||||
|
content, _ := json.Marshal(rec)
|
||||||
|
if len(content) > 500 {
|
||||||
|
content = content[:500]
|
||||||
|
}
|
||||||
|
return string(content)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
174
agent/search/handlers/db/handler_integration_test.go
Normal file
174
agent/search/handlers/db/handler_integration_test.go
Normal file
|
|
@ -0,0 +1,174 @@
|
||||||
|
package db_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/yaoapp/gou/model"
|
||||||
|
"github.com/yaoapp/gou/query/gou"
|
||||||
|
"github.com/yaoapp/yao/agent/context"
|
||||||
|
"github.com/yaoapp/yao/agent/search/handlers/db"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
|
"github.com/yaoapp/yao/agent/testutils"
|
||||||
|
oauthTypes "github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ============================================================================
|
||||||
|
// Integration Tests - Requires database and models
|
||||||
|
// ============================================================================
|
||||||
|
|
||||||
|
func TestHandler_Search_Integration(t *testing.T) {
|
||||||
|
// Skip if running short tests
|
||||||
|
if testing.Short() {
|
||||||
|
t.Skip("Skipping integration test")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initialize test environment (loads models, database, etc.)
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Create test context
|
||||||
|
ctx := newTestContext(t)
|
||||||
|
|
||||||
|
// Verify __yao.role model is loaded
|
||||||
|
mod := model.Select("__yao.role")
|
||||||
|
require.NotNil(t, mod, "__yao.role model should be loaded")
|
||||||
|
|
||||||
|
t.Run("search_role_model_with_results", func(t *testing.T) {
|
||||||
|
// First, ensure there's at least one role in the database
|
||||||
|
ensureTestRole(t, mod)
|
||||||
|
|
||||||
|
// Create handler with builtin QueryDSL generator
|
||||||
|
h := db.NewHandler("builtin", &types.DBConfig{
|
||||||
|
Models: []string{"__yao.role"},
|
||||||
|
MaxResults: 10,
|
||||||
|
})
|
||||||
|
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "查询所有角色",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Models: []string{"__yao.role"},
|
||||||
|
Scenario: types.ScenarioFilter,
|
||||||
|
Limit: 10,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := h.SearchWithContext(ctx, req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
// Verify result structure
|
||||||
|
assert.Equal(t, types.SearchTypeDB, result.Type)
|
||||||
|
assert.Equal(t, "查询所有角色", result.Query)
|
||||||
|
assert.Equal(t, types.SourceAuto, result.Source)
|
||||||
|
assert.GreaterOrEqual(t, result.Duration, int64(0))
|
||||||
|
|
||||||
|
// Should have results
|
||||||
|
if result.Error != "" {
|
||||||
|
t.Logf("Search error: %s", result.Error)
|
||||||
|
}
|
||||||
|
assert.Empty(t, result.Error, "Search should not return error")
|
||||||
|
assert.Greater(t, len(result.Items), 0, "Should have at least one result")
|
||||||
|
assert.Equal(t, len(result.Items), result.Total)
|
||||||
|
|
||||||
|
// Verify result items
|
||||||
|
for _, item := range result.Items {
|
||||||
|
assert.Equal(t, types.SearchTypeDB, item.Type)
|
||||||
|
assert.Equal(t, types.SourceAuto, item.Source)
|
||||||
|
assert.Equal(t, "__yao.role", item.Model)
|
||||||
|
assert.NotNil(t, item.Data, "Data should not be nil")
|
||||||
|
assert.NotNil(t, item.RecordID, "RecordID should not be nil")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("search_with_filter_scenario", func(t *testing.T) {
|
||||||
|
h := db.NewHandler("builtin", nil)
|
||||||
|
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "查询系统角色",
|
||||||
|
Source: types.SourceHook,
|
||||||
|
Models: []string{"__yao.role"},
|
||||||
|
Scenario: types.ScenarioFilter,
|
||||||
|
Limit: 5,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := h.SearchWithContext(ctx, req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
assert.Equal(t, types.SearchTypeDB, result.Type)
|
||||||
|
assert.Equal(t, types.SourceHook, result.Source)
|
||||||
|
assert.LessOrEqual(t, len(result.Items), 5, "Should respect limit")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("search_with_preset_wheres", func(t *testing.T) {
|
||||||
|
h := db.NewHandler("builtin", nil)
|
||||||
|
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "查询角色",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Models: []string{"__yao.role"},
|
||||||
|
Wheres: []gou.Where{
|
||||||
|
{Condition: gou.Condition{Field: &gou.Expression{Field: "is_active"}, Value: true, OP: "="}},
|
||||||
|
},
|
||||||
|
Limit: 10,
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := h.SearchWithContext(ctx, req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
|
||||||
|
// All results should have is_active = true (due to preset where)
|
||||||
|
for _, item := range result.Items {
|
||||||
|
if data, ok := item.Data["is_active"]; ok {
|
||||||
|
// is_active could be bool or int depending on driver
|
||||||
|
switch v := data.(type) {
|
||||||
|
case bool:
|
||||||
|
assert.True(t, v)
|
||||||
|
case int64:
|
||||||
|
assert.Equal(t, int64(1), v)
|
||||||
|
case float64:
|
||||||
|
assert.Equal(t, float64(1), v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// newTestContext creates a test context with required fields
|
||||||
|
func newTestContext(t *testing.T) *context.Context {
|
||||||
|
t.Helper()
|
||||||
|
authorized := &oauthTypes.AuthorizedInfo{
|
||||||
|
UserID: "test-user",
|
||||||
|
}
|
||||||
|
chatID := "test-chat-db-search"
|
||||||
|
ctx := context.New(t.Context(), authorized, chatID)
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensureTestRole ensures there's at least one role in the database for testing
|
||||||
|
func ensureTestRole(t *testing.T, mod *model.Model) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
// Try to find existing roles
|
||||||
|
rows, err := mod.Get(model.QueryParam{Limit: 1})
|
||||||
|
if err == nil && len(rows) > 0 {
|
||||||
|
return // Already have roles
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create a test role
|
||||||
|
_, err = mod.Create(map[string]interface{}{
|
||||||
|
"role_id": "test_role",
|
||||||
|
"name": "Test Role",
|
||||||
|
"description": "A test role for unit testing",
|
||||||
|
"is_active": true,
|
||||||
|
"is_system": false,
|
||||||
|
"level": 1,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("Note: Could not create test role: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -4,6 +4,8 @@ import (
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/gou/model"
|
||||||
|
"github.com/yaoapp/gou/query/gou"
|
||||||
"github.com/yaoapp/yao/agent/search/types"
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -38,14 +40,13 @@ func TestHandler_Type(t *testing.T) {
|
||||||
assert.Equal(t, types.SearchTypeDB, h.Type())
|
assert.Equal(t, types.SearchTypeDB, h.Type())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHandler_Search(t *testing.T) {
|
func TestHandler_Search_Validation(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
usesQueryDSL string
|
usesQueryDSL string
|
||||||
config *types.DBConfig
|
config *types.DBConfig
|
||||||
req *types.Request
|
req *types.Request
|
||||||
expectError string
|
expectError string
|
||||||
expectItems int
|
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "empty query",
|
name: "empty query",
|
||||||
|
|
@ -56,7 +57,6 @@ func TestHandler_Search(t *testing.T) {
|
||||||
Query: "",
|
Query: "",
|
||||||
},
|
},
|
||||||
expectError: "query is required",
|
expectError: "query is required",
|
||||||
expectItems: 0,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "no models in request or config",
|
name: "no models in request or config",
|
||||||
|
|
@ -66,92 +66,20 @@ func TestHandler_Search(t *testing.T) {
|
||||||
Type: types.SearchTypeDB,
|
Type: types.SearchTypeDB,
|
||||||
Query: "find products under $100",
|
Query: "find products under $100",
|
||||||
},
|
},
|
||||||
expectError: "",
|
expectError: "no models specified",
|
||||||
expectItems: 0,
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "models from config",
|
name: "context required for DB search",
|
||||||
usesQueryDSL: "builtin",
|
usesQueryDSL: "builtin",
|
||||||
config: &types.DBConfig{
|
config: &types.DBConfig{
|
||||||
Models: []string{"product"},
|
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{
|
req: &types.Request{
|
||||||
Type: types.SearchTypeDB,
|
Type: types.SearchTypeDB,
|
||||||
Query: "find products under $100",
|
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"},
|
Models: []string{"product"},
|
||||||
},
|
},
|
||||||
req: &types.Request{
|
expectError: "context is required for DB search",
|
||||||
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
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -163,16 +91,8 @@ func TestHandler_Search(t *testing.T) {
|
||||||
assert.NoError(t, err)
|
assert.NoError(t, err)
|
||||||
assert.NotNil(t, result)
|
assert.NotNil(t, result)
|
||||||
assert.Equal(t, types.SearchTypeDB, result.Type)
|
assert.Equal(t, types.SearchTypeDB, result.Type)
|
||||||
assert.Equal(t, tt.req.Query, result.Query)
|
assert.Equal(t, tt.expectError, result.Error)
|
||||||
assert.Equal(t, tt.expectItems, len(result.Items))
|
assert.Equal(t, 0, 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))
|
assert.GreaterOrEqual(t, result.Duration, int64(0))
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
@ -195,21 +115,374 @@ func TestHandler_Search_SourcePreserved(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestHandler_Search_MaxResultsFromConfig(t *testing.T) {
|
func TestHandler_BuildModelSchema(t *testing.T) {
|
||||||
cfg := &types.DBConfig{
|
h := NewHandler("builtin", nil)
|
||||||
Models: []string{"product"},
|
|
||||||
MaxResults: 50,
|
|
||||||
}
|
|
||||||
h := NewHandler("builtin", cfg)
|
|
||||||
|
|
||||||
req := &types.Request{
|
// Create a mock model for testing
|
||||||
Type: types.SearchTypeDB,
|
mod := &model.Model{
|
||||||
Query: "test",
|
MetaData: model.MetaData{
|
||||||
Models: []string{"product"},
|
Table: model.Table{
|
||||||
// No limit in request, should use config's MaxResults
|
Name: "test_products",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Columns: map[string]*model.Column{
|
||||||
|
"id": {
|
||||||
|
Name: "id",
|
||||||
|
Type: "ID",
|
||||||
|
Label: "ID",
|
||||||
|
},
|
||||||
|
"name": {
|
||||||
|
Name: "name",
|
||||||
|
Type: "string",
|
||||||
|
Label: "Name",
|
||||||
|
Description: "Product name",
|
||||||
|
},
|
||||||
|
"price": {
|
||||||
|
Name: "price",
|
||||||
|
Type: "decimal",
|
||||||
|
Label: "Price",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
schema := h.buildModelSchema(mod)
|
||||||
|
|
||||||
|
assert.NotNil(t, schema)
|
||||||
|
assert.Equal(t, "test_products", schema["name"])
|
||||||
|
|
||||||
|
columns, ok := schema["columns"].([]map[string]interface{})
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Len(t, columns, 3)
|
||||||
|
|
||||||
|
// Verify columns have required fields
|
||||||
|
for _, col := range columns {
|
||||||
|
assert.NotEmpty(t, col["name"])
|
||||||
|
assert.NotEmpty(t, col["type"])
|
||||||
}
|
}
|
||||||
result, err := h.Search(req)
|
}
|
||||||
assert.NoError(t, err)
|
|
||||||
assert.NotNil(t, result)
|
func TestHandler_BuildModelSchema_MultipleModels(t *testing.T) {
|
||||||
// Skeleton doesn't actually use maxResults yet, but the test ensures the handler runs
|
h := NewHandler("builtin", nil)
|
||||||
|
|
||||||
|
// Create mock models for testing joins
|
||||||
|
productMod := &model.Model{
|
||||||
|
MetaData: model.MetaData{
|
||||||
|
Table: model.Table{Name: "products"},
|
||||||
|
},
|
||||||
|
Columns: map[string]*model.Column{
|
||||||
|
"id": {Name: "id", Type: "ID"},
|
||||||
|
"name": {Name: "name", Type: "string"},
|
||||||
|
"category_id": {Name: "category_id", Type: "integer"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
categoryMod := &model.Model{
|
||||||
|
MetaData: model.MetaData{
|
||||||
|
Table: model.Table{Name: "categories"},
|
||||||
|
},
|
||||||
|
Columns: map[string]*model.Column{
|
||||||
|
"id": {Name: "id", Type: "ID"},
|
||||||
|
"name": {Name: "name", Type: "string"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
productSchema := h.buildModelSchema(productMod)
|
||||||
|
categorySchema := h.buildModelSchema(categoryMod)
|
||||||
|
|
||||||
|
assert.Equal(t, "products", productSchema["name"])
|
||||||
|
assert.Equal(t, "categories", categorySchema["name"])
|
||||||
|
|
||||||
|
// Verify both schemas can be combined into an array
|
||||||
|
schemas := []map[string]interface{}{productSchema, categorySchema}
|
||||||
|
assert.Len(t, schemas, 2)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_ExtractTitle(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", nil)
|
||||||
|
mod := &model.Model{}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
record map[string]interface{}
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "title field",
|
||||||
|
record: map[string]interface{}{"title": "Test Title", "id": 1},
|
||||||
|
expected: "Test Title",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "name field",
|
||||||
|
record: map[string]interface{}{"name": "Test Name", "id": 1},
|
||||||
|
expected: "Test Name",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "subject field",
|
||||||
|
record: map[string]interface{}{"subject": "Test Subject", "id": 1},
|
||||||
|
expected: "Test Subject",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "label field",
|
||||||
|
record: map[string]interface{}{"label": "Test Label", "id": 1},
|
||||||
|
expected: "Test Label",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no title field",
|
||||||
|
record: map[string]interface{}{"id": 1, "price": 100},
|
||||||
|
expected: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "empty title",
|
||||||
|
record: map[string]interface{}{"title": "", "name": "Fallback"},
|
||||||
|
expected: "Fallback",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
title := h.extractTitle(tt.record, mod)
|
||||||
|
assert.Equal(t, tt.expected, title)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_ExtractContent(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", nil)
|
||||||
|
mod := &model.Model{}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
record map[string]interface{}
|
||||||
|
expectEmpty bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "content field",
|
||||||
|
record: map[string]interface{}{"content": "Test Content"},
|
||||||
|
expectEmpty: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "description field",
|
||||||
|
record: map[string]interface{}{"description": "Test Description"},
|
||||||
|
expectEmpty: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "summary field",
|
||||||
|
record: map[string]interface{}{"summary": "Test Summary"},
|
||||||
|
expectEmpty: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "fallback to JSON",
|
||||||
|
record: map[string]interface{}{"id": 1, "price": 100},
|
||||||
|
expectEmpty: false, // Should return JSON representation
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
content := h.extractContent(tt.record, mod)
|
||||||
|
if tt.expectEmpty {
|
||||||
|
assert.Empty(t, content)
|
||||||
|
} else {
|
||||||
|
assert.NotEmpty(t, content)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_ConvertToResultItems(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", nil)
|
||||||
|
mod := &model.Model{
|
||||||
|
PrimaryKey: "id",
|
||||||
|
}
|
||||||
|
|
||||||
|
records := []map[string]interface{}{
|
||||||
|
{
|
||||||
|
"id": 1,
|
||||||
|
"name": "Product 1",
|
||||||
|
"description": "Description 1",
|
||||||
|
"price": 99.99,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": 2,
|
||||||
|
"title": "Product 2",
|
||||||
|
"content": "Content 2",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
items := h.convertToResultItems(records, "product", mod, types.SourceAuto)
|
||||||
|
|
||||||
|
assert.Len(t, items, 2)
|
||||||
|
|
||||||
|
// First item
|
||||||
|
assert.Equal(t, types.SearchTypeDB, items[0].Type)
|
||||||
|
assert.Equal(t, types.SourceAuto, items[0].Source)
|
||||||
|
assert.Equal(t, "product", items[0].Model)
|
||||||
|
assert.Equal(t, 1, items[0].RecordID)
|
||||||
|
assert.Equal(t, "Product 1", items[0].Title)
|
||||||
|
assert.Equal(t, "Description 1", items[0].Content)
|
||||||
|
assert.NotNil(t, items[0].Data)
|
||||||
|
|
||||||
|
// Second item
|
||||||
|
assert.Equal(t, 2, items[1].RecordID)
|
||||||
|
assert.Equal(t, "Product 2", items[1].Title)
|
||||||
|
assert.Equal(t, "Content 2", items[1].Content)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_ConvertToResultItems_NilModel(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", nil)
|
||||||
|
|
||||||
|
records := []map[string]interface{}{
|
||||||
|
{"id": 1, "name": "Test"},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should use default primary key "id" when model is nil
|
||||||
|
items := h.convertToResultItems(records, "test", nil, types.SourceHook)
|
||||||
|
|
||||||
|
assert.Len(t, items, 1)
|
||||||
|
assert.Equal(t, 1, items[0].RecordID)
|
||||||
|
assert.Equal(t, "Test", items[0].Title)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_Search_ScenarioTypes(t *testing.T) {
|
||||||
|
// Test that all scenario types are valid
|
||||||
|
scenarios := []types.ScenarioType{
|
||||||
|
types.ScenarioFilter,
|
||||||
|
types.ScenarioAggregation,
|
||||||
|
types.ScenarioJoin,
|
||||||
|
types.ScenarioComplex,
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, scenario := range scenarios {
|
||||||
|
t.Run(string(scenario), func(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", &types.DBConfig{Models: []string{"product"}})
|
||||||
|
req := &types.Request{
|
||||||
|
Type: types.SearchTypeDB,
|
||||||
|
Query: "test query",
|
||||||
|
Source: types.SourceAuto,
|
||||||
|
Models: []string{"product"},
|
||||||
|
Scenario: scenario,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Without context, should return error (but scenario should be preserved in request)
|
||||||
|
result, err := h.Search(req)
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NotNil(t, result)
|
||||||
|
// Verify request scenario is set correctly
|
||||||
|
assert.Equal(t, scenario, req.Scenario)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestScenarioTypeConstants(t *testing.T) {
|
||||||
|
// Verify scenario type constants match expected values
|
||||||
|
assert.Equal(t, types.ScenarioType("filter"), types.ScenarioFilter)
|
||||||
|
assert.Equal(t, types.ScenarioType("aggregation"), types.ScenarioAggregation)
|
||||||
|
assert.Equal(t, types.ScenarioType("join"), types.ScenarioJoin)
|
||||||
|
assert.Equal(t, types.ScenarioType("complex"), types.ScenarioComplex)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandler_MergeDSLConditions(t *testing.T) {
|
||||||
|
h := NewHandler("builtin", nil)
|
||||||
|
|
||||||
|
t.Run("merge wheres", func(t *testing.T) {
|
||||||
|
dsl := &gou.QueryDSL{
|
||||||
|
From: &gou.Table{Name: "users"},
|
||||||
|
Wheres: []gou.Where{
|
||||||
|
{Condition: gou.Condition{Field: &gou.Expression{Field: "status"}, Value: "active", OP: "="}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
req := &types.Request{
|
||||||
|
Wheres: []gou.Where{
|
||||||
|
{Condition: gou.Condition{Field: &gou.Expression{Field: "tenant_id"}, Value: 1, OP: "="}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
h.mergeDSLConditions(dsl, req)
|
||||||
|
|
||||||
|
// Preset wheres should be prepended
|
||||||
|
assert.Len(t, dsl.Wheres, 2)
|
||||||
|
assert.Equal(t, "tenant_id", dsl.Wheres[0].Field.Field)
|
||||||
|
assert.Equal(t, "status", dsl.Wheres[1].Field.Field)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("merge orders", func(t *testing.T) {
|
||||||
|
dsl := &gou.QueryDSL{
|
||||||
|
From: &gou.Table{Name: "products"},
|
||||||
|
Orders: gou.Orders{
|
||||||
|
{Field: &gou.Expression{Field: "name"}, Sort: "asc"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
req := &types.Request{
|
||||||
|
Orders: gou.Orders{
|
||||||
|
{Field: &gou.Expression{Field: "created_at"}, Sort: "desc"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
h.mergeDSLConditions(dsl, req)
|
||||||
|
|
||||||
|
// Preset orders should be prepended
|
||||||
|
assert.Len(t, dsl.Orders, 2)
|
||||||
|
assert.Equal(t, "created_at", dsl.Orders[0].Field.Field)
|
||||||
|
assert.Equal(t, "name", dsl.Orders[1].Field.Field)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("merge select fields", func(t *testing.T) {
|
||||||
|
dsl := &gou.QueryDSL{
|
||||||
|
From: &gou.Table{Name: "orders"},
|
||||||
|
Select: []gou.Expression{
|
||||||
|
{Field: "amount"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
req := &types.Request{
|
||||||
|
Select: []string{"id", "status"},
|
||||||
|
}
|
||||||
|
|
||||||
|
h.mergeDSLConditions(dsl, req)
|
||||||
|
|
||||||
|
// Preset select should be prepended
|
||||||
|
assert.Len(t, dsl.Select, 3)
|
||||||
|
assert.Equal(t, "id", dsl.Select[0].Field)
|
||||||
|
assert.Equal(t, "status", dsl.Select[1].Field)
|
||||||
|
assert.Equal(t, "amount", dsl.Select[2].Field)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("set limit from request", func(t *testing.T) {
|
||||||
|
dsl := &gou.QueryDSL{
|
||||||
|
From: &gou.Table{Name: "users"},
|
||||||
|
Limit: 0,
|
||||||
|
}
|
||||||
|
req := &types.Request{
|
||||||
|
Limit: 50,
|
||||||
|
}
|
||||||
|
|
||||||
|
h.mergeDSLConditions(dsl, req)
|
||||||
|
|
||||||
|
assert.Equal(t, 50, dsl.Limit)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("preserve dsl limit if set", func(t *testing.T) {
|
||||||
|
dsl := &gou.QueryDSL{
|
||||||
|
From: &gou.Table{Name: "users"},
|
||||||
|
Limit: 10,
|
||||||
|
}
|
||||||
|
req := &types.Request{
|
||||||
|
Limit: 50,
|
||||||
|
}
|
||||||
|
|
||||||
|
h.mergeDSLConditions(dsl, req)
|
||||||
|
|
||||||
|
// DSL limit should be preserved
|
||||||
|
assert.Equal(t, 10, dsl.Limit)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("nil dsl", func(t *testing.T) {
|
||||||
|
req := &types.Request{
|
||||||
|
Wheres: []gou.Where{
|
||||||
|
{Condition: gou.Condition{Field: &gou.Expression{Field: "id"}, Value: 1}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Should not panic
|
||||||
|
h.mergeDSLConditions(nil, req)
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
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"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -12,3 +13,12 @@ type Handler interface {
|
||||||
// Search executes the search and returns results
|
// Search executes the search and returns results
|
||||||
Search(req *types.Request) (*types.Result, error)
|
Search(req *types.Request) (*types.Result, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ContextHandler extends Handler with context support
|
||||||
|
// Handlers that need context (e.g., DB handler for QueryDSL generation) should implement this
|
||||||
|
type ContextHandler interface {
|
||||||
|
Handler
|
||||||
|
|
||||||
|
// SearchWithContext executes the search with context and returns results
|
||||||
|
SearchWithContext(ctx *context.Context, req *types.Request) (*types.Result, error)
|
||||||
|
}
|
||||||
|
|
|
||||||
190
agent/search/jsapi_db_test.go
Normal file
190
agent/search/jsapi_db_test.go
Normal file
|
|
@ -0,0 +1,190 @@
|
||||||
|
package search_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/yaoapp/gou/model"
|
||||||
|
"github.com/yaoapp/yao/agent/context"
|
||||||
|
"github.com/yaoapp/yao/agent/search"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
|
"github.com/yaoapp/yao/agent/testutils"
|
||||||
|
oauthTypes "github.com/yaoapp/yao/openapi/oauth/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ============================================================================
|
||||||
|
// DB Search JSAPI Integration Tests
|
||||||
|
// ============================================================================
|
||||||
|
|
||||||
|
func TestJSAPI_DB_Integration(t *testing.T) {
|
||||||
|
// Skip if running short tests
|
||||||
|
if testing.Short() {
|
||||||
|
t.Skip("Skipping integration test")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Initialize test environment
|
||||||
|
testutils.Prepare(t)
|
||||||
|
defer testutils.Clean(t)
|
||||||
|
|
||||||
|
// Create test context
|
||||||
|
ctx := newJSAPITestContext(t)
|
||||||
|
|
||||||
|
// Verify __yao.role model is loaded
|
||||||
|
mod := model.Select("__yao.role")
|
||||||
|
require.NotNil(t, mod, "__yao.role model should be loaded")
|
||||||
|
|
||||||
|
// Ensure test data exists
|
||||||
|
ensureJSAPITestRole(t, mod)
|
||||||
|
|
||||||
|
t.Run("db_search_with_context", func(t *testing.T) {
|
||||||
|
api := search.NewJSAPI(ctx, &types.Config{
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"__yao.role"},
|
||||||
|
MaxResults: 10,
|
||||||
|
},
|
||||||
|
}, &search.Uses{QueryDSL: "builtin"})
|
||||||
|
|
||||||
|
result := api.DB("查询所有角色", map[string]interface{}{
|
||||||
|
"models": []interface{}{"__yao.role"},
|
||||||
|
"limit": float64(10),
|
||||||
|
})
|
||||||
|
|
||||||
|
require.NotNil(t, result)
|
||||||
|
r, ok := result.(*types.Result)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
assert.Equal(t, types.SearchTypeDB, r.Type)
|
||||||
|
assert.Equal(t, "查询所有角色", r.Query)
|
||||||
|
assert.Equal(t, types.SourceHook, r.Source)
|
||||||
|
|
||||||
|
if r.Error != "" {
|
||||||
|
t.Logf("Search error: %s", r.Error)
|
||||||
|
}
|
||||||
|
assert.Empty(t, r.Error, "Should not have error")
|
||||||
|
assert.Greater(t, len(r.Items), 0, "Should have results")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("db_search_with_scenario", func(t *testing.T) {
|
||||||
|
api := search.NewJSAPI(ctx, &types.Config{
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"__yao.role"},
|
||||||
|
MaxResults: 5,
|
||||||
|
},
|
||||||
|
}, &search.Uses{QueryDSL: "builtin"})
|
||||||
|
|
||||||
|
result := api.DB("查询系统角色", map[string]interface{}{
|
||||||
|
"models": []interface{}{"__yao.role"},
|
||||||
|
"scenario": "filter",
|
||||||
|
"limit": float64(5),
|
||||||
|
})
|
||||||
|
|
||||||
|
require.NotNil(t, result)
|
||||||
|
r, ok := result.(*types.Result)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
assert.Equal(t, types.SearchTypeDB, r.Type)
|
||||||
|
assert.LessOrEqual(t, len(r.Items), 5, "Should respect limit")
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("db_search_with_select_fields", func(t *testing.T) {
|
||||||
|
api := search.NewJSAPI(ctx, &types.Config{
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"__yao.role"},
|
||||||
|
MaxResults: 10,
|
||||||
|
},
|
||||||
|
}, &search.Uses{QueryDSL: "builtin"})
|
||||||
|
|
||||||
|
result := api.DB("查询角色名称", map[string]interface{}{
|
||||||
|
"models": []interface{}{"__yao.role"},
|
||||||
|
"select": []interface{}{"id", "name", "description"},
|
||||||
|
"limit": float64(10),
|
||||||
|
})
|
||||||
|
|
||||||
|
require.NotNil(t, result)
|
||||||
|
r, ok := result.(*types.Result)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
assert.Equal(t, types.SearchTypeDB, r.Type)
|
||||||
|
if r.Error == "" && len(r.Items) > 0 {
|
||||||
|
// Verify items have data
|
||||||
|
for _, item := range r.Items {
|
||||||
|
assert.NotNil(t, item.Data)
|
||||||
|
assert.Equal(t, "__yao.role", item.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("db_search_all_with_multiple_types", func(t *testing.T) {
|
||||||
|
api := search.NewJSAPI(ctx, &types.Config{
|
||||||
|
KB: &types.KBConfig{Collections: []string{"docs"}},
|
||||||
|
DB: &types.DBConfig{
|
||||||
|
Models: []string{"__yao.role"},
|
||||||
|
MaxResults: 10,
|
||||||
|
},
|
||||||
|
}, &search.Uses{QueryDSL: "builtin"})
|
||||||
|
|
||||||
|
requests := []interface{}{
|
||||||
|
map[string]interface{}{
|
||||||
|
"type": "db",
|
||||||
|
"query": "查询角色",
|
||||||
|
"models": []interface{}{"__yao.role"},
|
||||||
|
"limit": float64(5),
|
||||||
|
},
|
||||||
|
map[string]interface{}{
|
||||||
|
"type": "kb",
|
||||||
|
"query": "知识库查询",
|
||||||
|
"collections": []interface{}{"docs"},
|
||||||
|
"limit": float64(5),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
results := api.All(requests)
|
||||||
|
require.Len(t, results, 2)
|
||||||
|
|
||||||
|
// DB result
|
||||||
|
r0, ok := results[0].(*types.Result)
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Equal(t, types.SearchTypeDB, r0.Type)
|
||||||
|
|
||||||
|
// KB result
|
||||||
|
r1, ok := results[1].(*types.Result)
|
||||||
|
require.True(t, ok)
|
||||||
|
assert.Equal(t, types.SearchTypeKB, r1.Type)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// newJSAPITestContext creates a test context for JSAPI tests
|
||||||
|
func newJSAPITestContext(t *testing.T) *context.Context {
|
||||||
|
t.Helper()
|
||||||
|
authorized := &oauthTypes.AuthorizedInfo{
|
||||||
|
UserID: "test-user-jsapi",
|
||||||
|
}
|
||||||
|
chatID := "test-chat-jsapi-db"
|
||||||
|
ctx := context.New(t.Context(), authorized, chatID)
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensureJSAPITestRole ensures there's at least one role in the database
|
||||||
|
func ensureJSAPITestRole(t *testing.T, mod *model.Model) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
// Try to find existing roles
|
||||||
|
rows, err := mod.Get(model.QueryParam{Limit: 1})
|
||||||
|
if err == nil && len(rows) > 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create a test role
|
||||||
|
_, err = mod.Create(map[string]interface{}{
|
||||||
|
"role_id": "jsapi_test_role",
|
||||||
|
"name": "JSAPI Test Role",
|
||||||
|
"description": "A test role for JSAPI unit testing",
|
||||||
|
"is_active": true,
|
||||||
|
"is_system": false,
|
||||||
|
"level": 1,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Logf("Note: Could not create test role: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -45,15 +45,14 @@ func (p *AgentProvider) Generate(ctx *agentContext.Context, input *Input) (*Resu
|
||||||
var lastLintErrors string
|
var lastLintErrors string
|
||||||
|
|
||||||
for attempt := 1; attempt <= MaxRetries; attempt++ {
|
for attempt := 1; attempt <= MaxRetries; attempt++ {
|
||||||
// Build the request message
|
// Build the request message in the format expected by querydsl agent
|
||||||
requestData := p.buildRequestData(input, attempt, lastLintErrors)
|
requestMessage := p.buildRequestMessage(input, attempt, lastLintErrors)
|
||||||
requestJSON, _ := json.Marshal(requestData)
|
|
||||||
|
|
||||||
// Create message for the agent
|
// Create message for the agent
|
||||||
messages := []agentContext.Message{
|
messages := []agentContext.Message{
|
||||||
{
|
{
|
||||||
Role: "user",
|
Role: "user",
|
||||||
Content: string(requestJSON),
|
Content: requestMessage,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -103,38 +102,36 @@ func (p *AgentProvider) Generate(ctx *agentContext.Context, input *Input) (*Resu
|
||||||
return nil, fmt.Errorf("QueryDSL generation failed after %d attempts: %w", MaxRetries, lastError)
|
return nil, fmt.Errorf("QueryDSL generation failed after %d attempts: %w", MaxRetries, lastError)
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildRequestData constructs the request data for the agent
|
// buildRequestMessage constructs the request message for the agent
|
||||||
func (p *AgentProvider) buildRequestData(input *Input, attempt int, lastLintErrors string) map[string]interface{} {
|
// Format follows the querydsl agent prompts.yml:
|
||||||
requestData := map[string]interface{}{
|
// "用户查询\nSchema:\n{schema JSON}"
|
||||||
"query": input.Query,
|
func (p *AgentProvider) buildRequestMessage(input *Input, attempt int, lastLintErrors string) string {
|
||||||
"models": input.ModelIDs,
|
// Build schema from extra params if provided
|
||||||
"limit": input.Limit,
|
var schemaJSON string
|
||||||
|
if input.ExtraParams != nil {
|
||||||
|
if schema, ok := input.ExtraParams["schema"]; ok {
|
||||||
|
schemaBytes, _ := json.Marshal(schema)
|
||||||
|
schemaJSON = string(schemaBytes)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add optional fields
|
// Build message in the expected format
|
||||||
if len(input.Wheres) > 0 {
|
message := input.Query
|
||||||
requestData["wheres"] = input.Wheres
|
if schemaJSON != "" {
|
||||||
|
message = fmt.Sprintf("%s\nSchema:\n%s", input.Query, schemaJSON)
|
||||||
}
|
}
|
||||||
if len(input.Orders) > 0 {
|
|
||||||
requestData["orders"] = input.Orders
|
// Add scenario hint if specified (filter, aggregation, join, complex)
|
||||||
}
|
if input.Scenario != "" {
|
||||||
if len(input.AllowedFields) > 0 {
|
message = fmt.Sprintf("%s\nScenario: %s", message, input.Scenario)
|
||||||
requestData["allowed_fields"] = input.AllowedFields
|
|
||||||
}
|
|
||||||
if len(input.ExtraParams) > 0 {
|
|
||||||
requestData["extra"] = input.ExtraParams
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add retry context if this is a retry attempt
|
// Add retry context if this is a retry attempt
|
||||||
if attempt > 1 && lastLintErrors != "" {
|
if attempt > 1 && lastLintErrors != "" {
|
||||||
requestData["retry"] = map[string]interface{}{
|
message = fmt.Sprintf("%s\n\nPrevious attempt failed with errors:\n%s\n\nPlease fix the errors and regenerate.", message, lastLintErrors)
|
||||||
"attempt": attempt,
|
|
||||||
"lint_errors": lastLintErrors,
|
|
||||||
"instructions": "The previous QueryDSL was invalid. Please fix the errors and regenerate.",
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return requestData
|
return message
|
||||||
}
|
}
|
||||||
|
|
||||||
// validateDSL validates the generated QueryDSL using the linter
|
// validateDSL validates the generated QueryDSL using the linter
|
||||||
|
|
@ -151,14 +148,25 @@ func (p *AgentProvider) validateDSL(dsl *gou.QueryDSL) *linter.LintResult {
|
||||||
}
|
}
|
||||||
|
|
||||||
// parseResult extracts QueryDSL from the agent's response
|
// parseResult extracts QueryDSL from the agent's response
|
||||||
// The agent should return data in NextHookResponse format: { data: { dsl: {...}, explain: "..." } }
|
// The querydsl agent returns QueryDSL JSON directly (not wrapped in {dsl: ...})
|
||||||
// The Stream() response wraps this in: { next: { data: { dsl: {...} } } }
|
// Or returns error JSON: {"error": "code", "message": "..."}
|
||||||
|
// Stream() returns *context.Response with QueryDSL in "next" field
|
||||||
func (p *AgentProvider) parseResult(result interface{}) (*Result, error) {
|
func (p *AgentProvider) parseResult(result interface{}) (*Result, error) {
|
||||||
if result == nil {
|
if result == nil {
|
||||||
return &Result{}, nil
|
return &Result{}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try to convert to map first (most common case)
|
// Handle *context.Response directly (most common case from Stream())
|
||||||
|
if resp, ok := result.(*agentContext.Response); ok {
|
||||||
|
genResult := &Result{}
|
||||||
|
if resp.Next != nil {
|
||||||
|
// Next contains the QueryDSL from hook
|
||||||
|
genResult.DSL = p.extractDSL(resp.Next)
|
||||||
|
}
|
||||||
|
return genResult, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try to convert to map first
|
||||||
var data map[string]interface{}
|
var data map[string]interface{}
|
||||||
|
|
||||||
switch v := result.(type) {
|
switch v := result.(type) {
|
||||||
|
|
@ -180,21 +188,79 @@ func (p *AgentProvider) parseResult(result interface{}) (*Result, error) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Check for "next" field (custom hook data from NextHookResponse)
|
genResult := &Result{}
|
||||||
// Stream() returns: { next: { data: { dsl: {...} } } }
|
|
||||||
if next, hasNext := data["next"]; hasNext && next != nil {
|
// Check for Stream() wrapper: { content: "...", next: {...} }
|
||||||
if nextMap, ok := next.(map[string]interface{}); ok {
|
// The actual response is in "content" field as a string
|
||||||
data = nextMap
|
if content, hasContent := data["content"]; hasContent && content != nil {
|
||||||
} else if nextStr, ok := next.(string); ok {
|
if contentStr, ok := content.(string); ok && contentStr != "" {
|
||||||
if err := json.Unmarshal([]byte(nextStr), &data); err != nil {
|
// Parse the content string as JSON
|
||||||
return &Result{}, nil
|
var contentData map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(contentStr), &contentData); err == nil {
|
||||||
|
data = contentData
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Extract QueryDSL from data
|
// Check for error response: {"error": "code", "message": "..."}
|
||||||
// Try common field names: "dsl", "data", "data.dsl"
|
if errCode, hasError := data["error"]; hasError {
|
||||||
genResult := &Result{}
|
errMsg := ""
|
||||||
|
if msg, ok := data["message"].(string); ok {
|
||||||
|
errMsg = msg
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("QueryDSL generation error [%v]: %s", errCode, errMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this is a direct QueryDSL (has "from" or "select" field)
|
||||||
|
// The querydsl agent returns QueryDSL directly, e.g., {"select": [...], "from": "table", ...}
|
||||||
|
if _, hasFrom := data["from"]; hasFrom {
|
||||||
|
genResult.DSL = p.extractDSL(data)
|
||||||
|
return genResult, nil
|
||||||
|
}
|
||||||
|
if _, hasSelect := data["select"]; hasSelect {
|
||||||
|
genResult.DSL = p.extractDSL(data)
|
||||||
|
return genResult, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fallback: check for wrapped formats
|
||||||
|
// Check for "next" field (custom hook data from NextHookResponse)
|
||||||
|
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 {
|
||||||
|
if err := json.Unmarshal([]byte(nextStr), &data); err == nil {
|
||||||
|
// Check if parsed data is a QueryDSL
|
||||||
|
if _, hasFrom := data["from"]; hasFrom {
|
||||||
|
genResult.DSL = p.extractDSL(data)
|
||||||
|
return genResult, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check for "dsl" field wrapper
|
||||||
|
if dsl, ok := data["dsl"]; ok {
|
||||||
|
genResult.DSL = p.extractDSL(dsl)
|
||||||
|
} else if d, ok := data["data"]; ok {
|
||||||
|
if dm, ok := d.(map[string]interface{}); ok {
|
||||||
|
// Check if data.data is a wrapped DSL: { dsl: {...} }
|
||||||
|
if dsl, ok := dm["dsl"]; ok {
|
||||||
|
genResult.DSL = p.extractDSL(dsl)
|
||||||
|
} else if _, hasFrom := dm["from"]; hasFrom {
|
||||||
|
// data.data is directly a QueryDSL (from __yao.querydsl Next hook)
|
||||||
|
genResult.DSL = p.extractDSL(dm)
|
||||||
|
} else if _, hasSelect := dm["select"]; hasSelect {
|
||||||
|
// data.data is directly a QueryDSL
|
||||||
|
genResult.DSL = p.extractDSL(dm)
|
||||||
|
}
|
||||||
|
if explain, ok := dm["explain"].(string); ok {
|
||||||
|
genResult.Explain = explain
|
||||||
|
}
|
||||||
|
if warnings, ok := dm["warnings"]; ok {
|
||||||
|
genResult.Warnings = p.extractWarnings(warnings)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Get explain if present
|
// Get explain if present
|
||||||
if explain, ok := data["explain"].(string); ok {
|
if explain, ok := data["explain"].(string); ok {
|
||||||
|
|
@ -206,23 +272,6 @@ func (p *AgentProvider) parseResult(result interface{}) (*Result, error) {
|
||||||
genResult.Warnings = p.extractWarnings(warnings)
|
genResult.Warnings = p.extractWarnings(warnings)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get DSL
|
|
||||||
if dsl, ok := data["dsl"]; ok {
|
|
||||||
genResult.DSL = p.extractDSL(dsl)
|
|
||||||
} else if d, ok := data["data"]; ok {
|
|
||||||
if dm, ok := d.(map[string]interface{}); ok {
|
|
||||||
if dsl, ok := dm["dsl"]; ok {
|
|
||||||
genResult.DSL = p.extractDSL(dsl)
|
|
||||||
}
|
|
||||||
if explain, ok := dm["explain"].(string); ok {
|
|
||||||
genResult.Explain = explain
|
|
||||||
}
|
|
||||||
if warnings, ok := dm["warnings"]; ok {
|
|
||||||
genResult.Warnings = p.extractWarnings(warnings)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return genResult, nil
|
return genResult, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,12 +2,14 @@ package querydsl
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"github.com/yaoapp/gou/query/gou"
|
"github.com/yaoapp/gou/query/gou"
|
||||||
|
"github.com/yaoapp/yao/agent/search/types"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Input contains all information needed to generate QueryDSL
|
// Input contains all information needed to generate QueryDSL
|
||||||
type Input struct {
|
type Input struct {
|
||||||
Query string // Natural language query
|
Query string // Natural language query
|
||||||
ModelIDs []string // Target model IDs (e.g., ["user", "order", "product"])
|
ModelIDs []string // Target model IDs (e.g., ["user", "order", "product"])
|
||||||
|
Scenario types.ScenarioType // QueryDSL scenario: "filter", "aggregation", "join", "complex"
|
||||||
Wheres []gou.Where // Pre-defined filters (optional)
|
Wheres []gou.Where // Pre-defined filters (optional)
|
||||||
Orders gou.Orders // Sort orders (optional)
|
Orders gou.Orders // Sort orders (optional)
|
||||||
AllowedFields []string // Allowed fields whitelist (optional, for security validation)
|
AllowedFields []string // Allowed fields whitelist (optional, for security validation)
|
||||||
|
|
|
||||||
|
|
@ -60,8 +60,14 @@ func (s *Searcher) Search(ctx *context.Context, req *types.Request) (*types.Resu
|
||||||
return &types.Result{Error: "unsupported search type"}, nil
|
return &types.Result{Error: "unsupported search type"}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute search
|
// Execute search - use context if handler supports it
|
||||||
result, err := handler.Search(req)
|
var result *types.Result
|
||||||
|
var err error
|
||||||
|
if ctxHandler, ok := handler.(interfaces.ContextHandler); ok {
|
||||||
|
result, err = ctxHandler.SearchWithContext(ctx, req)
|
||||||
|
} else {
|
||||||
|
result, err = handler.Search(req)
|
||||||
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return &types.Result{Error: err.Error()}, nil
|
return &types.Result{Error: err.Error()}, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -14,6 +14,17 @@ const (
|
||||||
SearchTypeDB SearchType = "db" // Database search (Yao Model/QueryDSL)
|
SearchTypeDB SearchType = "db" // Database search (Yao Model/QueryDSL)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// ScenarioType represents the QueryDSL generation scenario
|
||||||
|
type ScenarioType string
|
||||||
|
|
||||||
|
// ScenarioType constants for QueryDSL generation
|
||||||
|
const (
|
||||||
|
ScenarioFilter ScenarioType = "filter" // Simple filtering queries
|
||||||
|
ScenarioAggregation ScenarioType = "aggregation" // Aggregation/grouping queries
|
||||||
|
ScenarioJoin ScenarioType = "join" // Multi-table join queries
|
||||||
|
ScenarioComplex ScenarioType = "complex" // Complex queries combining multiple features
|
||||||
|
)
|
||||||
|
|
||||||
// SourceType represents where the search result came from
|
// SourceType represents where the search result came from
|
||||||
type SourceType string
|
type SourceType string
|
||||||
|
|
||||||
|
|
@ -42,10 +53,11 @@ type Request struct {
|
||||||
Graph bool `json:"graph,omitempty"` // Enable graph association
|
Graph bool `json:"graph,omitempty"` // Enable graph association
|
||||||
|
|
||||||
// Database search specific
|
// Database search specific
|
||||||
Models []string `json:"models,omitempty"` // Model IDs (e.g., "user", "agents.mybot.product")
|
Models []string `json:"models,omitempty"` // Model IDs (e.g., "user", "agents.mybot.product")
|
||||||
Wheres []gou.Where `json:"wheres,omitempty"` // Pre-defined filters (optional), uses GOU QueryDSL Where
|
Scenario ScenarioType `json:"scenario,omitempty"` // QueryDSL scenario: "filter", "aggregation", "join", "complex"
|
||||||
Orders gou.Orders `json:"orders,omitempty"` // Sort orders (optional), uses GOU QueryDSL Orders
|
Wheres []gou.Where `json:"wheres,omitempty"` // Pre-defined filters (optional), uses GOU QueryDSL Where
|
||||||
Select []string `json:"select,omitempty"` // Fields to return (optional)
|
Orders gou.Orders `json:"orders,omitempty"` // Sort orders (optional), uses GOU QueryDSL Orders
|
||||||
|
Select []string `json:"select,omitempty"` // Fields to return (optional)
|
||||||
|
|
||||||
// Reranking
|
// Reranking
|
||||||
Rerank *RerankOptions `json:"rerank,omitempty"`
|
Rerank *RerankOptions `json:"rerank,omitempty"`
|
||||||
|
|
|
||||||
|
|
@ -3,9 +3,12 @@ package testutils
|
||||||
import (
|
import (
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
_ "github.com/yaoapp/gou/encoding"
|
||||||
|
_ "github.com/yaoapp/gou/text"
|
||||||
"github.com/yaoapp/yao/agent"
|
"github.com/yaoapp/yao/agent"
|
||||||
"github.com/yaoapp/yao/config"
|
"github.com/yaoapp/yao/config"
|
||||||
"github.com/yaoapp/yao/kb"
|
"github.com/yaoapp/yao/kb"
|
||||||
|
"github.com/yaoapp/yao/query"
|
||||||
"github.com/yaoapp/yao/test"
|
"github.com/yaoapp/yao/test"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -17,8 +20,14 @@ import (
|
||||||
func Prepare(t *testing.T, opts ...interface{}) {
|
func Prepare(t *testing.T, opts ...interface{}) {
|
||||||
test.Prepare(t, config.Conf, opts...)
|
test.Prepare(t, config.Conf, opts...)
|
||||||
|
|
||||||
|
// Load Query Engine (required for DB search)
|
||||||
|
err := query.Load(config.Conf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
// Load KB (required for agent KB features)
|
// Load KB (required for agent KB features)
|
||||||
_, err := kb.Load(config.Conf)
|
_, err = kb.Load(config.Conf)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
326
data/bindata.go
326
data/bindata.go
File diff suppressed because it is too large
Load diff
|
|
@ -1,7 +1,6 @@
|
||||||
{
|
{
|
||||||
"name": "Query Builder",
|
"name": "Query Builder",
|
||||||
"description": "Build database queries",
|
"description": "Build database queries",
|
||||||
"connector": "deepseek.v3",
|
|
||||||
"type": "worker",
|
"type": "worker",
|
||||||
"uses": { "search": "disabled" },
|
"uses": { "search": "disabled" },
|
||||||
"options": {
|
"options": {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue