- Remove the `LoadModelCapabilities` test and associated model capabilities initialization from the agent, streamlining the loading process. - Update the LLM provider implementations to utilize a unified `Capabilities` structure, replacing references to `openai.Capabilities` with `llm.Capabilities`. - Enhance capability retrieval methods to simplify the extraction of connector capabilities, ensuring compatibility across different LLM providers. - Clean up unused functions and variables related to model capabilities, improving code maintainability.
602 lines
18 KiB
Go
602 lines
18 KiB
Go
package context
|
|
|
|
import (
|
|
"time"
|
|
|
|
"github.com/yaoapp/gou/llm"
|
|
"github.com/yaoapp/yao/agent/output"
|
|
"github.com/yaoapp/yao/agent/output/message"
|
|
)
|
|
|
|
// Send sends a message via the output module
|
|
// Automatically manages BlockID, ThreadID, lifecycle events, and metadata for delta operations
|
|
// - For delta operations: inherits BlockID and ThreadID from original message, increments chunk count
|
|
// - For new messages: auto-sets ThreadID from Stack, sends message_start event
|
|
// - Sends block_start event when a new BlockID is first encountered
|
|
// - Records metadata for all sent messages to enable delta inheritance
|
|
func (ctx *Context) Send(msg *message.Message) error {
|
|
// Call OnMessage callback if provided (for ctx.agent.Call with onChunk)
|
|
if ctx.Stack != nil && ctx.Stack.Options != nil && ctx.Stack.Options.OnMessage != nil {
|
|
if ret := ctx.Stack.Options.OnMessage(msg); ret != 0 {
|
|
return nil // Callback requested stop
|
|
}
|
|
}
|
|
|
|
out, err := ctx.getOutput()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
|
|
// Skip lifecycle events for event-type messages (prevent recursion)
|
|
isEventMessage := msg.Type == message.TypeEvent
|
|
|
|
// === Handle message_start event: record metadata for future delta chunks ===
|
|
if isEventMessage && msg.Props != nil {
|
|
if event, ok := msg.Props["event"].(string); ok && event == message.EventMessageStart {
|
|
if data, ok := msg.Props["data"].(message.EventMessageStartData); ok {
|
|
// Record metadata from message_start event
|
|
if data.MessageID != "" && ctx.messageMetadata != nil {
|
|
ctx.messageMetadata.setMessage(data.MessageID, &MessageMetadata{
|
|
MessageID: data.MessageID,
|
|
ThreadID: data.ThreadID,
|
|
Type: data.Type,
|
|
StartTime: time.Now(),
|
|
ChunkCount: 0, // Will be incremented by delta chunks
|
|
})
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// === Delta operations: Auto-inherit and update metadata ===
|
|
if msg.Delta && msg.MessageID != "" && ctx.messageMetadata != nil {
|
|
if metadata := ctx.getMessageMetadata(msg.MessageID); metadata != nil {
|
|
// Inherit BlockID if not specified
|
|
if msg.BlockID == "" {
|
|
msg.BlockID = metadata.BlockID
|
|
}
|
|
// Inherit ThreadID if not specified
|
|
if msg.ThreadID == "" {
|
|
msg.ThreadID = metadata.ThreadID
|
|
}
|
|
|
|
// Increment chunk count for this message
|
|
metadata.ChunkCount++
|
|
|
|
// Update Buffer content for streaming messages (for storage)
|
|
if ctx.Buffer != nil && msg.Props != nil {
|
|
if content, ok := msg.Props["content"].(string); ok {
|
|
ctx.Buffer.AppendMessageContent(msg.MessageID, content)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// === Auto-generate ChunkID (always) ===
|
|
if msg.ChunkID == "" && !isEventMessage {
|
|
if ctx.IDGenerator != nil {
|
|
msg.ChunkID = ctx.IDGenerator.GenerateChunkID()
|
|
} else {
|
|
msg.ChunkID = message.GenerateNanoID()
|
|
}
|
|
}
|
|
|
|
// === Non-delta operations: New message logic ===
|
|
if !msg.Delta && !isEventMessage {
|
|
// Auto-set ThreadID for non-root Stack (nested agent calls)
|
|
if msg.ThreadID == "" && ctx.Stack != nil && !ctx.Stack.IsRoot() {
|
|
msg.ThreadID = ctx.Stack.ID
|
|
}
|
|
|
|
// BlockID is NOT auto-generated by default (only manually specified in special cases)
|
|
// Example: Send a web card after LLM output, group them in the same Block
|
|
// Developers can specify via ctx.Send(message, blockId) or message.block_id
|
|
|
|
// === Send block_start event if this is a new block ===
|
|
if msg.BlockID != "" && ctx.messageMetadata != nil {
|
|
if ctx.messageMetadata.getBlock(msg.BlockID) == nil {
|
|
// New block, send block_start event
|
|
blockStartData := message.EventBlockStartData{
|
|
BlockID: msg.BlockID,
|
|
Type: "mixed", // Default type, can be enhanced later
|
|
Timestamp: time.Now().UnixMilli(),
|
|
}
|
|
blockStartEvent := output.NewEventMessage(message.EventBlockStart, "Block started", blockStartData)
|
|
if err := ctx.sendRaw(blockStartEvent); err != nil {
|
|
return err
|
|
}
|
|
|
|
// Record block metadata
|
|
ctx.messageMetadata.setBlock(msg.BlockID, &BlockMetadata{
|
|
BlockID: msg.BlockID,
|
|
Type: "mixed",
|
|
StartTime: time.Now(),
|
|
MessageCount: 0,
|
|
})
|
|
}
|
|
|
|
// Increment message count for this block
|
|
ctx.messageMetadata.updateBlock(msg.BlockID, func(block *BlockMetadata) {
|
|
block.MessageCount++
|
|
})
|
|
}
|
|
|
|
// === Generate MessageID if not provided ===
|
|
if msg.MessageID == "" {
|
|
if ctx.IDGenerator != nil {
|
|
msg.MessageID = ctx.IDGenerator.GenerateMessageID()
|
|
} else {
|
|
msg.MessageID = message.GenerateNanoID() // Use NanoID generator
|
|
}
|
|
}
|
|
|
|
// === Send message_start event ===
|
|
messageStartData := message.EventMessageStartData{
|
|
MessageID: msg.MessageID,
|
|
Type: msg.Type,
|
|
Timestamp: time.Now().UnixMilli(),
|
|
ThreadID: msg.ThreadID, // Include ThreadID for concurrent stream identification
|
|
}
|
|
messageStartEvent := output.NewEventMessage(message.EventMessageStart, "Message started", messageStartData)
|
|
if err := ctx.sendRaw(messageStartEvent); err != nil {
|
|
return err
|
|
}
|
|
|
|
// === Record message metadata with start time ===
|
|
if ctx.messageMetadata != nil {
|
|
ctx.messageMetadata.setMessage(msg.MessageID, &MessageMetadata{
|
|
MessageID: msg.MessageID,
|
|
BlockID: msg.BlockID,
|
|
ThreadID: msg.ThreadID,
|
|
Type: msg.Type,
|
|
StartTime: time.Now(),
|
|
ChunkCount: 1, // Initial chunk
|
|
})
|
|
}
|
|
}
|
|
|
|
// === Actually send the message ===
|
|
if err := out.Send(msg); err != nil {
|
|
return err
|
|
}
|
|
|
|
// === Buffer message for batch saving (non-delta, non-event messages only) ===
|
|
// Delta messages are streaming chunks; only final content should be saved
|
|
// Event messages are transient lifecycle signals, not stored
|
|
// Skip if History is disabled in options
|
|
if !msg.Delta && !isEventMessage && ctx.Buffer != nil && !ctx.shouldSkipHistory() {
|
|
assistantID := ""
|
|
if ctx.Stack != nil {
|
|
assistantID = ctx.Stack.AssistantID
|
|
}
|
|
ctx.Buffer.AddAssistantMessage(
|
|
msg.MessageID, // Use the same MessageID as sent to client
|
|
msg.Type,
|
|
msg.Props,
|
|
msg.BlockID,
|
|
msg.ThreadID,
|
|
assistantID,
|
|
nil, // metadata can be added if needed
|
|
)
|
|
}
|
|
|
|
// === Auto-send message_end for non-delta messages (complete messages) ===
|
|
if !msg.Delta && !isEventMessage && msg.MessageID != "" && ctx.messageMetadata != nil {
|
|
metadata := ctx.messageMetadata.getMessage(msg.MessageID)
|
|
if metadata != nil {
|
|
// Calculate duration
|
|
durationMs := time.Since(metadata.StartTime).Milliseconds()
|
|
|
|
// Extract content for the extra field
|
|
var content interface{}
|
|
if msg.Props != nil {
|
|
if c, ok := msg.Props["content"]; ok {
|
|
content = c
|
|
}
|
|
}
|
|
|
|
// Build message_end event data
|
|
endData := message.EventMessageEndData{
|
|
MessageID: msg.MessageID,
|
|
Type: msg.Type,
|
|
Timestamp: time.Now().UnixMilli(),
|
|
ThreadID: metadata.ThreadID, // Include ThreadID for concurrent stream identification
|
|
DurationMs: durationMs,
|
|
ChunkCount: metadata.ChunkCount,
|
|
Status: "completed",
|
|
}
|
|
|
|
// Add content to extra if available
|
|
if content != nil {
|
|
endData.Extra = map[string]interface{}{
|
|
"content": content,
|
|
}
|
|
}
|
|
|
|
// Send message_end event
|
|
messageEndEvent := output.NewEventMessage(message.EventMessageEnd, "Message completed", endData)
|
|
ctx.sendRaw(messageEndEvent)
|
|
}
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
// SendStream sends a streaming message that can be appended to later
|
|
// Unlike Send(), this does NOT automatically send message_end event
|
|
// Use ctx.Append() to add content, then ctx.End() to finalize
|
|
// Returns the message ID for use with Append/End
|
|
func (ctx *Context) SendStream(msg *message.Message) (string, error) {
|
|
out, err := ctx.getOutput()
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
|
|
// Skip lifecycle events for event-type messages
|
|
isEventMessage := msg.Type == message.TypeEvent
|
|
if isEventMessage {
|
|
// Event messages should use Send(), not SendStream()
|
|
return "", ctx.Send(msg)
|
|
}
|
|
|
|
// === Auto-generate ChunkID ===
|
|
if msg.ChunkID == "" {
|
|
if ctx.IDGenerator != nil {
|
|
msg.ChunkID = ctx.IDGenerator.GenerateChunkID()
|
|
} else {
|
|
msg.ChunkID = message.GenerateNanoID()
|
|
}
|
|
}
|
|
|
|
// === Auto-set ThreadID for non-root Stack ===
|
|
if msg.ThreadID == "" && ctx.Stack != nil && !ctx.Stack.IsRoot() {
|
|
msg.ThreadID = ctx.Stack.ID
|
|
}
|
|
|
|
// === Handle BlockID and block_start event ===
|
|
if msg.BlockID != "" && ctx.messageMetadata != nil {
|
|
if ctx.messageMetadata.getBlock(msg.BlockID) == nil {
|
|
blockStartData := message.EventBlockStartData{
|
|
BlockID: msg.BlockID,
|
|
Type: "mixed",
|
|
Timestamp: time.Now().UnixMilli(),
|
|
}
|
|
blockStartEvent := output.NewEventMessage(message.EventBlockStart, "Block started", blockStartData)
|
|
if err := ctx.sendRaw(blockStartEvent); err != nil {
|
|
return "", err
|
|
}
|
|
ctx.messageMetadata.setBlock(msg.BlockID, &BlockMetadata{
|
|
BlockID: msg.BlockID,
|
|
Type: "mixed",
|
|
StartTime: time.Now(),
|
|
MessageCount: 0,
|
|
})
|
|
}
|
|
ctx.messageMetadata.updateBlock(msg.BlockID, func(block *BlockMetadata) {
|
|
block.MessageCount++
|
|
})
|
|
}
|
|
|
|
// === Generate MessageID if not provided ===
|
|
if msg.MessageID == "" {
|
|
if ctx.IDGenerator != nil {
|
|
msg.MessageID = ctx.IDGenerator.GenerateMessageID()
|
|
} else {
|
|
msg.MessageID = message.GenerateNanoID()
|
|
}
|
|
}
|
|
|
|
// === Send message_start event ===
|
|
messageStartData := message.EventMessageStartData{
|
|
MessageID: msg.MessageID,
|
|
Type: msg.Type,
|
|
Timestamp: time.Now().UnixMilli(),
|
|
ThreadID: msg.ThreadID,
|
|
}
|
|
messageStartEvent := output.NewEventMessage(message.EventMessageStart, "Message started", messageStartData)
|
|
if err := ctx.sendRaw(messageStartEvent); err != nil {
|
|
return "", err
|
|
}
|
|
|
|
// === Record message metadata ===
|
|
if ctx.messageMetadata != nil {
|
|
ctx.messageMetadata.setMessage(msg.MessageID, &MessageMetadata{
|
|
MessageID: msg.MessageID,
|
|
BlockID: msg.BlockID,
|
|
ThreadID: msg.ThreadID,
|
|
Type: msg.Type,
|
|
StartTime: time.Now(),
|
|
ChunkCount: 1,
|
|
})
|
|
}
|
|
|
|
// === Actually send the message ===
|
|
if err := out.Send(msg); err != nil {
|
|
return "", err
|
|
}
|
|
|
|
// === Buffer streaming message (will be completed by End()) ===
|
|
if ctx.Buffer != nil && !ctx.shouldSkipHistory() {
|
|
assistantID := ""
|
|
if ctx.Stack != nil {
|
|
assistantID = ctx.Stack.AssistantID
|
|
}
|
|
ctx.Buffer.AddStreamingMessage(
|
|
msg.MessageID,
|
|
msg.Type,
|
|
msg.Props,
|
|
msg.BlockID,
|
|
msg.ThreadID,
|
|
assistantID,
|
|
nil,
|
|
)
|
|
}
|
|
|
|
// NOTE: No message_end event here - will be sent by End()
|
|
return msg.MessageID, nil
|
|
}
|
|
|
|
// End finalizes a streaming message started with SendStream
|
|
// Optionally appends final content before sending message_end event
|
|
// This also saves the complete message to the buffer for storage
|
|
func (ctx *Context) End(messageID string, finalContent ...string) error {
|
|
if messageID == "" {
|
|
return nil
|
|
}
|
|
|
|
// Append final content if provided
|
|
if len(finalContent) > 0 && finalContent[0] != "" {
|
|
// Create a delta message for the final content
|
|
deltaMsg := &message.Message{
|
|
MessageID: messageID,
|
|
Type: message.TypeText,
|
|
Delta: true,
|
|
DeltaAction: message.DeltaAppend,
|
|
Props: map[string]interface{}{
|
|
"content": finalContent[0],
|
|
},
|
|
}
|
|
if err := ctx.Send(deltaMsg); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
|
|
// Get complete content from buffer
|
|
var completeContent string
|
|
if ctx.Buffer != nil {
|
|
completeContent, _ = ctx.Buffer.CompleteStreamingMessage(messageID)
|
|
}
|
|
|
|
// Get metadata for duration calculation
|
|
var durationMs int64
|
|
var threadID string
|
|
var chunkCount int
|
|
var msgType string = message.TypeText
|
|
|
|
if ctx.messageMetadata != nil {
|
|
if metadata := ctx.messageMetadata.getMessage(messageID); metadata != nil {
|
|
durationMs = time.Since(metadata.StartTime).Milliseconds()
|
|
threadID = metadata.ThreadID
|
|
chunkCount = metadata.ChunkCount
|
|
msgType = metadata.Type
|
|
}
|
|
}
|
|
|
|
// Build message_end event data
|
|
endData := message.EventMessageEndData{
|
|
MessageID: messageID,
|
|
Type: msgType,
|
|
Timestamp: time.Now().UnixMilli(),
|
|
ThreadID: threadID,
|
|
DurationMs: durationMs,
|
|
ChunkCount: chunkCount,
|
|
Status: "completed",
|
|
}
|
|
|
|
// Add complete content to extra
|
|
if completeContent != "" {
|
|
endData.Extra = map[string]interface{}{
|
|
"content": completeContent,
|
|
}
|
|
}
|
|
|
|
// Send message_end event
|
|
messageEndEvent := output.NewEventMessage(message.EventMessageEnd, "Message completed", endData)
|
|
return ctx.sendRaw(messageEndEvent)
|
|
}
|
|
|
|
// EndMessage sends a message_end event for a completed message
|
|
// Note: For non-delta messages, message_end is automatically sent by Send()
|
|
// This method is primarily for delta streaming scenarios:
|
|
// - After all delta chunks are sent for a message, call EndMessage() to finalize it
|
|
// - For LLM streaming, this is typically called after receiving ChunkMessageEnd
|
|
func (ctx *Context) EndMessage(messageID string, content interface{}) error {
|
|
if messageID == "" || ctx.messageMetadata == nil {
|
|
return nil
|
|
}
|
|
|
|
metadata := ctx.messageMetadata.getMessage(messageID)
|
|
if metadata == nil {
|
|
return nil // Message not found, skip
|
|
}
|
|
|
|
// Calculate duration
|
|
durationMs := time.Since(metadata.StartTime).Milliseconds()
|
|
|
|
// Build message_end event data
|
|
endData := message.EventMessageEndData{
|
|
MessageID: messageID,
|
|
Type: metadata.Type,
|
|
Timestamp: time.Now().UnixMilli(),
|
|
ThreadID: metadata.ThreadID, // Include ThreadID for concurrent stream identification
|
|
DurationMs: durationMs,
|
|
ChunkCount: metadata.ChunkCount,
|
|
Status: "completed",
|
|
}
|
|
|
|
// Add content to extra if provided
|
|
if content != nil {
|
|
endData.Extra = map[string]interface{}{
|
|
"content": content,
|
|
}
|
|
}
|
|
|
|
// Send message_end event
|
|
messageEndEvent := output.NewEventMessage(message.EventMessageEnd, "Message completed", endData)
|
|
return ctx.sendRaw(messageEndEvent)
|
|
}
|
|
|
|
// EndBlock sends a block_end event for a completed block
|
|
// This should be called explicitly when all messages in a block are complete
|
|
func (ctx *Context) EndBlock(blockID string) error {
|
|
if blockID == "" || ctx.messageMetadata == nil {
|
|
return nil
|
|
}
|
|
|
|
blockMetadata := ctx.messageMetadata.getBlock(blockID)
|
|
if blockMetadata == nil {
|
|
return nil // Block not found, skip
|
|
}
|
|
|
|
// Calculate duration
|
|
durationMs := time.Since(blockMetadata.StartTime).Milliseconds()
|
|
|
|
// Build block_end event data
|
|
endData := message.EventBlockEndData{
|
|
BlockID: blockID,
|
|
Type: blockMetadata.Type,
|
|
Timestamp: time.Now().UnixMilli(),
|
|
DurationMs: durationMs,
|
|
MessageCount: blockMetadata.MessageCount,
|
|
Status: "completed",
|
|
}
|
|
|
|
// Send block_end event
|
|
blockEndEvent := output.NewEventMessage(message.EventBlockEnd, "Block completed", endData)
|
|
return ctx.sendRaw(blockEndEvent)
|
|
}
|
|
|
|
// SendGroup sends a group of messages via the output module
|
|
// Deprecated: This method is deprecated and will be removed in future versions
|
|
func (ctx *Context) SendGroup(group *message.Group) error {
|
|
output, err := ctx.getOutput()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return output.SendGroup(group)
|
|
}
|
|
|
|
// Flush flushes the output writer
|
|
func (ctx *Context) Flush() error {
|
|
output, err := ctx.getOutput()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return output.Flush()
|
|
}
|
|
|
|
// CloseOutput closes the output writer
|
|
func (ctx *Context) CloseOutput() error {
|
|
output, err := ctx.getOutput()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return output.Close()
|
|
}
|
|
|
|
// sendRaw sends a message directly without triggering lifecycle events
|
|
// Used internally to send event messages without recursion
|
|
func (ctx *Context) sendRaw(msg *message.Message) error {
|
|
out, err := ctx.getOutput()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return out.Send(msg)
|
|
}
|
|
|
|
// getWriter gets the effective Writer for the current context
|
|
// Priority: Skip.Output > Stack.Options.Writer > ctx.Writer
|
|
// Note: The Writer returned is always a SafeWriter (wrapped at context creation)
|
|
// to ensure thread-safe concurrent writes for SSE streaming.
|
|
func (ctx *Context) getWriter() Writer {
|
|
// Check if output is explicitly skipped (for internal A2A calls)
|
|
if ctx.Stack != nil && ctx.Stack.Options != nil && ctx.Stack.Options.Skip != nil && ctx.Stack.Options.Skip.Output {
|
|
return nil // Explicitly disable output
|
|
}
|
|
|
|
// Check if current Stack has a Writer override
|
|
if ctx.Stack != nil && ctx.Stack.Options != nil && ctx.Stack.Options.Writer != nil {
|
|
return ctx.Stack.Options.Writer
|
|
}
|
|
|
|
return ctx.Writer
|
|
}
|
|
|
|
// getOutput gets the output writer for the context
|
|
func (ctx *Context) getOutput() (*output.Output, error) {
|
|
// Check if current Stack has cached output
|
|
if ctx.Stack != nil && ctx.Stack.output != nil {
|
|
return ctx.Stack.output, nil
|
|
}
|
|
|
|
// Ensure Writer is wrapped in SafeWriter for concurrent-safe SSE writes
|
|
// This is essential for ctx.agent.All where multiple sub-agents
|
|
// write to the same SSE stream concurrently.
|
|
// We wrap once at the context level so all forked contexts share the same SafeWriter.
|
|
writer := ctx.getWriter()
|
|
if writer != nil {
|
|
// Check if it's already a SafeWriter
|
|
if _, ok := writer.(*output.SafeWriter); !ok {
|
|
// Wrap in SafeWriter with context for automatic cleanup on client disconnect
|
|
// This prevents goroutine leaks in enterprise applications
|
|
var safeWriter *output.SafeWriter
|
|
if ctx.Context != nil {
|
|
// Use request context to detect client disconnection
|
|
safeWriter = output.NewSafeWriterWithContext(ctx.Context, writer)
|
|
} else {
|
|
// Fallback to basic SafeWriter if no context available
|
|
safeWriter = output.NewSafeWriter(writer)
|
|
}
|
|
ctx.Writer = safeWriter
|
|
writer = safeWriter
|
|
}
|
|
}
|
|
|
|
trace, _ := ctx.Trace()
|
|
var options message.Options = message.Options{
|
|
BaseURL: "/",
|
|
Writer: writer,
|
|
Trace: trace,
|
|
Locale: ctx.Locale,
|
|
Accept: string(ctx.Accept),
|
|
}
|
|
|
|
if ctx.Capabilities != nil {
|
|
caps := llm.Capabilities(*ctx.Capabilities)
|
|
options.Capabilities = &caps
|
|
}
|
|
|
|
out, err := output.NewOutput(options)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// Cache to current Stack (each Stack has its own output with its own Writer)
|
|
if ctx.Stack != nil {
|
|
ctx.Stack.output = out
|
|
}
|
|
|
|
return out, nil
|
|
}
|
|
|
|
// CloseSafeWriter closes the SafeWriter if one was created
|
|
// This should be called at the end of the root request to flush any pending writes
|
|
// and stop the background goroutine.
|
|
func (ctx *Context) CloseSafeWriter() {
|
|
if ctx.Writer == nil {
|
|
return
|
|
}
|
|
if sw, ok := ctx.Writer.(*output.SafeWriter); ok {
|
|
sw.Close()
|
|
}
|
|
}
|