Merge pull request #1387 from trheyi/main

Enhance Search Functionality with Keyword Extraction and Intent Detection
This commit is contained in:
Max 2025-12-18 20:40:02 +08:00 committed by GitHub
commit 248579bc65
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
58 changed files with 3834 additions and 1659 deletions

View file

@ -16,7 +16,7 @@ import (
// Stream stream the agent // Stream stream the agent
// handler is optional, if not provided, a default handler will be used // handler is optional, if not provided, a default handler will be used
func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Message, options ...*context.Options) (interface{}, error) { func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Message, options ...*context.Options) (*context.Response, error) {
// Update logger with assistant ID and start logging // Update logger with assistant ID and start logging
ctx.Logger.SetAssistantID(ast.ID) ctx.Logger.SetAssistantID(ast.ID)
@ -205,8 +205,8 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
// ================================================ // ================================================
// Execute Auto Search (if enabled) // Execute Auto Search (if enabled)
// ================================================ // ================================================
if ast.shouldAutoSearch(ctx, completionMessages, createResponse, opts) { if intent := ast.shouldAutoSearch(ctx, completionMessages, createResponse, opts); intent != nil {
refCtx := ast.executeAutoSearch(ctx, completionMessages, createResponse, opts) refCtx := ast.executeAutoSearch(ctx, completionMessages, createResponse, intent, opts)
if refCtx != nil && len(refCtx.References) > 0 { if refCtx != nil && len(refCtx.References) > 0 {
completionMessages = ast.injectSearchContext(completionMessages, refCtx) completionMessages = ast.injectSearchContext(completionMessages, refCtx)
} }
@ -376,7 +376,7 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
// ================================================ // ================================================
// Execute Next Hook and Process Response // Execute Next Hook and Process Response
// ================================================ // ================================================
var finalResponse interface{} var finalResponse *context.Response
var nextResponse *context.NextHookResponse = nil var nextResponse *context.NextHookResponse = nil
if ast.HookScript != nil { if ast.HookScript != nil {

View file

@ -56,19 +56,18 @@ func TestAgentNextStandard(t *testing.T) {
assert.NoError(t, err) assert.NoError(t, err)
assert.NotNil(t, response) assert.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion)
assert.NotNil(t, resp.Completion) assert.Nil(t, response.Next)
assert.Nil(t, resp.Next)
// Verify response structure // Verify response structure
assert.Equal(t, "tests.realworld-next", resp.AssistantID) assert.Equal(t, "tests.realworld-next", response.AssistantID)
assert.NotEmpty(t, resp.ContextID) assert.NotEmpty(t, response.ContextID)
assert.NotEmpty(t, resp.RequestID) assert.NotEmpty(t, response.RequestID)
assert.NotEmpty(t, resp.TraceID) assert.NotEmpty(t, response.TraceID)
assert.NotEmpty(t, resp.ChatID) assert.NotEmpty(t, response.ChatID)
// Verify completion has content // Verify completion has content
assert.NotNil(t, resp.Completion.Content) assert.NotNil(t, response.Completion.Content)
t.Log("✓ Standard response test passed") t.Log("✓ Standard response test passed")
} }
@ -94,19 +93,18 @@ func TestAgentNextCustomData(t *testing.T) {
assert.NoError(t, err) assert.NoError(t, err)
assert.NotNil(t, response) assert.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion)
assert.NotNil(t, resp.Completion) assert.NotNil(t, response.Next)
assert.NotNil(t, resp.Next)
// Verify response structure // Verify response structure
assert.Equal(t, "tests.realworld-next", resp.AssistantID) assert.Equal(t, "tests.realworld-next", response.AssistantID)
assert.NotEmpty(t, resp.ContextID) assert.NotEmpty(t, response.ContextID)
assert.NotEmpty(t, resp.RequestID) assert.NotEmpty(t, response.RequestID)
assert.NotEmpty(t, resp.TraceID) assert.NotEmpty(t, response.TraceID)
// Verify custom data structure (from scenarioCustomData) // Verify custom data structure (from scenarioCustomData)
// resp.Next contains the "data" field value from NextHookResponse // response.Next contains the "data" field value from NextHookResponse
nextData, ok := resp.Next.(map[string]interface{}) nextData, ok := response.Next.(map[string]interface{})
assert.True(t, ok, "Next should be a map") assert.True(t, ok, "Next should be a map")
assert.Equal(t, "custom_response", nextData["type"]) assert.Equal(t, "custom_response", nextData["type"])
assert.Equal(t, "This is a custom response from Next Hook", nextData["message"]) assert.Equal(t, "This is a custom response from Next Hook", nextData["message"])
@ -137,21 +135,19 @@ func TestAgentNextDelegate(t *testing.T) {
assert.NoError(t, err) assert.NoError(t, err)
assert.NotNil(t, response) assert.NotNil(t, response)
resp := response.(*context.Response)
// Verify response structure // Verify response structure
assert.NotEmpty(t, resp.AssistantID) assert.NotEmpty(t, response.AssistantID)
assert.NotEmpty(t, resp.ContextID) assert.NotEmpty(t, response.ContextID)
assert.NotEmpty(t, resp.RequestID) assert.NotEmpty(t, response.RequestID)
assert.NotEmpty(t, resp.TraceID) assert.NotEmpty(t, response.TraceID)
// Verify completion (delegated agent should have returned completion) // Verify completion (delegated agent should have returned completion)
assert.NotNil(t, resp.Completion) assert.NotNil(t, response.Completion)
assert.NotNil(t, resp.Completion.Content) assert.NotNil(t, response.Completion.Content)
// Next should be from the delegated agent // Next should be from the delegated agent
// If delegated agent also has Next hook, it will be present // If delegated agent also has Next hook, it will be present
t.Logf("✓ Delegation test passed (delegated to: %s)", resp.AssistantID) t.Logf("✓ Delegation test passed (delegated to: %s)", response.AssistantID)
} }
// TestAgentNextConditional tests agent with conditional logic in Next Hook // TestAgentNextConditional tests agent with conditional logic in Next Hook
@ -177,18 +173,17 @@ func TestAgentNextConditional(t *testing.T) {
assert.NoError(t, err) assert.NoError(t, err)
assert.NotNil(t, response) assert.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Next)
assert.NotNil(t, resp.Next)
// Verify response structure // Verify response structure
assert.Equal(t, "tests.realworld-next", resp.AssistantID) assert.Equal(t, "tests.realworld-next", response.AssistantID)
assert.NotEmpty(t, resp.ContextID) assert.NotEmpty(t, response.ContextID)
assert.NotEmpty(t, resp.RequestID) assert.NotEmpty(t, response.RequestID)
assert.NotEmpty(t, resp.TraceID) assert.NotEmpty(t, response.TraceID)
// Verify conditional response structure (from scenarioConditional) // Verify conditional response structure (from scenarioConditional)
// resp.Next contains the "data" field value from NextHookResponse // response.Next contains the "data" field value from NextHookResponse
nextData, ok := resp.Next.(map[string]interface{}) nextData, ok := response.Next.(map[string]interface{})
assert.True(t, ok, "Next should be a map") assert.True(t, ok, "Next should be a map")
assert.Equal(t, "Conditional analysis complete", nextData["message"]) assert.Equal(t, "Conditional analysis complete", nextData["message"])
assert.Contains(t, nextData, "action") assert.Contains(t, nextData, "action")
@ -224,19 +219,18 @@ func TestAgentWithoutNextHook(t *testing.T) {
assert.NoError(t, err) assert.NoError(t, err)
assert.NotNil(t, response) assert.NotNil(t, response)
resp := response.(*context.Response) assert.Nil(t, response.Next)
assert.Nil(t, resp.Next)
// Verify response structure // Verify response structure
assert.Equal(t, "tests.create", resp.AssistantID) assert.Equal(t, "tests.create", response.AssistantID)
assert.NotEmpty(t, resp.ContextID) assert.NotEmpty(t, response.ContextID)
assert.NotEmpty(t, resp.RequestID) assert.NotEmpty(t, response.RequestID)
assert.NotEmpty(t, resp.TraceID) assert.NotEmpty(t, response.TraceID)
assert.NotEmpty(t, resp.ChatID) assert.NotEmpty(t, response.ChatID)
// Verify completion // Verify completion
assert.NotNil(t, resp.Completion) assert.NotNil(t, response.Completion)
assert.NotNil(t, resp.Completion.Content) assert.NotNil(t, response.Completion.Content)
t.Log("✓ No Next Hook test passed") t.Log("✓ No Next Hook test passed")
} }

View file

@ -51,7 +51,7 @@ type agentCallerWrapper struct {
ast *Assistant ast *Assistant
} }
func (w *agentCallerWrapper) Stream(ctx *agentContext.Context, messages []agentContext.Message, options ...*agentContext.Options) (interface{}, error) { func (w *agentCallerWrapper) Stream(ctx *agentContext.Context, messages []agentContext.Message, options ...*agentContext.Options) (*agentContext.Response, error) {
return w.ast.Stream(ctx, messages, options...) return w.ast.Stream(ctx, messages, options...)
} }

View file

@ -8,7 +8,7 @@ import (
) )
// processNextResponse processes the Next hook's response and handles agent delegation or custom data // processNextResponse processes the Next hook's response and handles agent delegation or custom data
func (ast *Assistant) processNextResponse(npc *NextProcessContext) (interface{}, error) { func (ast *Assistant) processNextResponse(npc *NextProcessContext) (*agentContext.Response, error) {
// If no Next hook response, return standard response // If no Next hook response, return standard response
if npc.NextResponse == nil { if npc.NextResponse == nil {
return ast.buildStandardResponse(npc), nil return ast.buildStandardResponse(npc), nil
@ -42,7 +42,7 @@ func (ast *Assistant) handleDelegation(
ctx *agentContext.Context, ctx *agentContext.Context,
delegate *agentContext.DelegateConfig, delegate *agentContext.DelegateConfig,
streamHandler func(message.StreamChunkType, []byte) int, streamHandler func(message.StreamChunkType, []byte) int,
) (interface{}, error) { ) (*agentContext.Response, error) {
// Load the target assistant // Load the target assistant
targetAssistant, err := Get(delegate.AgentID) targetAssistant, err := Get(delegate.AgentID)
if err != nil { if err != nil {
@ -62,7 +62,7 @@ func (ast *Assistant) handleDelegation(
} }
// buildStandardResponse builds the standard agent response when no custom Next hook processing is needed // buildStandardResponse builds the standard agent response when no custom Next hook processing is needed
func (ast *Assistant) buildStandardResponse(npc *NextProcessContext) interface{} { func (ast *Assistant) buildStandardResponse(npc *NextProcessContext) *agentContext.Response {
return &agentContext.Response{ return &agentContext.Response{
ContextID: npc.Context.ID, ContextID: npc.Context.ID,
RequestID: npc.Context.RequestID(), RequestID: npc.Context.RequestID(),

View file

@ -17,16 +17,32 @@ import (
) )
// shouldAutoSearch determines if auto search should be executed // shouldAutoSearch determines if auto search should be executed
// Returns false if: // Returns nil if search should be skipped, otherwise returns SearchIntent with types to search
// Search is skipped if:
// - opts.Skip.Search is true // - opts.Skip.Search is true
// - createResponse.Search is false
// - uses.search is "disabled" // - uses.search is "disabled"
// - assistant has no search configuration // - assistant has no search configuration
// - needsearch intent detection returns false // - needsearch intent detection returns false
func (ast *Assistant) shouldAutoSearch(ctx *context.Context, messages []context.Message, createResponse *context.HookCreateResponse, opts *context.Options) bool { func (ast *Assistant) shouldAutoSearch(ctx *context.Context, messages []context.Message, createResponse *context.HookCreateResponse, opts *context.Options) *SearchIntent {
// Check if search is skipped via options // Check if search is skipped via options
if opts != nil && opts.Skip != nil && opts.Skip.Search { if opts != nil && opts.Skip != nil && opts.Skip.Search {
ctx.Logger.Debug("Auto search skipped by opts.Skip.Search") ctx.Logger.Debug("Auto search skipped by opts.Skip.Search")
return false return nil
}
// Check createResponse.Search field (highest priority from Create hook)
// Supports: bool | SearchIntent | nil
if createResponse != nil && createResponse.Search != nil {
intent := parseSearchField(createResponse.Search)
if intent != nil {
if !intent.NeedSearch {
ctx.Logger.Info("Auto search disabled by createResponse.Search")
return nil
}
ctx.Logger.Info("Auto search controlled by createResponse.Search: types=%v", intent.SearchTypes)
return intent
}
} }
// Get merged uses configuration // Get merged uses configuration
@ -35,57 +51,123 @@ func (ast *Assistant) shouldAutoSearch(ctx *context.Context, messages []context.
// Check if search is explicitly disabled // Check if search is explicitly disabled
if uses != nil && uses.Search == "disabled" { if uses != nil && uses.Search == "disabled" {
ctx.Logger.Info("Auto search disabled by uses.search=disabled") ctx.Logger.Info("Auto search disabled by uses.search=disabled")
return false return nil
} }
// Check if assistant has search configuration // Check if assistant has search configuration
if ast.Search == nil && (uses == nil || uses.Search == "") { if ast.Search == nil && (uses == nil || uses.Search == "") {
return false return nil
} }
// Check search intent using __yao.needsearch agent // Check search intent using __yao.needsearch agent
if !ast.checkSearchIntent(ctx, messages) { intent := ast.checkSearchIntent(ctx, messages)
if intent == nil || !intent.NeedSearch {
ctx.Logger.Info("Auto search skipped: intent detection returned false") ctx.Logger.Info("Auto search skipped: intent detection returned false")
return false return nil
} }
// Check if search is enabled (builtin, agent, mcp, or empty means builtin) return intent
return true }
// parseSearchField parses the Search field from HookCreateResponse
// Supports: bool | SearchIntent | map[string]any | nil
func parseSearchField(search any) *SearchIntent {
if search == nil {
return nil
}
switch v := search.(type) {
case bool:
// bool: true = enable all, false = disable all
if v {
return &SearchIntent{
NeedSearch: true,
SearchTypes: []string{"web", "kb", "db"},
Confidence: 1.0,
Reason: "enabled by hook",
}
}
return &SearchIntent{
NeedSearch: false,
SearchTypes: []string{},
Confidence: 1.0,
Reason: "disabled by hook",
}
case *SearchIntent:
// SearchIntent is alias for context.SearchIntent, so this covers both
return v
case SearchIntent:
return &v
case map[string]any:
// Parse from map (e.g., from JSON)
intent := &SearchIntent{
NeedSearch: false,
SearchTypes: []string{},
Confidence: 0.5,
}
if needSearch, ok := v["need_search"].(bool); ok {
intent.NeedSearch = needSearch
}
if types, ok := v["search_types"].([]any); ok {
for _, t := range types {
if typeStr, ok := t.(string); ok {
intent.SearchTypes = append(intent.SearchTypes, typeStr)
}
}
}
if confidence, ok := v["confidence"].(float64); ok {
intent.Confidence = confidence
}
if reason, ok := v["reason"].(string); ok {
intent.Reason = reason
}
return intent
default:
return nil
}
} }
// checkSearchIntent uses __yao.needsearch agent to determine if search is needed // checkSearchIntent uses __yao.needsearch agent to determine if search is needed
// Returns true if search is needed, false otherwise // Returns SearchIntent with search types and confidence
func (ast *Assistant) checkSearchIntent(ctx *context.Context, messages []context.Message) bool { func (ast *Assistant) checkSearchIntent(ctx *context.Context, messages []context.Message) *SearchIntent {
// Get the last user message // Default intent: no search needed (fallback when agent unavailable or fails)
var userQuery string defaultIntent := &SearchIntent{
for i := len(messages) - 1; i >= 0; i-- { NeedSearch: false,
if messages[i].Role == "user" { SearchTypes: []string{},
if content, ok := messages[i].Content.(string); ok { Confidence: 0,
userQuery = content
break
} }
// Filter out system messages and pass full conversation context
var intentMessages []context.Message
for _, msg := range messages {
if msg.Role != "system" {
intentMessages = append(intentMessages, msg)
} }
} }
if userQuery == "" { if len(intentMessages) == 0 {
return true // No user message, proceed with search return defaultIntent // No messages, skip search
} }
// Try to get __yao.needsearch agent // Try to get __yao.needsearch agent
needsearchAst, err := Get("__yao.needsearch") needsearchAst, err := Get("__yao.needsearch")
if err != nil { if err != nil {
ctx.Logger.Debug("__yao.needsearch agent not available: %v, proceeding with search", err) ctx.Logger.Debug("__yao.needsearch agent not available: %v, skipping search", err)
return true // Agent not available, proceed with search return defaultIntent // Agent not available, skip search
} }
// === Output: Send loading message === // === Output: Send loading message ===
loadingID := ast.sendIntentLoading(ctx) loadingID := ast.sendIntentLoading(ctx)
// Build messages for intent detection
intentMessages := []context.Message{
{Role: "user", Content: userQuery},
}
// Call the needsearch agent (Stack will auto-track) // Call the needsearch agent (Stack will auto-track)
// IMPORTANT: Skip search to prevent infinite loop, skip output to prevent JSON showing in UI // IMPORTANT: Skip search to prevent infinite loop, skip output to prevent JSON showing in UI
opts := &context.Options{ opts := &context.Options{
@ -98,18 +180,17 @@ func (ast *Assistant) checkSearchIntent(ctx *context.Context, messages []context
result, err := needsearchAst.Stream(ctx, intentMessages, opts) result, err := needsearchAst.Stream(ctx, intentMessages, opts)
if err != nil { if err != nil {
ctx.Logger.Debug("__yao.needsearch failed: %v, proceeding with search", err) ctx.Logger.Debug("__yao.needsearch failed: %v, skipping search", err)
// === Output: Send done (error case, proceed with search) === // === Output: Send done (error case, skip search) ===
ast.sendIntentDone(ctx, loadingID, true, "") ast.sendIntentDone(ctx, loadingID, false, "")
return true // On error, proceed with search return defaultIntent // On error, skip search
} }
// Parse the result // Parse the result
// Next hook returns {data: {need_search: bool, search_types: [], confidence: float}} // Next hook returns {data: {need_search: bool, search_types: [], confidence: float}}
if response, ok := result.(*context.Response); ok {
// First try to get from Next hook response // First try to get from Next hook response
if response.Next != nil { if result.Next != nil {
if nextData, ok := response.Next.(map[string]interface{}); ok { if nextData, ok := result.Next.(map[string]interface{}); ok {
// Check for data field (from Next hook's {data: result}) // Check for data field (from Next hook's {data: result})
var intentData map[string]interface{} var intentData map[string]interface{}
if data, ok := nextData["data"].(map[string]interface{}); ok { if data, ok := nextData["data"].(map[string]interface{}); ok {
@ -118,38 +199,89 @@ func (ast *Assistant) checkSearchIntent(ctx *context.Context, messages []context
intentData = nextData intentData = nextData
} }
if needSearch, ok := intentData["need_search"].(bool); ok { intent := parseSearchIntent(intentData)
reason, _ := intentData["reason"].(string) if intent != nil {
ctx.Logger.Debug("Search intent (from Next): need_search=%v, reason=%s", needSearch, reason) ctx.Logger.Debug("Search intent (from Next): need_search=%v, types=%v, confidence=%.2f, reason=%s",
ast.sendIntentDone(ctx, loadingID, needSearch, reason) intent.NeedSearch, intent.SearchTypes, intent.Confidence, intent.Reason)
return needSearch ast.sendIntentDone(ctx, loadingID, intent.NeedSearch, intent.Reason)
return intent
} }
} }
} }
// Fallback: parse from Completion.Content if Next hook didn't process // Fallback: parse from Completion.Content if Next hook didn't process
if response.Completion != nil { if result.Completion != nil {
content, ok := response.Completion.Content.(string) content, ok := result.Completion.Content.(string)
if !ok || content == "" { if !ok || content == "" {
ast.sendIntentDone(ctx, loadingID, true, "") ast.sendIntentDone(ctx, loadingID, false, "")
return true return defaultIntent
}
needSearch, reason := parseNeedSearchFromContent(content)
ctx.Logger.Debug("Search intent (from Content): need_search=%v, reason=%s", needSearch, reason)
ast.sendIntentDone(ctx, loadingID, needSearch, reason)
return needSearch
} }
intent := parseSearchIntentFromContent(content)
ctx.Logger.Debug("Search intent (from Content): need_search=%v, types=%v, confidence=%.2f, reason=%s",
intent.NeedSearch, intent.SearchTypes, intent.Confidence, intent.Reason)
ast.sendIntentDone(ctx, loadingID, intent.NeedSearch, intent.Reason)
return intent
} }
// Default: proceed with search if we can't parse the result // Default: skip search if we can't parse the result
// === Output: Send done (default case) === // === Output: Send done (default case) ===
ast.sendIntentDone(ctx, loadingID, true, "") ast.sendIntentDone(ctx, loadingID, false, "")
return true return defaultIntent
} }
// parseNeedSearchFromContent parses need_search result from LLM completion content // parseSearchIntent parses SearchIntent from intent data map
func parseSearchIntent(intentData map[string]interface{}) *SearchIntent {
if intentData == nil {
return nil
}
needSearch, ok := intentData["need_search"].(bool)
if !ok {
return nil
}
intent := &SearchIntent{
NeedSearch: needSearch,
SearchTypes: []string{},
Confidence: 0.5, // Default confidence
}
// Parse search_types
if types, ok := intentData["search_types"].([]interface{}); ok {
for _, t := range types {
if typeStr, ok := t.(string); ok {
// Validate type
typeStr = strings.ToLower(typeStr)
if typeStr == "web" || typeStr == "kb" || typeStr == "db" {
intent.SearchTypes = append(intent.SearchTypes, typeStr)
}
}
}
}
// Parse confidence
if confidence, ok := intentData["confidence"].(float64); ok {
intent.Confidence = confidence
}
// Parse reason
if reason, ok := intentData["reason"].(string); ok {
intent.Reason = reason
}
return intent
}
// parseSearchIntentFromContent parses SearchIntent from LLM completion content
// Handles JSON wrapped in markdown code blocks // Handles JSON wrapped in markdown code blocks
func parseNeedSearchFromContent(content string) (bool, string) { func parseSearchIntentFromContent(content string) *SearchIntent {
// Default intent: no search needed
defaultIntent := &SearchIntent{
NeedSearch: false,
SearchTypes: []string{},
Confidence: 0,
}
// Remove markdown code block if present // Remove markdown code block if present
content = strings.TrimSpace(content) content = strings.TrimSpace(content)
if strings.HasPrefix(content, "```json") { if strings.HasPrefix(content, "```json") {
@ -165,17 +297,16 @@ func parseNeedSearchFromContent(content string) (bool, string) {
// Try to parse JSON // Try to parse JSON
var result map[string]interface{} var result map[string]interface{}
if err := json.Unmarshal([]byte(content), &result); err != nil { if err := json.Unmarshal([]byte(content), &result); err != nil {
// Failed to parse, default to search // Failed to parse, default to no search
return true, "" return defaultIntent
} }
needSearch, ok := result["need_search"].(bool) intent := parseSearchIntent(result)
if !ok { if intent == nil {
return true, "" return defaultIntent
} }
reason, _ := result["reason"].(string) return intent
return needSearch, reason
} }
// sendIntentLoading sends the initial intent detection loading message // sendIntentLoading sends the initial intent detection loading message
@ -271,10 +402,11 @@ func (ast *Assistant) getMergedSearchUses(createResponse *context.HookCreateResp
return uses return uses
} }
// executeAutoSearch executes auto search based on configuration // executeAutoSearch executes auto search based on configuration and intent
// Returns ReferenceContext with results and formatted context // Returns ReferenceContext with results and formatted context
// intent specifies which search types to execute (from needsearch agent)
// opts is optional, used to check Skip.Keyword // opts is optional, used to check Skip.Keyword
func (ast *Assistant) executeAutoSearch(ctx *context.Context, messages []context.Message, createResponse *context.HookCreateResponse, opts ...*context.Options) *searchTypes.ReferenceContext { func (ast *Assistant) executeAutoSearch(ctx *context.Context, messages []context.Message, createResponse *context.HookCreateResponse, intent *SearchIntent, opts ...*context.Options) *searchTypes.ReferenceContext {
ctx.Logger.Phase("Search") ctx.Logger.Phase("Search")
defer ctx.Logger.PhaseComplete("Search") defer ctx.Logger.PhaseComplete("Search")
@ -311,33 +443,23 @@ func (ast *Assistant) executeAutoSearch(ctx *context.Context, messages []context
skipKeyword = opts[0].Skip.Keyword skipKeyword = opts[0].Skip.Keyword
} }
// Extract keywords for web search if: // Build search requests based on configuration and intent
// 1. uses.keyword is configured (not empty) // Keyword extraction is done inside buildSearchRequests for web search
// 2. Skip.Keyword is not true buildOpts := &buildSearchRequestsOptions{
// 3. Web search is enabled skipKeyword: skipKeyword,
var extractedKeywords []string usesKeyword: searchUses.Keyword,
webSearchEnabled := searchConfig != nil && searchConfig.Web != nil
if webSearchEnabled && !skipKeyword && searchUses.Keyword != "" {
extractor := keyword.NewExtractor(searchUses.Keyword, searchConfig.Keyword)
keywords, err := extractor.Extract(ctx, query, nil)
if err != nil {
ctx.Logger.Warn("Keyword extraction failed, using original query: %v", err)
} else if len(keywords) > 0 {
extractedKeywords = keywords
// Use extracted keywords as the search query for web search
optimizedQuery := strings.Join(keywords, " ")
ctx.Logger.Info("Extracted keywords for web search: %s -> %s", truncateString(query, 30), optimizedQuery)
query = optimizedQuery
} }
} requests, extractedKeywords := ast.buildSearchRequests(ctx, query, searchConfig, intent, buildOpts)
// Build search requests based on configuration
requests := ast.buildSearchRequests(query, searchConfig)
if len(requests) == 0 { if len(requests) == 0 {
ctx.Logger.Info("No search requests to execute") ctx.Logger.Info("No search requests to execute")
return nil return nil
} }
// Update query if keywords were extracted (for web search)
if len(extractedKeywords) > 0 {
query = keywordsToQuery(extractedKeywords)
}
// === Output: Send loading message === // === Output: Send loading message ===
loadingID := ast.sendSearchLoading(ctx) loadingID := ast.sendSearchLoading(ctx)
@ -435,6 +557,52 @@ func (ast *Assistant) sendSearchLoading(ctx *context.Context) string {
return msgID return msgID
} }
// sendKeywordLoading sends the keyword extraction loading message
// Returns the message ID for later replacement
func (ast *Assistant) sendKeywordLoading(ctx *context.Context) string {
loadingMsg := i18n.T(ctx.Locale, "search.keyword.loading")
msg := &message.Message{
Type: "loading",
Props: map[string]any{
"message": loadingMsg,
},
}
// Send and get message ID
msgID, err := ctx.SendStream(msg)
if err != nil {
ctx.Logger.Warn("Failed to send keyword loading message: %v", err)
return ""
}
return msgID
}
// sendKeywordDone replaces keyword loading with done message
func (ast *Assistant) sendKeywordDone(ctx *context.Context, loadingID string, success bool) {
if loadingID == "" {
return
}
resultMsg := i18n.T(ctx.Locale, "search.keyword.done")
msg := &message.Message{
MessageID: loadingID,
Delta: true,
DeltaAction: message.DeltaReplace,
Type: "loading",
Props: map[string]any{
"message": resultMsg,
"done": true,
},
}
if err := ctx.Send(msg); err != nil {
ctx.Logger.Warn("Failed to send keyword done message: %v", err)
}
}
// sendSearchResult replaces loading with result message (without done flag) // sendSearchResult replaces loading with result message (without done flag)
func (ast *Assistant) sendSearchResult(ctx *context.Context, loadingID string, count int) { func (ast *Assistant) sendSearchResult(ctx *context.Context, loadingID string, count int) {
if loadingID == "" { if loadingID == "" {
@ -564,22 +732,67 @@ func (ast *Assistant) completeSearchTrace(node traceTypes.Node, resultCount int,
}) })
} }
// buildSearchRequests builds search requests based on assistant configuration // buildSearchRequestsOptions contains options for building search requests
func (ast *Assistant) buildSearchRequests(query string, config *searchTypes.Config) []*searchTypes.Request { type buildSearchRequestsOptions struct {
var requests []*searchTypes.Request skipKeyword bool // Skip keyword extraction
usesKeyword string // Keyword extractor config: "builtin", "<assistant-id>", "mcp:<server>.<tool>"
}
// buildSearchRequests builds search requests based on assistant configuration and intent
// intent specifies which search types to execute (from needsearch agent)
// Returns requests and extracted keywords (if any)
func (ast *Assistant) buildSearchRequests(ctx *context.Context, query string, config *searchTypes.Config, intent *SearchIntent, opts *buildSearchRequestsOptions) ([]*searchTypes.Request, []searchTypes.Keyword) {
var requests []*searchTypes.Request
var extractedKeywords []searchTypes.Keyword
// Helper to check if a search type is allowed by intent
isTypeAllowed := func(searchType string) bool {
if intent == nil || len(intent.SearchTypes) == 0 {
return true // No intent or empty types means all types allowed
}
for _, t := range intent.SearchTypes {
if t == searchType {
return true
}
}
return false
}
// Web search - check if web search is configured and allowed by intent
if config != nil && config.Web != nil && isTypeAllowed("web") {
webQuery := query
// Extract keywords for web search if configured
if opts != nil && !opts.skipKeyword && opts.usesKeyword != "" {
// === Output: Send keyword extraction loading ===
keywordLoadingID := ast.sendKeywordLoading(ctx)
extractor := keyword.NewExtractor(opts.usesKeyword, config.Keyword)
keywords, err := extractor.Extract(ctx, query, nil)
if err != nil {
ctx.Logger.Warn("Keyword extraction failed, using original query: %v", err)
ast.sendKeywordDone(ctx, keywordLoadingID, false)
} else if len(keywords) > 0 {
extractedKeywords = keywords
// Use extracted keywords as the search query for web search
webQuery = keywordsToQuery(keywords)
ctx.Logger.Info("Extracted keywords for web search: %s -> %s", truncateString(query, 30), webQuery)
ast.sendKeywordDone(ctx, keywordLoadingID, true)
} else {
ast.sendKeywordDone(ctx, keywordLoadingID, true)
}
}
// Web search - check if web search is configured
if config != nil && config.Web != nil {
requests = append(requests, &searchTypes.Request{ requests = append(requests, &searchTypes.Request{
Type: searchTypes.SearchTypeWeb, Type: searchTypes.SearchTypeWeb,
Query: query, Query: webQuery,
Source: searchTypes.SourceAuto, Source: searchTypes.SourceAuto,
Limit: config.Web.MaxResults, Limit: config.Web.MaxResults,
}) })
} }
// KB search - check if KB is configured // KB search - check if KB is configured and allowed by intent
if ast.KB != nil && len(ast.KB.Collections) > 0 { if ast.KB != nil && len(ast.KB.Collections) > 0 && isTypeAllowed("kb") {
limit := 10 limit := 10
threshold := 0.7 threshold := 0.7
if config != nil && config.KB != nil { if config != nil && config.KB != nil {
@ -589,7 +802,7 @@ func (ast *Assistant) buildSearchRequests(query string, config *searchTypes.Conf
} }
requests = append(requests, &searchTypes.Request{ requests = append(requests, &searchTypes.Request{
Type: searchTypes.SearchTypeKB, Type: searchTypes.SearchTypeKB,
Query: query, Query: query, // KB uses original query for semantic search
Source: searchTypes.SourceAuto, Source: searchTypes.SourceAuto,
Limit: limit, Limit: limit,
Collections: ast.KB.Collections, Collections: ast.KB.Collections,
@ -598,22 +811,22 @@ func (ast *Assistant) buildSearchRequests(query string, config *searchTypes.Conf
}) })
} }
// DB search - check if DB is configured // DB search - check if DB is configured and allowed by intent
if ast.DB != nil && len(ast.DB.Models) > 0 { if ast.DB != nil && len(ast.DB.Models) > 0 && isTypeAllowed("db") {
limit := 20 limit := 20
if config != nil && config.DB != nil && config.DB.MaxResults > 0 { if config != nil && config.DB != nil && config.DB.MaxResults > 0 {
limit = config.DB.MaxResults limit = config.DB.MaxResults
} }
requests = append(requests, &searchTypes.Request{ requests = append(requests, &searchTypes.Request{
Type: searchTypes.SearchTypeDB, Type: searchTypes.SearchTypeDB,
Query: query, Query: query, // DB uses original query for QueryDSL generation
Source: searchTypes.SourceAuto, Source: searchTypes.SourceAuto,
Limit: limit, Limit: limit,
Models: ast.DB.Models, Models: ast.DB.Models,
}) })
} }
return requests return requests, extractedKeywords
} }
// injectSearchContext injects search results into messages // injectSearchContext injects search results into messages
@ -709,7 +922,7 @@ func truncateString(s string, maxLen int) string {
// SearchExecutionResult holds all data from search execution for storage // SearchExecutionResult holds all data from search execution for storage
type SearchExecutionResult struct { type SearchExecutionResult struct {
Query string // Original query (before keyword optimization) Query string // Original query (before keyword optimization)
Keywords []string // Extracted keywords Keywords []searchTypes.Keyword // Extracted keywords with weights
Config map[string]any // Search config used Config map[string]any // Search config used
RefCtx *searchTypes.ReferenceContext // Reference context with results RefCtx *searchTypes.ReferenceContext // Reference context with results
Duration int64 // Search duration in ms Duration int64 // Search duration in ms
@ -717,6 +930,54 @@ type SearchExecutionResult struct {
SearchType string // "auto", "web", "kb", "db" SearchType string // "auto", "web", "kb", "db"
} }
// keywordsToQuery converts keywords with weights to a search query string
// Keywords are sorted by weight (descending) and joined with spaces
func keywordsToQuery(keywords []searchTypes.Keyword) string {
if len(keywords) == 0 {
return ""
}
// Sort by weight descending (higher weight first)
sorted := make([]searchTypes.Keyword, len(keywords))
copy(sorted, keywords)
for i := 0; i < len(sorted)-1; i++ {
for j := i + 1; j < len(sorted); j++ {
if sorted[j].W > sorted[i].W {
sorted[i], sorted[j] = sorted[j], sorted[i]
}
}
}
// Join keywords
parts := make([]string, len(sorted))
for i, kw := range sorted {
parts[i] = kw.K
}
return strings.Join(parts, " ")
}
// keywordsToStrings converts keywords to string slice for storage
func keywordsToStrings(keywords []searchTypes.Keyword) []string {
if len(keywords) == 0 {
return nil
}
result := make([]string, len(keywords))
for i, kw := range keywords {
result[i] = kw.K
}
return result
}
// containsSearchType checks if a search type is in the list
func containsSearchType(types []string, searchType string) bool {
for _, t := range types {
if t == searchType {
return true
}
}
return false
}
// saveSearch saves search results to storage // saveSearch saves search results to storage
// Called after search execution completes (success or failure) // Called after search execution completes (success or failure)
func (ast *Assistant) saveSearch(ctx *context.Context, execResult *SearchExecutionResult) { func (ast *Assistant) saveSearch(ctx *context.Context, execResult *SearchExecutionResult) {
@ -732,7 +993,7 @@ func (ast *Assistant) saveSearch(ctx *context.Context, execResult *SearchExecuti
RequestID: ctx.RequestID(), RequestID: ctx.RequestID(),
ChatID: ctx.ChatID, ChatID: ctx.ChatID,
Query: execResult.Query, Query: execResult.Query,
Keywords: execResult.Keywords, Keywords: keywordsToStrings(execResult.Keywords),
Config: execResult.Config, Config: execResult.Config,
Source: execResult.SearchType, Source: execResult.SearchType,
Duration: execResult.Duration, Duration: execResult.Duration,

View file

@ -77,8 +77,7 @@ func TestSearchAutoDisabled(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
require.NotNil(t, response) require.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion, "should have completion")
assert.NotNil(t, resp.Completion, "should have completion")
t.Logf("✓ Stream executed without search (disabled)") t.Logf("✓ Stream executed without search (disabled)")
}) })
} }

View file

@ -118,8 +118,7 @@ func TestSearchAutoFull(t *testing.T) {
} }
require.NotNil(t, response) require.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion, "should have completion")
assert.NotNil(t, resp.Completion, "should have completion")
t.Logf("✓ Stream executed with full search config (Web + KB + DB)") t.Logf("✓ Stream executed with full search config (Web + KB + DB)")
}) })
} }

View file

@ -101,8 +101,7 @@ func TestSearchAutoHookDisable(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
require.NotNil(t, response) require.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion, "should have completion")
assert.NotNil(t, resp.Completion, "should have completion")
t.Logf("✓ Stream executed with hook disabling search") t.Logf("✓ Stream executed with hook disabling search")
}) })
} }

View file

@ -87,8 +87,7 @@ func TestSearchAutoKeyword(t *testing.T) {
} }
require.NotNil(t, response) require.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion, "should have completion")
assert.NotNil(t, resp.Completion, "should have completion")
t.Logf("✓ Stream with keyword extraction executed successfully") t.Logf("✓ Stream with keyword extraction executed successfully")
}) })
@ -127,8 +126,7 @@ func TestSearchAutoKeyword(t *testing.T) {
} }
require.NotNil(t, response) require.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion, "should have completion")
assert.NotNil(t, resp.Completion, "should have completion")
t.Logf("✓ Stream with Skip.Keyword executed successfully") t.Logf("✓ Stream with Skip.Keyword executed successfully")
}) })
} }
@ -179,8 +177,7 @@ func TestSearchAutoKeywordNotConfigured(t *testing.T) {
} }
require.NotNil(t, response) require.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion, "should have completion")
assert.NotNil(t, resp.Completion, "should have completion")
t.Logf("✓ Stream without keyword config executed successfully") t.Logf("✓ Stream without keyword config executed successfully")
}) })
} }

View file

@ -96,8 +96,7 @@ func TestSearchAutoWeb(t *testing.T) {
} }
require.NotNil(t, response) require.NotNil(t, response)
resp := response.(*context.Response) assert.NotNil(t, response.Completion, "should have completion")
assert.NotNil(t, resp.Completion, "should have completion")
t.Logf("✓ Stream executed successfully with auto search") t.Logf("✓ Stream executed successfully with auto search")
}) })
} }

View file

@ -85,6 +85,10 @@ type NextProcessContext struct {
CreateResponse *chatctx.HookCreateResponse // Create hook response CreateResponse *chatctx.HookCreateResponse // Create hook response
} }
// SearchIntent is an alias for context.SearchIntent
// Used for search intent detection from __yao.needsearch agent
type SearchIntent = chatctx.SearchIntent
// ParsedContent extracts the actual tool return value from MCP ToolContent array // ParsedContent extracts the actual tool return value from MCP ToolContent array
// According to MCP protocol: // According to MCP protocol:
// - Content is []ToolContent array // - Content is []ToolContent array

View file

@ -9,7 +9,7 @@ import (
// AgentCaller interface for calling agents (to avoid circular dependency) // AgentCaller interface for calling agents (to avoid circular dependency)
// Used by content handlers (vision, audio, etc.) and search handlers (agent mode) // Used by content handlers (vision, audio, etc.) and search handlers (agent mode)
type AgentCaller interface { type AgentCaller interface {
Stream(ctx *agentContext.Context, messages []agentContext.Message, options ...*agentContext.Options) (interface{}, error) Stream(ctx *agentContext.Context, messages []agentContext.Message, options ...*agentContext.Options) (*agentContext.Response, error)
} }
// AgentGetterFunc is a function type that gets an agent by ID // AgentGetterFunc is a function type that gets an agent by ID

View file

@ -42,10 +42,10 @@ func CallAgent(ctx *agentContext.Context, agentID string, message agentContext.M
// Extract text from agent response // Extract text from agent response
// Two formats are supported: // Two formats are supported:
// 1. Custom Hook response (from Next hook) // 1. Custom Hook response (from Next hook) - response.Next
// 2. Standard Agent Stream response (LLM completion) // 2. Standard Agent Stream response (LLM completion) - response.Completion
return extractTextFromAgentResponse(response) return extractTextFromResponse(response)
} }
// CallAgentWithFileInfo calls an agent to process content with file metadata // CallAgentWithFileInfo calls an agent to process content with file metadata
@ -120,78 +120,38 @@ func CallAgentWithFileInfo(ctx *agentContext.Context, agentID string, message ag
return CallAgent(ctx, agentID, message) return CallAgent(ctx, agentID, message)
} }
// extractTextFromAgentResponse extracts text from agent response // extractTextFromResponse extracts text from agent response
// Handles two main response formats from agent.Stream(): // Now that agent.Stream() returns *agentContext.Response directly,
// // we can access fields without type assertions or JSON conversion.
// 1. Standard Response (No Next Hook or Next Hook returns nil):
// Structure: { completion: { content: "text" | [...ContentPart] } }
// Action: Extract text from completion.content field
//
// 2. Next Hook Response with Custom Data:
// Structure: { next: <any data from Next hook> }
// Action:
// - If next is string → return directly
// - If next is map/object → JSON stringify and return
// - This preserves the complete custom data structure from the hook
// //
// Priority: // Priority:
// 1. Check for "next" field (custom hook data) → return complete data // 1. Check response.Next (custom hook data) → return complete data
// 2. Check for "completion" field (standard LLM response) → extract text only // 2. Check response.Completion (standard LLM response) → extract text only
// 3. Fallback to direct string or JSON stringify func extractTextFromResponse(response *agentContext.Response) (string, error) {
func extractTextFromAgentResponse(response interface{}) (string, error) {
if response == nil { if response == nil {
return "", fmt.Errorf("agent returned nil response") return "", fmt.Errorf("agent returned nil response")
} }
// First, try to convert to map if it's a struct // Priority 1: Check Next field (custom hook data)
// agent.Stream() may return *agentContext.Response which needs to be converted // If Next hook returns custom data, return the complete structure
var responseMap map[string]interface{} if response.Next != nil {
// Check if it's already a map
if rm, ok := response.(map[string]interface{}); ok {
responseMap = rm
} else {
// Try to marshal and unmarshal to convert struct to map
jsonBytes, err := jsoniter.Marshal(response)
if err != nil {
// If it's a plain string, return directly
if responseStr, ok := response.(string); ok {
return responseStr, nil
}
return "", fmt.Errorf("failed to serialize agent response: %w", err)
}
// Unmarshal to map
if err := jsoniter.Unmarshal(jsonBytes, &responseMap); err != nil {
// If unmarshal fails, return the JSON string
return string(jsonBytes), nil
}
}
// Priority 1: Check for "next" field (custom hook data)
// If Next hook returns custom data, it's stored in the "next" field
// Return the complete custom data structure (preserve hook's intent)
if next, hasNext := responseMap["next"]; hasNext && next != nil {
// If next is a string, return directly // If next is a string, return directly
if nextStr, ok := next.(string); ok { if nextStr, ok := response.Next.(string); ok {
return nextStr, nil return nextStr, nil
} }
// Otherwise, JSON stringify to preserve complete structure // Otherwise, JSON stringify to preserve complete structure
jsonBytes, err := jsoniter.Marshal(next) jsonBytes, err := jsoniter.Marshal(response.Next)
if err != nil { if err != nil {
return "", fmt.Errorf("failed to serialize next hook data: %w", err) return "", fmt.Errorf("failed to serialize next hook data: %w", err)
} }
return string(jsonBytes), nil return string(jsonBytes), nil
} }
// Priority 2: Check for "completion" field (standard LLM response) // Priority 2: Check Completion field (standard LLM response)
// Extract text content from the LLM completion // Extract text content from the LLM completion
if completion, hasCompletion := responseMap["completion"]; hasCompletion && completion != nil { if response.Completion != nil {
if completionMap, ok := completion.(map[string]interface{}); ok {
// Extract content from completion
if content, hasContent := completionMap["content"]; hasContent {
// Content can be string or []ContentPart (multimodal) // Content can be string or []ContentPart (multimodal)
switch v := content.(type) { switch v := response.Completion.Content.(type) {
case string: case string:
// Simple text content // Simple text content
return v, nil return v, nil
@ -214,22 +174,9 @@ func extractTextFromAgentResponse(response interface{}) (string, error) {
return "", fmt.Errorf("no text content found in completion content parts") return "", fmt.Errorf("no text content found in completion content parts")
} }
} }
}
}
// Fallback: Try to find a "content" field directly (shouldn't happen normally) // No content found
if content, hasContent := responseMap["content"]; hasContent { return "", fmt.Errorf("no content found in agent response")
if contentStr, ok := content.(string); ok {
return contentStr, nil
}
}
// Last resort: JSON stringify the entire response
jsonBytes, err := jsoniter.Marshal(response)
if err != nil {
return "", fmt.Errorf("failed to serialize agent response: %w", err)
}
return string(jsonBytes), nil
} }
// CallMCPTool calls an MCP tool to process content // CallMCPTool calls an MCP tool to process content

View file

@ -16,7 +16,7 @@ func (opts *Options) ToMap() map[string]interface{} {
result["mode"] = opts.Mode result["mode"] = opts.Mode
} }
if opts.Search != nil { if opts.Search != nil {
result["search"] = *opts.Search result["search"] = opts.Search
} }
if opts.Skip != nil { if opts.Skip != nil {
result["skip"] = opts.Skip result["skip"] = opts.Skip
@ -25,6 +25,9 @@ func (opts *Options) ToMap() map[string]interface{} {
if opts.DisableGlobalPrompts { if opts.DisableGlobalPrompts {
result["disable_global_prompts"] = opts.DisableGlobalPrompts result["disable_global_prompts"] = opts.DisableGlobalPrompts
} }
if opts.Metadata != nil {
result["metadata"] = opts.Metadata
}
// Note: Runtime fields (Context, Writer) are not serialized (json:"-") // Note: Runtime fields (Context, Writer) are not serialized (json:"-")
// They should not be included in the map // They should not be included in the map
@ -47,8 +50,9 @@ func OptionsFromMap(m map[string]interface{}) *Options {
if mode, ok := m["mode"].(string); ok { if mode, ok := m["mode"].(string); ok {
opts.Mode = mode opts.Mode = mode
} }
if search, ok := m["search"].(bool); ok { // Search supports: bool | SearchIntent | map[string]any | nil
opts.Search = &search if search := m["search"]; search != nil {
opts.Search = search
} }
if skipMap, ok := m["skip"].(map[string]interface{}); ok { if skipMap, ok := m["skip"].(map[string]interface{}); ok {
skip := &Skip{} skip := &Skip{}
@ -66,6 +70,9 @@ func OptionsFromMap(m map[string]interface{}) *Options {
if disableGlobalPrompts, ok := m["disable_global_prompts"].(bool); ok { if disableGlobalPrompts, ok := m["disable_global_prompts"].(bool); ok {
opts.DisableGlobalPrompts = disableGlobalPrompts opts.DisableGlobalPrompts = disableGlobalPrompts
} }
if metadata, ok := m["metadata"].(map[string]interface{}); ok {
opts.Metadata = metadata
}
// Note: Context and Writer are runtime fields, not restored from map // Note: Context and Writer are runtime fields, not restored from map
// They should be set by the caller if needed // They should be set by the caller if needed

View file

@ -266,6 +266,15 @@ type Context struct {
Metadata map[string]interface{} `json:"metadata,omitempty"` // The metadata of the request, it will be used to pass data to the page Metadata map[string]interface{} `json:"metadata,omitempty"` // The metadata of the request, it will be used to pass data to the page
} }
// SearchIntent represents the result of search intent detection
// Used by Create hook to specify fine-grained search behavior
type SearchIntent struct {
NeedSearch bool `json:"need_search"` // Whether search is needed
SearchTypes []string `json:"search_types,omitempty"` // Types of search to perform: "web", "kb", "db"
Confidence float64 `json:"confidence,omitempty"` // Confidence level (0-1)
Reason string `json:"reason,omitempty"` // Reason for the decision
}
// Options represents the options for the context // Options represents the options for the context
type Options struct { type Options struct {
@ -284,11 +293,17 @@ type Options struct {
// Disable global prompts, default is false // Disable global prompts, default is false
DisableGlobalPrompts bool `json:"disable_global_prompts,omitempty"` // Temporarily disable global prompts for this request DisableGlobalPrompts bool `json:"disable_global_prompts,omitempty"` // Temporarily disable global prompts for this request
// Search mode, default is true // Search controls search behavior, supports multiple types:
Search *bool `json:"search,omitempty"` // Search mode, default is true // - bool: true = enable all search types, false = disable all search
// - SearchIntent: fine-grained control with specific types, confidence, etc.
// - nil: use default behavior (determined by __yao.needsearch agent)
Search any `json:"search,omitempty"` // Search mode: bool | SearchIntent | nil
// Agent mode, use to select the mode of the request, default is "chat" // Agent mode, use to select the mode of the request, default is "chat"
Mode string `json:"mode,omitempty"` // Agent mode, use to select the mode of the request, default is "chat" Mode string `json:"mode,omitempty"` // Agent mode, use to select the mode of the request, default is "chat"
// Metadata for passing custom data to hooks (e.g., scenario selection)
Metadata map[string]any `json:"metadata,omitempty"` // Custom metadata passed to Create/Next hooks
} }
// Stack represents the call stack node for tracing agent-to-agent calls // Stack represents the call stack node for tracing agent-to-agent calls
@ -369,6 +384,12 @@ type HookCreateResponse struct {
// ForceUses controls whether to force using Uses tools even when model has native capabilities // ForceUses controls whether to force using Uses tools even when model has native capabilities
ForceUses *bool `json:"force_uses,omitempty"` // Force using Uses tools regardless of model capabilities ForceUses *bool `json:"force_uses,omitempty"` // Force using Uses tools regardless of model capabilities
// Search controls search behavior, supports multiple types:
// - bool: true = enable all search types, false = disable all search
// - SearchIntent: fine-grained control with specific types, confidence, etc.
// - nil: use default behavior (determined by __yao.needsearch agent)
Search any `json:"search,omitempty"` // Search mode: bool | SearchIntent | nil
} }
// NextHookPayload payload for the next hook // NextHookPayload payload for the next hook

View file

@ -112,6 +112,10 @@ func init() {
"search.intent.need_search": "Searching for references...", "search.intent.need_search": "Searching for references...",
"search.intent.no_search": "No references needed", "search.intent.no_search": "No references needed",
// Keyword Extraction: assistant/search.go - Keyword extraction messages
"search.keyword.loading": "Analyzing conversation...",
"search.keyword.done": "Analysis complete",
// Search: assistant/search.go - Trace labels // Search: assistant/search.go - Trace labels
"search.trace.label": "Search", "search.trace.label": "Search",
"search.trace.description": "Search the web and knowledge base for relevant information", "search.trace.description": "Search the web and knowledge base for relevant information",
@ -201,6 +205,10 @@ func init() {
"search.intent.need_search": "正在查询相关资料...", "search.intent.need_search": "正在查询相关资料...",
"search.intent.no_search": "无需查询资料", "search.intent.no_search": "无需查询资料",
// Keyword Extraction: assistant/search.go - Keyword extraction messages
"search.keyword.loading": "正在分析对话内容...",
"search.keyword.done": "分析完成",
// Search: assistant/search.go - Trace labels // Search: assistant/search.go - Trace labels
"search.trace.label": "搜索", "search.trace.label": "搜索",
"search.trace.description": "搜索网络和知识库获取相关信息", "search.trace.description": "搜索网络和知识库获取相关信息",
@ -318,6 +326,10 @@ func init() {
"search.intent.need_search": "正在查询相关资料...", "search.intent.need_search": "正在查询相关资料...",
"search.intent.no_search": "无需查询资料", "search.intent.no_search": "无需查询资料",
// Keyword Extraction: assistant/search.go - Keyword extraction messages
"search.keyword.loading": "正在分析对话内容...",
"search.keyword.done": "分析完成",
// Search: assistant/search.go - Trace labels // Search: assistant/search.go - Trace labels
"search.trace.label": "搜索", "search.trace.label": "搜索",
"search.trace.description": "搜索网络和知识库获取相关信息", "search.trace.description": "搜索网络和知识库获取相关信息",

View file

@ -637,11 +637,25 @@ func (p *Provider) streamWithRetry(ctx *context.Context, messages []context.Mess
errorDetected := false errorDetected := false
// Wrap streamHandler to detect JSON error responses // Wrap streamHandler to detect JSON error responses
// Note: API error responses are raw JSON without "data: " prefix
// Normal SSE data always starts with "data: " prefix
wrappedHandler := func(data []byte) int { wrappedHandler := func(data []byte) int {
dataStr := string(data) dataStr := string(data)
trimmed := strings.TrimSpace(dataStr)
// Detect if this looks like a JSON error response (starts with "{" or contains "error") // Skip empty lines
if strings.Contains(dataStr, `"error"`) || (strings.TrimSpace(dataStr) == "{" && !errorDetected) { if trimmed == "" {
return http.HandlerReturnOk
}
// Normal SSE data starts with "data: " - pass to streamHandler
if strings.HasPrefix(dataStr, "data: ") {
return streamHandler(data)
}
// Detect if this looks like a JSON error response (raw JSON without "data: " prefix)
// API errors are returned as raw JSON: {"error": {...}}
if strings.HasPrefix(trimmed, "{") && strings.Contains(dataStr, `"error"`) {
errorDetected = true errorDetected = true
} }
@ -652,7 +666,7 @@ func (p *Provider) streamWithRetry(ctx *context.Context, messages []context.Mess
return http.HandlerReturnOk return http.HandlerReturnOk
} }
// Otherwise, use normal handler // Unknown format, pass to streamHandler (it will skip non-SSE data)
return streamHandler(data) return streamHandler(data)
} }

View file

@ -213,6 +213,10 @@ func initAssistant() error {
Audio: agentDSL.Uses.Audio, Audio: agentDSL.Uses.Audio,
Search: agentDSL.Uses.Search, Search: agentDSL.Uses.Search,
Fetch: agentDSL.Uses.Fetch, Fetch: agentDSL.Uses.Fetch,
Web: agentDSL.Uses.Web,
Keyword: agentDSL.Uses.Keyword,
QueryDSL: agentDSL.Uses.QueryDSL,
Rerank: agentDSL.Uses.Rerank,
} }
assistant.SetGlobalUses(globalUses) assistant.SetGlobalUses(globalUses)
} }

View file

@ -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
// - "<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, 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(), Duration: time.Since(start).Milliseconds(),
Error: "context is required for DB search",
}, nil
} }
// Store maxResults for later use // 1. Load all models and build combined schema
_ = maxResults models := make(map[string]*model.Model)
schemas := make([]map[string]interface{}, 0, len(modelIDs))
return result, nil for _, modelID := range modelIDs {
mod, err := model.Get(modelID)
if err != 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.Get(primaryModelID) // May be nil, that's ok
}
// 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,
Query: req.Query,
Source: req.Source,
Items: items,
Total: len(items),
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
}
// Merge preset Wheres (prepend to ensure they take priority)
if len(req.Wheres) > 0 {
dsl.Wheres = append(req.Wheres, dsl.Wheres...)
}
// 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)
} }

View file

@ -0,0 +1,221 @@
package db_test
import (
"testing"
"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"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// ============================================================================
// 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, query engine, 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)
}
}
}
})
t.Run("search_nonexistent_model_graceful", func(t *testing.T) {
h := db.NewHandler("builtin", nil)
req := &types.Request{
Type: types.SearchTypeDB,
Query: "查询文章",
Source: types.SourceAuto,
Models: []string{"nonexistent_model", "article", "fake_model"},
Limit: 10,
}
// Should NOT panic, should return gracefully with error
result, err := h.SearchWithContext(ctx, req)
require.NoError(t, err)
require.NotNil(t, result)
// Should have error message about no valid models
assert.Equal(t, types.SearchTypeDB, result.Type)
assert.Equal(t, "no valid models found", result.Error)
assert.Empty(t, result.Items)
})
t.Run("search_mixed_models_partial_exist", func(t *testing.T) {
h := db.NewHandler("builtin", nil)
req := &types.Request{
Type: types.SearchTypeDB,
Query: "查询角色",
Source: types.SourceAuto,
Models: []string{"nonexistent_model", "__yao.role", "fake_model"}, // Only __yao.role exists
Limit: 10,
}
// Should NOT panic, should work with the existing model
result, err := h.SearchWithContext(ctx, req)
require.NoError(t, err)
require.NotNil(t, result)
// Should succeed with partial models
assert.Equal(t, types.SearchTypeDB, result.Type)
if result.Error == "" {
// If no error, should have results from __yao.role
assert.GreaterOrEqual(t, len(result.Items), 0)
}
})
}
// 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)
}
}

View file

@ -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{ req: &types.Request{
Type: types.SearchTypeDB, Type: types.SearchTypeDB,
Query: "find products under $100", 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"}, 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.expectItems, len(result.Items))
if tt.expectError != "" {
assert.Equal(t, tt.expectError, result.Error) assert.Equal(t, tt.expectError, result.Error)
} else { assert.Equal(t, 0, len(result.Items))
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)
// Create a mock model for testing
mod := &model.Model{
MetaData: model.MetaData{
Table: model.Table{
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"])
}
}
func TestHandler_BuildModelSchema_MultipleModels(t *testing.T) {
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{ req := &types.Request{
Type: types.SearchTypeDB, Type: types.SearchTypeDB,
Query: "test", Query: "test query",
Source: types.SourceAuto,
Models: []string{"product"}, Models: []string{"product"},
// No limit in request, should use config's MaxResults Scenario: scenario,
} }
// Without context, should return error (but scenario should be preserved in request)
result, err := h.Search(req) result, err := h.Search(req)
assert.NoError(t, err) assert.NoError(t, err)
assert.NotNil(t, result) assert.NotNil(t, result)
// Skeleton doesn't actually use maxResults yet, but the test ensures the handler runs // 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)
})
} }

View file

@ -150,46 +150,20 @@ func (p *AgentProvider) Search(ctx *agentContext.Context, req *types.Request) (*
}, nil }, nil
} }
// parseAgentResponse parses the agent response into search result items // parseAgentResponse parses the agent's *context.Response into search result items
// The agent should return a JSON structure with search results // Now that agent.Stream() returns *context.Response directly,
func (p *AgentProvider) parseAgentResponse(response interface{}, source types.SourceType) ([]*types.ResultItem, int, string) { // we can access fields without type assertions.
if response == nil { //
// The agent returns search results in response.Next field
func (p *AgentProvider) parseAgentResponse(response *agentContext.Response, source types.SourceType) ([]*types.ResultItem, int, string) {
if response == nil || response.Next == nil {
return nil, 0, "Agent returned nil response" return nil, 0, "Agent returned nil response"
} }
// Try to extract data from response // Extract data from Next field
var data map[string]interface{} data := extractNextData(response.Next)
if data == nil {
// Handle different response types return nil, 0, "Failed to extract data from agent response"
switch v := response.(type) {
case map[string]interface{}:
data = v
case string:
// Try to parse as JSON
if err := json.Unmarshal([]byte(v), &data); err != nil {
return nil, 0, fmt.Sprintf("Failed to parse agent response as JSON: %v", err)
}
default:
// Try to marshal and unmarshal
jsonBytes, err := json.Marshal(response)
if err != nil {
return nil, 0, fmt.Sprintf("Failed to serialize agent response: %v", err)
}
if err := json.Unmarshal(jsonBytes, &data); err != nil {
return nil, 0, fmt.Sprintf("Failed to parse agent response: %v", err)
}
}
// Check for "next" field (custom hook data)
if next, hasNext := data["next"]; hasNext && next != nil {
if nextMap, ok := next.(map[string]interface{}); ok {
data = nextMap
} else if nextStr, ok := next.(string); ok {
// Try to parse as JSON
if err := json.Unmarshal([]byte(nextStr), &data); err != nil {
return nil, 0, fmt.Sprintf("Failed to parse next hook data: %v", err)
}
}
} }
// Extract items from data // Extract items from data
@ -230,3 +204,34 @@ func (p *AgentProvider) parseAgentResponse(response interface{}, source types.So
return items, total, "" return items, total, ""
} }
// extractNextData extracts the actual data from response.Next field
// Handles nested structures like { "data": { ... } }
func extractNextData(next interface{}) map[string]interface{} {
if next == nil {
return nil
}
switch v := next.(type) {
case map[string]interface{}:
// Check for "data" wrapper
if data, ok := v["data"].(map[string]interface{}); ok {
return data
}
return v
case string:
// Try to parse as JSON
var data map[string]interface{}
if err := json.Unmarshal([]byte(v), &data); err == nil {
return extractNextData(data)
}
}
// Try to handle other types by converting to JSON and back
if bytes, err := json.Marshal(next); err == nil {
var data map[string]interface{}
if err := json.Unmarshal(bytes, &data); err == nil {
return extractNextData(data)
}
}
return nil
}

View file

@ -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)
}

View file

@ -0,0 +1,191 @@
package search_test
import (
"testing"
"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"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// ============================================================================
// 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 (loads models, database, query engine, etc.)
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)
}
}

View file

@ -23,8 +23,8 @@ func NewAgentProvider(agentID string) *AgentProvider {
} }
// Extract extracts keywords by calling the target agent // Extract extracts keywords by calling the target agent
// The agent receives the content and returns extracted keywords // The agent receives the content and returns extracted keywords with weights
func (p *AgentProvider) Extract(ctx *agentContext.Context, content string, opts *types.KeywordOptions) ([]string, error) { func (p *AgentProvider) Extract(ctx *agentContext.Context, content string, opts *types.KeywordOptions) ([]types.Keyword, error) {
if ctx == nil { if ctx == nil {
return nil, fmt.Errorf("context is required for agent keyword extraction") return nil, fmt.Errorf("context is required for agent keyword extraction")
} }
@ -64,73 +64,64 @@ func (p *AgentProvider) Extract(ctx *agentContext.Context, content string, opts
}, },
} }
result, err := agent.Stream(ctx, messages, options) response, err := agent.Stream(ctx, messages, options)
if err != nil { if err != nil {
return nil, fmt.Errorf("agent call failed: %w", err) return nil, fmt.Errorf("agent call failed: %w", err)
} }
// Debug: log the result type and value // Parse the result from response.Next
// fmt.Printf("DEBUG Agent result type: %T, value: %+v\n", result, result) return p.parseResponse(response)
// Parse the result
return p.parseResult(result)
} }
// parseResult extracts keywords from the agent's response // parseResponse extracts keywords from the agent's *context.Response
// The agent should return data in NextHookResponse format: { data: { keywords: [...] } } // Now that agent.Stream() returns *context.Response directly,
// The Stream() response wraps this in: { next: { data: { keywords: [...] } } } // we can access fields without type assertions.
func (p *AgentProvider) parseResult(result interface{}) ([]string, error) { //
if result == nil { // The agent returns keywords in response.Next field as {data: {keywords: [{k, w}, ...]}}
return []string{}, nil func (p *AgentProvider) parseResponse(response *agentContext.Response) ([]types.Keyword, error) {
if response == nil || response.Next == nil {
return []types.Keyword{}, nil
}
return p.parseNextData(response.Next)
}
// parseNextData extracts keywords from Next hook data
// Expected format: {data: {keywords: [{k: "keyword", w: 0.9}, ...]}}
func (p *AgentProvider) parseNextData(next interface{}) ([]types.Keyword, error) {
if next == nil {
return []types.Keyword{}, nil
} }
// Try to convert to map first (most common case) // Try to convert to map first (most common case)
var data map[string]interface{} var data map[string]interface{}
switch v := result.(type) { switch v := next.(type) {
case map[string]interface{}: case map[string]interface{}:
data = v data = v
case string: case string:
// Try to parse as JSON // Try to parse as JSON
if err := json.Unmarshal([]byte(v), &data); err != nil { if err := json.Unmarshal([]byte(v), &data); err != nil {
// Not a JSON object, try as array // Not a JSON object, try as array of keywords
var keywords []string var keywords []types.Keyword
if err := json.Unmarshal([]byte(v), &keywords); err == nil { if err := json.Unmarshal([]byte(v), &keywords); err == nil {
return keywords, nil return keywords, nil
} }
// Return as single keyword // Return as single keyword with default weight
return []string{v}, nil return []types.Keyword{{K: v, W: 0.5}}, nil
} }
case []string: case []types.Keyword:
return v, nil return v, nil
case []interface{}: case []interface{}:
keywords := make([]string, 0, len(v)) return p.extractKeywordsFromArray(v)
for _, item := range v {
if s, ok := item.(string); ok {
keywords = append(keywords, s)
}
}
return keywords, nil
default: default:
// Try to marshal and unmarshal // Try to marshal and unmarshal
jsonBytes, err := json.Marshal(result) jsonBytes, err := json.Marshal(next)
if err != nil { if err != nil {
return []string{}, nil return []types.Keyword{}, nil
} }
if err := json.Unmarshal(jsonBytes, &data); err != nil { if err := json.Unmarshal(jsonBytes, &data); err != nil {
return []string{}, nil return []types.Keyword{}, nil
}
}
// Check for "next" field (custom hook data from NextHookResponse)
// Stream() returns: { next: { data: { keywords: [...] } } }
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 {
return []string{}, nil
}
} }
} }
@ -148,28 +139,48 @@ func (p *AgentProvider) parseResult(result interface{}) ([]string, error) {
return p.extractKeywordsFromValue(d) return p.extractKeywordsFromValue(d)
} }
return []string{}, nil return []types.Keyword{}, nil
} }
// extractKeywordsFromValue extracts string array from various types // extractKeywordsFromValue extracts Keyword array from various types
func (p *AgentProvider) extractKeywordsFromValue(v interface{}) ([]string, error) { func (p *AgentProvider) extractKeywordsFromValue(v interface{}) ([]types.Keyword, error) {
switch kw := v.(type) { switch kw := v.(type) {
case []string: case []types.Keyword:
return kw, nil return kw, nil
case []interface{}: case []interface{}:
keywords := make([]string, 0, len(kw)) return p.extractKeywordsFromArray(kw)
for _, item := range kw {
if s, ok := item.(string); ok {
keywords = append(keywords, s)
}
}
return keywords, nil
case string: case string:
var keywords []string var keywords []types.Keyword
if err := json.Unmarshal([]byte(kw), &keywords); err == nil { if err := json.Unmarshal([]byte(kw), &keywords); err == nil {
return keywords, nil return keywords, nil
} }
return []string{kw}, nil return []types.Keyword{{K: kw, W: 0.5}}, nil
} }
return []string{}, nil return []types.Keyword{}, nil
}
// extractKeywordsFromArray extracts keywords from []interface{}
// Handles both {k, w} objects and plain strings
func (p *AgentProvider) extractKeywordsFromArray(items []interface{}) ([]types.Keyword, error) {
keywords := make([]types.Keyword, 0, len(items))
for _, item := range items {
switch v := item.(type) {
case map[string]interface{}:
// Handle {k: "keyword", w: 0.9} format
k, _ := v["k"].(string)
w, _ := v["w"].(float64)
if k != "" {
if w == 0 {
w = 0.5 // Default weight
}
keywords = append(keywords, types.Keyword{K: k, W: w})
}
case string:
// Plain string, use default weight
if v != "" {
keywords = append(keywords, types.Keyword{K: v, W: 0.5})
}
}
}
return keywords, nil
} }

View file

@ -1,243 +0,0 @@
package keyword
import (
"regexp"
"sort"
"strings"
"unicode"
)
// BuiltinExtractor implements simple frequency-based keyword extraction
// This is a lightweight implementation with no external dependencies.
//
// Algorithm:
// 1. Tokenize text (split by whitespace and punctuation)
// 2. Normalize (lowercase, trim)
// 3. Filter stop words and short words
// 4. Count word frequency
// 5. Return top N words by frequency
//
// Limitations:
// - No semantic understanding
// - No phrase extraction (single words only)
// - Basic Chinese support (splits by punctuation, no proper segmentation)
//
// For better results, use Agent or MCP mode with LLM-based extraction.
type BuiltinExtractor struct {
stopWords map[string]bool
minLength int // minimum word length to consider
}
// Result represents an extracted keyword with its score
type Result struct {
Word string `json:"word"`
Score float64 `json:"score"` // frequency-based score (0-1)
}
// NewBuiltinExtractor creates a new builtin keyword extractor
func NewBuiltinExtractor() *BuiltinExtractor {
return &BuiltinExtractor{
stopWords: defaultStopWords,
minLength: 2,
}
}
// Extract extracts keywords from text using frequency-based algorithm
func (e *BuiltinExtractor) Extract(text string, limit int) []Result {
if text == "" || limit <= 0 {
return []Result{}
}
// Step 1: Tokenize
tokens := e.tokenize(text)
// Step 2 & 3: Normalize and filter
var words []string
for _, token := range tokens {
word := e.normalize(token)
if e.shouldKeep(word) {
words = append(words, word)
}
}
if len(words) == 0 {
return []Result{}
}
// Step 4: Count frequency
freq := make(map[string]int)
for _, word := range words {
freq[word]++
}
// Step 5: Sort by frequency and return top N
type wordFreq struct {
word string
freq int
}
var sorted []wordFreq
for word, count := range freq {
sorted = append(sorted, wordFreq{word, count})
}
sort.Slice(sorted, func(i, j int) bool {
return sorted[i].freq > sorted[j].freq
})
// Calculate max frequency for normalization
maxFreq := 1
if len(sorted) > 0 {
maxFreq = sorted[0].freq
}
// Build result with normalized scores
result := make([]Result, 0, limit)
for i := 0; i < len(sorted) && i < limit; i++ {
result = append(result, Result{
Word: sorted[i].word,
Score: float64(sorted[i].freq) / float64(maxFreq),
})
}
return result
}
// ExtractAsStrings is a convenience method that returns just the keyword strings
func (e *BuiltinExtractor) ExtractAsStrings(text string, limit int) []string {
results := e.Extract(text, limit)
words := make([]string, len(results))
for i, r := range results {
words[i] = r.Word
}
return words
}
// tokenize splits text into tokens
// Handles both English (space-separated) and Chinese (character-based with punctuation splits)
func (e *BuiltinExtractor) tokenize(text string) []string {
// Split by whitespace and common punctuation
splitter := regexp.MustCompile(`[\s\p{P}\p{S}]+`)
tokens := splitter.Split(text, -1)
// Further split mixed Chinese/English text
var result []string
for _, token := range tokens {
if token == "" {
continue
}
// Split Chinese characters as individual tokens (basic approach)
// For proper Chinese segmentation, use Agent/MCP mode
subTokens := e.splitMixedText(token)
result = append(result, subTokens...)
}
return result
}
// splitMixedText handles mixed Chinese/English text
// Chinese characters are grouped together, English words stay as-is
func (e *BuiltinExtractor) splitMixedText(text string) []string {
var result []string
var current strings.Builder
var lastType int // 0=none, 1=chinese, 2=other
for _, r := range text {
currentType := 0
if unicode.Is(unicode.Han, r) {
currentType = 1
} else if unicode.IsLetter(r) || unicode.IsDigit(r) {
currentType = 2
}
if currentType == 0 {
// Non-word character, flush current
if current.Len() > 0 {
result = append(result, current.String())
current.Reset()
}
lastType = 0
continue
}
if lastType != 0 && lastType != currentType {
// Type changed, flush current
if current.Len() > 0 {
result = append(result, current.String())
current.Reset()
}
}
current.WriteRune(r)
lastType = currentType
}
// Flush remaining
if current.Len() > 0 {
result = append(result, current.String())
}
return result
}
// normalize converts word to lowercase and trims whitespace
func (e *BuiltinExtractor) normalize(word string) string {
return strings.ToLower(strings.TrimSpace(word))
}
// shouldKeep checks if a word should be kept (not a stop word, meets length requirement)
func (e *BuiltinExtractor) shouldKeep(word string) bool {
if len(word) < e.minLength {
return false
}
if e.stopWords[word] {
return false
}
// Keep if it contains at least one letter or Chinese character
for _, r := range word {
if unicode.IsLetter(r) {
return true
}
}
return false
}
// defaultStopWords contains common stop words for English and Chinese
// This is a minimal set to keep the implementation lightweight.
// For comprehensive stop word filtering, use Agent/MCP mode.
var defaultStopWords = map[string]bool{
// English stop words (most common ~100)
"a": true, "an": true, "the": true, "and": true, "or": true, "but": true,
"is": true, "are": true, "was": true, "were": true, "be": true, "been": true, "being": true,
"have": true, "has": true, "had": true, "do": true, "does": true, "did": true,
"will": true, "would": true, "could": true, "should": true, "may": true, "might": true,
"must": true, "shall": true, "can": true, "need": true, "dare": true,
"i": true, "you": true, "he": true, "she": true, "it": true, "we": true, "they": true,
"me": true, "him": true, "her": true, "us": true, "them": true,
"my": true, "your": true, "his": true, "its": true, "our": true, "their": true,
"mine": true, "yours": true, "hers": true, "ours": true, "theirs": true,
"this": true, "that": true, "these": true, "those": true,
"what": true, "which": true, "who": true, "whom": true, "whose": true,
"where": true, "when": true, "why": true, "how": true,
"all": true, "each": true, "every": true, "both": true, "few": true, "more": true,
"most": true, "other": true, "some": true, "such": true, "no": true, "not": true,
"only": true, "same": true, "so": true, "than": true, "too": true, "very": true,
"just": true, "also": true, "now": true, "here": true, "there": true,
"in": true, "on": true, "at": true, "by": true, "for": true, "with": true,
"about": true, "against": true, "between": true, "into": true, "through": true,
"during": true, "before": true, "after": true, "above": true, "below": true,
"to": true, "from": true, "up": true, "down": true, "out": true, "off": true,
"over": true, "under": true, "again": true, "further": true, "then": true, "once": true,
"as": true, "if": true, "because": true, "until": true, "while": true,
// Chinese stop words (most common ~50)
"的": true, "了": true, "和": true, "是": true, "就": true,
"都": true, "而": true, "及": true, "与": true, "着": true,
"或": true, "一个": true, "没有": true, "我们": true, "你们": true,
"他们": true, "它们": true, "这个": true, "那个": true, "这些": true,
"那些": true, "这里": true, "那里": true, "什么": true, "怎么": true,
"为什么": true, "哪里": true, "谁": true, "哪个": true, "多少": true,
"在": true, "有": true, "个": true, "中": true, "为": true,
"以": true, "于": true, "上": true, "下": true, "不": true,
"也": true, "很": true, "到": true, "说": true, "要": true,
"会": true, "可以": true, "这": true, "那": true, "但": true,
"如果": true, "因为": true, "所以": true, "虽然": true, "但是": true,
}

View file

@ -1,137 +0,0 @@
package keyword
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestBuiltinExtractor_Extract(t *testing.T) {
extractor := NewBuiltinExtractor()
tests := []struct {
name string
text string
limit int
minCount int // minimum expected keywords
}{
{
name: "English text",
text: "The quick brown fox jumps over the lazy dog. The fox is very quick.",
limit: 5,
minCount: 3, // fox, quick, etc.
},
{
name: "Chinese text",
text: "人工智能技术正在快速发展,机器学习和深度学习是人工智能的核心技术",
limit: 5,
minCount: 2,
},
{
name: "Mixed text",
text: "AI人工智能 machine learning 机器学习 deep learning 深度学习",
limit: 10,
minCount: 3,
},
{
name: "Empty text",
text: "",
limit: 5,
minCount: 0,
},
{
name: "Only stop words",
text: "the a an is are was were",
limit: 5,
minCount: 0,
},
{
name: "Technical query",
text: "How to implement a search engine with Elasticsearch and Redis caching?",
limit: 5,
minCount: 3, // search, engine, elasticsearch, redis, caching
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
results := extractor.Extract(tt.text, tt.limit)
assert.GreaterOrEqual(t, len(results), tt.minCount, "Expected at least %d keywords", tt.minCount)
assert.LessOrEqual(t, len(results), tt.limit, "Should not exceed limit")
// Check scores are valid
for _, r := range results {
assert.NotEmpty(t, r.Word)
assert.GreaterOrEqual(t, r.Score, 0.0)
assert.LessOrEqual(t, r.Score, 1.0)
}
})
}
}
func TestBuiltinExtractor_ExtractAsStrings(t *testing.T) {
extractor := NewBuiltinExtractor()
text := "Machine learning and deep learning are subfields of artificial intelligence"
keywords := extractor.ExtractAsStrings(text, 5)
assert.NotEmpty(t, keywords)
assert.LessOrEqual(t, len(keywords), 5)
// Check that common ML terms are extracted
keywordSet := make(map[string]bool)
for _, k := range keywords {
keywordSet[k] = true
}
assert.True(t, keywordSet["learning"] || keywordSet["machine"] || keywordSet["artificial"],
"Expected at least one relevant keyword")
}
func TestBuiltinExtractor_StopWords(t *testing.T) {
extractor := NewBuiltinExtractor()
// Test that stop words are filtered
text := "the quick brown fox is very lazy"
results := extractor.Extract(text, 10)
for _, r := range results {
assert.NotEqual(t, "the", r.Word)
assert.NotEqual(t, "is", r.Word)
assert.NotEqual(t, "very", r.Word)
}
}
func TestBuiltinExtractor_Frequency(t *testing.T) {
extractor := NewBuiltinExtractor()
// Word "search" appears 3 times, should rank higher
text := "search engine optimization, search ranking, search results"
results := extractor.Extract(text, 3)
assert.NotEmpty(t, results)
// "search" should be the top keyword
assert.Equal(t, "search", results[0].Word)
assert.Equal(t, 1.0, results[0].Score) // highest frequency = 1.0
}
func TestBuiltinExtractor_ZeroLimit(t *testing.T) {
extractor := NewBuiltinExtractor()
results := extractor.Extract("some text here", 0)
assert.Empty(t, results)
}
func TestBuiltinExtractor_ChineseStopWords(t *testing.T) {
extractor := NewBuiltinExtractor()
// Test that Chinese stop words are filtered
text := "这是一个关于人工智能的文章"
results := extractor.Extract(text, 10)
for _, r := range results {
assert.NotEqual(t, "这", r.Word)
assert.NotEqual(t, "是", r.Word)
assert.NotEqual(t, "一个", r.Word)
assert.NotEqual(t, "的", r.Word)
}
}

View file

@ -1,19 +1,21 @@
// Package keyword provides keyword extraction for web search optimization // Package keyword provides keyword extraction for web search optimization
// Supports three modes via uses.keyword configuration: // Supports three modes via uses.keyword configuration:
// - "builtin": Simple frequency-based extraction (no external dependencies) // - "builtin" or "": Uses __yao.keyword system agent (LLM-powered)
// - "<assistant-id>": Delegate to an LLM-powered assistant for high-quality extraction // - "<assistant-id>": Delegate to a custom LLM-powered assistant
// - "mcp:<server>.<tool>": Call external MCP tool // - "mcp:<server>.<tool>": Call external MCP tool
//
// For production use cases requiring high accuracy, use Agent or MCP mode.
package keyword package keyword
import ( import (
"fmt"
"strings" "strings"
"github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/agent/context"
"github.com/yaoapp/yao/agent/search/types" "github.com/yaoapp/yao/agent/search/types"
) )
// SystemKeywordAgent is the default system agent for keyword extraction
const SystemKeywordAgent = "__yao.keyword"
// Extractor extracts keywords from text // Extractor extracts keywords from text
// Mode is determined by uses.keyword configuration // Mode is determined by uses.keyword configuration
type Extractor struct { type Extractor struct {
@ -32,19 +34,20 @@ func NewExtractor(usesKeyword string, cfg *types.KeywordConfig) *Extractor {
} }
// Extract extracts keywords from content based on configured mode // Extract extracts keywords from content based on configured mode
// Returns a list of keywords optimized for search queries // Returns a list of keywords with weights optimized for search queries
func (e *Extractor) Extract(ctx *context.Context, content string, opts *types.KeywordOptions) ([]string, error) { func (e *Extractor) Extract(ctx *context.Context, content string, opts *types.KeywordOptions) ([]types.Keyword, error) {
// Merge options with config defaults // Merge options with config defaults
mergedOpts := e.mergeOptions(opts) mergedOpts := e.mergeOptions(opts)
switch { switch {
case e.usesKeyword == "builtin" || e.usesKeyword == "": case e.usesKeyword == "builtin" || e.usesKeyword == "":
return e.builtinExtract(content, mergedOpts) // Use system keyword agent
return e.agentExtract(ctx, content, SystemKeywordAgent, mergedOpts)
case strings.HasPrefix(e.usesKeyword, "mcp:"): case strings.HasPrefix(e.usesKeyword, "mcp:"):
return e.mcpExtract(ctx, content, mergedOpts) return e.mcpExtract(ctx, content, mergedOpts)
default: default:
// Assume it's an assistant ID for Agent mode // Assume it's an assistant ID for Agent mode
return e.agentExtract(ctx, content, mergedOpts) return e.agentExtract(ctx, content, e.usesKeyword, mergedOpts)
} }
} }
@ -78,29 +81,24 @@ func (e *Extractor) mergeOptions(opts *types.KeywordOptions) *types.KeywordOptio
return result return result
} }
// builtinExtract uses simple frequency-based extraction
// This is a lightweight implementation with no external dependencies.
// For better results, use Agent or MCP mode.
func (e *Extractor) builtinExtract(content string, opts *types.KeywordOptions) ([]string, error) {
extractor := NewBuiltinExtractor()
return extractor.ExtractAsStrings(content, opts.MaxKeywords), nil
}
// agentExtract delegates to an LLM-powered assistant // agentExtract delegates to an LLM-powered assistant
// The assistant can understand context and extract semantically relevant keywords // The assistant can understand context and extract semantically relevant keywords
func (e *Extractor) agentExtract(ctx *context.Context, content string, opts *types.KeywordOptions) ([]string, error) { func (e *Extractor) agentExtract(ctx *context.Context, content string, agentID string, opts *types.KeywordOptions) ([]types.Keyword, error) {
provider := NewAgentProvider(e.usesKeyword) if ctx == nil {
return nil, fmt.Errorf("context is required for keyword extraction")
}
provider := NewAgentProvider(agentID)
return provider.Extract(ctx, content, opts) return provider.Extract(ctx, content, opts)
} }
// mcpExtract calls an external MCP tool // mcpExtract calls an external MCP tool
// Format: "mcp:<server>.<tool>" // Format: "mcp:<server>.<tool>"
func (e *Extractor) mcpExtract(ctx *context.Context, content string, opts *types.KeywordOptions) ([]string, error) { func (e *Extractor) mcpExtract(ctx *context.Context, content string, opts *types.KeywordOptions) ([]types.Keyword, error) {
mcpRef := strings.TrimPrefix(e.usesKeyword, "mcp:") mcpRef := strings.TrimPrefix(e.usesKeyword, "mcp:")
provider, err := NewMCPProvider(mcpRef) provider, err := NewMCPProvider(mcpRef)
if err != nil { if err != nil {
// Fallback to builtin on invalid MCP format // Fallback to system agent on invalid MCP format
return e.builtinExtract(content, opts) return e.agentExtract(ctx, content, SystemKeywordAgent, e.mergeOptions(nil))
} }
return provider.Extract(ctx, content, opts) return provider.Extract(ctx, content, opts)
} }

View file

@ -8,56 +8,48 @@ import (
"github.com/yaoapp/yao/agent/search/types" "github.com/yaoapp/yao/agent/search/types"
) )
func TestExtractor_BuiltinMode(t *testing.T) { func TestExtractor_BuiltinMode_RequiresContext(t *testing.T) {
// Test builtin mode (no external dependencies) // Test builtin mode requires context (now uses __yao.keyword agent)
extractor := keyword.NewExtractor("builtin", &types.KeywordConfig{ extractor := keyword.NewExtractor("builtin", &types.KeywordConfig{
MaxKeywords: 5, MaxKeywords: 5,
Language: "auto", Language: "auto",
}) })
keywords, err := extractor.Extract(nil, "How to build a search engine with Elasticsearch?", nil) // Without context, should return error
assert.NoError(t, err) _, err := extractor.Extract(nil, "How to build a search engine with Elasticsearch?", nil)
assert.NotEmpty(t, keywords) assert.Error(t, err)
assert.LessOrEqual(t, len(keywords), 5) assert.Contains(t, err.Error(), "context is required")
} }
func TestExtractor_EmptyUsesKeyword(t *testing.T) { func TestExtractor_EmptyUsesKeyword_RequiresContext(t *testing.T) {
// Empty uses.keyword should default to builtin // Empty uses.keyword should default to __yao.keyword agent
extractor := keyword.NewExtractor("", nil) extractor := keyword.NewExtractor("", nil)
keywords, err := extractor.Extract(nil, "Machine learning algorithms", nil) // Without context, should return error
assert.NoError(t, err) _, err := extractor.Extract(nil, "Machine learning algorithms", nil)
assert.NotEmpty(t, keywords) assert.Error(t, err)
assert.Contains(t, err.Error(), "context is required")
} }
func TestExtractor_RuntimeOptionsOverride(t *testing.T) { func TestExtractor_AgentMode_RequiresContext(t *testing.T) {
// Config has max_keywords=10, but runtime opts override to 3 // Custom agent mode requires context
extractor := keyword.NewExtractor("builtin", &types.KeywordConfig{ extractor := keyword.NewExtractor("custom.keyword.agent", nil)
MaxKeywords: 10,
})
keywords, err := extractor.Extract(nil, "one two three four five six seven eight nine ten", &types.KeywordOptions{ _, err := extractor.Extract(nil, "Test query", nil)
MaxKeywords: 3, assert.Error(t, err)
}) assert.Contains(t, err.Error(), "context is required")
assert.NoError(t, err)
assert.LessOrEqual(t, len(keywords), 3)
} }
func TestExtractor_ConfigDefaults(t *testing.T) { func TestExtractor_MCPMode_InvalidFormat(t *testing.T) {
// No config, should use defaults // Invalid MCP format should fallback to system agent (which requires context)
extractor := keyword.NewExtractor("builtin", nil)
keywords, err := extractor.Extract(nil, "Test query for keyword extraction", nil)
assert.NoError(t, err)
assert.NotEmpty(t, keywords)
assert.LessOrEqual(t, len(keywords), 10) // default max_keywords is 10
}
func TestExtractor_InvalidMCPFormat(t *testing.T) {
// Invalid MCP format should fallback to builtin
extractor := keyword.NewExtractor("mcp:invalid", nil) extractor := keyword.NewExtractor("mcp:invalid", nil)
keywords, err := extractor.Extract(nil, "Test query", nil) _, err := extractor.Extract(nil, "Test query", nil)
assert.NoError(t, err) assert.Error(t, err)
assert.NotEmpty(t, keywords) assert.Contains(t, err.Error(), "context is required")
}
func TestExtractor_SystemKeywordAgentConstant(t *testing.T) {
// Verify the system keyword agent constant
assert.Equal(t, "__yao.keyword", keyword.SystemKeywordAgent)
} }

View file

@ -31,7 +31,7 @@ func NewMCPProvider(mcpRef string) (*MCPProvider, error) {
} }
// Extract extracts keywords by calling the MCP tool // Extract extracts keywords by calling the MCP tool
func (p *MCPProvider) Extract(ctx *agentContext.Context, content string, opts *types.KeywordOptions) ([]string, error) { func (p *MCPProvider) Extract(ctx *agentContext.Context, content string, opts *types.KeywordOptions) ([]types.Keyword, error) {
// Get MCP client // Get MCP client
client, err := mcp.Select(p.serverID) client, err := mcp.Select(p.serverID)
if err != nil { if err != nil {
@ -56,9 +56,9 @@ func (p *MCPProvider) Extract(ctx *agentContext.Context, content string, opts *t
} }
// parseResult extracts keywords from the MCP tool response // parseResult extracts keywords from the MCP tool response
func (p *MCPProvider) parseResult(result *gouMCPTypes.CallToolResponse) ([]string, error) { func (p *MCPProvider) parseResult(result *gouMCPTypes.CallToolResponse) ([]types.Keyword, error) {
if result == nil { if result == nil {
return []string{}, nil return []types.Keyword{}, nil
} }
// Check for errors in result // Check for errors in result
@ -72,7 +72,7 @@ func (p *MCPProvider) parseResult(result *gouMCPTypes.CallToolResponse) ([]strin
// Parse content - expect JSON data with "keywords" field // Parse content - expect JSON data with "keywords" field
if len(result.Content) == 0 { if len(result.Content) == 0 {
return []string{}, nil return []types.Keyword{}, nil
} }
// Try to extract keywords from content // Try to extract keywords from content
@ -88,36 +88,50 @@ func (p *MCPProvider) parseResult(result *gouMCPTypes.CallToolResponse) ([]strin
} }
} }
// Try to parse as direct array // Try to parse as direct array of keywords
var keywords []string var keywords []types.Keyword
if err := json.Unmarshal([]byte(content.Text), &keywords); err == nil { if err := json.Unmarshal([]byte(content.Text), &keywords); err == nil {
return keywords, nil return keywords, nil
} }
} }
} }
return []string{}, nil return []types.Keyword{}, nil
} }
// extractKeywordsFromValue extracts string array from various types // extractKeywordsFromValue extracts Keyword array from various types
func (p *MCPProvider) extractKeywordsFromValue(v interface{}) ([]string, error) { func (p *MCPProvider) extractKeywordsFromValue(v interface{}) ([]types.Keyword, error) {
switch kw := v.(type) { switch kw := v.(type) {
case []string: case []types.Keyword:
return kw, nil return kw, nil
case []interface{}: case []interface{}:
keywords := make([]string, 0, len(kw)) keywords := make([]types.Keyword, 0, len(kw))
for _, item := range kw { for _, item := range kw {
if s, ok := item.(string); ok { switch v := item.(type) {
keywords = append(keywords, s) case map[string]interface{}:
// Handle {k: "keyword", w: 0.9} format
k, _ := v["k"].(string)
w, _ := v["w"].(float64)
if k != "" {
if w == 0 {
w = 0.5 // Default weight
}
keywords = append(keywords, types.Keyword{K: k, W: w})
}
case string:
// Plain string, use default weight
if v != "" {
keywords = append(keywords, types.Keyword{K: v, W: 0.5})
}
} }
} }
return keywords, nil return keywords, nil
case string: case string:
var keywords []string var keywords []types.Keyword
if err := json.Unmarshal([]byte(kw), &keywords); err == nil { if err := json.Unmarshal([]byte(kw), &keywords); err == nil {
return keywords, nil return keywords, nil
} }
return []string{kw}, nil return []types.Keyword{{K: kw, W: 0.5}}, nil
} }
return []string{}, nil return []types.Keyword{}, nil
} }

View file

@ -72,13 +72,13 @@ func TestMCPProviderWithCustomOptions(t *testing.T) {
} }
func TestMCPProviderInvalidFormat(t *testing.T) { func TestMCPProviderInvalidFormat(t *testing.T) {
// Test invalid MCP format fallback to builtin // Test invalid MCP format fallback to system agent (requires context)
extractor := keyword.NewExtractor("mcp:invalid", nil) extractor := keyword.NewExtractor("mcp:invalid", nil)
// Should fallback to builtin (no error) // Should fallback to system agent which requires context
keywords, err := extractor.Extract(nil, "test content for keyword extraction", nil) _, err := extractor.Extract(nil, "test content for keyword extraction", nil)
assert.NoError(t, err) assert.Error(t, err)
assert.NotEmpty(t, keywords, "Should fallback to builtin and extract keywords") assert.Contains(t, err.Error(), "context is required")
} }
func TestMCPProviderServerNotFound(t *testing.T) { func TestMCPProviderServerNotFound(t *testing.T) {

View file

@ -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,
}, },
} }
@ -71,8 +70,8 @@ func (p *AgentProvider) Generate(ctx *agentContext.Context, input *Input) (*Resu
continue continue
} }
// Parse the result // Parse the result from response
genResult, err := p.parseResult(result) genResult, err := p.parseResponse(result)
if err != nil { if err != nil {
lastError = err lastError = err
continue continue
@ -103,27 +102,32 @@ 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{} { // Returns JSON format for structured communication with the agent
func (p *AgentProvider) buildRequestMessage(input *Input, attempt int, lastLintErrors string) string {
// Build request data as JSON
requestData := map[string]interface{}{ requestData := map[string]interface{}{
"query": input.Query, "query": input.Query,
"models": input.ModelIDs, "models": input.ModelIDs,
"limit": input.Limit, "limit": input.Limit,
} }
// Add optional fields // Add schema from extra params if provided
if len(input.Wheres) > 0 { if input.ExtraParams != nil {
requestData["wheres"] = input.Wheres if schema, ok := input.ExtraParams["schema"]; ok {
requestData["schema"] = schema
} }
if len(input.Orders) > 0 {
requestData["orders"] = input.Orders
} }
// Add scenario hint if specified (filter, aggregation, join, complex)
if input.Scenario != "" {
requestData["scenario"] = string(input.Scenario)
}
// Add allowed fields if specified
if len(input.AllowedFields) > 0 { if len(input.AllowedFields) > 0 {
requestData["allowed_fields"] = input.AllowedFields 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 != "" {
@ -134,7 +138,8 @@ func (p *AgentProvider) buildRequestData(input *Input, attempt int, lastLintErro
} }
} }
return requestData jsonBytes, _ := json.Marshal(requestData)
return string(jsonBytes)
} }
// validateDSL validates the generated QueryDSL using the linter // validateDSL validates the generated QueryDSL using the linter
@ -150,18 +155,36 @@ func (p *AgentProvider) validateDSL(dsl *gou.QueryDSL) *linter.LintResult {
return lintResult return lintResult
} }
// parseResult extracts QueryDSL from the agent's response // parseResponse extracts QueryDSL from the agent's *context.Response
// The agent should return data in NextHookResponse format: { data: { dsl: {...}, explain: "..." } } // Now that agent.Stream() returns *context.Response directly,
// The Stream() response wraps this in: { next: { data: { dsl: {...} } } } // we can access fields without type assertions.
func (p *AgentProvider) parseResult(result interface{}) (*Result, error) { //
if result == nil { // The querydsl agent returns QueryDSL in response.Next field
// Or returns error JSON: {"error": "code", "message": "..."}
func (p *AgentProvider) parseResponse(response *agentContext.Response) (*Result, error) {
if response == nil {
return &Result{}, nil return &Result{}, nil
} }
// Try to convert to map first (most common case) // Check Next field first (custom hook data)
if response.Next != nil {
return p.parseNextData(response.Next)
}
// No Next data, return empty result
return &Result{}, nil
}
// parseNextData extracts QueryDSL from Next hook data
func (p *AgentProvider) parseNextData(next interface{}) (*Result, error) {
if next == nil {
return &Result{}, nil
}
// Try to convert to map first
var data map[string]interface{} var data map[string]interface{}
switch v := result.(type) { switch v := next.(type) {
case map[string]interface{}: case map[string]interface{}:
data = v data = v
case string: case string:
@ -171,7 +194,7 @@ func (p *AgentProvider) parseResult(result interface{}) (*Result, error) {
} }
default: default:
// Try to marshal and unmarshal // Try to marshal and unmarshal
jsonBytes, err := json.Marshal(result) jsonBytes, err := json.Marshal(next)
if err != nil { if err != nil {
return &Result{}, nil return &Result{}, nil
} }
@ -180,49 +203,72 @@ func (p *AgentProvider) parseResult(result interface{}) (*Result, error) {
} }
} }
// Check for "next" field (custom hook data from NextHookResponse)
// Stream() returns: { next: { data: { dsl: {...} } } }
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 {
return &Result{}, nil
}
}
}
// Extract QueryDSL from data
// Try common field names: "dsl", "data", "data.dsl"
genResult := &Result{} genResult := &Result{}
// Get explain if present // Check for error response: {"error": "code", "message": "..."}
if errCode, hasError := data["error"]; hasError {
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
}
// Check for "dsl" field wrapper: { dsl: {...} }
if dsl, ok := data["dsl"]; ok {
genResult.DSL = p.extractDSL(dsl)
if explain, ok := data["explain"].(string); ok { if explain, ok := data["explain"].(string); ok {
genResult.Explain = explain genResult.Explain = explain
} }
// Get warnings if present
if warnings, ok := data["warnings"]; ok { if warnings, ok := data["warnings"]; ok {
genResult.Warnings = p.extractWarnings(warnings) genResult.Warnings = p.extractWarnings(warnings)
} }
return genResult, nil
}
// Get DSL // Check for "data" field wrapper: { data: { dsl: {...}, explain: "...", warnings: [] } }
if dsl, ok := data["dsl"]; ok { if d, ok := data["data"]; ok {
genResult.DSL = p.extractDSL(dsl)
} else if d, ok := data["data"]; ok {
if dm, ok := d.(map[string]interface{}); ok { if dm, ok := d.(map[string]interface{}); ok {
// Check if data.data contains dsl field: { data: { dsl: {...} } }
if dsl, ok := dm["dsl"]; ok { if dsl, ok := dm["dsl"]; ok {
genResult.DSL = p.extractDSL(dsl) 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)
} }
// Extract explain and warnings from data.data
if explain, ok := dm["explain"].(string); ok { if explain, ok := dm["explain"].(string); ok {
genResult.Explain = explain genResult.Explain = explain
} }
if warnings, ok := dm["warnings"]; ok { if warnings, ok := dm["warnings"]; ok {
genResult.Warnings = p.extractWarnings(warnings) genResult.Warnings = p.extractWarnings(warnings)
} }
return genResult, nil
} }
} }
// Fallback: Get explain and warnings from top level
if explain, ok := data["explain"].(string); ok {
genResult.Explain = explain
}
if warnings, ok := data["warnings"]; ok {
genResult.Warnings = p.extractWarnings(warnings)
}
return genResult, nil return genResult, nil
} }

View file

@ -1,124 +0,0 @@
package querydsl
import (
"github.com/yaoapp/gou/model"
"github.com/yaoapp/gou/query/gou"
)
// BuiltinGenerator implements template-based QueryDSL generation
// This is a placeholder implementation that returns a basic QueryDSL.
//
// TODO: Implement actual template-based generation:
// - Parse natural language query
// - Match against model schema
// - Generate appropriate where clauses
// - Handle common query patterns (search, filter, sort)
//
// For production use cases requiring high accuracy, use Agent or MCP mode.
type BuiltinGenerator struct{}
// NewBuiltinGenerator creates a new builtin QueryDSL generator
func NewBuiltinGenerator() *BuiltinGenerator {
return &BuiltinGenerator{}
}
// Generate generates QueryDSL from natural language
// Currently returns a placeholder QueryDSL that searches all searchable fields
func (g *BuiltinGenerator) Generate(input *Input) (*Result, error) {
if input == nil || input.Query == "" {
return &Result{
Warnings: []string{"empty query, returning empty DSL"},
}, nil
}
// Build a basic QueryDSL
dsl := &gou.QueryDSL{}
// Set limit
limit := input.Limit
if limit <= 0 {
limit = 20
}
dsl.Limit = limit
// Apply pre-defined wheres if provided
if len(input.Wheres) > 0 {
dsl.Wheres = input.Wheres
}
// Apply orders if provided
if len(input.Orders) > 0 {
dsl.Orders = input.Orders
}
// Load models and try to generate basic search conditions
// Use the first model as the primary table, others can be joined
if len(input.ModelIDs) > 0 {
primaryModelID := input.ModelIDs[0]
// Check if model exists before selecting
if !model.Exists(primaryModelID) {
return &Result{
DSL: dsl,
Explain: "Generated basic QueryDSL (model not found)",
Warnings: []string{
"model '" + primaryModelID + "' not found, returning basic DSL without search conditions",
},
}, nil
}
primaryModel := model.Select(primaryModelID)
if primaryModel != nil && len(primaryModel.MetaData.Columns) > 0 {
// Find searchable text columns (string/text types with index)
var searchableColumns []string
for _, col := range primaryModel.MetaData.Columns {
// Use Index as a proxy for searchable, and check for text types
if col.Index && (col.Type == "string" || col.Type == "text" || col.Type == "longText") {
searchableColumns = append(searchableColumns, col.Name)
}
}
// If we have searchable columns and no pre-defined wheres, add a basic search
if len(searchableColumns) > 0 && len(input.Wheres) == 0 {
// Build OR conditions for searchable columns
orWheres := make([]gou.Where, 0, len(searchableColumns))
for _, col := range searchableColumns {
orWheres = append(orWheres, gou.Where{
Condition: gou.Condition{
Field: &gou.Expression{Field: col},
OP: "match",
Value: input.Query,
},
})
}
// Wrap in OR group if multiple columns
if len(orWheres) > 1 {
// Mark all but the first as OR conditions
for i := 1; i < len(orWheres); i++ {
orWheres[i].OR = true
}
dsl.Wheres = []gou.Where{
{
Wheres: orWheres,
},
}
} else if len(orWheres) == 1 {
dsl.Wheres = orWheres
}
}
}
// TODO: For multi-model queries, generate joins based on model relations
// This requires analyzing the relations between models and generating
// appropriate JOIN clauses in the QueryDSL
}
return &Result{
DSL: dsl,
Explain: "Generated basic search QueryDSL using builtin template (placeholder implementation)",
Warnings: []string{
"builtin generator is a placeholder, consider using Agent or MCP mode for production",
},
}, nil
}

View file

@ -1,13 +1,12 @@
// Package querydsl provides QueryDSL generation from natural language for DB search // Package querydsl provides QueryDSL generation from natural language for DB search
// Supports three modes via uses.querydsl configuration: // Supports three modes via uses.querydsl configuration:
// - "builtin": Template-based generation (no external dependencies) // - "builtin" or "": Uses __yao.querydsl system agent (LLM-powered)
// - "<assistant-id>": Delegate to an LLM-powered assistant for high-quality generation // - "<assistant-id>": Delegate to a custom LLM-powered assistant
// - "mcp:<server>.<tool>": Call external MCP tool // - "mcp:<server>.<tool>": Call external MCP tool
//
// For production use cases requiring high accuracy, use Agent or MCP mode.
package querydsl package querydsl
import ( import (
"fmt"
"strings" "strings"
"github.com/yaoapp/gou/query/gou" "github.com/yaoapp/gou/query/gou"
@ -15,6 +14,9 @@ import (
"github.com/yaoapp/yao/agent/search/types" "github.com/yaoapp/yao/agent/search/types"
) )
// SystemQueryDSLAgent is the default system agent for QueryDSL generation
const SystemQueryDSLAgent = "__yao.querydsl"
// Generator generates QueryDSL from natural language // Generator generates QueryDSL from natural language
// Mode is determined by uses.querydsl configuration // Mode is determined by uses.querydsl configuration
type Generator struct { type Generator struct {
@ -40,12 +42,13 @@ func (g *Generator) Generate(ctx *context.Context, input *Input) (*Result, error
switch { switch {
case g.usesQueryDSL == "builtin" || g.usesQueryDSL == "": case g.usesQueryDSL == "builtin" || g.usesQueryDSL == "":
result, err = g.builtinGenerate(input) // Use system querydsl agent
result, err = g.agentGenerate(ctx, input, SystemQueryDSLAgent)
case strings.HasPrefix(g.usesQueryDSL, "mcp:"): case strings.HasPrefix(g.usesQueryDSL, "mcp:"):
result, err = g.mcpGenerate(ctx, input) result, err = g.mcpGenerate(ctx, input)
default: default:
// Assume it's an assistant ID for Agent mode // Assume it's an assistant ID for Agent mode
result, err = g.agentGenerate(ctx, input) result, err = g.agentGenerate(ctx, input, g.usesQueryDSL)
} }
if err != nil { if err != nil {
@ -60,18 +63,13 @@ func (g *Generator) Generate(ctx *context.Context, input *Input) (*Result, error
return result, nil return result, nil
} }
// builtinGenerate uses template-based generation
// This is a lightweight implementation with no external dependencies.
// For better results, use Agent or MCP mode.
func (g *Generator) builtinGenerate(input *Input) (*Result, error) {
generator := NewBuiltinGenerator()
return generator.Generate(input)
}
// agentGenerate delegates to an LLM-powered assistant // agentGenerate delegates to an LLM-powered assistant
// The assistant can understand context and generate semantically correct QueryDSL // The assistant can understand context and generate semantically correct QueryDSL
func (g *Generator) agentGenerate(ctx *context.Context, input *Input) (*Result, error) { func (g *Generator) agentGenerate(ctx *context.Context, input *Input, agentID string) (*Result, error) {
provider := NewAgentProvider(g.usesQueryDSL) if ctx == nil {
return nil, fmt.Errorf("context is required for QueryDSL generation")
}
provider := NewAgentProvider(agentID)
return provider.Generate(ctx, input) return provider.Generate(ctx, input)
} }
@ -81,8 +79,8 @@ func (g *Generator) mcpGenerate(ctx *context.Context, input *Input) (*Result, er
mcpRef := strings.TrimPrefix(g.usesQueryDSL, "mcp:") mcpRef := strings.TrimPrefix(g.usesQueryDSL, "mcp:")
provider, err := NewMCPProvider(mcpRef) provider, err := NewMCPProvider(mcpRef)
if err != nil { if err != nil {
// Fallback to builtin on invalid MCP format // Fallback to system agent on invalid MCP format
return g.builtinGenerate(input) return g.agentGenerate(ctx, input, SystemQueryDSLAgent)
} }
return provider.Generate(ctx, input) return provider.Generate(ctx, input)
} }

View file

@ -46,27 +46,24 @@ func TestNewGenerator(t *testing.T) {
} }
} }
func TestGenerator_Generate_Builtin(t *testing.T) { func TestGenerator_Generate_Builtin_RequiresContext(t *testing.T) {
// Builtin mode now uses __yao.querydsl agent which requires context
gen := NewGenerator("builtin", nil) gen := NewGenerator("builtin", nil)
// Note: In real usage, models are loaded internally via model.Select()
// For this test, we just verify the basic flow works without models
input := &Input{ input := &Input{
Query: "find all active users", Query: "find all active users",
ModelIDs: []string{"user"}, ModelIDs: []string{"user"},
Limit: 10, Limit: 10,
} }
result, err := gen.Generate(nil, input) // Without context, should return error
assert.NoError(t, err) _, err := gen.Generate(nil, input)
assert.NotNil(t, result) assert.Error(t, err)
assert.NotNil(t, result.DSL) assert.Contains(t, err.Error(), "context is required")
assert.NotEmpty(t, result.Explain)
assert.NotEmpty(t, result.Warnings)
} }
func TestGenerator_Generate_EmptyMode(t *testing.T) { func TestGenerator_Generate_EmptyMode_RequiresContext(t *testing.T) {
// Empty mode should default to builtin // Empty mode defaults to __yao.querydsl agent which requires context
gen := NewGenerator("", nil) gen := NewGenerator("", nil)
input := &Input{ input := &Input{
@ -75,116 +72,45 @@ func TestGenerator_Generate_EmptyMode(t *testing.T) {
Limit: 5, Limit: 5,
} }
result, err := gen.Generate(nil, input) // Without context, should return error
assert.NoError(t, err) _, err := gen.Generate(nil, input)
assert.NotNil(t, result) assert.Error(t, err)
assert.Contains(t, err.Error(), "context is required")
} }
func TestBuiltinGenerator_Generate(t *testing.T) { func TestGenerator_Generate_AgentMode_RequiresContext(t *testing.T) {
gen := NewBuiltinGenerator() // Custom agent mode requires context
gen := NewGenerator("custom.querydsl.agent", nil)
t.Run("empty query", func(t *testing.T) { input := &Input{
result, err := gen.Generate(&Input{})
assert.NoError(t, err)
assert.NotNil(t, result)
assert.Nil(t, result.DSL)
assert.Contains(t, result.Warnings, "empty query, returning empty DSL")
})
t.Run("nil input", func(t *testing.T) {
result, err := gen.Generate(nil)
assert.NoError(t, err)
assert.NotNil(t, result)
assert.Nil(t, result.DSL)
})
t.Run("basic query without models loaded", func(t *testing.T) {
// Models are loaded internally via model.Select()
// When model is not found, it still generates basic DSL
result, err := gen.Generate(&Input{
Query: "find users", Query: "find users",
ModelIDs: []string{"user"}, ModelIDs: []string{"user"},
Limit: 10, Limit: 10,
})
assert.NoError(t, err)
assert.NotNil(t, result)
assert.NotNil(t, result.DSL)
assert.Equal(t, 10, result.DSL.Limit)
})
t.Run("query with pre-defined wheres", func(t *testing.T) {
preWheres := []gou.Where{
{
Condition: gou.Condition{
Field: &gou.Expression{Field: "status"},
OP: "=",
Value: "active",
},
},
} }
result, err := gen.Generate(&Input{
Query: "find users",
ModelIDs: []string{"user"},
Wheres: preWheres,
Limit: 10,
})
assert.NoError(t, err)
assert.NotNil(t, result)
assert.NotNil(t, result.DSL)
// Should use pre-defined wheres
assert.Equal(t, preWheres, result.DSL.Wheres)
})
t.Run("query with orders", func(t *testing.T) { _, err := gen.Generate(nil, input)
orders := gou.Orders{ assert.Error(t, err)
{Field: &gou.Expression{Field: "created_at"}, Sort: "desc"}, assert.Contains(t, err.Error(), "context is required")
} }
result, err := gen.Generate(&Input{
func TestGenerator_Generate_MCPMode_InvalidFormat(t *testing.T) {
// Invalid MCP format should fallback to system agent (which requires context)
gen := NewGenerator("mcp:invalid", nil)
input := &Input{
Query: "find users", Query: "find users",
ModelIDs: []string{"user"}, ModelIDs: []string{"user"},
Orders: orders,
Limit: 10, Limit: 10,
}) }
assert.NoError(t, err)
assert.NotNil(t, result)
assert.NotNil(t, result.DSL)
assert.Equal(t, orders, result.DSL.Orders)
})
t.Run("query with allowed fields", func(t *testing.T) { _, err := gen.Generate(nil, input)
result, err := gen.Generate(&Input{ assert.Error(t, err)
Query: "find users", assert.Contains(t, err.Error(), "context is required")
ModelIDs: []string{"user"}, }
AllowedFields: []string{"id", "name", "email"},
Limit: 10,
})
assert.NoError(t, err)
assert.NotNil(t, result)
assert.NotNil(t, result.DSL)
})
t.Run("default limit", func(t *testing.T) { func TestSystemQueryDSLAgentConstant(t *testing.T) {
result, err := gen.Generate(&Input{ // Verify the system querydsl agent constant
Query: "find users", assert.Equal(t, "__yao.querydsl", SystemQueryDSLAgent)
ModelIDs: []string{"user"},
})
assert.NoError(t, err)
assert.NotNil(t, result)
assert.NotNil(t, result.DSL)
assert.Equal(t, 20, result.DSL.Limit)
})
t.Run("multi-model query", func(t *testing.T) {
// Models are loaded internally via model.Select()
result, err := gen.Generate(&Input{
Query: "find user orders",
ModelIDs: []string{"user", "order"},
Limit: 10,
})
assert.NoError(t, err)
assert.NotNil(t, result)
assert.NotNil(t, result.DSL)
})
} }
func TestResult(t *testing.T) { func TestResult(t *testing.T) {
@ -201,3 +127,80 @@ func TestResult(t *testing.T) {
assert.NotEmpty(t, result.Explain) assert.NotEmpty(t, result.Explain)
assert.Len(t, result.Warnings, 1) assert.Len(t, result.Warnings, 1)
} }
func TestGenerator_ValidateFields(t *testing.T) {
gen := NewGenerator("", nil)
t.Run("validate select fields", func(t *testing.T) {
result := &Result{
DSL: &gou.QueryDSL{
Select: []gou.Expression{
{Field: "id"},
{Field: "name"},
{Field: "secret_field"},
},
},
}
allowedFields := []string{"id", "name", "email"}
validated := gen.validateFields(result, allowedFields)
assert.NotNil(t, validated)
assert.Len(t, validated.DSL.Select, 2)
assert.Contains(t, validated.Warnings[0], "secret_field")
})
t.Run("validate where fields", func(t *testing.T) {
result := &Result{
DSL: &gou.QueryDSL{
Wheres: []gou.Where{
{
Condition: gou.Condition{
Field: &gou.Expression{Field: "status"},
OP: "=",
Value: "active",
},
},
{
Condition: gou.Condition{
Field: &gou.Expression{Field: "secret"},
OP: "=",
Value: "hidden",
},
},
},
},
}
allowedFields := []string{"status", "name"}
validated := gen.validateFields(result, allowedFields)
assert.NotNil(t, validated)
assert.Len(t, validated.DSL.Wheres, 1)
assert.Contains(t, validated.Warnings[0], "secret")
})
t.Run("validate order fields", func(t *testing.T) {
result := &Result{
DSL: &gou.QueryDSL{
Orders: gou.Orders{
{Field: &gou.Expression{Field: "created_at"}, Sort: "desc"},
{Field: &gou.Expression{Field: "secret_sort"}, Sort: "asc"},
},
},
}
allowedFields := []string{"created_at", "updated_at"}
validated := gen.validateFields(result, allowedFields)
assert.NotNil(t, validated)
assert.Len(t, validated.DSL.Orders, 1)
assert.Contains(t, validated.Warnings[0], "secret_sort")
})
t.Run("nil DSL", func(t *testing.T) {
result := &Result{DSL: nil}
allowedFields := []string{"id", "name"}
validated := gen.validateFields(result, allowedFields)
assert.NotNil(t, validated)
assert.Nil(t, validated.DSL)
})
}

View file

@ -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)

View file

@ -59,20 +59,24 @@ func (p *AgentProvider) Rerank(ctx *context.Context, query string, items []*type
}, },
} }
result, err := agent.Stream(ctx, messages, options) response, err := agent.Stream(ctx, messages, options)
if err != nil { if err != nil {
return nil, fmt.Errorf("agent stream failed: %w", err) return nil, fmt.Errorf("agent stream failed: %w", err)
} }
// Parse response // Parse response from response.Next
return p.parseResponse(result, items, opts) return p.parseAgentResponse(response, items, opts)
} }
// parseResponse extracts reranked items from agent response // parseAgentResponse extracts reranked items from agent's *context.Response
// The response format from agent.Stream is typically: // Now that agent.Stream() returns *context.Response directly,
// { "next": { "data": { "order": [...] } } } // we can access fields without type assertions.
func (p *AgentProvider) parseResponse(result interface{}, originalItems []*types.ResultItem, opts *types.RerankOptions) ([]*types.ResultItem, error) { //
if result == nil { // Expected response.Next format:
// { "order": ["ref_001", "ref_003", "ref_002"] }
// Or: { "items": [{ "citation_id": "ref_001", ... }, ...] }
func (p *AgentProvider) parseAgentResponse(response *context.Response, originalItems []*types.ResultItem, opts *types.RerankOptions) ([]*types.ResultItem, error) {
if response == nil || response.Next == nil {
return originalItems, nil return originalItems, nil
} }
@ -84,20 +88,20 @@ func (p *AgentProvider) parseResponse(result interface{}, originalItems []*types
} }
} }
// Extract response data // Extract response data from Next field
response := extractResponseData(result) data := extractNextData(response.Next)
if response == nil { if data == nil {
return originalItems, nil return originalItems, nil
} }
// Try to get reranked order from response // Try to get reranked order from data
// Expected format: { "order": ["ref_001", "ref_003", "ref_002"] } // Expected format: { "order": ["ref_001", "ref_003", "ref_002"] }
// Or: { "items": [{ "citation_id": "ref_001", ... }, ...] } // Or: { "items": [{ "citation_id": "ref_001", ... }, ...] }
var reranked []*types.ResultItem var reranked []*types.ResultItem
// Try "order" field (list of citation IDs) // Try "order" field (list of citation IDs)
if order, ok := response["order"]; ok { if order, ok := data["order"]; ok {
if orderList := toStringSlice(order); len(orderList) > 0 { if orderList := toStringSlice(order); len(orderList) > 0 {
for _, id := range orderList { for _, id := range orderList {
if item, exists := itemMap[id]; exists { if item, exists := itemMap[id]; exists {
@ -116,7 +120,7 @@ func (p *AgentProvider) parseResponse(result interface{}, originalItems []*types
// Try "items" field (full items or items with citation_id) // Try "items" field (full items or items with citation_id)
if len(reranked) == 0 { if len(reranked) == 0 {
if items, ok := response["items"]; ok { if items, ok := data["items"]; ok {
if itemsList := toItemsList(items); len(itemsList) > 0 { if itemsList := toItemsList(items); len(itemsList) > 0 {
for _, respItem := range itemsList { for _, respItem := range itemsList {
// Check if it's just a reference or full item // Check if it's just a reference or full item
@ -150,20 +154,16 @@ func (p *AgentProvider) parseResponse(result interface{}, originalItems []*types
return reranked, nil return reranked, nil
} }
// extractResponseData extracts the actual response data from agent.Stream result // extractNextData extracts the actual data from response.Next field
// Handles nested structures like { "next": { "data": { ... } } } // Handles nested structures like { "data": { ... } }
func extractResponseData(result interface{}) map[string]interface{} { func extractNextData(next interface{}) map[string]interface{} {
switch v := result.(type) { if next == nil {
return nil
}
switch v := next.(type) {
case map[string]interface{}: case map[string]interface{}:
// Check for "next" wrapper (from NextHookResponse) // Check for "data" wrapper
if next, ok := v["next"].(map[string]interface{}); ok {
// Check for "data" inside next
if data, ok := next["data"].(map[string]interface{}); ok {
return data
}
return next
}
// Check for direct "data" wrapper
if data, ok := v["data"].(map[string]interface{}); ok { if data, ok := v["data"].(map[string]interface{}); ok {
return data return data
} }
@ -172,16 +172,14 @@ func extractResponseData(result interface{}) map[string]interface{} {
// Try to parse as JSON // Try to parse as JSON
var data map[string]interface{} var data map[string]interface{}
if err := json.Unmarshal([]byte(v), &data); err == nil { if err := json.Unmarshal([]byte(v), &data); err == nil {
return extractResponseData(data) return extractNextData(data)
} }
} }
// Try to handle other types by converting to JSON and back // Try to handle other types by converting to JSON and back
if result != nil { if bytes, err := json.Marshal(next); err == nil {
if bytes, err := json.Marshal(result); err == nil {
var data map[string]interface{} var data map[string]interface{}
if err := json.Unmarshal(bytes, &data); err == nil { if err := json.Unmarshal(bytes, &data); err == nil {
return extractResponseData(data) return extractNextData(data)
}
} }
} }
return nil return nil

View file

@ -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
} }

View file

@ -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
@ -43,6 +54,7 @@ type Request struct {
// 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")
Scenario ScenarioType `json:"scenario,omitempty"` // QueryDSL scenario: "filter", "aggregation", "join", "complex"
Wheres []gou.Where `json:"wheres,omitempty"` // Pre-defined filters (optional), uses GOU QueryDSL Where Wheres []gou.Where `json:"wheres,omitempty"` // Pre-defined filters (optional), uses GOU QueryDSL Where
Orders gou.Orders `json:"orders,omitempty"` // Sort orders (optional), uses GOU QueryDSL Orders Orders gou.Orders `json:"orders,omitempty"` // Sort orders (optional), uses GOU QueryDSL Orders
Select []string `json:"select,omitempty"` // Fields to return (optional) Select []string `json:"select,omitempty"` // Fields to return (optional)
@ -108,6 +120,12 @@ type ProcessedQuery struct {
DSL *gou.QueryDSL `json:"dsl,omitempty"` // For DB search, uses GOU QueryDSL DSL *gou.QueryDSL `json:"dsl,omitempty"` // For DB search, uses GOU QueryDSL
} }
// Keyword represents an extracted keyword with weight
type Keyword struct {
K string `json:"k"` // Keyword text
W float64 `json:"w"` // Weight (0.1-1.0), higher means more relevant
}
// Note: For QueryDSL and Model types, use GOU types directly: // Note: For QueryDSL and Model types, use GOU types directly:
// - github.com/yaoapp/gou/query/gou.QueryDSL // - github.com/yaoapp/gou/query/gou.QueryDSL
// - github.com/yaoapp/gou/model.Model // - github.com/yaoapp/gou/model.Model

View file

@ -110,10 +110,58 @@ Each line is a JSON object:
| `expected` | any | No | Expected output (exact match) | | `expected` | any | No | Expected output (exact match) |
| `user` | string | No | Override user ID | | `user` | string | No | Override user ID |
| `team` | string | No | Override team ID | | `team` | string | No | Override team ID |
| `options` | Options | No | Context options (see below) |
| `timeout` | string | No | Override timeout (e.g., "30s") | | `timeout` | string | No | Override timeout (e.g., "30s") |
| `skip` | bool | No | Skip this test | | `skip` | bool | No | Skip this test |
| `metadata` | map | No | Additional metadata | | `metadata` | map | No | Additional metadata |
### Options
The `options` field allows per-test-case configuration that maps to `context.Options`:
| Field | Type | Description |
| ------------------------ | ------- | -------------------------------------------------- |
| `connector` | string | Override connector (e.g., `"deepseek.v3"`) |
| `mode` | string | Agent mode (default: `"chat"`) |
| `search` | bool | Enable/disable search mode (default: `true`) |
| `disable_global_prompts` | bool | Temporarily disable global prompts |
| `metadata` | map | Custom data passed to hooks (e.g., scenario) |
| `skip` | object | Skip configuration (see below) |
#### Options.skip
| Field | Type | Description |
| --------- | ---- | ------------------------ |
| `history` | bool | Skip history loading |
| `trace` | bool | Skip trace logging |
| `output` | bool | Skip output to client |
| `keyword` | bool | Skip keyword extraction |
| `search` | bool | Skip auto search |
**Example with options:**
```jsonl
{
"id": "T001",
"input": "Query users with status active",
"options": {
"connector": "deepseek.v3",
"metadata": {"scenario": "filter"},
"skip": {"trace": true}
},
"assert": {"type": "json_path", "path": "from", "value": "users"}
}
```
**Using metadata for hook scenarios:**
The `options.metadata` field is passed to agent hooks. For example, a Create Hook can read `options.metadata.scenario` to select different prompt presets:
```jsonl
{"id": "T001", "input": "...", "options": {"metadata": {"scenario": "aggregation"}}}
{"id": "T002", "input": "...", "options": {"metadata": {"scenario": "join"}}}
```
### Input Types ### Input Types
| Type | Description | Example | | Type | Description | Example |
@ -225,9 +273,28 @@ return { pass: true, message: "Validation passed" };
### JSON Path Notes ### JSON Path Notes
- Supports dot-notation: `$.field.subfield` or `field.subfield` - Supports dot-notation: `$.field.subfield` or `field.subfield`
- Supports array indexing: `field[0]`, `field[0].subfield`, `field[0].nested[1]`
- Supports multiple expected values (OR logic): `"value": ["a", "b"]` - passes if actual matches any
- Auto-extracts JSON from markdown code blocks (` ```json ... ``` `) - Auto-extracts JSON from markdown code blocks (` ```json ... ``` `)
- Works with both string output and structured objects - Works with both string output and structured objects
**Array index examples:**
```jsonl
{"id": "T001", "assert": {"type": "json_path", "path": "wheres[0].like", "value": "%test%"}}
{"id": "T002", "assert": {"type": "json_path", "path": "wheres[0].in[0]", "value": "pending"}}
{"id": "T003", "assert": {"type": "json_path", "path": "joins[0].from", "value": "users"}}
{"id": "T004", "assert": {"type": "json_path", "path": "groups[0]", "value": "category"}}
```
**Multiple expected values (OR logic):**
```jsonl
{"id": "T005", "assert": {"type": "json_path", "path": "error", "value": ["missing_schema", "missing_query"]}}
```
This passes if `error` equals either `"missing_schema"` or `"missing_query"`.
## Output Formats ## Output Formats
Determined by `-o` file extension: Determined by `-o` file extension:

View file

@ -4,6 +4,7 @@ import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"regexp" "regexp"
"strconv"
"strings" "strings"
jsoniter "github.com/json-iterator/go" jsoniter "github.com/json-iterator/go"
@ -244,7 +245,7 @@ func (a *Asserter) assertJSONPath(assertion *Assertion, output interface{}) *Ass
jsonData = v jsonData = v
default: default:
result.Passed = false result.Passed = false
result.Message = "output is not a JSON object or array" result.Message = fmt.Sprintf("output is not a JSON object or array, got: %T = %v", output, truncateOutput(output, 200))
return result return result
} }
@ -253,6 +254,8 @@ func (a *Asserter) assertJSONPath(assertion *Assertion, output interface{}) *Ass
actual := a.extractPath(jsonData, path) actual := a.extractPath(jsonData, path)
result.Actual = actual result.Actual = actual
// Compare expected value with actual value
// First, try direct comparison (handles both primitive values and arrays)
if validateOutput(actual, assertion.Value) { if validateOutput(actual, assertion.Value) {
result.Passed = true result.Passed = true
result.Message = fmt.Sprintf("path '%s' equals expected value", assertion.Path) result.Message = fmt.Sprintf("path '%s' equals expected value", assertion.Path)
@ -264,27 +267,113 @@ func (a *Asserter) assertJSONPath(assertion *Assertion, output interface{}) *Ass
return result return result
} }
// extractPath extracts a value from JSON using a simple dot-notation path // truncateOutput truncates output for error messages
func truncateOutput(output interface{}, maxLen int) string {
var s string
switch v := output.(type) {
case string:
s = v
case nil:
return "<nil>"
default:
bytes, err := jsoniter.Marshal(v)
if err != nil {
s = fmt.Sprintf("%v", v)
} else {
s = string(bytes)
}
}
if len(s) > maxLen {
return s[:maxLen] + "..."
}
return s
}
// extractPath extracts a value from JSON using dot-notation path with array index support
// Supports: "field", "field.nested", "field[0]", "field[0].nested", "field.nested[0].value"
func (a *Asserter) extractPath(data interface{}, path string) interface{} { func (a *Asserter) extractPath(data interface{}, path string) interface{} {
parts := strings.Split(path, ".")
current := data current := data
for _, part := range parts { // Parse path into segments, handling both dots and array indices
if part == "" { // e.g., "wheres[0].like" -> ["wheres", "[0]", "like"]
segments := parsePathSegments(path)
for _, segment := range segments {
if segment == "" {
continue continue
} }
// Check if this is an array index like "[0]"
if strings.HasPrefix(segment, "[") && strings.HasSuffix(segment, "]") {
indexStr := segment[1 : len(segment)-1]
index, err := strconv.Atoi(indexStr)
if err != nil {
return nil
}
arr, ok := current.([]interface{})
if !ok {
return nil
}
if index < 0 || index >= len(arr) {
return nil
}
current = arr[index]
} else {
// Regular field access
switch v := current.(type) { switch v := current.(type) {
case map[string]interface{}: case map[string]interface{}:
current = v[part] current = v[segment]
default: default:
return nil return nil
} }
} }
}
return current return current
} }
// parsePathSegments splits a path like "wheres[0].like" into ["wheres", "[0]", "like"]
func parsePathSegments(path string) []string {
var segments []string
var current strings.Builder
for i := 0; i < len(path); i++ {
ch := path[i]
switch ch {
case '.':
if current.Len() > 0 {
segments = append(segments, current.String())
current.Reset()
}
case '[':
if current.Len() > 0 {
segments = append(segments, current.String())
current.Reset()
}
// Find the closing bracket
j := i + 1
for j < len(path) && path[j] != ']' {
j++
}
if j < len(path) {
segments = append(segments, path[i:j+1]) // Include "[" and "]"
i = j
}
default:
current.WriteByte(ch)
}
}
if current.Len() > 0 {
segments = append(segments, current.String())
}
return segments
}
// assertRegex checks if output matches a regex pattern // assertRegex checks if output matches a regex pattern
func (a *Asserter) assertRegex(assertion *Assertion, output interface{}) *AssertionResult { func (a *Asserter) assertRegex(assertion *Assertion, output interface{}) *AssertionResult {
result := &AssertionResult{ result := &AssertionResult{

View file

@ -595,13 +595,13 @@ func (r *AgentReporter) Write(report *Report, w io.Writer) error {
}, },
} }
result, err := agent.Stream(ctx, messages, options) response, err := agent.Stream(ctx, messages, options)
if err != nil { if err != nil {
return fmt.Errorf("reporter agent call failed: %w", err) return fmt.Errorf("reporter agent call failed: %w", err)
} }
// Extract content from result // Extract content from response
content, err := r.extractContent(result) content, err := r.extractContent(response)
if err != nil { if err != nil {
return fmt.Errorf("failed to extract report content: %w", err) return fmt.Errorf("failed to extract report content: %w", err)
} }
@ -615,55 +615,25 @@ func (r *AgentReporter) Write(report *Report, w io.Writer) error {
return nil return nil
} }
// extractContent extracts the report content from the agent's response // extractContent extracts the report content from the agent's *context.Response
func (r *AgentReporter) extractContent(result interface{}) (string, error) { // Now that agent.Stream() returns *context.Response directly,
if result == nil { // we can access fields without type assertions.
return "", fmt.Errorf("agent returned nil result") func (r *AgentReporter) extractContent(response *context.Response) (string, error) {
if response == nil {
return "", fmt.Errorf("agent returned nil response")
} }
// Try to convert to map first (context.Response) // Priority 1: Check Next field (custom hook data)
switch v := result.(type) { if response.Next != nil {
case string: return r.contentToString(response.Next)
return v, nil }
case *context.Response: // Priority 2: Extract from completion content
// Extract from completion content if response.Completion != nil && response.Completion.Content != nil {
if v.Completion != nil && v.Completion.Content != nil { return r.contentToString(response.Completion.Content)
return r.contentToString(v.Completion.Content)
}
// Try next field
if v.Next != nil {
return r.contentToString(v.Next)
} }
return "", fmt.Errorf("no content in response") return "", fmt.Errorf("no content in response")
case map[string]interface{}:
// Check for completion.content
if completion, ok := v["completion"].(map[string]interface{}); ok {
if content, ok := completion["content"]; ok {
return r.contentToString(content)
}
}
// Check for next
if next, ok := v["next"]; ok {
return r.contentToString(next)
}
// Check for content directly
if content, ok := v["content"]; ok {
return r.contentToString(content)
}
// Marshal the whole thing
jsonBytes, _ := jsoniter.Marshal(v)
return string(jsonBytes), nil
default:
// Try to marshal as JSON
jsonBytes, err := jsoniter.Marshal(result)
if err != nil {
return fmt.Sprintf("%v", result), nil
}
return string(jsonBytes), nil
}
} }
// contentToString converts various content types to string // contentToString converts various content types to string

View file

@ -64,15 +64,8 @@ func (r *Executor) RunDirect() (*Report, error) {
ctx := NewTestContextFromOptions(chatID, agentInfo.ID, r.opts, tc) ctx := NewTestContextFromOptions(chatID, agentInfo.ID, r.opts, tc)
defer ctx.Release() defer ctx.Release()
// Set options: skip history (input already contains conversation), connector override // Build context options
opts := &context.Options{ opts := buildContextOptions(tc, r.opts)
Skip: &context.Skip{
History: true, // Skip history loading - input already contains full conversation
},
}
if r.opts.Connector != "" {
opts.Connector = r.opts.Connector
}
// Create timeout context // Create timeout context
timeout := tc.GetTimeout(r.opts.Timeout) timeout := tc.GetTimeout(r.opts.Timeout)
@ -103,12 +96,19 @@ func (r *Executor) RunDirect() (*Report, error) {
output := extractOutput(response) output := extractOutput(response)
r.output.DirectOutput(output) r.output.DirectOutput(output)
// Determine connector: user-specified > agent default
connector := r.opts.Connector
if connector == "" {
connector = agentInfo.Connector
}
// Return minimal report (for exit code handling) // Return minimal report (for exit code handling)
return &Report{ return &Report{
Summary: &Summary{ Summary: &Summary{
Total: 1, Total: 1,
Passed: 1, Passed: 1,
AgentID: agentInfo.ID, AgentID: agentInfo.ID,
Connector: connector,
}, },
}, nil }, nil
} }
@ -165,13 +165,19 @@ func (r *Executor) RunTests() (*Report, error) {
return nil, fmt.Errorf("failed to get assistant: %w", err) return nil, fmt.Errorf("failed to get assistant: %w", err)
} }
// Determine connector: user-specified > agent default
connector := r.opts.Connector
if connector == "" {
connector = agentInfo.Connector
}
// Create report // Create report
report := &Report{ report := &Report{
Summary: &Summary{ Summary: &Summary{
Total: len(testCases), Total: len(testCases),
AgentID: agentInfo.ID, AgentID: agentInfo.ID,
AgentPath: agentInfo.Path, AgentPath: agentInfo.Path,
Connector: r.opts.Connector, Connector: connector,
RunsPerCase: r.opts.Runs, RunsPerCase: r.opts.Runs,
}, },
Environment: NewEnvironment(r.opts.UserID, r.opts.TeamID), Environment: NewEnvironment(r.opts.UserID, r.opts.TeamID),
@ -279,6 +285,7 @@ func (r *Executor) runSingleTest(ast *assistant.Assistant, tc *Case, agentID str
ID: tc.ID, ID: tc.ID,
Input: tc.Input, Input: tc.Input,
Expected: tc.Expected, Expected: tc.Expected,
Options: tc.Options,
} }
// Parse input to messages // Parse input to messages
@ -297,15 +304,8 @@ func (r *Executor) runSingleTest(ast *assistant.Assistant, tc *Case, agentID str
ctx := NewTestContextFromOptions(chatID, agentID, r.opts, tc) ctx := NewTestContextFromOptions(chatID, agentID, r.opts, tc)
defer ctx.Release() defer ctx.Release()
// Set options: skip history (input already contains conversation), connector override // Build context options from test case and runner options
opts := &context.Options{ opts := buildContextOptions(tc, r.opts)
Skip: &context.Skip{
History: true, // Skip history loading - input already contains full conversation
},
}
if r.opts.Connector != "" {
opts.Connector = r.opts.Connector
}
// Create timeout context // Create timeout context
timeout := tc.GetTimeout(r.opts.Timeout) timeout := tc.GetTimeout(r.opts.Timeout)
@ -477,23 +477,96 @@ func writeJSONLine(writer *bufio.Writer, data interface{}) error {
return err return err
} }
// buildContextOptions builds context.Options from test case and runner options
// Priority: test case options > runner options > defaults
func buildContextOptions(tc *Case, runnerOpts *Options) *context.Options {
opts := &context.Options{
Skip: &context.Skip{
History: true, // Default: skip history loading - input already contains full conversation
},
}
// Apply test case options if specified
if tc.Options != nil {
// Connector: test case > runner
if tc.Options.Connector != "" {
opts.Connector = tc.Options.Connector
}
// Mode
if tc.Options.Mode != "" {
opts.Mode = tc.Options.Mode
}
// DisableGlobalPrompts
if tc.Options.DisableGlobalPrompts {
opts.DisableGlobalPrompts = true
}
// Search (pointer to distinguish unset from false)
if tc.Options.Search != nil {
opts.Search = tc.Options.Search
}
// Metadata for hooks
if tc.Options.Metadata != nil {
opts.Metadata = tc.Options.Metadata
}
// Skip options from test case
if tc.Options.Skip != nil {
opts.Skip.Trace = tc.Options.Skip.Trace
opts.Skip.Output = tc.Options.Skip.Output
opts.Skip.Keyword = tc.Options.Skip.Keyword
opts.Skip.Search = tc.Options.Skip.Search
// Note: History defaults to true for tests
}
}
// Runner connector override (highest priority)
if runnerOpts != nil && runnerOpts.Connector != "" {
opts.Connector = runnerOpts.Connector
}
return opts
}
// extractOutput extracts the output from the agent response // extractOutput extracts the output from the agent response
func extractOutput(response interface{}) interface{} { // Priority: Next hook data (if non-empty) > Completion content > nil
func extractOutput(response *context.Response) interface{} {
if response == nil { if response == nil {
return nil return nil
} }
// Try to get completion content from context.Response // Prefer Next hook data if available and non-empty
if resp, ok := response.(*context.Response); ok { // response.Next is already the Data value (not NextHookResponse struct)
if resp.Completion != nil { if response.Next != nil && !isEmptyValue(response.Next) {
return resp.Completion.Content return response.Next
}
if resp.Next != nil {
return resp.Next
} }
// Fall back to raw completion content
if response.Completion != nil {
return response.Completion.Content
} }
return response return nil
}
// isEmptyValue checks if a value is considered "empty" for output purposes
func isEmptyValue(v interface{}) bool {
if v == nil {
return true
}
switch val := v.(type) {
case string:
return val == ""
case map[string]interface{}:
return len(val) == 0
case []interface{}:
return len(val) == 0
}
return false
} }
// validateOutput validates the actual output against expected // validateOutput validates the actual output against expected

View file

@ -205,8 +205,13 @@ type Case struct {
TeamID string `json:"team,omitempty"` TeamID string `json:"team,omitempty"`
// Metadata contains additional metadata for the test case // Metadata contains additional metadata for the test case
// This is passed to ctx.Metadata and can be used by Create Hook
Metadata map[string]interface{} `json:"metadata,omitempty"` Metadata map[string]interface{} `json:"metadata,omitempty"`
// Options contains context options for this test case
// Supports: connector, skip (history, trace, output, keyword, search), mode
Options *CaseOptions `json:"options,omitempty"`
// Skip indicates whether to skip this test case // Skip indicates whether to skip this test case
Skip bool `json:"skip,omitempty"` Skip bool `json:"skip,omitempty"`
@ -215,6 +220,38 @@ type Case struct {
Timeout string `json:"timeout,omitempty"` Timeout string `json:"timeout,omitempty"`
} }
// CaseOptions represents per-test-case context options
// Maps to context.Options fields
type CaseOptions struct {
// Connector overrides the agent's default connector
Connector string `json:"connector,omitempty"`
// Skip configuration
Skip *CaseSkipOptions `json:"skip,omitempty"`
// DisableGlobalPrompts temporarily disables global prompts for this request
DisableGlobalPrompts bool `json:"disable_global_prompts,omitempty"`
// Search mode, default is true (use pointer to distinguish unset from false)
Search *bool `json:"search,omitempty"`
// Mode is the agent mode (default: "chat")
Mode string `json:"mode,omitempty"`
// Metadata for passing custom data to hooks (e.g., scenario selection)
Metadata map[string]interface{} `json:"metadata,omitempty"`
}
// CaseSkipOptions represents skip configuration for a test case
// Maps to context.Skip fields
type CaseSkipOptions struct {
History bool `json:"history,omitempty"` // Skip history loading
Trace bool `json:"trace,omitempty"` // Skip trace logging
Output bool `json:"output,omitempty"` // Skip output to client
Keyword bool `json:"keyword,omitempty"` // Skip keyword extraction
Search bool `json:"search,omitempty"` // Skip auto search
}
// Assertion represents a single assertion rule // Assertion represents a single assertion rule
type Assertion struct { type Assertion struct {
// Type is the assertion type: // Type is the assertion type:
@ -334,6 +371,9 @@ type Result struct {
// Error contains the error message if status is failed/error/timeout // Error contains the error message if status is failed/error/timeout
Error string `json:"error,omitempty"` Error string `json:"error,omitempty"`
// Options contains the context options used for this test case
Options *CaseOptions `json:"options,omitempty"`
// Metadata contains additional result metadata // Metadata contains additional result metadata
Metadata map[string]interface{} `json:"metadata,omitempty"` Metadata map[string]interface{} `json:"metadata,omitempty"`
} }

View file

@ -3,6 +3,12 @@ package testutils
import ( import (
"testing" "testing"
_ "github.com/yaoapp/gou/encoding"
"github.com/yaoapp/gou/model"
"github.com/yaoapp/gou/query"
"github.com/yaoapp/gou/query/gou"
_ "github.com/yaoapp/gou/text"
"github.com/yaoapp/xun/capsule"
"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"
@ -28,6 +34,21 @@ func Prepare(t *testing.T, opts ...interface{}) {
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
// Register default query engine (required for DB search)
// capsule.Global is initialized by test.Prepare
if _, has := query.Engines["default"]; !has && capsule.Global != nil {
query.Register("default", &gou.Query{
Query: capsule.Query(),
GetTableName: func(s string) string {
if mod, has := model.Models[s]; has {
return mod.MetaData.Table.Name
}
return s
},
AESKey: config.Conf.DB.AESKey,
})
}
} }
// Clean clean the test environment // Clean clean the test environment

File diff suppressed because one or more lines are too long

2
go.mod
View file

@ -13,7 +13,7 @@ require (
github.com/elazarl/go-bindata-assetfs v1.0.1 github.com/elazarl/go-bindata-assetfs v1.0.1
github.com/emersion/go-imap v1.2.1 github.com/emersion/go-imap v1.2.1
github.com/evanw/esbuild v0.25.4 github.com/evanw/esbuild v0.25.4
github.com/expr-lang/expr v1.17.3 github.com/expr-lang/expr v1.17.7
github.com/fatih/color v1.18.0 github.com/fatih/color v1.18.0
github.com/fsnotify/fsnotify v1.9.0 github.com/fsnotify/fsnotify v1.9.0
github.com/gin-gonic/gin v1.10.1 github.com/gin-gonic/gin v1.10.1

2
go.sum
View file

@ -81,6 +81,8 @@ github.com/evanw/esbuild v0.25.4 h1:k1bTSim+usBG27w7BfOCorhgx3tO+6bAfMj5pR+6SKg=
github.com/evanw/esbuild v0.25.4/go.mod h1:D2vIQZqV/vIf/VRHtViaUtViZmG7o+kKmlBfVQuRi48= github.com/evanw/esbuild v0.25.4/go.mod h1:D2vIQZqV/vIf/VRHtViaUtViZmG7o+kKmlBfVQuRi48=
github.com/expr-lang/expr v1.17.3 h1:myeTTuDFz7k6eFe/JPlep/UsiIjVhG61FMHFu63U7j0= github.com/expr-lang/expr v1.17.3 h1:myeTTuDFz7k6eFe/JPlep/UsiIjVhG61FMHFu63U7j0=
github.com/expr-lang/expr v1.17.3/go.mod h1:8/vRC7+7HBzESEqt5kKpYXxrxkr31SaO8r40VO/1IT4= github.com/expr-lang/expr v1.17.3/go.mod h1:8/vRC7+7HBzESEqt5kKpYXxrxkr31SaO8r40VO/1IT4=
github.com/expr-lang/expr v1.17.7 h1:Q0xY/e/2aCIp8g9s/LGvMDCC5PxYlvHgDZRQ4y16JX8=
github.com/expr-lang/expr v1.17.7/go.mod h1:8/vRC7+7HBzESEqt5kKpYXxrxkr31SaO8r40VO/1IT4=
github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk= github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk=
github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM= github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM=
github.com/fatih/color v1.18.0/go.mod h1:4FelSpRwEGDpQ12mAdzqdOukCy4u8WUtOY6lkT/6HfU= github.com/fatih/color v1.18.0/go.mod h1:4FelSpRwEGDpQ12mAdzqdOukCy4u8WUtOY6lkT/6HfU=

View file

@ -1,21 +1,25 @@
- role: system - role: system
content: | content: |
Extract keywords from text content. You are a keyword extraction tool, NOT a chatbot. Do NOT answer questions or provide explanations.
Your ONLY job: analyze text and output a JSON array of keywords with weights.
Task: Output format: ["keyword:weight", ...]
1. Analyze input text
2. Extract important keywords
3. Return JSON format
4. Match input language
Response Format (JSON only): Weight:
```json - 1.0: Core topic
{"keywords": ["keyword1", "keyword2", ...]} - 0.8-0.9: Key concepts
``` - 0.6-0.7: Supporting themes
- 0.4-0.5: Peripheral concepts
Guidelines: Examples:
- Extract 5-15 keywords based on content length - Input(EN): "Developers frustrated with callback hell. ES2017 async/await improved readability."
- Prioritize nouns, proper nouns, key concepts - Output: ["async/await:1", "asynchronous programming:0.9", "ES2017:0.8", "code readability:0.7"]
- Include single words and short phrases
- Exclude common stop words - Input(中文): "用户反馈APP启动慢、页面卡顿。需要优化首屏加载和内存占用。"
- Keywords MUST be in the same language as input - Output: ["性能优化:1", "启动速度:0.9", "内存管理:0.8", "用户体验:0.7"]
Rules:
- ONLY output JSON array, nothing else
- Max 5 keywords, sorted by weight
- Summarize related concepts
- Match input language (EN→EN, 中文→中文)

View file

@ -1,13 +1,20 @@
/** /**
* Keyword Extraction Agent - Next Hook * Keyword Extraction Agent - Next Hook
* Parses LLM response and extracts keywords with error tolerance * Parses LLM response and extracts keywords with weight
* Format: ["keyword:weight", ...] -> [{k, w}, ...]
*/ */
// @ts-nocheck // @ts-nocheck
/** Keyword with weight */
interface Keyword {
k: string; // keyword
w: number; // weight (0.1-1.0)
}
/** /**
* Next hook - processes keyword extraction response * Next hook - processes keyword extraction response
* Uses text.ExtractJSON for fault-tolerant JSON extraction from LLM output * Parses format: ["keyword1:0.9", "keyword2:0.8", ...]
*/ */
function Next( function Next(
ctx: agent.Context, ctx: agent.Context,
@ -21,22 +28,17 @@ function Next(
} }
const content = completion.content; const content = completion.content;
let keywords: string[] = []; let keywords: Keyword[] = [];
try { try {
// Use text.ExtractJSON for fault-tolerant extraction // Extract JSON array from response
// Handles markdown code blocks, broken JSON, etc. const parsed = Process("text.ExtractJSON", content) as string[] | null;
const parsed = Process("text.ExtractJSON", content) as {
keywords?: string[];
} | null;
if (parsed && Array.isArray(parsed.keywords)) { if (parsed && Array.isArray(parsed)) {
keywords = parsed.keywords.filter( keywords = parseKeywordArray(parsed);
(k) => typeof k === "string" && k.trim().length > 0
);
} }
} catch (e) { } catch (e) {
// If extraction fails, try to extract keywords from text // If extraction fails, try to extract from text
keywords = extractKeywordsFromText(content); keywords = extractKeywordsFromText(content);
} }
@ -45,6 +47,9 @@ function Next(
keywords = extractKeywordsFromText(content); keywords = extractKeywordsFromText(content);
} }
// Sort by weight descending and limit to 5
keywords = keywords.sort((a, b) => b.w - a.w).slice(0, 5);
// Return parsed keywords // Return parsed keywords
return { return {
data: { data: {
@ -54,49 +59,99 @@ function Next(
} }
/** /**
* Extract keywords from plain text when JSON parsing fails * Parse keyword array format: ["keyword:weight", ...]
* Handles formats like: * Examples: ["AI:0.9", "机器学习:0.8", "deep learning:0.7"]
* - Comma-separated: "keyword1, keyword2, keyword3"
* - Line-separated: "keyword1\nkeyword2\nkeyword3"
* - Bullet points: "- keyword1\n- keyword2"
* - Numbered: "1. keyword1\n2. keyword2"
*/ */
function extractKeywordsFromText(text: string): string[] { function parseKeywordArray(items: (string | any)[]): Keyword[] {
const keywords: string[] = []; const keywords: Keyword[] = [];
// Remove common prefixes/suffixes for (const item of items) {
let cleaned = text if (typeof item === "string") {
.replace(/^[\s\S]*?keywords?[\s:]*\[?/i, "") // Remove "keywords:" prefix const parsed = parseKeywordString(item);
.replace(/\][\s\S]*$/, "") // Remove trailing ] if (parsed) {
.trim(); keywords.push(parsed);
}
} else if (item && typeof item === "object" && item.k) {
// Fallback: handle {k, w} format
const k = String(item.k).trim();
const w =
typeof item.w === "number" ? Math.min(1.0, Math.max(0.1, item.w)) : 0.5;
if (k.length > 0) {
keywords.push({ k, w });
}
}
}
// Try line-by-line extraction return keywords;
const lines = cleaned.split(/[\n\r]+/); }
/**
* Parse single keyword string: "keyword:weight" or "keyword"
*/
function parseKeywordString(str: string): Keyword | null {
const trimmed = str.trim().replace(/^["']+|["']+$/g, ""); // Remove quotes
if (!trimmed) return null;
// Try to split by last colon (keyword may contain colons)
const lastColonIdx = trimmed.lastIndexOf(":");
if (lastColonIdx > 0) {
const keyword = trimmed.substring(0, lastColonIdx).trim();
const weightStr = trimmed.substring(lastColonIdx + 1).trim();
const weight = parseFloat(weightStr);
if (keyword && !isNaN(weight)) {
return {
k: keyword,
w: Math.min(1.0, Math.max(0.1, weight)),
};
}
}
// No weight found, return with default weight
return { k: trimmed, w: 0.5 };
}
/**
* Extract keywords from plain text when JSON parsing fails
*/
function extractKeywordsFromText(text: string): Keyword[] {
const keywords: Keyword[] = [];
// Try to find array-like content
const arrayMatch = text.match(/\[([^\]]+)\]/);
if (arrayMatch) {
const items = arrayMatch[1].split(",");
for (const item of items) {
const parsed = parseKeywordString(item);
if (parsed) {
keywords.push(parsed);
}
}
if (keywords.length > 0) return keywords;
}
// Fallback: line-by-line extraction
const lines = text.split(/[\n\r,]+/);
let defaultWeight = 1.0;
for (const line of lines) { for (const line of lines) {
// Remove bullet points, numbers, quotes let cleaned = line
let keyword = line .replace(/^[\s\-\*\•\d\.\[\]"'`]+/, "") // Remove prefixes
.replace(/^[\s\-\*\•\d\.]+/, "") // Remove bullets/numbers .replace(/[\]"'`]+$/, "") // Remove suffixes
.replace(/^["'`]+|["'`]+$/g, "") // Remove quotes
.replace(/,\s*$/, "") // Remove trailing comma
.trim(); .trim();
// Skip empty or too long if (cleaned.length > 0 && cleaned.length < 100) {
if (keyword.length > 0 && keyword.length < 100) { const parsed = parseKeywordString(cleaned);
// Split by comma if contains multiple if (parsed) {
if (keyword.includes(",")) { // Use parsed weight or assign decreasing default
const parts = keyword.split(",").map((p) => p.trim()); if (parsed.w === 0.5) {
for (const part of parts) { parsed.w = Math.max(0.1, defaultWeight);
if (part.length > 0 && part.length < 100) { defaultWeight -= 0.1;
keywords.push(part);
} }
} keywords.push(parsed);
} else {
keywords.push(keyword);
} }
} }
} }
// Deduplicate return keywords;
return [...new Set(keywords)];
} }

View file

@ -1,20 +1,57 @@
# Need Search Agent # Need Search Agent
- role: system - role: system
content: | content: |
Classify if user query needs external search. You are a search intent classifier. Analyze user input and classify whether external search is needed.
## Rules ## Your Task
NO SEARCH: greetings, chitchat, math, code generation, text processing, general knowledge - Classify the user's query into search categories
WEB: real-time data (weather, news, prices), current events, recent info - Output MUST be a JSON with exactly these 3 fields: need_search, search_types, confidence
KB: documentation, how-to, configuration, FAQ - DO NOT extract keywords, DO NOT answer the question, DO NOT add explanations
DB: user data, orders, records, business data
## Response (JSON only) ## Classification Rules
{"need_search": bool, "search_types": ["web"|"kb"|"db"], "confidence": 0-1}
### need_search=false (No search needed)
Use when the question can be answered from LLM's internal knowledge:
- Greetings & chitchat: "hello", "how are you", casual conversation
- Math & calculations: arithmetic, equations, formulas
- Code generation: write code, debug, explain code, algorithms
- Text processing: translate, summarize, rewrite, format
- General knowledge: history, science, concepts (not time-sensitive)
- Creative tasks: write stories, poems, brainstorm ideas
- Reasoning & logic: philosophy, opinions, hypothetical questions
### need_search=true with search_types=["web"] (Web search)
Use when real-time or frequently changing information is needed:
- Current events: news, breaking stories, recent happenings
- Time-sensitive data: weather, stock prices, exchange rates, sports scores
- Live information: event schedules, store hours, availability
- Recent updates: latest versions, new releases, current status
- Location-based: nearby places, local info, addresses
### need_search=true with search_types=["kb"] (Knowledge base)
Use when querying internal documentation or product knowledge:
- Documentation: how-to guides, tutorials, setup instructions
- Configuration: settings, parameters, options explained
- Product info: features, specifications, capabilities
- Policies: terms, rules, guidelines, compliance
- FAQ: common questions about the system/product
- Troubleshooting: error messages, known issues, solutions
### need_search=true with search_types=["db"] (Database)
Use when querying user-specific or transactional data:
- Personal data: "my orders", "my profile", "my history"
- Account info: balance, subscription, membership status
- Business records: invoices, transactions, payments
- User preferences: settings, saved items, favorites
- Keywords: "my", "mine", specific order/ID numbers
## Required Output Format (JSON only, no markdown)
{"need_search": true/false, "search_types": [], "confidence": 0.0-1.0}
## Examples ## Examples
"Hello" → {"need_search": false, "search_types": [], "confidence": 0.99} "Hello" → {"need_search": false, "search_types": [], "confidence": 0.99}
"Today's weather" → {"need_search": true, "search_types": ["web"], "confidence": 0.95} "Today's weather" → {"need_search": true, "search_types": ["web"], "confidence": 0.95}
"Write a sort function" → {"need_search": false, "search_types": [], "confidence": 0.90} "Write a bubble sort in JS" → {"need_search": false, "search_types": [], "confidence": 0.95}
"用JavaScript写冒泡排序" → {"need_search": false, "search_types": [], "confidence": 0.95}
"How to config DB" → {"need_search": true, "search_types": ["kb"], "confidence": 0.85} "How to config DB" → {"need_search": true, "search_types": ["kb"], "confidence": 0.85}
"My orders" → {"need_search": true, "search_types": ["db"], "confidence": 0.95} "My orders" → {"need_search": true, "search_types": ["db"], "confidence": 0.95}

View file

@ -4,7 +4,7 @@
"type": "worker", "type": "worker",
"uses": { "search": "disabled" }, "uses": { "search": "disabled" },
"options": { "options": {
"max_tokens": 2000, "max_tokens": 8192,
"temperature": 0.2 "temperature": 0.2
} }
} }

View file

@ -1,43 +1,141 @@
# QueryDSL Generator Agent Prompts # QueryDSL Generator Agent - Main Prompt (Default/Basic Queries)
- role: system - role: system
content: | content: |
You are a QueryDSL generator. Your task is to convert natural language queries into Yao QueryDSL format. You are a QueryDSL generator. Convert natural language queries into Yao QueryDSL JSON format.
## QueryDSL Structure ## QueryDSL JSON Schema
```json ```json
{ {
"select": ["field1", "field2"], "$schema": "http://json-schema.org/draft-07/schema#",
"from": "table_name", "title": "QueryDSL",
"wheres": [ "description": "Gou Query Domain Specific Language for database queries",
{"field": "name", "op": "=", "value": "test"}, "type": "object",
{"field": "status", "op": "in", "value": ["active", "pending"]} "definitions": {
], "expression": {
"orders": [ "type": "string",
{"field": "created_at", "sort": "desc"} "description": "Field expression. Syntax: field, table.field, :FUNC(args), field as alias"
], },
"limit": 20 "condition": {
"type": "object",
"description": "Query condition",
"properties": {
"field": { "type": "string" },
"op": { "type": "string", "description": "=, >, >=, <, <=, <>, like, match, in, is" },
"value": { "description": "Compare value" },
"or": { "type": "boolean", "default": false },
"=": { "description": "Shorthand for op='='" },
">": {}, ">=": {}, "<": {}, "<=": {}, "<>": {},
"like": { "description": "Shorthand for op='like'" },
"in": { "type": "array", "description": "Shorthand for op='in'" },
"is": { "type": "string", "enum": ["null", "not null"] }
}
},
"where": {
"allOf": [
{ "$ref": "#/definitions/condition" },
{ "properties": { "wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } } } }
]
},
"order": {
"oneOf": [
{ "type": "string", "description": "'field desc', 'field asc'" },
{ "type": "object", "properties": { "field": {}, "sort": { "enum": ["asc", "desc"] } } }
]
},
"group": {
"oneOf": [
{ "type": "string", "description": "'field', 'field rollup 合计'" },
{ "type": "object", "properties": { "field": {}, "rollup": { "type": "string" } } }
]
},
"join": {
"type": "object",
"properties": {
"from": { "description": "Table to join" },
"key": { "description": "Join key field" },
"foreign": { "description": "Foreign key field" },
"left": { "type": "boolean" },
"right": { "type": "boolean" }
},
"required": ["from", "key", "foreign"]
}
},
"properties": {
"select": { "type": "array", "items": { "$ref": "#/definitions/expression" } },
"from": { "type": "string", "description": "Table name" },
"wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } },
"orders": { "description": "ORDER BY" },
"groups": { "description": "GROUP BY" },
"havings": { "type": "array", "description": "HAVING conditions" },
"joins": { "type": "array", "items": { "$ref": "#/definitions/join" } },
"limit": { "type": "integer", "description": "Max records" },
"offset": { "type": "integer", "description": "Skip records" },
"page": { "type": "integer", "description": "Page number (1-based)" },
"pagesize": { "type": "integer", "description": "Records per page" },
"first": { "description": "Return first record(s)" }
}
} }
``` ```
## Supported Operators ## Condition Format
- Comparison: =, !=, >, >=, <, <= Conditions use operator as JSON key: `{"field": "xxx", "OPERATOR": VALUE}`
- Pattern: like, not like - `"="` : `{"field": "status", "=": "active"}`
- Range: in, not in, between - `">"` : `{"field": "price", ">": 100}`
- Null check: is null, is not null - `">="` : `{"field": "age", ">=": 18}`
- `"like"` : `{"field": "name", "like": "%test%"}`
## Basic Examples
Input: "查询所有用户"
Schema:
```json
{"name": "users", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "姓名"},
{"name": "email", "type": "string", "label": "邮箱"},
{"name": "status", "type": "string", "label": "状态"}
]}
```
Output:
{"select": ["id", "name", "email", "status"], "from": "users", "limit": 20}
Input: "Find active users sorted by name"
Schema:
```json
{"name": "users", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "Name"},
{"name": "status", "type": "string", "label": "Status"}
]}
```
Output:
{"select": ["id", "name", "status"], "from": "users", "wheres": [{"field": "status", "=": "active"}], "orders": ["name asc"], "limit": 20}
## Response Format ## Response Format
Always respond with valid JSON: Output JSON only. No markdown, no explanation.
```json
{ ### Success Response
"dsl": { ... }, Return QueryDSL directly:
"explain": "Brief explanation of the query", {"select": [...], "from": "table", "wheres": [...], "limit": 20}
"warnings": ["any warnings or notes"]
} ### Error Response
``` When input is insufficient or invalid, return error JSON:
{"error": "error_code", "message": "Error description"}
Error codes:
- `missing_schema`: No schema provided
- `missing_query`: No query/requirement provided
- `invalid_field`: Referenced field not in schema
- `ambiguous_query`: Query intent unclear, need more details
Error examples:
{"error": "missing_schema", "message": "Schema is required"}
{"error": "missing_query", "message": "Query requirement is required"}
{"error": "invalid_field", "message": "Field 'xxx' does not exist in schema"}
{"error": "ambiguous_query", "message": "Query is ambiguous, please provide more details"}
## Guidelines ## Guidelines
- Generate valid QueryDSL based on the provided schema 1. Only use fields from the provided schema (use column.name)
- Use appropriate operators for the query intent 2. Default limit to 20 if not specified
- Include only fields that exist in the schema 3. Return error JSON if input is insufficient
- Add helpful explanations for complex queries 4. IMPORTANT: Verify your JSON syntax before output. Ensure all key-value pairs use colon (:), e.g. {"field": "price", ">": 100} NOT {"field": "price", ">", 100}

View file

@ -0,0 +1,172 @@
# QueryDSL Generator - Aggregation/Statistics Scenario
- role: system
content: |
You are a QueryDSL generator. Convert natural language queries into Yao QueryDSL JSON format.
This scenario focuses on AGGREGATION and STATISTICS queries.
## QueryDSL JSON Schema
```json
{
"$schema": "http://json-schema.org/draft-07/schema#",
"title": "QueryDSL",
"description": "Gou Query Domain Specific Language for database queries",
"type": "object",
"definitions": {
"expression": {
"type": "string",
"description": "Field expression. Syntax: field, table.field, :FUNC(args), field as alias"
},
"condition": {
"type": "object",
"description": "Query condition",
"properties": {
"field": { "type": "string" },
"op": { "type": "string", "description": "=, >, >=, <, <=, <>, like, match, in, is" },
"value": { "description": "Compare value" },
"or": { "type": "boolean", "default": false },
"=": { "description": "Shorthand for op='='" },
">": {}, ">=": {}, "<": {}, "<=": {}, "<>": {},
"like": { "description": "Shorthand for op='like'" },
"in": { "type": "array", "description": "Shorthand for op='in'" },
"is": { "type": "string", "enum": ["null", "not null"] }
}
},
"where": {
"allOf": [
{ "$ref": "#/definitions/condition" },
{ "properties": { "wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } } } }
]
},
"order": {
"oneOf": [
{ "type": "string", "description": "'field desc', 'field asc'" },
{ "type": "object", "properties": { "field": {}, "sort": { "enum": ["asc", "desc"] } } }
]
},
"group": {
"oneOf": [
{ "type": "string", "description": "'field', 'field rollup 合计'" },
{ "type": "object", "properties": { "field": {}, "rollup": { "type": "string" } } }
]
},
"join": {
"type": "object",
"properties": {
"from": { "description": "Table to join" },
"key": { "description": "Join key field" },
"foreign": { "description": "Foreign key field" },
"left": { "type": "boolean" },
"right": { "type": "boolean" }
},
"required": ["from", "key", "foreign"]
}
},
"properties": {
"select": { "type": "array", "items": { "$ref": "#/definitions/expression" } },
"from": { "type": "string", "description": "Table name" },
"wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } },
"orders": { "description": "ORDER BY" },
"groups": { "description": "GROUP BY" },
"havings": { "type": "array", "description": "HAVING conditions" },
"joins": { "type": "array", "items": { "$ref": "#/definitions/join" } },
"limit": { "type": "integer", "description": "Max records" },
"offset": { "type": "integer", "description": "Skip records" },
"page": { "type": "integer", "description": "Page number (1-based)" },
"pagesize": { "type": "integer", "description": "Records per page" },
"first": { "description": "Return first record(s)" }
}
}
```
## Aggregate Functions
- `:COUNT(field)` - Count records
- `:SUM(field)` - Sum values
- `:AVG(field)` - Average
- `:MAX(field)` - Maximum
- `:MIN(field)` - Minimum
- `:DATE(field)` - Extract date from datetime
- `:YEAR(field)`, `:MONTH(field)` - Extract year/month
## Groups Syntax
- String: `"category"` or with rollup `"category rollup 合计"`
- Object: `{"field": "category", "rollup": "Total"}`
## Condition Format
Conditions use operator as JSON key: `{"field": "xxx", "OPERATOR": VALUE}`
- `">"` : `{"field": "price", ">": 100}`
- `">="` : `{"field": "count", ">=": 10}`
- `"="` : `{"field": "status", "=": "active"}`
## Havings (filter aggregated results)
Use after GROUP BY to filter aggregated values:
- `{"field": ":SUM(amount)", ">": 1000}`
- `{"field": ":COUNT(id)", ">=": 10}`
## Examples
Input: "按状态统计订单数量"
Schema:
```json
{"name": "orders", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "status", "type": "string", "label": "状态"},
{"name": "amount", "type": "decimal", "label": "金额"}
]}
```
Output:
{"select": ["status", ":COUNT(id) as count"], "from": "orders", "groups": ["status"]}
Input: "各分类销售总额只显示超过10000的"
Schema:
```json
{"name": "products", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "category", "type": "string", "label": "分类"},
{"name": "sales", "type": "decimal", "label": "销售额"}
]}
```
Output:
{"select": ["category", ":SUM(sales) as total"], "from": "products", "groups": ["category"], "havings": [{"field": ":SUM(sales)", ">": 10000}], "orders": ["total desc"]}
Input: "Monthly order count"
Schema:
```json
{"name": "orders", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "amount", "type": "decimal", "label": "Amount"},
{"name": "created_at", "type": "datetime", "label": "Created At"}
]}
```
Output:
{"select": [":YEAR(created_at) as year", ":MONTH(created_at) as month", ":COUNT(id) as count"], "from": "orders", "groups": [":YEAR(created_at)", ":MONTH(created_at)"], "orders": ["year desc", "month desc"]}
Input: "每个用户的平均消费和最大单笔订单"
Schema:
```json
{"name": "orders", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "user_id", "type": "integer", "label": "用户ID"},
{"name": "amount", "type": "decimal", "label": "金额"}
]}
```
Output:
{"select": ["user_id", ":AVG(amount) as avg_amount", ":MAX(amount) as max_amount"], "from": "orders", "groups": ["user_id"]}
## Response Format
Output JSON only. No markdown, no explanation.
### Success Response
{"select": [...], "from": "table", "groups": [...]}
### Error Response
{"error": "error_code", "message": "Error description"}
- `missing_schema`: No schema provided
- `missing_query`: No query/requirement provided
- `invalid_field`: Referenced field not in schema
- `ambiguous_query`: Query intent unclear
## Guidelines
1. Only use fields from the provided schema (use column.name)
2. Default limit to 20 if not specified
3. Return error JSON if input is insufficient
4. IMPORTANT: Verify your JSON syntax before output. Ensure all key-value pairs use colon (:), e.g. {"field": "price", ">": 100} NOT {"field": "price", ">", 100}

View file

@ -0,0 +1,163 @@
# QueryDSL Generator - Complex Query Scenario (Filter + Aggregation)
- role: system
content: |
You are a QueryDSL generator. Convert natural language queries into Yao QueryDSL JSON format.
This scenario focuses on COMPLEX queries combining filters, aggregations, and sorting.
## QueryDSL JSON Schema
```json
{
"$schema": "http://json-schema.org/draft-07/schema#",
"title": "QueryDSL",
"description": "Gou Query Domain Specific Language for database queries",
"type": "object",
"definitions": {
"expression": {
"type": "string",
"description": "Field expression. Syntax: field, table.field, :FUNC(args), field as alias"
},
"condition": {
"type": "object",
"description": "Query condition",
"properties": {
"field": { "type": "string" },
"op": { "type": "string", "description": "=, >, >=, <, <=, <>, like, match, in, is" },
"value": { "description": "Compare value" },
"or": { "type": "boolean", "default": false },
"=": { "description": "Shorthand for op='='" },
">": {}, ">=": {}, "<": {}, "<=": {}, "<>": {},
"like": { "description": "Shorthand for op='like'" },
"in": { "type": "array", "description": "Shorthand for op='in'" },
"is": { "type": "string", "enum": ["null", "not null"] }
}
},
"where": {
"allOf": [
{ "$ref": "#/definitions/condition" },
{ "properties": { "wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } } } }
]
},
"order": {
"oneOf": [
{ "type": "string", "description": "'field desc', 'field asc'" },
{ "type": "object", "properties": { "field": {}, "sort": { "enum": ["asc", "desc"] } } }
]
},
"group": {
"oneOf": [
{ "type": "string", "description": "'field', 'field rollup 合计'" },
{ "type": "object", "properties": { "field": {}, "rollup": { "type": "string" } } }
]
},
"join": {
"type": "object",
"properties": {
"from": { "description": "Table to join" },
"key": { "description": "Join key field" },
"foreign": { "description": "Foreign key field" },
"left": { "type": "boolean" },
"right": { "type": "boolean" }
},
"required": ["from", "key", "foreign"]
}
},
"properties": {
"select": { "type": "array", "items": { "$ref": "#/definitions/expression" } },
"from": { "type": "string", "description": "Table name" },
"wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } },
"orders": { "description": "ORDER BY" },
"groups": { "description": "GROUP BY" },
"havings": { "type": "array", "description": "HAVING conditions" },
"joins": { "type": "array", "items": { "$ref": "#/definitions/join" } },
"limit": { "type": "integer", "description": "Max records" },
"offset": { "type": "integer", "description": "Skip records" },
"page": { "type": "integer", "description": "Page number (1-based)" },
"pagesize": { "type": "integer", "description": "Records per page" },
"first": { "description": "Return first record(s)" }
}
}
```
## Condition Format
Conditions use operator as JSON key with value: `{"field": "xxx", "OPERATOR": VALUE}`
Operators (used as JSON keys):
- `"="` : `{"field": "status", "=": "active"}`
- `">"` : `{"field": "price", ">": 100}`
- `">="` : `{"field": "age", ">=": 18}`
- `"<"` : `{"field": "stock", "<": 10}`
- `"<="` : `{"field": "score", "<=": 60}`
- `"like"` : `{"field": "name", "like": "%test%"}`
- `"in"` : `{"field": "status", "in": ["a", "b"]}`
- `"is"` : `{"field": "deleted_at", "is": "null"}`
- OR: `{"or": true, "field": "name", "=": "test"}`
- Nested: `{"wheres": [cond1, {"or": true, ...cond2}]}`
## Aggregate Functions
- `:COUNT(field)`, `:SUM(field)`, `:AVG(field)`, `:MAX(field)`, `:MIN(field)`
- `:DATE(field)`, `:YEAR(field)`, `:MONTH(field)`
## Havings (filter aggregated results)
- `{"field": ":SUM(amount)", ">": 1000}`
## Examples
Input: "统计今年每月的活跃订单数和总金额"
Schema:
```json
{"name": "orders", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "status", "type": "string", "label": "状态"},
{"name": "amount", "type": "decimal", "label": "金额"},
{"name": "created_at", "type": "datetime", "label": "创建时间"}
]}
```
Output:
{"select": [":MONTH(created_at) as month", ":COUNT(id) as count", ":SUM(amount) as total"], "from": "orders", "wheres": [{"field": "status", "=": "active"}, {"field": "created_at", ">=": "2024-01-01"}], "groups": [":MONTH(created_at)"], "orders": ["month asc"]}
Input: "Find top 5 categories by sales where price > 100, only show categories with total > 10000"
Schema:
```json
{"name": "products", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "Name"},
{"name": "category", "type": "string", "label": "Category"},
{"name": "price", "type": "decimal", "label": "Price"},
{"name": "sales", "type": "integer", "label": "Sales"}
]}
```
Output:
{"select": ["category", ":SUM(sales) as total_sales"], "from": "products", "wheres": [{"field": "price", ">": 100}], "groups": ["category"], "havings": [{"field": ":SUM(sales)", ">": 10000}], "orders": ["total_sales desc"], "limit": 5}
Input: "按地区统计VIP用户的消费总额只显示消费超过5000的地区"
Schema:
```json
{"name": "users", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "姓名"},
{"name": "region", "type": "string", "label": "地区"},
{"name": "is_vip", "type": "boolean", "label": "VIP"},
{"name": "total_spent", "type": "decimal", "label": "消费总额"}
]}
```
Output:
{"select": ["region", ":COUNT(id) as user_count", ":SUM(total_spent) as total"], "from": "users", "wheres": [{"field": "is_vip", "=": true}], "groups": ["region"], "havings": [{"field": ":SUM(total_spent)", ">": 5000}], "orders": ["total desc"]}
## Response Format
Output JSON only. No markdown, no explanation.
### Success Response
{"select": [...], "from": "table", "wheres": [...], "groups": [...], "havings": [...]}
### Error Response
{"error": "error_code", "message": "Error description"}
- `missing_schema`: No schema provided
- `missing_query`: No query/requirement provided
- `invalid_field`: Referenced field not in schema
- `ambiguous_query`: Query intent unclear
## Guidelines
1. Only use fields from the provided schema (use column.name)
2. Default limit to 20 if not specified
3. Return error JSON if input is insufficient
4. IMPORTANT: Verify your JSON syntax before output. Ensure all key-value pairs use colon (:), e.g. {"field": "price", ">": 100} NOT {"field": "price", ">", 100}

View file

@ -0,0 +1,174 @@
# QueryDSL Generator - Filter/Where Conditions Scenario
- role: system
content: |
You are a QueryDSL generator. Convert natural language queries into Yao QueryDSL JSON format.
This scenario focuses on FILTER and WHERE condition queries.
## QueryDSL JSON Schema
```json
{
"$schema": "http://json-schema.org/draft-07/schema#",
"title": "QueryDSL",
"description": "Gou Query Domain Specific Language for database queries",
"type": "object",
"definitions": {
"expression": {
"type": "string",
"description": "Field expression. Syntax: field, table.field, :FUNC(args), field as alias"
},
"condition": {
"type": "object",
"description": "Query condition",
"properties": {
"field": { "type": "string" },
"op": { "type": "string", "description": "=, >, >=, <, <=, <>, like, match, in, is" },
"value": { "description": "Compare value" },
"or": { "type": "boolean", "default": false },
"=": { "description": "Shorthand for op='='" },
">": {}, ">=": {}, "<": {}, "<=": {}, "<>": {},
"like": { "description": "Shorthand for op='like'" },
"in": { "type": "array", "description": "Shorthand for op='in'" },
"is": { "type": "string", "enum": ["null", "not null"] }
}
},
"where": {
"allOf": [
{ "$ref": "#/definitions/condition" },
{ "properties": { "wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } } } }
]
},
"order": {
"oneOf": [
{ "type": "string", "description": "'field desc', 'field asc'" },
{ "type": "object", "properties": { "field": {}, "sort": { "enum": ["asc", "desc"] } } }
]
},
"group": {
"oneOf": [
{ "type": "string", "description": "'field', 'field rollup 合计'" },
{ "type": "object", "properties": { "field": {}, "rollup": { "type": "string" } } }
]
},
"join": {
"type": "object",
"properties": {
"from": { "description": "Table to join" },
"key": { "description": "Join key field" },
"foreign": { "description": "Foreign key field" },
"left": { "type": "boolean" },
"right": { "type": "boolean" }
},
"required": ["from", "key", "foreign"]
}
},
"properties": {
"select": { "type": "array", "items": { "$ref": "#/definitions/expression" } },
"from": { "type": "string", "description": "Table name" },
"wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } },
"orders": { "description": "ORDER BY" },
"groups": { "description": "GROUP BY" },
"havings": { "type": "array", "description": "HAVING conditions" },
"joins": { "type": "array", "items": { "$ref": "#/definitions/join" } },
"limit": { "type": "integer", "description": "Max records" },
"offset": { "type": "integer", "description": "Skip records" },
"page": { "type": "integer", "description": "Page number (1-based)" },
"pagesize": { "type": "integer", "description": "Records per page" },
"first": { "description": "Return first record(s)" }
}
}
```
## Condition Format
Conditions use operator as JSON key with value: `{"field": "xxx", "OPERATOR": VALUE}`
Operators (used as JSON keys):
- `"="` : `{"field": "status", "=": "active"}`
- `">"` : `{"field": "price", ">": 100}`
- `">="` : `{"field": "age", ">=": 18}`
- `"<"` : `{"field": "stock", "<": 10}`
- `"<="` : `{"field": "score", "<=": 60}`
- `"<>"` : `{"field": "type", "<>": "deleted"}`
- `"like"` : `{"field": "name", "like": "%test%"}`
- `"in"` : `{"field": "status", "in": ["a", "b"]}`
- `"is"` : `{"field": "deleted_at", "is": "null"}`
## When to use = vs like
- Use `=` for: ID, status, type, boolean, enum, exact values
- Use `like` for: name search, title search, content search
- `%keyword%` : contains
- `keyword%` : starts with
- `%keyword` : ends with
## OR and Nested Conditions
- OR: `{"or": true, "field": "name", "=": "test"}`
- Nested (grouping): `{"wheres": [cond1, {"or": true, ...cond2}]}`
- Example: (A AND B) OR C → `[{"wheres": [A, B]}, {"or": true, ...C}]`
## Examples
Input: "查询状态为active的用户"
Schema:
```json
{"name": "users", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "姓名"},
{"name": "status", "type": "string", "label": "状态"}
]}
```
Output:
{"select": ["id", "name", "status"], "from": "users", "wheres": [{"field": "status", "=": "active"}], "limit": 20}
Input: "Search products containing iPhone"
Schema:
```json
{"name": "products", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "Name"},
{"name": "price", "type": "decimal", "label": "Price"}
]}
```
Output:
{"select": ["id", "name", "price"], "from": "products", "wheres": [{"field": "name", "like": "%iPhone%"}], "limit": 20}
Input: "价格100-500的商品"
Schema:
```json
{"name": "products", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "名称"},
{"name": "price", "type": "decimal", "label": "价格"}
]}
```
Output:
{"select": ["id", "name", "price"], "from": "products", "wheres": [{"field": "price", ">=": 100}, {"field": "price", "<=": 500}], "limit": 20}
Input: "状态为pending或processing的订单"
Schema:
```json
{"name": "orders", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "status", "type": "string", "label": "状态"},
{"name": "amount", "type": "decimal", "label": "金额"}
]}
```
Output:
{"select": ["id", "status", "amount"], "from": "orders", "wheres": [{"field": "status", "in": ["pending", "processing"]}], "limit": 20}
## Response Format
Output JSON only. No markdown, no explanation.
### Success Response
{"select": [...], "from": "table", "wheres": [...], "limit": 20}
### Error Response
{"error": "error_code", "message": "Error description"}
- `missing_schema`: No schema provided
- `missing_query`: No query/requirement provided
- `invalid_field`: Referenced field not in schema
- `ambiguous_query`: Query intent unclear
## Guidelines
1. Only use fields from the provided schema (use column.name)
2. Default limit to 20 if not specified
3. Return error JSON if input is insufficient
4. IMPORTANT: Verify your JSON syntax before output. Ensure all key-value pairs use colon (:), e.g. {"field": "price", ">": 100} NOT {"field": "price", ">", 100}

View file

@ -0,0 +1,197 @@
# QueryDSL Generator - Multi-table Join Scenario
- role: system
content: |
You are a QueryDSL generator. Convert natural language queries into Yao QueryDSL JSON format.
This scenario focuses on MULTI-TABLE JOIN queries.
## QueryDSL JSON Schema
```json
{
"$schema": "http://json-schema.org/draft-07/schema#",
"title": "QueryDSL",
"description": "Gou Query Domain Specific Language for database queries",
"type": "object",
"definitions": {
"expression": {
"type": "string",
"description": "Field expression. Syntax: field, table.field, :FUNC(args), field as alias"
},
"condition": {
"type": "object",
"description": "Query condition",
"properties": {
"field": { "type": "string" },
"op": { "type": "string", "description": "=, >, >=, <, <=, <>, like, match, in, is" },
"value": { "description": "Compare value" },
"or": { "type": "boolean", "default": false },
"=": { "description": "Shorthand for op='='" },
">": {}, ">=": {}, "<": {}, "<=": {}, "<>": {},
"like": { "description": "Shorthand for op='like'" },
"in": { "type": "array", "description": "Shorthand for op='in'" },
"is": { "type": "string", "enum": ["null", "not null"] }
}
},
"where": {
"allOf": [
{ "$ref": "#/definitions/condition" },
{ "properties": { "wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } } } }
]
},
"order": {
"oneOf": [
{ "type": "string", "description": "'field desc', 'field asc'" },
{ "type": "object", "properties": { "field": {}, "sort": { "enum": ["asc", "desc"] } } }
]
},
"group": {
"oneOf": [
{ "type": "string", "description": "'field', 'field rollup 合计'" },
{ "type": "object", "properties": { "field": {}, "rollup": { "type": "string" } } }
]
},
"join": {
"type": "object",
"properties": {
"from": { "description": "Table to join" },
"key": { "description": "Join key field" },
"foreign": { "description": "Foreign key field" },
"left": { "type": "boolean" },
"right": { "type": "boolean" }
},
"required": ["from", "key", "foreign"]
}
},
"properties": {
"select": { "type": "array", "items": { "$ref": "#/definitions/expression" } },
"from": { "type": "string", "description": "Table name" },
"wheres": { "type": "array", "items": { "$ref": "#/definitions/where" } },
"orders": { "description": "ORDER BY" },
"groups": { "description": "GROUP BY" },
"havings": { "type": "array", "description": "HAVING conditions" },
"joins": { "type": "array", "items": { "$ref": "#/definitions/join" } },
"limit": { "type": "integer", "description": "Max records" },
"offset": { "type": "integer", "description": "Skip records" },
"page": { "type": "integer", "description": "Page number (1-based)" },
"pagesize": { "type": "integer", "description": "Records per page" },
"first": { "description": "Return first record(s)" }
}
}
```
## Join Syntax
```json
{"from": "table_to_join", "key": "foreign_key_field", "foreign": "primary_key_field", "left": true}
```
- `from`: Table to join
- `key`: Field in main table (foreign key)
- `foreign`: Field in joined table (usually id)
- `left`: true for LEFT JOIN (keep all main table records)
- `right`: true for RIGHT JOIN
- Omit left/right for INNER JOIN (only matching records)
## Condition Format
Conditions use operator as JSON key: `{"field": "xxx", "OPERATOR": VALUE}`
- `"="` : `{"field": "status", "=": "active"}`
- `">"` : `{"field": "amount", ">": 100}`
## Important Rules
1. Always prefix fields with table name: `orders.id`, `users.name`
2. Use alias for clarity: `users.name as user_name`
3. Use LEFT JOIN when you want all main records even without matches
## Examples
Input: "查询订单及用户信息"
Schema:
```json
[
{"name": "orders", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "user_id", "type": "integer", "label": "用户ID"},
{"name": "amount", "type": "decimal", "label": "金额"}
]},
{"name": "users", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "姓名"},
{"name": "email", "type": "string", "label": "邮箱"}
]}
]
```
Output:
{"select": ["orders.id", "orders.amount", "users.name", "users.email"], "from": "orders", "joins": [{"from": "users", "key": "user_id", "foreign": "id", "left": true}], "limit": 20}
Input: "Products with category names"
Schema:
```json
[
{"name": "products", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "Name"},
{"name": "category_id", "type": "integer", "label": "Category ID"},
{"name": "price", "type": "decimal", "label": "Price"}
]},
{"name": "categories", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "Name"}
]}
]
```
Output:
{"select": ["products.id", "products.name as product_name", "products.price", "categories.name as category_name"], "from": "products", "joins": [{"from": "categories", "key": "category_id", "foreign": "id", "left": true}], "limit": 20}
Input: "查询VIP用户的订单"
Schema:
```json
[
{"name": "orders", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "user_id", "type": "integer", "label": "用户ID"},
{"name": "amount", "type": "decimal", "label": "金额"}
]},
{"name": "users", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "姓名"},
{"name": "is_vip", "type": "boolean", "label": "VIP"}
]}
]
```
Output:
{"select": ["orders.id", "orders.amount", "users.name"], "from": "orders", "joins": [{"from": "users", "key": "user_id", "foreign": "id"}], "wheres": [{"field": "users.is_vip", "=": true}], "limit": 20}
Input: "每个用户的订单总额"
Schema:
```json
[
{"name": "users", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "name", "type": "string", "label": "姓名"}
]},
{"name": "orders", "columns": [
{"name": "id", "type": "ID", "label": "ID"},
{"name": "user_id", "type": "integer", "label": "用户ID"},
{"name": "amount", "type": "decimal", "label": "金额"}
]}
]
```
Output:
{"select": ["users.id", "users.name", ":SUM(orders.amount) as total"], "from": "users", "joins": [{"from": "orders", "key": "id", "foreign": "user_id", "left": true}], "groups": ["users.id", "users.name"]}
## Response Format
Output JSON only. No markdown, no explanation.
### Success Response
{"select": [...], "from": "table", "joins": [...], "limit": 20}
### Error Response
{"error": "error_code", "message": "Error description"}
- `missing_schema`: No schema provided
- `missing_query`: No query/requirement provided
- `invalid_field`: Referenced field not in schema
- `missing_relation`: Cannot determine join relationship between tables
- `ambiguous_query`: Query intent unclear
## Guidelines
1. Only use fields from the provided schema (use column.name)
2. Default limit to 20 if not specified
3. Return error JSON if input is insufficient
4. IMPORTANT: Verify your JSON syntax before output. Ensure all key-value pairs use colon (:), e.g. {"field": "price", ">": 100} NOT {"field": "price", ">", 100}

View file

@ -0,0 +1,69 @@
/**
* QueryDSL Generator Agent - Hooks
*
* Scenarios (via metadata.scenario):
* - "filter" : WHERE conditions (=, like, in, OR, nested)
* - "aggregation" : GROUP BY, COUNT, SUM, AVG, HAVING
* - "join" : Multi-table JOIN queries
*
* If not specified, uses default prompts.yml (basic queries)
*/
// @ts-nocheck
// Valid scenario names that map to prompt presets in prompts/ directory
const VALID_SCENARIOS = ["filter", "aggregation", "join", "complex"];
/**
* Create hook - selects prompt preset based on metadata.scenario
*/
function Create(
ctx: agent.Context,
messages: agent.Message[],
options?: Record<string, any>
): agent.HookCreateResponse | null {
// Get scenario from metadata
const scenario = options.metadata?.scenario || ctx.metadata?.scenario;
// If valid scenario specified, return the corresponding preset
if (typeof scenario === "string" && VALID_SCENARIOS.includes(scenario)) {
return {
prompt_preset: scenario,
};
}
// No preset - use default prompts.yml
return null;
}
/**
* Next hook - extracts QueryDSL JSON from LLM response
*/
function Next(
ctx: agent.Context,
payload: agent.NextHookPayload
): agent.NextHookResponse | null {
const completion = payload.completion;
if (!completion || !completion.content) {
return {
data: { error: "empty_response", message: "LLM returned empty content" },
};
}
const content = completion.content;
// Use text.ExtractJSON for fault-tolerant extraction
const dsl = Process("text.ExtractJSON", content);
if (dsl && typeof dsl === "object" && Object.keys(dsl).length > 0) {
return { data: dsl };
}
// Extraction failed, return error with original content
return {
data: {
error: "extraction_failed",
message: "Failed to extract JSON from LLM response",
raw: content,
},
};
}