Refactor Assistant streaming logic and remove MCP hook
- Simplified the Stream method by extracting capability initialization into a separate function, improving readability and maintainability. - Enhanced error handling by centralizing trace logging for failures and history tracking. - Removed the MCP hook implementation as it was no longer needed, streamlining the codebase. - Updated agent trace handling to ensure consistent logging of agent output and errors.
This commit is contained in:
parent
e3fde0aa6c
commit
a0483bac5d
4 changed files with 206 additions and 115 deletions
|
|
@ -10,9 +10,7 @@ import (
|
||||||
"github.com/yaoapp/yao/agent/assistant/handlers"
|
"github.com/yaoapp/yao/agent/assistant/handlers"
|
||||||
"github.com/yaoapp/yao/agent/context"
|
"github.com/yaoapp/yao/agent/context"
|
||||||
"github.com/yaoapp/yao/agent/i18n"
|
"github.com/yaoapp/yao/agent/i18n"
|
||||||
"github.com/yaoapp/yao/agent/llm"
|
|
||||||
"github.com/yaoapp/yao/agent/output/message"
|
"github.com/yaoapp/yao/agent/output/message"
|
||||||
"github.com/yaoapp/yao/trace/types"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// Stream stream the agent
|
// Stream stream the agent
|
||||||
|
|
@ -35,61 +33,38 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
}
|
}
|
||||||
|
|
||||||
// Initialize stack and auto-handle completion/failure/restore
|
// Initialize stack and auto-handle completion/failure/restore
|
||||||
_, traceID, done := context.EnterStack(ctx, ast.ID, ctx.Referer)
|
_, _, done := context.EnterStack(ctx, ast.ID, ctx.Referer)
|
||||||
defer done()
|
defer done()
|
||||||
|
|
||||||
_ = traceID // traceID is available for trace logging
|
|
||||||
|
|
||||||
// Get connector and capabilities early (before sending stream_start)
|
|
||||||
// so that output adapters can use them when converting stream_start event
|
|
||||||
if ast.Prompts != nil || ast.MCP != nil {
|
|
||||||
_, capabilities, err := ast.GetConnector(ctx)
|
|
||||||
if err != nil {
|
|
||||||
streamHandler := ast.getStreamHandler(ctx, handler...)
|
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Set capabilities in context for output adapters to use
|
|
||||||
if capabilities != nil {
|
|
||||||
ctx.Capabilities = capabilities
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Determine stream handler
|
// Determine stream handler
|
||||||
streamHandler := ast.getStreamHandler(ctx, handler...)
|
streamHandler := ast.getStreamHandler(ctx, handler...)
|
||||||
|
|
||||||
|
// Get connector and capabilities early (before sending stream_start)
|
||||||
|
// so that output adapters can use them when converting stream_start event
|
||||||
|
err = ast.initializeCapabilities(ctx)
|
||||||
|
if err != nil {
|
||||||
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
// Send ChunkStreamStart only for root stack (agent-level stream start)
|
// Send ChunkStreamStart only for root stack (agent-level stream start)
|
||||||
// Now ctx.Capabilities is set, so output adapters can use it
|
// Now ctx.Capabilities is set, so output adapters can use it
|
||||||
ast.sendAgentStreamStart(ctx, streamHandler, streamStartTime)
|
ast.sendAgentStreamStart(ctx, streamHandler, streamStartTime)
|
||||||
|
|
||||||
// Trace Add
|
// Initialize agent trace node
|
||||||
trace, _ := ctx.Trace()
|
agentNode := ast.initAgentTraceNode(ctx, inputMessages)
|
||||||
var agentNode types.Node = nil
|
|
||||||
if trace != nil {
|
|
||||||
agentNode, _ = trace.Add(inputMessages, types.TraceNodeOption{
|
|
||||||
Label: i18n.Tr(ast.ID, ctx.Locale, "assistant.agent.stream.label"), // "Assistant {{name}}"
|
|
||||||
Type: "agent",
|
|
||||||
Icon: "assistant",
|
|
||||||
Description: i18n.Tr(ast.ID, ctx.Locale, "assistant.agent.stream.description"), // "Assistant {{name}} is processing the request"
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Full input messages with chat history
|
// Full input messages with chat history
|
||||||
fullMessages, err := ast.WithHistory(ctx, inputMessages)
|
fullMessages, err := ast.WithHistory(ctx, inputMessages)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if agentNode != nil {
|
ast.traceAgentFail(agentNode, err)
|
||||||
agentNode.Fail(err)
|
|
||||||
}
|
|
||||||
// Send error stream_end for root stack
|
// Send error stream_end for root stack
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Log the chat history
|
// Log the chat history
|
||||||
if agentNode != nil {
|
ast.traceAgentHistory(ctx, agentNode, fullMessages)
|
||||||
agentNode.Info(i18n.Tr(ast.ID, ctx.Locale, "assistant.agent.stream.history"), map[string]any{"messages": fullMessages}) // "Get Chat History"
|
|
||||||
}
|
|
||||||
|
|
||||||
// Request Create hook ( Optional )
|
// Request Create hook ( Optional )
|
||||||
var createResponse *context.HookCreateResponse
|
var createResponse *context.HookCreateResponse
|
||||||
|
|
@ -97,95 +72,35 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
var err error
|
var err error
|
||||||
createResponse, err = ast.Script.Create(ctx, fullMessages)
|
createResponse, err = ast.Script.Create(ctx, fullMessages)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if agentNode != nil {
|
ast.traceAgentFail(agentNode, err)
|
||||||
agentNode.Fail(err)
|
|
||||||
}
|
|
||||||
// Send error stream_end for root stack
|
// Send error stream_end for root stack
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Log the create response
|
// Log the create response
|
||||||
if agentNode != nil {
|
ast.traceCreateHook(agentNode, createResponse)
|
||||||
agentNode.Debug("Call Create Hook", map[string]any{"response": createResponse})
|
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
var completionOptions *context.CompletionOptions // default is nil
|
|
||||||
|
|
||||||
// LLM Call Stream ( Optional )
|
// LLM Call Stream ( Optional )
|
||||||
var completionMessages []context.Message
|
|
||||||
var completionResponse *context.CompletionResponse
|
var completionResponse *context.CompletionResponse
|
||||||
if ast.Prompts != nil || ast.MCP != nil {
|
if ast.Prompts != nil || ast.MCP != nil {
|
||||||
// Build the LLM request first
|
// Build the LLM request first
|
||||||
completionMessages, completionOptions, err = ast.BuildRequest(ctx, inputMessages, createResponse)
|
completionMessages, completionOptions, err := ast.BuildRequest(ctx, inputMessages, createResponse)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if agentNode != nil {
|
ast.traceAgentFail(agentNode, err)
|
||||||
agentNode.Fail(err)
|
|
||||||
}
|
|
||||||
// Send error stream_end for root stack
|
// Send error stream_end for root stack
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Get connector object (capabilities were already set above, before stream_start)
|
// Execute the LLM streaming call
|
||||||
conn, capabilities, err := ast.GetConnector(ctx)
|
completionResponse, err = ast.executeLLMStream(ctx, completionMessages, completionOptions, agentNode, streamHandler)
|
||||||
if err != nil {
|
|
||||||
if agentNode != nil {
|
|
||||||
agentNode.Fail(err)
|
|
||||||
}
|
|
||||||
// Send error stream_end for root stack
|
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Set capabilities in options if not already set
|
|
||||||
if completionOptions.Capabilities == nil && capabilities != nil {
|
|
||||||
completionOptions.Capabilities = capabilities
|
|
||||||
}
|
|
||||||
|
|
||||||
// Log the capabilities
|
|
||||||
if agentNode != nil {
|
|
||||||
agentNode.Debug("Get Connector Capabilities", map[string]any{"capabilities": capabilities})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Trace Add
|
|
||||||
if trace != nil {
|
|
||||||
trace.Add(
|
|
||||||
map[string]any{"messages": completionMessages, "options": completionOptions},
|
|
||||||
types.TraceNodeOption{
|
|
||||||
Label: fmt.Sprintf(i18n.Tr(ast.ID, ctx.Locale, "llm.openai.stream.label"), conn.ID()), // "LLM %s"
|
|
||||||
Type: "llm",
|
|
||||||
Icon: "psychology",
|
|
||||||
Description: fmt.Sprintf(i18n.Tr(ast.ID, ctx.Locale, "llm.openai.stream.description"), conn.ID()), // "LLM %s is processing the request"
|
|
||||||
},
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Create LLM instance with connector and options
|
|
||||||
llmInstance, err := llm.New(conn, completionOptions)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Send error stream_end for root stack
|
// Send error stream_end for root stack
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Call the LLM Completion Stream (streamHandler was set earlier)
|
|
||||||
log.Trace("[AGENT] Calling LLM Stream: assistant=%s", ast.ID)
|
|
||||||
completionResponse, err = llmInstance.Stream(ctx, completionMessages, completionOptions, streamHandler)
|
|
||||||
log.Trace("[AGENT] LLM Stream returned: assistant=%s, err=%v", ast.ID, err)
|
|
||||||
if err != nil {
|
|
||||||
// Send error stream_end for root stack
|
|
||||||
log.Trace("[AGENT] Calling sendStreamEndOnError")
|
|
||||||
ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err)
|
|
||||||
log.Trace("[AGENT] sendStreamEndOnError returned")
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Mark LLM Request Complete
|
|
||||||
if trace != nil {
|
|
||||||
trace.Complete(completionResponse)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Request MCP hook ( Optional )
|
// Request MCP hook ( Optional )
|
||||||
|
|
@ -211,9 +126,7 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa
|
||||||
_ = doneResponse // doneResponse is available for further processing
|
_ = doneResponse // doneResponse is available for further processing
|
||||||
|
|
||||||
// Set the output of the agent node
|
// Set the output of the agent node
|
||||||
if agentNode != nil {
|
ast.traceAgentOutput(agentNode, createResponse, doneResponse, completionResponse)
|
||||||
agentNode.SetOutput(context.Response{Create: createResponse, Done: doneResponse, Completion: completionResponse})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Only close output and send stream_end if this is the root call (entry point)
|
// Only close output and send stream_end if this is the root call (entry point)
|
||||||
// Nested calls (from MCP, hooks, etc.) should not close the output or send stream_end
|
// Nested calls (from MCP, hooks, etc.) should not close the output or send stream_end
|
||||||
|
|
@ -464,3 +377,24 @@ func (ast *Assistant) handleInterrupt(ctx *context.Context, signal *context.Inte
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// initializeCapabilities gets connector and capabilities, then sets them in context
|
||||||
|
// This should be called early (before sending stream_start) so that output adapters
|
||||||
|
// can use capabilities when converting stream_start event
|
||||||
|
func (ast *Assistant) initializeCapabilities(ctx *context.Context) error {
|
||||||
|
if ast.Prompts == nil && ast.MCP == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
_, capabilities, err := ast.GetConnector(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set capabilities in context for output adapters to use
|
||||||
|
if capabilities != nil {
|
||||||
|
ctx.Capabilities = capabilities
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,8 +0,0 @@
|
||||||
package hook
|
|
||||||
|
|
||||||
import "github.com/yaoapp/yao/agent/context"
|
|
||||||
|
|
||||||
// MCP MCP hook
|
|
||||||
func (s *Script) MCP(ctx *context.Context, messages []context.Message) (*context.ResponseHookMCP, error) {
|
|
||||||
return &context.ResponseHookMCP{}, nil
|
|
||||||
}
|
|
||||||
59
agent/assistant/llm.go
Normal file
59
agent/assistant/llm.go
Normal file
|
|
@ -0,0 +1,59 @@
|
||||||
|
package assistant
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/yaoapp/kun/log"
|
||||||
|
"github.com/yaoapp/yao/agent/context"
|
||||||
|
"github.com/yaoapp/yao/agent/llm"
|
||||||
|
"github.com/yaoapp/yao/agent/output/message"
|
||||||
|
"github.com/yaoapp/yao/trace/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// executeLLMStream executes the LLM streaming call with pre-built request
|
||||||
|
// Returns completionResponse and error
|
||||||
|
func (ast *Assistant) executeLLMStream(
|
||||||
|
ctx *context.Context,
|
||||||
|
completionMessages []context.Message,
|
||||||
|
completionOptions *context.CompletionOptions,
|
||||||
|
agentNode types.Node,
|
||||||
|
streamHandler message.StreamFunc,
|
||||||
|
) (*context.CompletionResponse, error) {
|
||||||
|
|
||||||
|
// Get connector object (capabilities were already set above, before stream_start)
|
||||||
|
conn, capabilities, err := ast.GetConnector(ctx)
|
||||||
|
if err != nil {
|
||||||
|
ast.traceAgentFail(agentNode, err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set capabilities in options if not already set
|
||||||
|
if completionOptions.Capabilities == nil && capabilities != nil {
|
||||||
|
completionOptions.Capabilities = capabilities
|
||||||
|
}
|
||||||
|
|
||||||
|
// Log the capabilities
|
||||||
|
ast.traceConnectorCapabilities(agentNode, capabilities)
|
||||||
|
|
||||||
|
// Trace Add LLM request
|
||||||
|
ast.traceLLMRequest(ctx, conn.ID(), completionMessages, completionOptions)
|
||||||
|
|
||||||
|
// Create LLM instance with connector and options
|
||||||
|
llmInstance, err := llm.New(conn, completionOptions)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call the LLM Completion Stream (streamHandler was set earlier)
|
||||||
|
log.Trace("[AGENT] Calling LLM Stream: assistant=%s", ast.ID)
|
||||||
|
completionResponse, err := llmInstance.Stream(ctx, completionMessages, completionOptions, streamHandler)
|
||||||
|
log.Trace("[AGENT] LLM Stream returned: assistant=%s, err=%v", ast.ID, err)
|
||||||
|
if err != nil {
|
||||||
|
log.Trace("[AGENT] Calling sendStreamEndOnError")
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Mark LLM Request Complete
|
||||||
|
ast.traceLLMComplete(ctx, completionResponse)
|
||||||
|
|
||||||
|
return completionResponse, nil
|
||||||
|
}
|
||||||
|
|
||||||
106
agent/assistant/trace.go
Normal file
106
agent/assistant/trace.go
Normal file
|
|
@ -0,0 +1,106 @@
|
||||||
|
package assistant
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/yaoapp/yao/agent/context"
|
||||||
|
"github.com/yaoapp/yao/agent/i18n"
|
||||||
|
"github.com/yaoapp/yao/trace/types"
|
||||||
|
)
|
||||||
|
|
||||||
|
// initAgentTraceNode creates and returns the agent trace node
|
||||||
|
func (ast *Assistant) initAgentTraceNode(ctx *context.Context, inputMessages []context.Message) types.Node {
|
||||||
|
trace, _ := ctx.Trace()
|
||||||
|
if trace == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
agentNode, _ := trace.Add(inputMessages, types.TraceNodeOption{
|
||||||
|
Label: i18n.Tr(ast.ID, ctx.Locale, "assistant.agent.stream.label"), // "Assistant {{name}}"
|
||||||
|
Type: "agent",
|
||||||
|
Icon: "assistant",
|
||||||
|
Description: i18n.Tr(ast.ID, ctx.Locale, "assistant.agent.stream.description"), // "Assistant {{name}} is processing the request"
|
||||||
|
})
|
||||||
|
|
||||||
|
return agentNode
|
||||||
|
}
|
||||||
|
|
||||||
|
// traceAgentHistory logs the chat history to the agent trace node
|
||||||
|
func (ast *Assistant) traceAgentHistory(ctx *context.Context, agentNode types.Node, fullMessages []context.Message) {
|
||||||
|
if agentNode == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
agentNode.Info(
|
||||||
|
i18n.Tr(ast.ID, ctx.Locale, "assistant.agent.stream.history"), // "Get Chat History"
|
||||||
|
map[string]any{"messages": fullMessages},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// traceCreateHook logs the create hook response to the agent trace node
|
||||||
|
func (ast *Assistant) traceCreateHook(agentNode types.Node, createResponse *context.HookCreateResponse) {
|
||||||
|
if agentNode == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
agentNode.Debug("Call Create Hook", map[string]any{"response": createResponse})
|
||||||
|
}
|
||||||
|
|
||||||
|
// traceConnectorCapabilities logs the connector capabilities to the agent trace node
|
||||||
|
func (ast *Assistant) traceConnectorCapabilities(agentNode types.Node, capabilities *context.ModelCapabilities) {
|
||||||
|
if agentNode == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
agentNode.Debug("Get Connector Capabilities", map[string]any{"capabilities": capabilities})
|
||||||
|
}
|
||||||
|
|
||||||
|
// traceLLMRequest adds a LLM trace node to the trace
|
||||||
|
func (ast *Assistant) traceLLMRequest(ctx *context.Context, connID string, completionMessages []context.Message, completionOptions *context.CompletionOptions) {
|
||||||
|
trace, _ := ctx.Trace()
|
||||||
|
if trace == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
trace.Add(
|
||||||
|
map[string]any{"messages": completionMessages, "options": completionOptions},
|
||||||
|
types.TraceNodeOption{
|
||||||
|
Label: fmt.Sprintf(i18n.Tr(ast.ID, ctx.Locale, "llm.openai.stream.label"), connID), // "LLM %s"
|
||||||
|
Type: "llm",
|
||||||
|
Icon: "psychology",
|
||||||
|
Description: fmt.Sprintf(i18n.Tr(ast.ID, ctx.Locale, "llm.openai.stream.description"), connID), // "LLM %s is processing the request"
|
||||||
|
},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// traceLLMComplete marks the LLM request as complete in the trace
|
||||||
|
func (ast *Assistant) traceLLMComplete(ctx *context.Context, completionResponse *context.CompletionResponse) {
|
||||||
|
trace, _ := ctx.Trace()
|
||||||
|
if trace == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
trace.Complete(completionResponse)
|
||||||
|
}
|
||||||
|
|
||||||
|
// traceAgentOutput sets the output of the agent trace node
|
||||||
|
func (ast *Assistant) traceAgentOutput(agentNode types.Node, createResponse *context.HookCreateResponse, doneResponse *context.ResponseHookDone, completionResponse *context.CompletionResponse) {
|
||||||
|
if agentNode == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
agentNode.SetOutput(context.Response{
|
||||||
|
Create: createResponse,
|
||||||
|
Done: doneResponse,
|
||||||
|
Completion: completionResponse,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// traceAgentFail marks the agent trace node as failed
|
||||||
|
func (ast *Assistant) traceAgentFail(agentNode types.Node, err error) {
|
||||||
|
if agentNode == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
agentNode.Fail(err)
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue