From a0483bac5dbfa8e957711e3ee352a69e46058c2d Mon Sep 17 00:00:00 2001 From: Max Date: Sat, 29 Nov 2025 09:13:18 +0800 Subject: [PATCH] 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. --- agent/assistant/agent.go | 148 ++++++++++-------------------------- agent/assistant/hook/mcp.go | 8 -- agent/assistant/llm.go | 59 ++++++++++++++ agent/assistant/trace.go | 106 ++++++++++++++++++++++++++ 4 files changed, 206 insertions(+), 115 deletions(-) delete mode 100644 agent/assistant/hook/mcp.go create mode 100644 agent/assistant/llm.go create mode 100644 agent/assistant/trace.go diff --git a/agent/assistant/agent.go b/agent/assistant/agent.go index 0b8bc776..f6385097 100644 --- a/agent/assistant/agent.go +++ b/agent/assistant/agent.go @@ -10,9 +10,7 @@ import ( "github.com/yaoapp/yao/agent/assistant/handlers" "github.com/yaoapp/yao/agent/context" "github.com/yaoapp/yao/agent/i18n" - "github.com/yaoapp/yao/agent/llm" "github.com/yaoapp/yao/agent/output/message" - "github.com/yaoapp/yao/trace/types" ) // 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 - _, traceID, done := context.EnterStack(ctx, ast.ID, ctx.Referer) + _, _, done := context.EnterStack(ctx, ast.ID, ctx.Referer) 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 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) // Now ctx.Capabilities is set, so output adapters can use it ast.sendAgentStreamStart(ctx, streamHandler, streamStartTime) - // Trace Add - trace, _ := ctx.Trace() - 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" - }) - } + // Initialize agent trace node + agentNode := ast.initAgentTraceNode(ctx, inputMessages) // Full input messages with chat history fullMessages, err := ast.WithHistory(ctx, inputMessages) if err != nil { - if agentNode != nil { - agentNode.Fail(err) - } + ast.traceAgentFail(agentNode, err) // Send error stream_end for root stack ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err) return nil, err } // Log the chat history - if agentNode != nil { - agentNode.Info(i18n.Tr(ast.ID, ctx.Locale, "assistant.agent.stream.history"), map[string]any{"messages": fullMessages}) // "Get Chat History" - } + ast.traceAgentHistory(ctx, agentNode, fullMessages) // Request Create hook ( Optional ) var createResponse *context.HookCreateResponse @@ -97,95 +72,35 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa var err error createResponse, err = ast.Script.Create(ctx, fullMessages) if err != nil { - if agentNode != nil { - agentNode.Fail(err) - } + ast.traceAgentFail(agentNode, err) // Send error stream_end for root stack ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err) return nil, err } // Log the create response - if agentNode != nil { - agentNode.Debug("Call Create Hook", map[string]any{"response": createResponse}) - } + ast.traceCreateHook(agentNode, createResponse) } - var completionOptions *context.CompletionOptions // default is nil - // LLM Call Stream ( Optional ) - var completionMessages []context.Message var completionResponse *context.CompletionResponse if ast.Prompts != nil || ast.MCP != nil { // 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 agentNode != nil { - agentNode.Fail(err) - } + ast.traceAgentFail(agentNode, err) // Send error stream_end for root stack ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, err) return nil, err } - // Get connector object (capabilities were already set above, before stream_start) - conn, capabilities, err := ast.GetConnector(ctx) - 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) + // Execute the LLM streaming call + completionResponse, err = ast.executeLLMStream(ctx, completionMessages, completionOptions, agentNode, streamHandler) if err != nil { // Send error stream_end for root stack ast.sendStreamEndOnError(ctx, streamHandler, streamStartTime, 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 ) @@ -211,9 +126,7 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa _ = doneResponse // doneResponse is available for further processing // Set the output of the agent node - if agentNode != nil { - agentNode.SetOutput(context.Response{Create: createResponse, Done: doneResponse, Completion: completionResponse}) - } + ast.traceAgentOutput(agentNode, createResponse, doneResponse, completionResponse) // 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 @@ -464,3 +377,24 @@ func (ast *Assistant) handleInterrupt(ctx *context.Context, signal *context.Inte 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 +} diff --git a/agent/assistant/hook/mcp.go b/agent/assistant/hook/mcp.go deleted file mode 100644 index 0e0ba6cc..00000000 --- a/agent/assistant/hook/mcp.go +++ /dev/null @@ -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 -} diff --git a/agent/assistant/llm.go b/agent/assistant/llm.go new file mode 100644 index 00000000..125d4ff1 --- /dev/null +++ b/agent/assistant/llm.go @@ -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 +} + diff --git a/agent/assistant/trace.go b/agent/assistant/trace.go new file mode 100644 index 00000000..4a9274cb --- /dev/null +++ b/agent/assistant/trace.go @@ -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) +}