From 31ce461162e20cb6166caeca849a2e5c8804da0c Mon Sep 17 00:00:00 2001 From: Max Date: Sat, 27 Dec 2025 19:36:31 +0800 Subject: [PATCH] Refactor Assistant Stream and Content Handling for Improved Context Management - Updated the Stream method to include options for handling message history and context more effectively, ensuring original messages are preserved for autoSearch and delegation. - Introduced a new buildContextMessage function to consolidate conversation context, filtering out system messages and limiting to the last five user messages for efficiency. - Enhanced content processing by adding a convertToContentParts function to handle different content formats, improving compatibility with historical data. - Improved logging and error handling in the executeLLMStream method to ensure clarity in LLM request tracing and response handling. - Added new utility functions for extracting text content and building context messages, enhancing overall code clarity and maintainability. --- Makefile | 4 +- agent/assistant/agent.go | 15 +-- agent/assistant/build_content.go | 3 +- agent/assistant/llm.go | 34 +++++-- agent/assistant/search.go | 158 +++++++++++++++++++++++++++++-- agent/content/content.go | 108 ++++++++++++++++++++- 6 files changed, 290 insertions(+), 32 deletions(-) diff --git a/Makefile b/Makefile index 9d56165e..b1cba35d 100644 --- a/Makefile +++ b/Makefile @@ -13,8 +13,8 @@ OS := $(shell uname) TESTFOLDER := $(shell $(GO) list ./... | grep -vE 'examples|openai|aigc|neo|twilio|share*' | awk '!/\/tests\// || /openapi\/tests/') # Core tests (exclude AI-related: agent, aigc, openai, and KB) TESTFOLDER_CORE := $(shell $(GO) list ./... | grep -vE 'examples|openai|aigc|neo|twilio|share*|agent|kb' | awk '!/\/tests\// || /openapi\/tests/') -# AI tests (agent, aigc) -TESTFOLDER_AI := $(shell $(GO) list ./agent/... ./aigc/...) +# AI tests (agent, aigc) - exclude agent/search/handlers/web (requires external API keys) +TESTFOLDER_AI := $(shell $(GO) list ./agent/... ./aigc/... | grep -v 'agent/search/handlers/web') # KB tests (kb) TESTFOLDER_KB := $(shell $(GO) list ./kb/...) TESTTAGS ?= "" diff --git a/agent/assistant/agent.go b/agent/assistant/agent.go index 6803da38..4dd3acda 100644 --- a/agent/assistant/agent.go +++ b/agent/assistant/agent.go @@ -123,7 +123,7 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa // Get Full Messages with chat history // ================================================ ctx.Logger.Phase("History") - historyResult, err := ast.WithHistory(ctx, inputMessages, agentNode) + historyResult, err := ast.WithHistory(ctx, inputMessages, agentNode, opts) if err != nil { ast.traceAgentFail(agentNode, err) ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err) @@ -213,6 +213,9 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa ctx.Logger.Phase("LLM") // Build the LLM request first (use fullMessages which includes history) + // Note: completionMessages here are still in original format (with __yao.attachment:// URLs) + // Content conversion (BuildContent) happens inside executeLLMStream, right before LLM call + // This ensures autoSearch and delegate receive original messages, not converted ones completionMessages, completionOptions, err = ast.BuildRequest(ctx, fullMessages, createResponse) if err != nil { finalStatus = context.ResumeStatusFailed @@ -223,16 +226,6 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa return nil, err } - // Build content - convert extended types (file, data) to standard LLM types (text, image_url, input_audio) - completionMessages, err = ast.BuildContent(ctx, completionMessages, completionOptions, opts) - if err != nil { - finalStatus = context.ResumeStatusFailed - finalError = err - ast.traceAgentFail(agentNode, err) - ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err) - return nil, err - } - // ================================================ // Execute Auto Search (if enabled) // ================================================ diff --git a/agent/assistant/build_content.go b/agent/assistant/build_content.go index 9a23eb7e..ce2f939c 100644 --- a/agent/assistant/build_content.go +++ b/agent/assistant/build_content.go @@ -19,10 +19,11 @@ func (ast *Assistant) BuildContent(ctx *context.Context, messages []context.Mess } // Get connector and capabilities - _, capabilities, err := ast.GetConnector(ctx, opts) + connector, capabilities, err := ast.GetConnector(ctx, opts) if err != nil { return nil, fmt.Errorf("failed to get connector: %w", err) } + _ = connector // unused but needed for GetConnector call // Get Uses configuration from options (already merged in BuildRequest) uses := options.Uses diff --git a/agent/assistant/llm.go b/agent/assistant/llm.go index 415f3ee7..4c34f6ff 100644 --- a/agent/assistant/llm.go +++ b/agent/assistant/llm.go @@ -33,11 +33,22 @@ func (ast *Assistant) executeLLMStream( // Log the capabilities ast.traceConnectorCapabilities(agentNode, capabilities) - // Trace Add LLM request - ast.traceLLMRequest(ctx, conn.ID(), completionMessages, completionOptions) + // Build content - convert extended types (file, data, __yao.attachment://) to standard LLM types + // This is done here (right before LLM call) to ensure: + // 1. autoSearch receives original messages (not converted) + // 2. delegate receives original messages (not converted) + // 3. Only the actual LLM call sees converted messages + llmMessages, err := ast.BuildContent(ctx, completionMessages, completionOptions, opts) + if err != nil { + ast.traceAgentFail(agentNode, err) + return nil, err + } + + // Trace Add LLM request (use converted messages for trace) + ast.traceLLMRequest(ctx, conn.ID(), llmMessages, completionOptions) // Log LLM call start - ctx.Logger.LLMStart(conn.ID(), "", len(completionMessages)) + ctx.Logger.LLMStart(conn.ID(), "", len(llmMessages)) // Create LLM instance with connector and options llmInstance, err := llm.New(conn, completionOptions) @@ -48,7 +59,8 @@ func (ast *Assistant) executeLLMStream( } // Call the LLM Completion Stream (streamHandler was set earlier) - completionResponse, err := llmInstance.Stream(ctx, completionMessages, completionOptions, streamHandler) + // Use llmMessages (converted) instead of completionMessages (original) + completionResponse, err := llmInstance.Stream(ctx, llmMessages, completionOptions, streamHandler) if err != nil { // Mark LLM Request as failed in trace @@ -86,11 +98,18 @@ func (ast *Assistant) executeLLMForToolRetry( completionOptions.Capabilities = capabilities } + // Build content - convert extended types for LLM call + llmMessages, err := ast.BuildContent(ctx, completionMessages, completionOptions, opts) + if err != nil { + ast.traceAgentFail(agentNode, err) + return nil, err + } + // Trace Add LLM retry request - ast.traceLLMRetryRequest(ctx, conn.ID(), completionMessages, completionOptions) + ast.traceLLMRetryRequest(ctx, conn.ID(), llmMessages, completionOptions) // Log LLM call start (retry) - ctx.Logger.LLMStart(conn.ID(), "", len(completionMessages)) + ctx.Logger.LLMStart(conn.ID(), "", len(llmMessages)) // Create LLM instance with connector and options llmInstance, err := llm.New(conn, completionOptions) @@ -101,7 +120,8 @@ func (ast *Assistant) executeLLMForToolRetry( } // Call the LLM Completion Stream (still streaming for tool retry) - completionResponse, err := llmInstance.Stream(ctx, completionMessages, completionOptions, streamHandler) + // Use llmMessages (converted) instead of completionMessages (original) + completionResponse, err := llmInstance.Stream(ctx, llmMessages, completionOptions, streamHandler) if err != nil { // Mark LLM Retry Request as failed in trace ast.traceLLMFail(ctx, err) diff --git a/agent/assistant/search.go b/agent/assistant/search.go index 764baf14..af6e2982 100644 --- a/agent/assistant/search.go +++ b/agent/assistant/search.go @@ -146,14 +146,8 @@ func (ast *Assistant) checkSearchIntent(ctx *context.Context, messages []context Confidence: 0, } - // 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) - } - } - + // Build a single text message with conversation context + intentMessages := buildContextMessage(messages) if len(intentMessages) == 0 { return defaultIntent // No messages, skip search } @@ -435,7 +429,16 @@ func (ast *Assistant) executeAutoSearch(ctx *context.Context, messages []context ctx.Logger.Info("No query found in messages, skipping auto search") return nil } + + // Build query with conversation context for better keyword extraction + // This helps the keyword extractor understand the full context + contextMessages := buildContextMessage(messages) query := originalQuery + if len(contextMessages) > 0 { + if contextStr, ok := contextMessages[0].Content.(string); ok { + query = contextStr + } + } // Check if keyword extraction should be skipped skipKeyword := false @@ -467,7 +470,9 @@ func (ast *Assistant) executeAutoSearch(ctx *context.Context, messages []context searchNode := ast.createSearchTrace(ctx, query, requests) // Execute searches in parallel - ctx.Logger.Info("Executing %d search requests for query: %s", len(requests), truncateString(query, 50)) + // Build provider info for logging + providerInfo := ast.getSearchProviderInfo(searchConfig, searchUses) + ctx.Logger.Info("Executing %d search requests via %s for query: %s", len(requests), providerInfo, truncateString(query, 50)) startTime := time.Now() results, err := searcher.All(ctx, requests) @@ -899,6 +904,105 @@ func (ast *Assistant) injectSearchContext(messages []context.Message, refCtx *se return result } +// extractTextContent extracts text-only content from a message +// For multimodal messages, concatenates all text parts +// Returns empty string if no text content found +func extractTextContent(msg context.Message) string { + content := msg.Content + // Handle string content + if str, ok := content.(string); ok { + return str + } + // Handle content parts (array of objects) - extract only text parts + if parts, ok := content.([]interface{}); ok { + var texts []string + for _, part := range parts { + if partMap, ok := part.(map[string]interface{}); ok { + if partMap["type"] == "text" { + if text, ok := partMap["text"].(string); ok { + texts = append(texts, text) + } + } + } + } + if len(texts) > 0 { + return strings.Join(texts, "\n") + } + } + // Handle []context.ContentPart + if parts, ok := content.([]context.ContentPart); ok { + var texts []string + for _, part := range parts { + if part.Type == context.ContentText && part.Text != "" { + texts = append(texts, part.Text) + } + } + if len(texts) > 0 { + return strings.Join(texts, "\n") + } + } + return "" +} + +// buildContextMessage builds a single user message with conversation context +// Filters out system messages and extracts text-only content +// Only takes the last 5 messages for efficiency +// Returns a slice with one message containing the full context, or empty slice if no content +func buildContextMessage(messages []context.Message) []context.Message { + const maxMessages = 5 + + // Take only the last maxMessages (excluding system messages) + var recentMessages []context.Message + for i := len(messages) - 1; i >= 0 && len(recentMessages) < maxMessages; i-- { + if messages[i].Role != "system" { + recentMessages = append(recentMessages, messages[i]) + } + } + // Reverse to maintain chronological order + for i, j := 0, len(recentMessages)-1; i < j; i, j = i+1, j-1 { + recentMessages[i], recentMessages[j] = recentMessages[j], recentMessages[i] + } + + var contextParts []string + var lastUserMessage string + + for _, msg := range recentMessages { + textContent := extractTextContent(msg) + if textContent == "" { + continue + } + + // Format message with role label + switch msg.Role { + case "user": + contextParts = append(contextParts, "[User]: "+textContent) + lastUserMessage = textContent + case "assistant": + contextParts = append(contextParts, "[Assistant]: "+textContent) + default: + contextParts = append(contextParts, "["+string(msg.Role)+"]: "+textContent) + } + } + + // Build single message with context + var result []context.Message + if len(contextParts) > 1 { + // Multiple messages: include conversation context + fullContext := "=== Conversation Context ===\n" + strings.Join(contextParts, "\n\n") + "\n=== End Context ===\n\nCurrent user request: " + lastUserMessage + result = append(result, context.Message{ + Role: "user", + Content: fullContext, + }) + } else if lastUserMessage != "" { + // Single user message: just use it directly + result = append(result, context.Message{ + Role: "user", + Content: lastUserMessage, + }) + } + return result +} + // extractQueryFromMessages extracts the search query from messages // Uses the last user message as the query func extractQueryFromMessages(messages []context.Message) string { @@ -1129,3 +1233,39 @@ func (ast *Assistant) configToMap(config *searchTypes.Config) map[string]any { return result } + +// getSearchProviderInfo returns a human-readable string describing the search provider(s) +func (ast *Assistant) getSearchProviderInfo(config *searchTypes.Config, uses *search.Uses) string { + var parts []string + + // Web search provider - always show when web search is being executed + webMode := "" + if uses != nil { + webMode = uses.Web + } + + if webMode == "" || webMode == "builtin" { + // Builtin mode: show the actual provider (tavily/serper/serpapi) + provider := "tavily" // default + if config != nil && config.Web != nil && config.Web.Provider != "" { + provider = config.Web.Provider + } + parts = append(parts, "web:"+provider) + } else if strings.HasPrefix(webMode, "mcp:") { + parts = append(parts, "web:"+webMode) + } else { + parts = append(parts, "web:agent:"+webMode) + } + + // KB search + if config != nil && config.KB != nil && len(config.KB.Collections) > 0 { + parts = append(parts, "kb") + } + + // DB search + if config != nil && config.DB != nil && len(config.DB.Models) > 0 { + parts = append(parts, "db") + } + + return strings.Join(parts, ", ") +} diff --git a/agent/content/content.go b/agent/content/content.go index 6dc6daf2..4686b134 100644 --- a/agent/content/content.go +++ b/agent/content/content.go @@ -99,10 +99,14 @@ func processMessage( return *msg, nil } - // Get content parts + // Get content parts - try typed first, then convert from interface{} parts, ok := msg.GetContentAsParts() if !ok { - return *msg, nil + // Try to convert from []interface{} (common when loaded from history/JSON) + parts, ok = convertToContentParts(msg.Content) + if !ok { + return *msg, nil + } } // Note: File information will be collected and stored in Space by CallAgentWithFileInfo @@ -293,6 +297,28 @@ func processImageURLContent( }, nil } + // Check model capabilities + supportsVision := false + if capabilities != nil && capabilities.Vision != nil { + // Vision can be bool or string (format) + switch v := capabilities.Vision.(type) { + case bool: + supportsVision = v + case string: + supportsVision = v != "" && v != "false" && v != "none" + } + } + + // If model supports vision AND we're not forcing uses, pass through + if supportsVision && !forceUses { + return &Result{ + ContentPart: part, + }, nil + } + + // Model doesn't support vision OR forceUses is true + // Need to convert image to text + // If it's uploader wrapper or HTTP URL, need to process // Check cache first cachedText, found, err := tryGetCachedText(ctx, url, processedFiles) @@ -441,6 +467,7 @@ func getToolForProcessing(uses *agentContext.Uses, fileType FileType) string { func tryGetCachedText(ctx *agentContext.Context, url string, processedFiles map[string]string) (string, bool, error) { // Parse URL to check if it's an uploader wrapper uploaderName, fileID, isWrapper := attachment.Parse(url) + if !isWrapper { return "", false, nil // Not an uploader wrapper, no cache } @@ -485,3 +512,80 @@ func cacheProcessedText(ctx *agentContext.Context, url string, text string, proc return nil } + +// convertToContentParts converts []interface{} to []ContentPart +// This is needed when content is loaded from JSON/history and is []interface{} instead of []ContentPart +func convertToContentParts(content interface{}) ([]agentContext.ContentPart, bool) { + // Check if it's []interface{} + arr, ok := content.([]interface{}) + if !ok { + return nil, false + } + + parts := make([]agentContext.ContentPart, 0, len(arr)) + for _, item := range arr { + // Each item should be a map + m, ok := item.(map[string]interface{}) + if !ok { + continue + } + + // Get type field + typeStr, _ := m["type"].(string) + if typeStr == "" { + continue + } + + part := agentContext.ContentPart{ + Type: agentContext.ContentPartType(typeStr), + } + + switch typeStr { + case "text": + if text, ok := m["text"].(string); ok { + part.Text = text + } + + case "image_url": + if imgData, ok := m["image_url"].(map[string]interface{}); ok { + part.ImageURL = &agentContext.ImageURL{} + if url, ok := imgData["url"].(string); ok { + part.ImageURL.URL = url + } + if detail, ok := imgData["detail"].(string); ok { + part.ImageURL.Detail = agentContext.ImageDetailLevel(detail) + } + } + + case "file": + if fileData, ok := m["file"].(map[string]interface{}); ok { + part.File = &agentContext.FileAttachment{} + if url, ok := fileData["url"].(string); ok { + part.File.URL = url + } + if filename, ok := fileData["filename"].(string); ok { + part.File.Filename = filename + } + } + + case "input_audio": + if audioData, ok := m["input_audio"].(map[string]interface{}); ok { + part.InputAudio = &agentContext.InputAudio{} + if data, ok := audioData["data"].(string); ok { + part.InputAudio.Data = data + } + if format, ok := audioData["format"].(string); ok { + part.InputAudio.Format = format + } + } + } + + parts = append(parts, part) + } + + if len(parts) == 0 { + return nil, false + } + + return parts, true +}