diff --git a/agent/assistant/agent.go b/agent/assistant/agent.go index 204d87b1..7aff7782 100644 --- a/agent/assistant/agent.go +++ b/agent/assistant/agent.go @@ -16,7 +16,7 @@ import ( // Stream stream the agent // 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 ctx.Logger.SetAssistantID(ast.ID) @@ -376,7 +376,7 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa // ================================================ // Execute Next Hook and Process Response // ================================================ - var finalResponse interface{} + var finalResponse *context.Response var nextResponse *context.NextHookResponse = nil if ast.HookScript != nil { diff --git a/agent/assistant/agent_next_test.go b/agent/assistant/agent_next_test.go index 997b19b2..bf3de999 100644 --- a/agent/assistant/agent_next_test.go +++ b/agent/assistant/agent_next_test.go @@ -56,19 +56,18 @@ func TestAgentNextStandard(t *testing.T) { assert.NoError(t, err) assert.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion) - assert.Nil(t, resp.Next) + assert.NotNil(t, response.Completion) + assert.Nil(t, response.Next) // Verify response structure - assert.Equal(t, "tests.realworld-next", resp.AssistantID) - assert.NotEmpty(t, resp.ContextID) - assert.NotEmpty(t, resp.RequestID) - assert.NotEmpty(t, resp.TraceID) - assert.NotEmpty(t, resp.ChatID) + assert.Equal(t, "tests.realworld-next", response.AssistantID) + assert.NotEmpty(t, response.ContextID) + assert.NotEmpty(t, response.RequestID) + assert.NotEmpty(t, response.TraceID) + assert.NotEmpty(t, response.ChatID) // Verify completion has content - assert.NotNil(t, resp.Completion.Content) + assert.NotNil(t, response.Completion.Content) t.Log("✓ Standard response test passed") } @@ -94,19 +93,18 @@ func TestAgentNextCustomData(t *testing.T) { assert.NoError(t, err) assert.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion) - assert.NotNil(t, resp.Next) + assert.NotNil(t, response.Completion) + assert.NotNil(t, response.Next) // Verify response structure - assert.Equal(t, "tests.realworld-next", resp.AssistantID) - assert.NotEmpty(t, resp.ContextID) - assert.NotEmpty(t, resp.RequestID) - assert.NotEmpty(t, resp.TraceID) + assert.Equal(t, "tests.realworld-next", response.AssistantID) + assert.NotEmpty(t, response.ContextID) + assert.NotEmpty(t, response.RequestID) + assert.NotEmpty(t, response.TraceID) // Verify custom data structure (from scenarioCustomData) - // resp.Next contains the "data" field value from NextHookResponse - nextData, ok := resp.Next.(map[string]interface{}) + // response.Next contains the "data" field value from NextHookResponse + nextData, ok := response.Next.(map[string]interface{}) assert.True(t, ok, "Next should be a map") assert.Equal(t, "custom_response", nextData["type"]) 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.NotNil(t, response) - resp := response.(*context.Response) - // Verify response structure - assert.NotEmpty(t, resp.AssistantID) - assert.NotEmpty(t, resp.ContextID) - assert.NotEmpty(t, resp.RequestID) - assert.NotEmpty(t, resp.TraceID) + assert.NotEmpty(t, response.AssistantID) + assert.NotEmpty(t, response.ContextID) + assert.NotEmpty(t, response.RequestID) + assert.NotEmpty(t, response.TraceID) // Verify completion (delegated agent should have returned completion) - assert.NotNil(t, resp.Completion) - assert.NotNil(t, resp.Completion.Content) + assert.NotNil(t, response.Completion) + assert.NotNil(t, response.Completion.Content) // Next should be from the delegated agent // 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 @@ -177,18 +173,17 @@ func TestAgentNextConditional(t *testing.T) { assert.NoError(t, err) assert.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Next) + assert.NotNil(t, response.Next) // Verify response structure - assert.Equal(t, "tests.realworld-next", resp.AssistantID) - assert.NotEmpty(t, resp.ContextID) - assert.NotEmpty(t, resp.RequestID) - assert.NotEmpty(t, resp.TraceID) + assert.Equal(t, "tests.realworld-next", response.AssistantID) + assert.NotEmpty(t, response.ContextID) + assert.NotEmpty(t, response.RequestID) + assert.NotEmpty(t, response.TraceID) // Verify conditional response structure (from scenarioConditional) - // resp.Next contains the "data" field value from NextHookResponse - nextData, ok := resp.Next.(map[string]interface{}) + // response.Next contains the "data" field value from NextHookResponse + nextData, ok := response.Next.(map[string]interface{}) assert.True(t, ok, "Next should be a map") assert.Equal(t, "Conditional analysis complete", nextData["message"]) assert.Contains(t, nextData, "action") @@ -224,19 +219,18 @@ func TestAgentWithoutNextHook(t *testing.T) { assert.NoError(t, err) assert.NotNil(t, response) - resp := response.(*context.Response) - assert.Nil(t, resp.Next) + assert.Nil(t, response.Next) // Verify response structure - assert.Equal(t, "tests.create", resp.AssistantID) - assert.NotEmpty(t, resp.ContextID) - assert.NotEmpty(t, resp.RequestID) - assert.NotEmpty(t, resp.TraceID) - assert.NotEmpty(t, resp.ChatID) + assert.Equal(t, "tests.create", response.AssistantID) + assert.NotEmpty(t, response.ContextID) + assert.NotEmpty(t, response.RequestID) + assert.NotEmpty(t, response.TraceID) + assert.NotEmpty(t, response.ChatID) // Verify completion - assert.NotNil(t, resp.Completion) - assert.NotNil(t, resp.Completion.Content) + assert.NotNil(t, response.Completion) + assert.NotNil(t, response.Completion.Content) t.Log("✓ No Next Hook test passed") } diff --git a/agent/assistant/assistant.go b/agent/assistant/assistant.go index bd6413e5..7a6825be 100644 --- a/agent/assistant/assistant.go +++ b/agent/assistant/assistant.go @@ -51,7 +51,7 @@ type agentCallerWrapper struct { 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...) } diff --git a/agent/assistant/next.go b/agent/assistant/next.go index 940d0739..fc6c7c7a 100644 --- a/agent/assistant/next.go +++ b/agent/assistant/next.go @@ -8,7 +8,7 @@ import ( ) // 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 npc.NextResponse == nil { return ast.buildStandardResponse(npc), nil @@ -42,7 +42,7 @@ func (ast *Assistant) handleDelegation( ctx *agentContext.Context, delegate *agentContext.DelegateConfig, streamHandler func(message.StreamChunkType, []byte) int, -) (interface{}, error) { +) (*agentContext.Response, error) { // Load the target assistant targetAssistant, err := Get(delegate.AgentID) 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 -func (ast *Assistant) buildStandardResponse(npc *NextProcessContext) interface{} { +func (ast *Assistant) buildStandardResponse(npc *NextProcessContext) *agentContext.Response { return &agentContext.Response{ ContextID: npc.Context.ID, RequestID: npc.Context.RequestID(), diff --git a/agent/assistant/search.go b/agent/assistant/search.go index 3685eba7..dcbbf8c3 100644 --- a/agent/assistant/search.go +++ b/agent/assistant/search.go @@ -98,39 +98,37 @@ func (ast *Assistant) checkSearchIntent(ctx *context.Context, messages []context // Parse the result // 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 - if response.Next != nil { - if nextData, ok := response.Next.(map[string]interface{}); ok { - // Check for data field (from Next hook's {data: result}) - var intentData map[string]interface{} - if data, ok := nextData["data"].(map[string]interface{}); ok { - intentData = data - } else { - intentData = nextData - } + // First try to get from Next hook response + if result.Next != nil { + if nextData, ok := result.Next.(map[string]interface{}); ok { + // Check for data field (from Next hook's {data: result}) + var intentData map[string]interface{} + if data, ok := nextData["data"].(map[string]interface{}); ok { + intentData = data + } else { + intentData = nextData + } - if needSearch, ok := intentData["need_search"].(bool); ok { - reason, _ := intentData["reason"].(string) - ctx.Logger.Debug("Search intent (from Next): need_search=%v, reason=%s", needSearch, reason) - ast.sendIntentDone(ctx, loadingID, needSearch, reason) - return needSearch - } + if needSearch, ok := intentData["need_search"].(bool); ok { + reason, _ := intentData["reason"].(string) + ctx.Logger.Debug("Search intent (from Next): need_search=%v, reason=%s", needSearch, reason) + ast.sendIntentDone(ctx, loadingID, needSearch, reason) + return needSearch } } + } - // Fallback: parse from Completion.Content if Next hook didn't process - if response.Completion != nil { - content, ok := response.Completion.Content.(string) - if !ok || content == "" { - ast.sendIntentDone(ctx, loadingID, true, "") - return true - } - 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 + // Fallback: parse from Completion.Content if Next hook didn't process + if result.Completion != nil { + content, ok := result.Completion.Content.(string) + if !ok || content == "" { + ast.sendIntentDone(ctx, loadingID, true, "") + return true } + 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 } // Default: proceed with search if we can't parse the result diff --git a/agent/assistant/search_auto_disabled_test.go b/agent/assistant/search_auto_disabled_test.go index 5ab7ced6..cc5688c8 100644 --- a/agent/assistant/search_auto_disabled_test.go +++ b/agent/assistant/search_auto_disabled_test.go @@ -77,8 +77,7 @@ func TestSearchAutoDisabled(t *testing.T) { require.NoError(t, err) require.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion, "should have completion") + assert.NotNil(t, response.Completion, "should have completion") t.Logf("✓ Stream executed without search (disabled)") }) } diff --git a/agent/assistant/search_auto_full_test.go b/agent/assistant/search_auto_full_test.go index 204d46d9..e4bc8e84 100644 --- a/agent/assistant/search_auto_full_test.go +++ b/agent/assistant/search_auto_full_test.go @@ -118,8 +118,7 @@ func TestSearchAutoFull(t *testing.T) { } require.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion, "should have completion") + assert.NotNil(t, response.Completion, "should have completion") t.Logf("✓ Stream executed with full search config (Web + KB + DB)") }) } diff --git a/agent/assistant/search_auto_hook_disable_test.go b/agent/assistant/search_auto_hook_disable_test.go index fed6c8cd..23c8285f 100644 --- a/agent/assistant/search_auto_hook_disable_test.go +++ b/agent/assistant/search_auto_hook_disable_test.go @@ -101,8 +101,7 @@ func TestSearchAutoHookDisable(t *testing.T) { require.NoError(t, err) require.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion, "should have completion") + assert.NotNil(t, response.Completion, "should have completion") t.Logf("✓ Stream executed with hook disabling search") }) } diff --git a/agent/assistant/search_auto_keyword_test.go b/agent/assistant/search_auto_keyword_test.go index d5e2a1b0..c6c3d76b 100644 --- a/agent/assistant/search_auto_keyword_test.go +++ b/agent/assistant/search_auto_keyword_test.go @@ -87,8 +87,7 @@ func TestSearchAutoKeyword(t *testing.T) { } require.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion, "should have completion") + assert.NotNil(t, response.Completion, "should have completion") t.Logf("✓ Stream with keyword extraction executed successfully") }) @@ -127,8 +126,7 @@ func TestSearchAutoKeyword(t *testing.T) { } require.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion, "should have completion") + assert.NotNil(t, response.Completion, "should have completion") t.Logf("✓ Stream with Skip.Keyword executed successfully") }) } @@ -179,8 +177,7 @@ func TestSearchAutoKeywordNotConfigured(t *testing.T) { } require.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion, "should have completion") + assert.NotNil(t, response.Completion, "should have completion") t.Logf("✓ Stream without keyword config executed successfully") }) } diff --git a/agent/assistant/search_auto_web_test.go b/agent/assistant/search_auto_web_test.go index 22ec0896..630151e3 100644 --- a/agent/assistant/search_auto_web_test.go +++ b/agent/assistant/search_auto_web_test.go @@ -96,8 +96,7 @@ func TestSearchAutoWeb(t *testing.T) { } require.NotNil(t, response) - resp := response.(*context.Response) - assert.NotNil(t, resp.Completion, "should have completion") + assert.NotNil(t, response.Completion, "should have completion") t.Logf("✓ Stream executed successfully with auto search") }) } diff --git a/agent/caller/caller.go b/agent/caller/caller.go index fb4d2049..229e5c3f 100644 --- a/agent/caller/caller.go +++ b/agent/caller/caller.go @@ -9,7 +9,7 @@ import ( // AgentCaller interface for calling agents (to avoid circular dependency) // Used by content handlers (vision, audio, etc.) and search handlers (agent mode) type AgentCaller interface { - Stream(ctx *agentContext.Context, messages []agentContext.Message, options ...*agentContext.Options) (interface{}, error) + Stream(ctx *agentContext.Context, messages []agentContext.Message, options ...*agentContext.Options) (*agentContext.Response, error) } // AgentGetterFunc is a function type that gets an agent by ID diff --git a/agent/content/tools.go b/agent/content/tools.go index cc9d7787..b4fc6468 100644 --- a/agent/content/tools.go +++ b/agent/content/tools.go @@ -42,10 +42,10 @@ func CallAgent(ctx *agentContext.Context, agentID string, message agentContext.M // Extract text from agent response // Two formats are supported: - // 1. Custom Hook response (from Next hook) - // 2. Standard Agent Stream response (LLM completion) + // 1. Custom Hook response (from Next hook) - response.Next + // 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 @@ -120,116 +120,63 @@ func CallAgentWithFileInfo(ctx *agentContext.Context, agentID string, message ag return CallAgent(ctx, agentID, message) } -// extractTextFromAgentResponse extracts text from agent response -// Handles two main response formats from agent.Stream(): -// -// 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: } -// 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 +// extractTextFromResponse extracts text from agent response +// Now that agent.Stream() returns *agentContext.Response directly, +// we can access fields without type assertions or JSON conversion. // // Priority: -// 1. Check for "next" field (custom hook data) → return complete data -// 2. Check for "completion" field (standard LLM response) → extract text only -// 3. Fallback to direct string or JSON stringify -func extractTextFromAgentResponse(response interface{}) (string, error) { +// 1. Check response.Next (custom hook data) → return complete data +// 2. Check response.Completion (standard LLM response) → extract text only +func extractTextFromResponse(response *agentContext.Response) (string, error) { if response == nil { return "", fmt.Errorf("agent returned nil response") } - // First, try to convert to map if it's a struct - // agent.Stream() may return *agentContext.Response which needs to be converted - var responseMap map[string]interface{} - - // 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 { + // Priority 1: Check Next field (custom hook data) + // If Next hook returns custom data, return the complete structure + if response.Next != nil { // If next is a string, return directly - if nextStr, ok := next.(string); ok { + if nextStr, ok := response.Next.(string); ok { return nextStr, nil } // Otherwise, JSON stringify to preserve complete structure - jsonBytes, err := jsoniter.Marshal(next) + jsonBytes, err := jsoniter.Marshal(response.Next) if err != nil { return "", fmt.Errorf("failed to serialize next hook data: %w", err) } 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 - if completion, hasCompletion := responseMap["completion"]; hasCompletion && 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) - switch v := content.(type) { - case string: - // Simple text content - return v, nil - case []interface{}: - // Multimodal content array - extract all text parts - var text string - for _, part := range v { - if partMap, ok := part.(map[string]interface{}); ok { - if partType, _ := partMap["type"].(string); partType == "text" { - if textContent, ok := partMap["text"].(string); ok { - text += textContent - } - } + if response.Completion != nil { + // Content can be string or []ContentPart (multimodal) + switch v := response.Completion.Content.(type) { + case string: + // Simple text content + return v, nil + case []interface{}: + // Multimodal content array - extract all text parts + var text string + for _, part := range v { + if partMap, ok := part.(map[string]interface{}); ok { + if partType, _ := partMap["type"].(string); partType == "text" { + if textContent, ok := partMap["text"].(string); ok { + text += textContent } } - if text != "" { - return text, nil - } - // No text found in content parts - return "", fmt.Errorf("no text content found in completion content parts") } } + if text != "" { + return text, nil + } + // No text found in 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) - if content, hasContent := responseMap["content"]; hasContent { - 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 + // No content found + return "", fmt.Errorf("no content found in agent response") } // CallMCPTool calls an MCP tool to process content diff --git a/agent/search/handlers/web/agent.go b/agent/search/handlers/web/agent.go index e23c2bc6..0daa84d7 100644 --- a/agent/search/handlers/web/agent.go +++ b/agent/search/handlers/web/agent.go @@ -150,46 +150,20 @@ func (p *AgentProvider) Search(ctx *agentContext.Context, req *types.Request) (* }, nil } -// parseAgentResponse parses the agent response into search result items -// The agent should return a JSON structure with search results -func (p *AgentProvider) parseAgentResponse(response interface{}, source types.SourceType) ([]*types.ResultItem, int, string) { - if response == nil { +// parseAgentResponse parses the agent's *context.Response into search result items +// Now that agent.Stream() returns *context.Response directly, +// we can access fields without type assertions. +// +// 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" } - // Try to extract data from response - var data map[string]interface{} - - // Handle different response types - switch v := response.(type) { - case map[string]interface{}: - data = v - case string: - // Try to parse as JSON - if err := json.Unmarshal([]byte(v), &data); err != nil { - return nil, 0, fmt.Sprintf("Failed to parse agent response as JSON: %v", err) - } - default: - // Try to marshal and unmarshal - jsonBytes, err := json.Marshal(response) - if err != nil { - return nil, 0, fmt.Sprintf("Failed to serialize agent response: %v", err) - } - if err := json.Unmarshal(jsonBytes, &data); err != nil { - return nil, 0, fmt.Sprintf("Failed to parse agent response: %v", err) - } - } - - // Check for "next" field (custom hook data) - if next, hasNext := data["next"]; hasNext && next != nil { - if nextMap, ok := next.(map[string]interface{}); ok { - data = nextMap - } else if nextStr, ok := next.(string); ok { - // Try to parse as JSON - if err := json.Unmarshal([]byte(nextStr), &data); err != nil { - return nil, 0, fmt.Sprintf("Failed to parse next hook data: %v", err) - } - } + // Extract data from Next field + data := extractNextData(response.Next) + if data == nil { + return nil, 0, "Failed to extract data from agent response" } // Extract items from data @@ -230,3 +204,34 @@ func (p *AgentProvider) parseAgentResponse(response interface{}, source types.So 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 +} diff --git a/agent/search/nlp/keyword/agent.go b/agent/search/nlp/keyword/agent.go index 22f72c08..903bdde6 100644 --- a/agent/search/nlp/keyword/agent.go +++ b/agent/search/nlp/keyword/agent.go @@ -64,30 +64,38 @@ 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 { return nil, fmt.Errorf("agent call failed: %w", err) } - // Debug: log the result type and value - // fmt.Printf("DEBUG Agent result type: %T, value: %+v\n", result, result) - - // Parse the result - return p.parseResult(result) + // Parse the result from response.Next + return p.parseResponse(response) } -// parseResult extracts keywords from the agent's response -// The agent should return data in NextHookResponse format: { data: { keywords: [...] } } -// The Stream() response wraps this in: { next: { data: { keywords: [...] } } } -func (p *AgentProvider) parseResult(result interface{}) ([]string, error) { - if result == nil { +// parseResponse extracts keywords from the agent's *context.Response +// Now that agent.Stream() returns *context.Response directly, +// we can access fields without type assertions. +// +// The agent returns keywords in response.Next field +func (p *AgentProvider) parseResponse(response *agentContext.Response) ([]string, error) { + if response == nil || response.Next == nil { + return []string{}, nil + } + + return p.parseNextData(response.Next) +} + +// parseNextData extracts keywords from Next hook data +func (p *AgentProvider) parseNextData(next interface{}) ([]string, error) { + if next == nil { return []string{}, nil } // Try to convert to map first (most common case) var data map[string]interface{} - switch v := result.(type) { + switch v := next.(type) { case map[string]interface{}: data = v case string: @@ -113,7 +121,7 @@ func (p *AgentProvider) parseResult(result interface{}) ([]string, error) { return keywords, nil default: // Try to marshal and unmarshal - jsonBytes, err := json.Marshal(result) + jsonBytes, err := json.Marshal(next) if err != nil { return []string{}, nil } @@ -122,18 +130,6 @@ func (p *AgentProvider) parseResult(result interface{}) ([]string, error) { } } - // 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 - } - } - } - // Extract keywords from data // Try common field names: "keywords", "data", "data.keywords" if kw, ok := data["keywords"]; ok { diff --git a/agent/search/nlp/querydsl/agent.go b/agent/search/nlp/querydsl/agent.go index b4b5531b..85d6aa60 100644 --- a/agent/search/nlp/querydsl/agent.go +++ b/agent/search/nlp/querydsl/agent.go @@ -70,8 +70,8 @@ func (p *AgentProvider) Generate(ctx *agentContext.Context, input *Input) (*Resu continue } - // Parse the result - genResult, err := p.parseResult(result) + // Parse the result from response + genResult, err := p.parseResponse(result) if err != nil { lastError = err continue @@ -155,28 +155,36 @@ func (p *AgentProvider) validateDSL(dsl *gou.QueryDSL) *linter.LintResult { return lintResult } -// parseResult extracts QueryDSL from the agent's response -// The querydsl agent returns QueryDSL JSON directly (not wrapped in {dsl: ...}) +// parseResponse extracts QueryDSL from the agent's *context.Response +// Now that agent.Stream() returns *context.Response directly, +// we can access fields without type assertions. +// +// The querydsl agent returns QueryDSL in response.Next field // Or returns error JSON: {"error": "code", "message": "..."} -// Stream() returns *context.Response with QueryDSL in "next" field -func (p *AgentProvider) parseResult(result interface{}) (*Result, error) { - if result == nil { +func (p *AgentProvider) parseResponse(response *agentContext.Response) (*Result, error) { + if response == nil { return &Result{}, nil } - // Handle *context.Response directly (most common case from Stream()) - if resp, ok := result.(*agentContext.Response); ok { - if resp.Next != nil { - // Next contains the hook response, recursively parse it - return p.parseResult(resp.Next) - } + // 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{} - switch v := result.(type) { + switch v := next.(type) { case map[string]interface{}: data = v case string: @@ -186,7 +194,7 @@ func (p *AgentProvider) parseResult(result interface{}) (*Result, error) { } default: // Try to marshal and unmarshal - jsonBytes, err := json.Marshal(result) + jsonBytes, err := json.Marshal(next) if err != nil { return &Result{}, nil } @@ -197,18 +205,6 @@ func (p *AgentProvider) parseResult(result interface{}) (*Result, error) { genResult := &Result{} - // Check for Stream() wrapper: { content: "...", next: {...} } - // The actual response is in "content" field as a string - if content, hasContent := data["content"]; hasContent && content != nil { - if contentStr, ok := content.(string); ok && contentStr != "" { - // Parse the content string as JSON - var contentData map[string]interface{} - if err := json.Unmarshal([]byte(contentStr), &contentData); err == nil { - data = contentData - } - } - } - // Check for error response: {"error": "code", "message": "..."} if errCode, hasError := data["error"]; hasError { errMsg := "" @@ -229,22 +225,6 @@ func (p *AgentProvider) parseResult(result interface{}) (*Result, error) { return genResult, nil } - // Fallback: check for wrapped formats - // Check for "next" field (custom hook data from NextHookResponse) - if next, hasNext := data["next"]; hasNext && next != nil { - if nextMap, ok := next.(map[string]interface{}); ok { - data = nextMap - } else if nextStr, ok := next.(string); ok { - if err := json.Unmarshal([]byte(nextStr), &data); err == nil { - // Check if parsed data is a QueryDSL - if _, hasFrom := data["from"]; hasFrom { - genResult.DSL = p.extractDSL(data) - return genResult, nil - } - } - } - } - // Check for "dsl" field wrapper: { dsl: {...} } if dsl, ok := data["dsl"]; ok { genResult.DSL = p.extractDSL(dsl) diff --git a/agent/search/rerank/agent.go b/agent/search/rerank/agent.go index 0a5137e6..994e50dd 100644 --- a/agent/search/rerank/agent.go +++ b/agent/search/rerank/agent.go @@ -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 { return nil, fmt.Errorf("agent stream failed: %w", err) } - // Parse response - return p.parseResponse(result, items, opts) + // Parse response from response.Next + return p.parseAgentResponse(response, items, opts) } -// parseResponse extracts reranked items from agent response -// The response format from agent.Stream is typically: -// { "next": { "data": { "order": [...] } } } -func (p *AgentProvider) parseResponse(result interface{}, originalItems []*types.ResultItem, opts *types.RerankOptions) ([]*types.ResultItem, error) { - if result == nil { +// parseAgentResponse extracts reranked items from agent's *context.Response +// Now that agent.Stream() returns *context.Response directly, +// we can access fields without type assertions. +// +// 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 } @@ -84,20 +88,20 @@ func (p *AgentProvider) parseResponse(result interface{}, originalItems []*types } } - // Extract response data - response := extractResponseData(result) - if response == nil { + // Extract response data from Next field + data := extractNextData(response.Next) + if data == 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"] } // Or: { "items": [{ "citation_id": "ref_001", ... }, ...] } var reranked []*types.ResultItem // 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 { for _, id := range orderList { 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) if len(reranked) == 0 { - if items, ok := response["items"]; ok { + if items, ok := data["items"]; ok { if itemsList := toItemsList(items); len(itemsList) > 0 { for _, respItem := range itemsList { // 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 } -// extractResponseData extracts the actual response data from agent.Stream result -// Handles nested structures like { "next": { "data": { ... } } } -func extractResponseData(result interface{}) map[string]interface{} { - switch v := result.(type) { +// 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 "next" wrapper (from NextHookResponse) - 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 + // Check for "data" wrapper if data, ok := v["data"].(map[string]interface{}); ok { return data } @@ -172,16 +172,14 @@ func extractResponseData(result interface{}) map[string]interface{} { // Try to parse as JSON var data map[string]interface{} 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 - if result != nil { - if bytes, err := json.Marshal(result); err == nil { - var data map[string]interface{} - if err := json.Unmarshal(bytes, &data); err == nil { - return extractResponseData(data) - } + 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 diff --git a/agent/test/runner.go b/agent/test/runner.go index dedc00b2..8065c81c 100644 --- a/agent/test/runner.go +++ b/agent/test/runner.go @@ -532,26 +532,23 @@ func buildContextOptions(tc *Case, runnerOpts *Options) *context.Options { } // extractOutput extracts the output from the agent response -// Priority: Next hook data (if non-empty) > Completion content > raw 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 { return nil } - // Try to get data from context.Response - if resp, ok := response.(*context.Response); ok { - // Prefer Next hook data if available and non-empty - // resp.Next is already the Data value (not NextHookResponse struct) - if resp.Next != nil && !isEmptyValue(resp.Next) { - return resp.Next - } - // Fall back to raw completion content - if resp.Completion != nil { - return resp.Completion.Content - } + // Prefer Next hook data if available and non-empty + // response.Next is already the Data value (not NextHookResponse struct) + if response.Next != nil && !isEmptyValue(response.Next) { + return response.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