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.
This commit is contained in:
Max 2025-12-27 19:36:31 +08:00
parent 401f22eeb1
commit 31ce461162
6 changed files with 290 additions and 32 deletions

View file

@ -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 ?= ""

View file

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

View file

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

View file

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

View file

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

View file

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