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:
parent
401f22eeb1
commit
31ce461162
6 changed files with 290 additions and 32 deletions
4
Makefile
4
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 ?= ""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
// ================================================
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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, ", ")
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue