From 19b4c7354f8e70d6a9b11040f3f1d0caca2b8d3d Mon Sep 17 00:00:00 2001 From: Max Date: Thu, 20 Nov 2025 11:36:44 +0800 Subject: [PATCH] Enhance output handling and trace management in Assistant's Stream method - Updated the Assistant's Stream method to set output for the agent node and conditionally close the output writer based on call depth, improving resource management. - Integrated detailed trace logging for output closure events, distinguishing between root and nested calls to enhance debugging capabilities. - Refactored trace creation in context to utilize authorized information for better trace context management. - Introduced a lightweight persistNode structure for efficient storage of trace nodes, optimizing the save and load operations in local and store drivers. --- agent/assistant/agent.go | 41 ++++++++++++--- agent/context/context.go | 14 ++++-- trace/local/driver.go | 106 +++++++++++++++++++++++++++++++++++++-- trace/store/driver.go | 106 +++++++++++++++++++++++++++++++++++++-- 4 files changed, 249 insertions(+), 18 deletions(-) diff --git a/agent/assistant/agent.go b/agent/assistant/agent.go index bc03e26f..2ce9f6fe 100644 --- a/agent/assistant/agent.go +++ b/agent/assistant/agent.go @@ -166,15 +166,42 @@ func (ast *Assistant) Stream(ctx *context.Context, inputMessages []context.Messa _ = doneResponse // doneResponse is available for further processing - // Close the output writer to send [DONE] marker and flush data - if err := output.Close(ctx); err != nil { - // Log error but don't fail the request - fmt.Printf("Warning: Failed to close output writer: %v\n", err) + // Set the output of the agent node + if agentNode != nil { + agentNode.SetOutput(context.Response{Create: createResponse, Done: doneResponse, Completion: completionResponse}) } - // Flush any remaining data to the client - if err := output.Flush(ctx); err != nil { - fmt.Printf("Warning: Failed to flush output: %v\n", err) + // Only close output if this is the root call (entry point) + // Nested calls (from MCP, hooks, etc.) should not close the output + // Note: Flush is already handled by the stream handler (handleStreamEnd) + if ctx.Stack != nil && ctx.Stack.IsRoot() { + // Log closing output for root call + if trace, _ := ctx.Trace(); trace != nil { + trace.Debug("Agent: Closing output (root call)", map[string]any{ + "stack_id": ctx.Stack.ID, + "depth": ctx.Stack.Depth, + "assistant_id": ctx.Stack.AssistantID, + }) + } + + // Close the output writer to send [DONE] marker + if err := output.Close(ctx); err != nil { + if trace, _ := ctx.Trace(); trace != nil { + trace.Error("Agent: Failed to close output", map[string]any{ + "error": err.Error(), + }) + } + } + } else { + // Log skipping close for nested call + if trace, _ := ctx.Trace(); trace != nil && ctx.Stack != nil { + trace.Debug("Agent: Skipping output close (nested call)", map[string]any{ + "stack_id": ctx.Stack.ID, + "depth": ctx.Stack.Depth, + "parent_id": ctx.Stack.ParentID, + "assistant_id": ctx.Stack.AssistantID, + }) + } } return &context.Response{Create: createResponse, Done: doneResponse, Completion: completionResponse}, nil diff --git a/agent/context/context.go b/agent/context/context.go index 214046f5..6d2b14fe 100644 --- a/agent/context/context.go +++ b/agent/context/context.go @@ -176,10 +176,18 @@ func (ctx *Context) Trace() (traceTypes.Manager, error) { return nil, fmt.Errorf("unsupported trace driver: %s", cfg.Trace.Driver) } + // Prepare trace options + traceOption := &traceTypes.TraceOption{ID: traceID} + + // Set trace options from authorized information + if ctx.Authorized != nil { + traceOption.CreatedBy = ctx.Authorized.UserID + traceOption.TeamID = ctx.Authorized.TeamID + traceOption.TenantID = ctx.Authorized.TenantID + } + // Create trace using trace.New (handles registry) - createdTraceID, manager, err := trace.New(ctx.Context, driverType, &traceTypes.TraceOption{ - ID: traceID, // Use existing ID from Stack or empty to generate new one - }, driverOptions...) + createdTraceID, manager, err := trace.New(ctx.Context, driverType, traceOption, driverOptions...) if err != nil { return nil, fmt.Errorf("failed to create trace: %w", err) } diff --git a/trace/local/driver.go b/trace/local/driver.go index 1d58661f..45eeae14 100644 --- a/trace/local/driver.go +++ b/trace/local/driver.go @@ -12,6 +12,83 @@ import ( "github.com/yaoapp/yao/trace/types" ) +// persistNode is a lightweight version of TraceNode for storage +// Only stores IDs of children instead of full child nodes +type persistNode struct { + ID string `json:"ID"` + ParentID string `json:"ParentID"` + ChildrenIDs []string `json:"ChildrenIDs,omitempty"` + Label string `json:"Label,omitempty"` + Icon string `json:"Icon,omitempty"` + Description string `json:"Description,omitempty"` + Metadata map[string]any `json:"Metadata,omitempty"` + Status types.NodeStatus `json:"Status"` + Input types.TraceInput `json:"Input,omitempty"` + Output types.TraceOutput `json:"Output,omitempty"` + CreatedAt int64 `json:"CreatedAt"` + StartTime int64 `json:"StartTime"` + EndTime int64 `json:"EndTime,omitempty"` + UpdatedAt int64 `json:"UpdatedAt"` +} + +// toPersistNode converts TraceNode to persistNode for storage +func toPersistNode(node *types.TraceNode) *persistNode { + if node == nil { + return nil + } + + // Extract children IDs + childrenIDs := make([]string, 0, len(node.Children)) + for _, child := range node.Children { + if child != nil { + childrenIDs = append(childrenIDs, child.ID) + } + } + + return &persistNode{ + ID: node.ID, + ParentID: node.ParentID, + ChildrenIDs: childrenIDs, + Label: node.Label, + Icon: node.Icon, + Description: node.Description, + Metadata: node.Metadata, + Status: node.Status, + Input: node.Input, + Output: node.Output, + CreatedAt: node.CreatedAt, + StartTime: node.StartTime, + EndTime: node.EndTime, + UpdatedAt: node.UpdatedAt, + } +} + +// fromPersistNode converts persistNode to TraceNode +func fromPersistNode(pn *persistNode) *types.TraceNode { + if pn == nil { + return nil + } + + return &types.TraceNode{ + ID: pn.ID, + ParentID: pn.ParentID, + Children: nil, // Children will be loaded separately if needed + TraceNodeOption: types.TraceNodeOption{ + Label: pn.Label, + Icon: pn.Icon, + Description: pn.Description, + Metadata: pn.Metadata, + }, + Status: pn.Status, + Input: pn.Input, + Output: pn.Output, + CreatedAt: pn.CreatedAt, + StartTime: pn.StartTime, + EndTime: pn.EndTime, + UpdatedAt: pn.UpdatedAt, + } +} + // Driver the local disk storage driver implementation type Driver struct { basePath string // Base directory for storing trace files @@ -66,9 +143,12 @@ func (d *Driver) SaveNode(ctx context.Context, traceID string, node *types.Trace return fmt.Errorf("failed to create nodes directory: %w", err) } + // Convert to persist format (only store children IDs) + persistData := toPersistNode(node) + // Save node as JSON filePath := filepath.Join(nodesDir, node.ID+".json") - data, err := json.MarshalIndent(node, "", " ") + data, err := json.MarshalIndent(persistData, "", " ") if err != nil { return fmt.Errorf("failed to marshal node: %w", err) } @@ -92,12 +172,30 @@ func (d *Driver) LoadNode(ctx context.Context, traceID string, nodeID string) (* return nil, fmt.Errorf("failed to read node file: %w", err) } - var node types.TraceNode - if err := json.Unmarshal(data, &node); err != nil { + var pn persistNode + if err := json.Unmarshal(data, &pn); err != nil { return nil, fmt.Errorf("failed to unmarshal node: %w", err) } - return &node, nil + // Convert to TraceNode + node := fromPersistNode(&pn) + + // Load children if needed + if len(pn.ChildrenIDs) > 0 { + children := make([]*types.TraceNode, 0, len(pn.ChildrenIDs)) + for _, childID := range pn.ChildrenIDs { + child, err := d.LoadNode(ctx, traceID, childID) + if err != nil { + return nil, fmt.Errorf("failed to load child node %s: %w", childID, err) + } + if child != nil { + children = append(children, child) + } + } + node.Children = children + } + + return node, nil } // LoadTrace loads the entire trace tree from disk diff --git a/trace/store/driver.go b/trace/store/driver.go index d159bcfb..066e5290 100644 --- a/trace/store/driver.go +++ b/trace/store/driver.go @@ -11,6 +11,83 @@ import ( "github.com/yaoapp/yao/trace/types" ) +// persistNode is a lightweight version of TraceNode for storage +// Only stores IDs of children instead of full child nodes +type persistNode struct { + ID string `json:"ID"` + ParentID string `json:"ParentID"` + ChildrenIDs []string `json:"ChildrenIDs,omitempty"` + Label string `json:"Label,omitempty"` + Icon string `json:"Icon,omitempty"` + Description string `json:"Description,omitempty"` + Metadata map[string]any `json:"Metadata,omitempty"` + Status types.NodeStatus `json:"Status"` + Input types.TraceInput `json:"Input,omitempty"` + Output types.TraceOutput `json:"Output,omitempty"` + CreatedAt int64 `json:"CreatedAt"` + StartTime int64 `json:"StartTime"` + EndTime int64 `json:"EndTime,omitempty"` + UpdatedAt int64 `json:"UpdatedAt"` +} + +// toPersistNode converts TraceNode to persistNode for storage +func toPersistNode(node *types.TraceNode) *persistNode { + if node == nil { + return nil + } + + // Extract children IDs + childrenIDs := make([]string, 0, len(node.Children)) + for _, child := range node.Children { + if child != nil { + childrenIDs = append(childrenIDs, child.ID) + } + } + + return &persistNode{ + ID: node.ID, + ParentID: node.ParentID, + ChildrenIDs: childrenIDs, + Label: node.Label, + Icon: node.Icon, + Description: node.Description, + Metadata: node.Metadata, + Status: node.Status, + Input: node.Input, + Output: node.Output, + CreatedAt: node.CreatedAt, + StartTime: node.StartTime, + EndTime: node.EndTime, + UpdatedAt: node.UpdatedAt, + } +} + +// fromPersistNode converts persistNode to TraceNode +func fromPersistNode(pn *persistNode) *types.TraceNode { + if pn == nil { + return nil + } + + return &types.TraceNode{ + ID: pn.ID, + ParentID: pn.ParentID, + Children: nil, // Children will be loaded separately if needed + TraceNodeOption: types.TraceNodeOption{ + Label: pn.Label, + Icon: pn.Icon, + Description: pn.Description, + Metadata: pn.Metadata, + }, + Status: pn.Status, + Input: pn.Input, + Output: pn.Output, + CreatedAt: pn.CreatedAt, + StartTime: pn.StartTime, + EndTime: pn.EndTime, + UpdatedAt: pn.UpdatedAt, + } +} + // Driver the gou store storage driver implementation type Driver struct { storeName string // Store name in gou @@ -54,7 +131,10 @@ func (d *Driver) getKey(traceID string, parts ...string) string { func (d *Driver) SaveNode(ctx context.Context, traceID string, node *types.TraceNode) error { key := d.getKey(traceID, "node", node.ID) - data, err := json.Marshal(node) + // Convert to persist format (only store children IDs) + persistData := toPersistNode(node) + + data, err := json.Marshal(persistData) if err != nil { return fmt.Errorf("failed to marshal node: %w", err) } @@ -80,12 +160,30 @@ func (d *Driver) LoadNode(ctx context.Context, traceID string, nodeID string) (* return nil, fmt.Errorf("invalid data type in store") } - var node types.TraceNode - if err := json.Unmarshal([]byte(dataStr), &node); err != nil { + var pn persistNode + if err := json.Unmarshal([]byte(dataStr), &pn); err != nil { return nil, fmt.Errorf("failed to unmarshal node: %w", err) } - return &node, nil + // Convert to TraceNode + node := fromPersistNode(&pn) + + // Load children if needed + if len(pn.ChildrenIDs) > 0 { + children := make([]*types.TraceNode, 0, len(pn.ChildrenIDs)) + for _, childID := range pn.ChildrenIDs { + child, err := d.LoadNode(ctx, traceID, childID) + if err != nil { + return nil, fmt.Errorf("failed to load child node %s: %w", childID, err) + } + if child != nil { + children = append(children, child) + } + } + node.Children = children + } + + return node, nil } // LoadTrace loads the entire trace tree from store