diff --git a/trace/local/driver.go b/trace/local/driver.go index e14c7183..1d58661f 100644 --- a/trace/local/driver.go +++ b/trace/local/driver.go @@ -478,6 +478,71 @@ func (d *Driver) DeleteTrace(ctx context.Context, traceID string) error { return nil } +// SaveUpdate persists a trace update event to disk (append-only) +func (d *Driver) SaveUpdate(ctx context.Context, traceID string, update *types.TraceUpdate) error { + if err := d.ensureTraceDir(traceID); err != nil { + return err + } + + filePath := filepath.Join(d.getTracePath(traceID), "updates.jsonl") + + // Marshal update to JSON + data, err := json.Marshal(update) + if err != nil { + return fmt.Errorf("failed to marshal update: %w", err) + } + + // Append to file (create if not exists) + f, err := os.OpenFile(filePath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) + if err != nil { + return fmt.Errorf("failed to open updates file: %w", err) + } + defer f.Close() + + if _, err := f.Write(append(data, '\n')); err != nil { + return fmt.Errorf("failed to write update: %w", err) + } + + return nil +} + +// LoadUpdates loads trace update events from disk +func (d *Driver) LoadUpdates(ctx context.Context, traceID string, since int64) ([]*types.TraceUpdate, error) { + filePath := filepath.Join(d.getTracePath(traceID), "updates.jsonl") + + // Read file + data, err := os.ReadFile(filePath) + if err != nil { + if os.IsNotExist(err) { + return []*types.TraceUpdate{}, nil + } + return nil, fmt.Errorf("failed to read updates file: %w", err) + } + + // Parse line by line + lines := strings.Split(string(data), "\n") + updates := make([]*types.TraceUpdate, 0, len(lines)) + + for _, line := range lines { + if strings.TrimSpace(line) == "" { + continue + } + + var update types.TraceUpdate + if err := json.Unmarshal([]byte(line), &update); err != nil { + // Skip malformed lines + continue + } + + // Filter by timestamp + if update.Timestamp >= since { + updates = append(updates, &update) + } + } + + return updates, nil +} + // Close closes the local driver func (d *Driver) Close() error { // No cleanup needed for local file system diff --git a/trace/manager.go b/trace/manager.go index 04ff921e..9481e708 100644 --- a/trace/manager.go +++ b/trace/manager.go @@ -3,56 +3,23 @@ package trace import ( "context" "fmt" - "sync" "time" gonanoid "github.com/matoous/go-nanoid/v2" "github.com/yaoapp/yao/trace/types" ) -// manager implements the Manager interface with unified business logic +// manager implements the Manager interface with channel-based state management type manager struct { ctx context.Context - cancel context.CancelFunc // Cancel function to stop background goroutines + cancel context.CancelFunc traceID string driver types.Driver - rootNode *types.TraceNode - currentNodes []*types.TraceNode - spaces map[string]*types.TraceSpace - spaceLocks map[string]*sync.RWMutex // Per-space locks for concurrent safety - mu sync.RWMutex // Protects currentNodes and spaces - - // Subscription mechanism - updates []*types.TraceUpdate // Update history (all events) - updatesMu sync.RWMutex // Protects updates - subscribers map[string]chan *types.TraceUpdate // Active subscribers - subMu sync.RWMutex // Protects subscribers - completed bool // Trace completion status + stateCmdChan chan stateCommand // Single channel for all state mutations } // NewManager creates a new trace manager instance func NewManager(ctx context.Context, traceID string, driver types.Driver) (types.Manager, error) { - // Create root node - now := time.Now().Unix() - rootNode := &types.TraceNode{ - ID: genNodeID(), - ParentID: "", - Children: []*types.TraceNode{}, - Status: types.StatusRunning, - CreatedAt: now, - StartTime: now, - UpdatedAt: now, - TraceNodeOption: types.TraceNodeOption{ - Label: "Root", - Icon: "root", - }, - } - - // Save root node - if err := driver.SaveNode(ctx, traceID, rootNode); err != nil { - return nil, fmt.Errorf("failed to save root node: %w", err) - } - // Create a cancellable context for the manager managerCtx, cancel := context.WithCancel(ctx) @@ -61,31 +28,35 @@ func NewManager(ctx context.Context, traceID string, driver types.Driver) (types cancel: cancel, traceID: traceID, driver: driver, - rootNode: rootNode, - currentNodes: []*types.TraceNode{rootNode}, - spaces: make(map[string]*types.TraceSpace), - spaceLocks: make(map[string]*sync.RWMutex), - updates: make([]*types.TraceUpdate, 0, 100), - subscribers: make(map[string]chan *types.TraceUpdate), - completed: false, + stateCmdChan: make(chan stateCommand, 100), // Buffered channel for performance } - // Broadcast init event - m.addUpdate(&types.TraceUpdate{ - Type: types.UpdateTypeInit, - TraceID: traceID, - Timestamp: now, - Data: types.NewTraceInitData(traceID, rootNode), - }) + // Start state worker goroutine + go m.startStateWorker() - // Broadcast root node start event - m.addUpdate(&types.TraceUpdate{ - Type: types.UpdateTypeNodeStart, - TraceID: traceID, - NodeID: rootNode.ID, - Timestamp: now, - Data: rootNode.ToStartData(), - }) + // Try to load existing updates from driver (for resumed traces) + if existingUpdates, err := driver.LoadUpdates(ctx, traceID, 0); err == nil && len(existingUpdates) > 0 { + m.stateSetUpdates(existingUpdates) + // Check if trace was already completed + for _, update := range existingUpdates { + if update.Type == types.UpdateTypeComplete { + m.stateMarkCompleted() + if data, ok := update.Data.(*types.TraceCompleteData); ok { + m.stateSetTraceStatus(data.Status) + } + break + } + } + } else { + // New trace - create and broadcast init event + now := time.Now().Unix() + m.addUpdateAndBroadcast(&types.TraceUpdate{ + Type: types.UpdateTypeInit, + TraceID: traceID, + Timestamp: now, + Data: types.NewTraceInitData(traceID, nil), + }) + } return m, nil } @@ -96,16 +67,94 @@ func genNodeID() string { return id } +// addUpdateAndBroadcast persists, adds to history, and broadcasts an update +func (m *manager) addUpdateAndBroadcast(update *types.TraceUpdate) { + // Persist to driver (synchronous - no race) + _ = m.driver.SaveUpdate(context.Background(), m.traceID, update) + + // Add to in-memory history + m.stateAddUpdate(update) + + // Broadcast to subscribers + m.stateBroadcast(update) +} + // checkContext checks if context is cancelled func (m *manager) checkContext() error { select { case <-m.ctx.Done(): + // Context cancelled - just return the error + // Don't call handleCancellation here to avoid deadlock + // handleCancellation should be called explicitly when needed return m.ctx.Err() default: return nil } } +// handleCancellation marks nodes and trace as cancelled (called when context is done) +func (m *manager) handleCancellation() { + // Mark as completed first - this will trigger state worker to exit + // IMPORTANT: Must mark completed before any state queries to prevent deadlock + if !m.stateMarkCompleted() { + return // Already completed + } + + now := time.Now().Unix() + + // Get current nodes (state worker will process this before exiting) + nodes := m.stateGetCurrentNodes() + + // Mark only running/pending nodes as cancelled + for _, node := range nodes { + if node.Status == types.StatusRunning || node.Status == types.StatusPending { + node.Status = types.StatusCancelled + node.EndTime = now + node.UpdatedAt = now + + // Save node with background context (ignore errors) + _ = m.driver.SaveNode(context.Background(), m.traceID, node) + + // Broadcast cancelled event + m.addUpdateAndBroadcast(&types.TraceUpdate{ + Type: types.UpdateTypeNodeFailed, + TraceID: m.traceID, + NodeID: node.ID, + Timestamp: now, + Data: &types.NodeFailedData{ + NodeID: node.ID, + Status: types.CompleteStatusCancelled, + EndTime: now, + Duration: (now - node.StartTime) * 1000, + Error: "context cancelled", + }, + }) + } + } + + // Update trace status + m.stateSetTraceStatus(types.TraceStatusCancelled) + + // Calculate total duration + totalDuration := int64(0) + rootNode := m.stateGetRoot() + if rootNode != nil && rootNode.CreatedAt > 0 { + totalDuration = (now - rootNode.CreatedAt) * 1000 + } + + // Broadcast trace cancelled event (this will be processed even after state worker starts draining) + m.addUpdateAndBroadcast(&types.TraceUpdate{ + Type: types.UpdateTypeComplete, + TraceID: m.traceID, + Timestamp: now, + Data: &types.TraceCompleteData{ + TraceID: m.traceID, + Status: types.TraceStatusCancelled, + TotalDuration: totalDuration, + }, + }) +} + // newNode creates a node instance that broadcasts events (for external use) func (m *manager) newNode(data *types.TraceNode) types.Node { return &node{ @@ -114,68 +163,63 @@ func (m *manager) newNode(data *types.TraceNode) types.Node { } } -// Helper functions for thread-safe access - -// getCurrentNodes returns a copy of current nodes (thread-safe read) -func (m *manager) getCurrentNodes() []*types.TraceNode { - m.mu.RLock() - defer m.mu.RUnlock() - nodes := make([]*types.TraceNode, len(m.currentNodes)) - copy(nodes, m.currentNodes) - return nodes -} - -// getSpace returns a space by ID (thread-safe read) -func (m *manager) getSpace(id string) (*types.TraceSpace, bool) { - m.mu.RLock() - defer m.mu.RUnlock() - space, ok := m.spaces[id] - return space, ok -} - -// setSpace stores a space (thread-safe write) -func (m *manager) setSpace(id string, space *types.TraceSpace) { - m.mu.Lock() - defer m.mu.Unlock() - m.spaces[id] = space -} - -// deleteSpace removes a space (thread-safe write) -func (m *manager) deleteSpace(id string) { - m.mu.Lock() - defer m.mu.Unlock() - delete(m.spaces, id) -} - -// getAllSpaces returns all spaces (thread-safe read) -func (m *manager) getAllSpaces() []*types.TraceSpace { - m.mu.RLock() - defer m.mu.RUnlock() - spaces := make([]*types.TraceSpace, 0, len(m.spaces)) - for _, space := range m.spaces { - spaces = append(spaces, space) - } - return spaces -} - // Add creates next sequential node - auto-joins if currently in parallel state func (m *manager) Add(input types.TraceInput, option types.TraceNodeOption) (types.Node, error) { if err := m.checkContext(); err != nil { return nil, err } - m.mu.Lock() - defer m.mu.Unlock() - now := time.Now().Unix() - // If in parallel state (multiple current nodes), auto-join first + // Check if root exists + rootNode := m.stateGetRoot() + + if rootNode == nil { + // Create root node + rootNode = &types.TraceNode{ + ID: genNodeID(), + ParentID: "", + Children: []*types.TraceNode{}, + TraceNodeOption: option, + Status: types.StatusRunning, + Input: input, + CreatedAt: now, + StartTime: now, + UpdatedAt: now, + } + + // Save root node + if err := m.driver.SaveNode(m.ctx, m.traceID, rootNode); err != nil { + return nil, fmt.Errorf("failed to save root node: %w", err) + } + + // Update state + m.stateUpdateRootAndCurrent(rootNode, []*types.TraceNode{rootNode}) + + // Update trace status + m.stateSetTraceStatus(types.TraceStatusRunning) + + // Broadcast event + m.addUpdateAndBroadcast(&types.TraceUpdate{ + Type: types.UpdateTypeNodeStart, + TraceID: m.traceID, + NodeID: rootNode.ID, + Timestamp: now, + Data: rootNode.ToStartData(), + }) + + return &node{manager: m, data: rootNode}, nil + } + + // Get current nodes + currentNodes := m.stateGetCurrentNodes() + var parentNode *types.TraceNode - if len(m.currentNodes) > 1 { + if len(currentNodes) > 1 { // Auto-join: create join node parentNode = &types.TraceNode{ ID: genNodeID(), - ParentID: m.currentNodes[0].ParentID, // Same parent as parallel nodes + ParentID: currentNodes[0].ParentID, Children: []*types.TraceNode{}, Status: types.StatusCompleted, CreatedAt: now, @@ -184,15 +228,16 @@ func (m *manager) Add(input types.TraceInput, option types.TraceNodeOption) (typ UpdatedAt: now, TraceNodeOption: types.TraceNodeOption{Label: "Join", Icon: "join"}, } + // Save join node if err := m.driver.SaveNode(m.ctx, m.traceID, parentNode); err != nil { return nil, err } } else { - parentNode = m.currentNodes[0] + parentNode = currentNodes[0] } - // Create new node data + // Create new node newNodeData := &types.TraceNode{ ID: genNodeID(), ParentID: parentNode.ID, @@ -205,7 +250,7 @@ func (m *manager) Add(input types.TraceInput, option types.TraceNodeOption) (typ UpdatedAt: now, } - // Add to parent's children + // Update parent's children parentNode.Children = append(parentNode.Children, newNodeData) // Save nodes @@ -216,11 +261,11 @@ func (m *manager) Add(input types.TraceInput, option types.TraceNodeOption) (typ return nil, err } - // Set as current node - m.currentNodes = []*types.TraceNode{newNodeData} + // Update current nodes + m.stateSetCurrentNodes([]*types.TraceNode{newNodeData}) - // Broadcast node start event - m.addUpdate(&types.TraceUpdate{ + // Broadcast event + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeNodeStart, TraceID: m.traceID, NodeID: newNodeData.ID, @@ -228,11 +273,7 @@ func (m *manager) Add(input types.TraceInput, option types.TraceNodeOption) (typ Data: newNodeData.ToStartData(), }) - // Return Node interface - return &node{ - manager: m, - data: newNodeData, - }, nil + return &node{manager: m, data: newNodeData}, nil } // Parallel creates multiple concurrent child nodes, returns Node interfaces for direct control @@ -241,11 +282,21 @@ func (m *manager) Parallel(parallelInputs []types.TraceParallelInput) ([]types.N return nil, err } - m.mu.Lock() - defer m.mu.Unlock() + if len(parallelInputs) == 0 { + return nil, fmt.Errorf("parallel inputs cannot be empty") + } + + // Check if root exists + if m.stateGetRoot() == nil { + return nil, fmt.Errorf("root node does not exist, please call Add first before using Parallel") + } now := time.Now().Unix() - parentNode := m.currentNodes[0] + + // Get current nodes + currentNodes := m.stateGetCurrentNodes() + parentNode := currentNodes[0] + nodeData := make([]*types.TraceNode, 0, len(parallelInputs)) nodeInterfaces := make([]types.Node, 0, len(parallelInputs)) @@ -265,11 +316,6 @@ func (m *manager) Parallel(parallelInputs []types.TraceParallelInput) ([]types.N nodeData = append(nodeData, data) parentNode.Children = append(parentNode.Children, data) - // Save node - if err := m.driver.SaveNode(m.ctx, m.traceID, data); err != nil { - return nil, err - } - // Create Node interface wrapper nodeInterfaces = append(nodeInterfaces, &node{ manager: m, @@ -277,16 +323,29 @@ func (m *manager) Parallel(parallelInputs []types.TraceParallelInput) ([]types.N }) } - // Save parent node - if err := m.driver.SaveNode(m.ctx, m.traceID, parentNode); err != nil { - return nil, err + // Save all nodes in batch - collect errors + var saveErrors []error + for _, data := range nodeData { + if err := m.driver.SaveNode(m.ctx, m.traceID, data); err != nil { + saveErrors = append(saveErrors, fmt.Errorf("failed to save node %s: %w", data.ID, err)) + } } - // Set all as current nodes (parallel state) - m.currentNodes = nodeData + // Return error if any node failed to save + if len(saveErrors) > 0 { + return nil, fmt.Errorf("failed to save %d node(s): %v", len(saveErrors), saveErrors) + } - // Broadcast parallel nodes as batch (frontend supports data.nodes[]) - m.addUpdate(&types.TraceUpdate{ + // Save parent node + if err := m.driver.SaveNode(m.ctx, m.traceID, parentNode); err != nil { + return nil, fmt.Errorf("failed to save parent node: %w", err) + } + + // Set all as current nodes + m.stateSetCurrentNodes(nodeData) + + // Broadcast parallel nodes as batch + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeNodeStart, TraceID: m.traceID, Timestamp: now, @@ -325,8 +384,8 @@ func (m *manager) log(level string, format string, args ...any) { message := fmt.Sprintf(format, args...) now := time.Now().Unix() - // Get current nodes safely - nodes := m.getCurrentNodes() + // Get current nodes + nodes := m.stateGetCurrentNodes() // Log to all current nodes for _, node := range nodes { @@ -340,7 +399,7 @@ func (m *manager) log(level string, format string, args ...any) { _ = m.driver.SaveLog(m.ctx, m.traceID, log) // Broadcast log event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeLogAdded, TraceID: m.traceID, NodeID: node.ID, @@ -357,7 +416,7 @@ func (m *manager) SetOutput(output types.TraceOutput) error { } now := time.Now().Unix() - nodes := m.getCurrentNodes() + nodes := m.stateGetCurrentNodes() for _, node := range nodes { node.Output = output node.UpdatedAt = now @@ -366,7 +425,7 @@ func (m *manager) SetOutput(output types.TraceOutput) error { } // Broadcast node update event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeNodeUpdated, TraceID: m.traceID, NodeID: node.ID, @@ -384,7 +443,7 @@ func (m *manager) SetMetadata(key string, value any) error { } now := time.Now().Unix() - nodes := m.getCurrentNodes() + nodes := m.stateGetCurrentNodes() for _, node := range nodes { if node.Metadata == nil { node.Metadata = make(map[string]any) @@ -396,7 +455,7 @@ func (m *manager) SetMetadata(key string, value any) error { } // Broadcast node update event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeNodeUpdated, TraceID: m.traceID, NodeID: node.ID, @@ -415,16 +474,30 @@ func (m *manager) Complete(output ...types.TraceOutput) error { } now := time.Now().Unix() - nodes := m.getCurrentNodes() + nodes := m.stateGetCurrentNodes() - // Set output if provided + // Determine output value once + var nodeOutput types.TraceOutput if len(output) > 0 { - for _, node := range nodes { - node.Output = output[0] - } + nodeOutput = output[0] } for _, node := range nodes { + // Set output if provided + if len(output) > 0 { + node.Output = nodeOutput + } + + // Create complete data BEFORE modifying other fields to avoid race + completeData := &types.NodeCompleteData{ + NodeID: node.ID, + Status: types.CompleteStatusSuccess, + EndTime: now, + Duration: (now - node.StartTime) * 1000, + Output: node.Output, + } + + // Now modify node status node.Status = types.StatusCompleted node.EndTime = now node.UpdatedAt = now @@ -432,13 +505,13 @@ func (m *manager) Complete(output ...types.TraceOutput) error { return err } - // Broadcast node complete event - m.addUpdate(&types.TraceUpdate{ + // Broadcast node complete event with pre-created data + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeNodeComplete, TraceID: m.traceID, NodeID: node.ID, Timestamp: now, - Data: node.ToCompleteData(), + Data: completeData, }) } return nil @@ -454,7 +527,7 @@ func (m *manager) Fail(err error) error { // Log error first m.Error("Node failed: %v", err) - nodes := m.getCurrentNodes() + nodes := m.stateGetCurrentNodes() for _, node := range nodes { node.Status = types.StatusFailed node.EndTime = now @@ -464,14 +537,14 @@ func (m *manager) Fail(err error) error { } // Broadcast node failed event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeNodeFailed, TraceID: m.traceID, NodeID: node.ID, Timestamp: now, Data: &types.NodeFailedData{ NodeID: node.ID, - Status: "failed", + Status: types.CompleteStatusFailed, EndTime: now, Duration: (node.EndTime - node.StartTime) * 1000, // Convert to milliseconds Error: err.Error(), @@ -483,7 +556,7 @@ func (m *manager) Fail(err error) error { // GetRootNode returns the root node func (m *manager) GetRootNode() (*types.TraceNode, error) { - return m.rootNode, nil + return m.stateGetRoot(), nil } // GetNode returns a node by ID @@ -493,28 +566,29 @@ func (m *manager) GetNode(id string) (*types.TraceNode, error) { // GetCurrentNodes returns current active nodes func (m *manager) GetCurrentNodes() ([]*types.TraceNode, error) { - return m.getCurrentNodes(), nil + return m.stateGetCurrentNodes(), nil } // MarkComplete marks the entire trace as completed func (m *manager) MarkComplete() error { - m.updatesMu.Lock() - if m.completed { - m.updatesMu.Unlock() + // Try to mark as completed + if !m.stateMarkCompleted() { return nil // Already completed } - m.completed = true - m.updatesMu.Unlock() + + // Update trace status + m.stateSetTraceStatus(types.TraceStatusCompleted) // Calculate total duration from root node now := time.Now().Unix() totalDuration := int64(0) - if m.rootNode != nil && m.rootNode.CreatedAt > 0 { - totalDuration = (now - m.rootNode.CreatedAt) * 1000 // Convert to milliseconds + rootNode := m.stateGetRoot() + if rootNode != nil && rootNode.CreatedAt > 0 { + totalDuration = (now - rootNode.CreatedAt) * 1000 // Convert to milliseconds } // Broadcast completion event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeComplete, TraceID: m.traceID, Timestamp: now, @@ -545,11 +619,11 @@ func (m *manager) CreateSpace(option types.TraceSpaceOption) (*types.TraceSpace, return nil, err } - // Cache in memory (thread-safe) - m.setSpace(space.ID, space) + // Cache in memory + m.stateSetSpace(space.ID, space) // Broadcast space created event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeSpaceCreated, TraceID: m.traceID, SpaceID: space.ID, @@ -562,8 +636,8 @@ func (m *manager) CreateSpace(option types.TraceSpaceOption) (*types.TraceSpace, // GetSpace returns a space by ID func (m *manager) GetSpace(id string) (*types.TraceSpace, error) { - // Check cache first (thread-safe) - if space, ok := m.getSpace(id); ok { + // Check cache first + if space, ok := m.stateGetSpace(id); ok { return space, nil } @@ -573,9 +647,9 @@ func (m *manager) GetSpace(id string) (*types.TraceSpace, error) { return nil, err } - // Cache it (thread-safe) + // Cache it if space != nil { - m.setSpace(id, space) + m.stateSetSpace(id, space) } return space, nil @@ -583,8 +657,8 @@ func (m *manager) GetSpace(id string) (*types.TraceSpace, error) { // HasSpace checks if a space exists func (m *manager) HasSpace(id string) bool { - // Check cache (thread-safe) - if _, ok := m.getSpace(id); ok { + // Check cache + if _, ok := m.stateGetSpace(id); ok { return true } @@ -601,8 +675,8 @@ func (m *manager) DeleteSpace(id string) error { now := time.Now().Unix() - // Remove from cache (thread-safe) - m.deleteSpace(id) + // Remove from cache + m.stateDeleteSpace(id) // Delete from driver if err := m.driver.DeleteSpace(m.ctx, m.traceID, id); err != nil { @@ -610,7 +684,7 @@ func (m *manager) DeleteSpace(id string) error { } // Broadcast space deleted event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeSpaceDeleted, TraceID: m.traceID, SpaceID: id, @@ -626,8 +700,8 @@ func (m *manager) ListSpaces() []*types.TraceSpace { // Load from driver to ensure we have all spaces spaceIDs, err := m.driver.ListSpaces(m.ctx, m.traceID) if err != nil { - // Fallback to cached spaces (thread-safe) - return m.getAllSpaces() + // Fallback to cached spaces + return m.stateGetAllSpaces() } // Load all spaces @@ -648,11 +722,6 @@ func (m *manager) SetSpaceValue(spaceID, key string, value any) error { return err } - // Lock this specific space for concurrent safety - spaceLock := m.getSpaceLock(spaceID) - spaceLock.Lock() - defer spaceLock.Unlock() - now := time.Now().Unix() // Get space @@ -661,19 +730,27 @@ func (m *manager) SetSpaceValue(spaceID, key string, value any) error { return fmt.Errorf("space not found: %s", spaceID) } - // Set value in driver - if err := m.driver.SetSpaceKey(m.ctx, m.traceID, spaceID, key, value); err != nil { - return err - } + // Set value in driver (through state worker for concurrent safety) + err = m.stateExecuteSpaceOp(spaceID, func() error { + if err := m.driver.SetSpaceKey(m.ctx, m.traceID, spaceID, key, value); err != nil { + return err + } - // Update space timestamp - space.UpdatedAt = now - if err := m.driver.SaveSpace(m.ctx, m.traceID, space); err != nil { + // Update space timestamp + space.UpdatedAt = now + if err := m.driver.SaveSpace(m.ctx, m.traceID, space); err != nil { + return err + } + + return nil + }) + + if err != nil { return err } // Broadcast memory_add event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeMemoryAdd, TraceID: m.traceID, SpaceID: spaceID, @@ -686,12 +763,23 @@ func (m *manager) SetSpaceValue(spaceID, key string, value any) error { // GetSpaceValue gets a value from a space func (m *manager) GetSpaceValue(spaceID, key string) (any, error) { - return m.driver.GetSpaceKey(m.ctx, m.traceID, spaceID, key) + var result any + err := m.stateExecuteSpaceOp(spaceID, func() error { + var err error + result, err = m.driver.GetSpaceKey(m.ctx, m.traceID, spaceID, key) + return err + }) + return result, err } // HasSpaceValue checks if a key exists in a space func (m *manager) HasSpaceValue(spaceID, key string) bool { - return m.driver.HasSpaceKey(m.ctx, m.traceID, spaceID, key) + var result bool + _ = m.stateExecuteSpaceOp(spaceID, func() error { + result = m.driver.HasSpaceKey(m.ctx, m.traceID, spaceID, key) + return nil + }) + return result } // DeleteSpaceValue deletes a value from a space and broadcasts memory_delete event @@ -700,20 +788,19 @@ func (m *manager) DeleteSpaceValue(spaceID, key string) error { return err } - // Lock this specific space for concurrent safety - spaceLock := m.getSpaceLock(spaceID) - spaceLock.Lock() - defer spaceLock.Unlock() - now := time.Now().Unix() - // Delete value from driver - if err := m.driver.DeleteSpaceKey(m.ctx, m.traceID, spaceID, key); err != nil { + // Delete value from driver (through state worker for concurrent safety) + err := m.stateExecuteSpaceOp(spaceID, func() error { + return m.driver.DeleteSpaceKey(m.ctx, m.traceID, spaceID, key) + }) + + if err != nil { return err } // Broadcast memory_delete event - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeMemoryDelete, TraceID: m.traceID, SpaceID: spaceID, @@ -730,20 +817,19 @@ func (m *manager) ClearSpaceValues(spaceID string) error { return err } - // Lock this specific space for concurrent safety - spaceLock := m.getSpaceLock(spaceID) - spaceLock.Lock() - defer spaceLock.Unlock() - now := time.Now().Unix() - // Clear values from driver - if err := m.driver.ClearSpaceKeys(m.ctx, m.traceID, spaceID); err != nil { + // Clear values from driver (through state worker for concurrent safety) + err := m.stateExecuteSpaceOp(spaceID, func() error { + return m.driver.ClearSpaceKeys(m.ctx, m.traceID, spaceID) + }) + + if err != nil { return err } // Broadcast memory_delete event (for all keys) - m.addUpdate(&types.TraceUpdate{ + m.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeMemoryDelete, TraceID: m.traceID, SpaceID: spaceID, @@ -756,24 +842,16 @@ func (m *manager) ClearSpaceValues(spaceID string) error { // ListSpaceKeys returns all keys in a space func (m *manager) ListSpaceKeys(spaceID string) []string { - keys, err := m.driver.ListSpaceKeys(m.ctx, m.traceID, spaceID) - if err != nil { - return nil - } + var keys []string + _ = m.stateExecuteSpaceOp(spaceID, func() error { + var err error + keys, err = m.driver.ListSpaceKeys(m.ctx, m.traceID, spaceID) + return err + }) return keys } -// getSpaceLock gets or creates a lock for a specific space (thread-safe) -func (m *manager) getSpaceLock(spaceID string) *sync.RWMutex { - m.mu.Lock() - defer m.mu.Unlock() - - if lock, exists := m.spaceLocks[spaceID]; exists { - return lock - } - - // Create new lock for this space - lock := &sync.RWMutex{} - m.spaceLocks[spaceID] = lock - return lock +// IsComplete returns whether the trace is completed +func (m *manager) IsComplete() bool { + return m.stateIsCompleted() } diff --git a/trace/node.go b/trace/node.go index 34d56874..3245cf26 100644 --- a/trace/node.go +++ b/trace/node.go @@ -42,7 +42,7 @@ func (n *node) logWithBroadcast(level string, format string, args ...any) { log := n.log(level, format, args...) // Broadcast event - n.manager.addUpdate(&types.TraceUpdate{ + n.manager.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeLogAdded, TraceID: n.manager.traceID, NodeID: n.data.ID, @@ -193,7 +193,7 @@ func (n *node) SetMetadata(key string, value any) error { // SetStatus sets the node status func (n *node) SetStatus(status string) error { - n.data.Status = status + n.data.Status = types.NodeStatus(status) n.data.UpdatedAt = time.Now().Unix() return n.manager.driver.SaveNode(n.manager.ctx, n.manager.traceID, n.data) } @@ -206,7 +206,7 @@ func (n *node) Complete(output ...types.TraceOutput) error { } // Broadcast event - n.manager.addUpdate(&types.TraceUpdate{ + n.manager.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeNodeComplete, TraceID: n.manager.traceID, NodeID: n.data.ID, @@ -242,7 +242,7 @@ func (n *node) Fail(err error) error { } // Broadcast event - n.manager.addUpdate(&types.TraceUpdate{ + n.manager.addUpdateAndBroadcast(&types.TraceUpdate{ Type: types.UpdateTypeNodeFailed, TraceID: n.manager.traceID, NodeID: n.data.ID, diff --git a/trace/state.go b/trace/state.go new file mode 100644 index 00000000..a4b26519 --- /dev/null +++ b/trace/state.go @@ -0,0 +1,420 @@ +package trace + +import ( + "time" + + "github.com/yaoapp/yao/trace/types" +) + +// State management using channel-based serialization (no locks needed) +// All state mutations go through a single worker goroutine + +// managerState holds all mutable state (accessed only by state worker) +type managerState struct { + rootNode *types.TraceNode + currentNodes []*types.TraceNode + spaces map[string]*types.TraceSpace + traceStatus types.TraceStatus + completed bool + updates []*types.TraceUpdate + subscribers map[string]chan *types.TraceUpdate +} + +// State command interface - all commands are processed serially +type stateCommand interface { + execute(s *managerState) +} + +// Commands with response channels for synchronous operations + +// --- Root Node Commands --- + +type cmdSetRoot struct { + node *types.TraceNode +} + +func (c *cmdSetRoot) execute(s *managerState) { + s.rootNode = c.node +} + +type cmdGetRoot struct { + resp chan *types.TraceNode +} + +func (c *cmdGetRoot) execute(s *managerState) { + c.resp <- s.rootNode +} + +// --- Current Nodes Commands --- + +type cmdSetCurrentNodes struct { + nodes []*types.TraceNode +} + +func (c *cmdSetCurrentNodes) execute(s *managerState) { + s.currentNodes = c.nodes +} + +type cmdGetCurrentNodes struct { + resp chan []*types.TraceNode +} + +func (c *cmdGetCurrentNodes) execute(s *managerState) { + // Return a copy to prevent external mutation + nodes := make([]*types.TraceNode, len(s.currentNodes)) + copy(nodes, s.currentNodes) + c.resp <- nodes +} + +type cmdUpdateRootAndCurrent struct { + root *types.TraceNode + current []*types.TraceNode +} + +func (c *cmdUpdateRootAndCurrent) execute(s *managerState) { + s.rootNode = c.root + s.currentNodes = c.current +} + +// --- Space Commands --- + +type cmdGetSpace struct { + id string + resp chan *types.TraceSpace +} + +func (c *cmdGetSpace) execute(s *managerState) { + c.resp <- s.spaces[c.id] +} + +type cmdSetSpace struct { + id string + space *types.TraceSpace +} + +func (c *cmdSetSpace) execute(s *managerState) { + s.spaces[c.id] = c.space +} + +type cmdDeleteSpace struct { + id string +} + +func (c *cmdDeleteSpace) execute(s *managerState) { + delete(s.spaces, c.id) +} + +type cmdGetAllSpaces struct { + resp chan []*types.TraceSpace +} + +func (c *cmdGetAllSpaces) execute(s *managerState) { + spaces := make([]*types.TraceSpace, 0, len(s.spaces)) + for _, space := range s.spaces { + spaces = append(spaces, space) + } + c.resp <- spaces +} + +// --- Trace Status Commands --- + +type cmdSetTraceStatus struct { + status types.TraceStatus +} + +func (c *cmdSetTraceStatus) execute(s *managerState) { + s.traceStatus = c.status +} + +type cmdGetTraceStatus struct { + resp chan types.TraceStatus +} + +func (c *cmdGetTraceStatus) execute(s *managerState) { + c.resp <- s.traceStatus +} + +// --- Completion Commands --- + +type cmdMarkCompleted struct { + resp chan bool // Returns true if marked, false if already completed +} + +func (c *cmdMarkCompleted) execute(s *managerState) { + if s.completed { + c.resp <- false + } else { + s.completed = true + c.resp <- true + } +} + +type cmdIsCompleted struct { + resp chan bool +} + +func (c *cmdIsCompleted) execute(s *managerState) { + c.resp <- s.completed +} + +// --- Update Commands --- + +type cmdAddUpdate struct { + update *types.TraceUpdate +} + +func (c *cmdAddUpdate) execute(s *managerState) { + s.updates = append(s.updates, c.update) +} + +type cmdGetUpdates struct { + since int64 + resp chan []*types.TraceUpdate +} + +func (c *cmdGetUpdates) execute(s *managerState) { + filtered := make([]*types.TraceUpdate, 0) + for _, update := range s.updates { + if update.Timestamp >= c.since { + filtered = append(filtered, update) + } + } + c.resp <- filtered +} + +type cmdSetUpdates struct { + updates []*types.TraceUpdate +} + +func (c *cmdSetUpdates) execute(s *managerState) { + s.updates = c.updates +} + +// --- Subscriber Commands --- + +type cmdAddSubscriber struct { + id string + ch chan *types.TraceUpdate +} + +func (c *cmdAddSubscriber) execute(s *managerState) { + s.subscribers[c.id] = c.ch +} + +type cmdRemoveSubscriber struct { + id string +} + +func (c *cmdRemoveSubscriber) execute(s *managerState) { + delete(s.subscribers, c.id) +} + +type cmdGetSubscribers struct { + resp chan map[string]chan *types.TraceUpdate +} + +func (c *cmdGetSubscribers) execute(s *managerState) { + // Return a copy of the map + subs := make(map[string]chan *types.TraceUpdate, len(s.subscribers)) + for id, ch := range s.subscribers { + subs[id] = ch + } + c.resp <- subs +} + +// --- Broadcast Command (special - sends to all subscribers) --- + +type cmdBroadcast struct { + update *types.TraceUpdate +} + +func (c *cmdBroadcast) execute(s *managerState) { + // Send to all subscribers (non-blocking, with panic recovery) + for _, ch := range s.subscribers { + func(channel chan *types.TraceUpdate) { + defer func() { + // Recover from panic if channel is closed + if r := recover(); r != nil { + // Channel was closed, ignore + } + }() + select { + case channel <- c.update: + default: + // Subscriber is slow, skip (non-blocking) + } + }(ch) + } +} + +// --- Space KV Commands (for concurrent safety) --- +// These ensure all operations on a space are serialized through state worker + +type cmdSpaceKVOp struct { + spaceID string + fn func() error + resp chan error +} + +func (c *cmdSpaceKVOp) execute(s *managerState) { + // Execute the operation (typically a driver call) + // The function is provided by caller and executed serially here + err := c.fn() + c.resp <- err +} + +// State worker - processes all commands serially in a single goroutine +func (m *manager) startStateWorker() { + // Initialize state + state := &managerState{ + rootNode: nil, + currentNodes: []*types.TraceNode{}, + spaces: make(map[string]*types.TraceSpace), + traceStatus: types.TraceStatusPending, + completed: false, + updates: make([]*types.TraceUpdate, 0, 100), + subscribers: make(map[string]chan *types.TraceUpdate), + } + + // Process commands until context is cancelled or trace is completed + for { + select { + case cmd, ok := <-m.stateCmdChan: + if !ok { + // Channel closed + return + } + cmd.execute(state) + + // Exit after processing completion + if state.completed { + // Drain remaining commands with timeout + drainTimer := time.NewTimer(100 * time.Millisecond) + defer drainTimer.Stop() + drainLoop: + for { + select { + case cmd := <-m.stateCmdChan: + cmd.execute(state) + case <-drainTimer.C: + break drainLoop + } + } + return + } + case <-m.ctx.Done(): + // Context cancelled - continue processing for a short time to handle cancellation + // Then exit to prevent deadlock + time.Sleep(10 * time.Millisecond) + return + } + } +} + +// Helper methods for manager to send commands + +func (m *manager) stateSetRoot(node *types.TraceNode) { + m.stateCmdChan <- &cmdSetRoot{node: node} +} + +func (m *manager) stateGetRoot() *types.TraceNode { + resp := make(chan *types.TraceNode, 1) + m.stateCmdChan <- &cmdGetRoot{resp: resp} + return <-resp +} + +func (m *manager) stateSetCurrentNodes(nodes []*types.TraceNode) { + m.stateCmdChan <- &cmdSetCurrentNodes{nodes: nodes} +} + +func (m *manager) stateGetCurrentNodes() []*types.TraceNode { + resp := make(chan []*types.TraceNode, 1) + m.stateCmdChan <- &cmdGetCurrentNodes{resp: resp} + return <-resp +} + +func (m *manager) stateUpdateRootAndCurrent(root *types.TraceNode, current []*types.TraceNode) { + m.stateCmdChan <- &cmdUpdateRootAndCurrent{root: root, current: current} +} + +func (m *manager) stateGetSpace(id string) (*types.TraceSpace, bool) { + resp := make(chan *types.TraceSpace, 1) + m.stateCmdChan <- &cmdGetSpace{id: id, resp: resp} + space := <-resp + return space, space != nil +} + +func (m *manager) stateSetSpace(id string, space *types.TraceSpace) { + m.stateCmdChan <- &cmdSetSpace{id: id, space: space} +} + +func (m *manager) stateDeleteSpace(id string) { + m.stateCmdChan <- &cmdDeleteSpace{id: id} +} + +func (m *manager) stateGetAllSpaces() []*types.TraceSpace { + resp := make(chan []*types.TraceSpace, 1) + m.stateCmdChan <- &cmdGetAllSpaces{resp: resp} + return <-resp +} + +func (m *manager) stateSetTraceStatus(status types.TraceStatus) { + m.stateCmdChan <- &cmdSetTraceStatus{status: status} +} + +func (m *manager) stateGetTraceStatus() types.TraceStatus { + resp := make(chan types.TraceStatus, 1) + m.stateCmdChan <- &cmdGetTraceStatus{resp: resp} + return <-resp +} + +func (m *manager) stateMarkCompleted() bool { + resp := make(chan bool, 1) + m.stateCmdChan <- &cmdMarkCompleted{resp: resp} + return <-resp +} + +func (m *manager) stateIsCompleted() bool { + resp := make(chan bool, 1) + m.stateCmdChan <- &cmdIsCompleted{resp: resp} + return <-resp +} + +func (m *manager) stateAddUpdate(update *types.TraceUpdate) { + m.stateCmdChan <- &cmdAddUpdate{update: update} +} + +func (m *manager) stateGetUpdates(since int64) []*types.TraceUpdate { + resp := make(chan []*types.TraceUpdate, 1) + m.stateCmdChan <- &cmdGetUpdates{since: since, resp: resp} + return <-resp +} + +func (m *manager) stateSetUpdates(updates []*types.TraceUpdate) { + m.stateCmdChan <- &cmdSetUpdates{updates: updates} +} + +func (m *manager) stateAddSubscriber(id string, ch chan *types.TraceUpdate) { + m.stateCmdChan <- &cmdAddSubscriber{id: id, ch: ch} +} + +func (m *manager) stateRemoveSubscriber(id string) { + m.stateCmdChan <- &cmdRemoveSubscriber{id: id} +} + +func (m *manager) stateGetSubscribers() map[string]chan *types.TraceUpdate { + resp := make(chan map[string]chan *types.TraceUpdate, 1) + m.stateCmdChan <- &cmdGetSubscribers{resp: resp} + return <-resp +} + +func (m *manager) stateBroadcast(update *types.TraceUpdate) { + m.stateCmdChan <- &cmdBroadcast{update: update} +} + +// stateExecuteSpaceOp executes a space operation serially through state worker +func (m *manager) stateExecuteSpaceOp(spaceID string, fn func() error) error { + resp := make(chan error, 1) + m.stateCmdChan <- &cmdSpaceKVOp{spaceID: spaceID, fn: fn, resp: resp} + return <-resp +} diff --git a/trace/store/driver.go b/trace/store/driver.go index 1fea76a5..d159bcfb 100644 --- a/trace/store/driver.go +++ b/trace/store/driver.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "strings" + "sync" "github.com/yaoapp/gou/store" "github.com/yaoapp/yao/trace/types" @@ -15,6 +16,7 @@ type Driver struct { storeName string // Store name in gou store store.Store // Gou store instance prefix string // Key prefix for isolation + updatesMu sync.Mutex // Protects concurrent updates } // New creates a new store driver @@ -428,6 +430,66 @@ func (d *Driver) DeleteTrace(ctx context.Context, traceID string) error { return nil } +// SaveUpdate persists a trace update event to store (append to list) +func (d *Driver) SaveUpdate(ctx context.Context, traceID string, update *types.TraceUpdate) error { + key := d.getKey(traceID, "updates") + + // Lock to prevent concurrent updates + d.updatesMu.Lock() + defer d.updatesMu.Unlock() + + // Load existing updates + existingUpdates, _ := d.LoadUpdates(ctx, traceID, 0) + + // Append new update + existingUpdates = append(existingUpdates, update) + + // Marshal all updates + data, err := json.Marshal(existingUpdates) + if err != nil { + return fmt.Errorf("failed to marshal updates: %w", err) + } + + // Save back to store + if err := d.store.Set(key, string(data), 0); err != nil { + return fmt.Errorf("failed to save updates to store: %w", err) + } + + return nil +} + +// LoadUpdates loads trace update events from store +func (d *Driver) LoadUpdates(ctx context.Context, traceID string, since int64) ([]*types.TraceUpdate, error) { + key := d.getKey(traceID, "updates") + + // Get data from store + value, ok := d.store.Get(key) + if !ok { + return []*types.TraceUpdate{}, nil + } + + dataStr, ok := value.(string) + if !ok { + return []*types.TraceUpdate{}, nil + } + + // Unmarshal updates array + var allUpdates []*types.TraceUpdate + if err := json.Unmarshal([]byte(dataStr), &allUpdates); err != nil { + return []*types.TraceUpdate{}, nil + } + + // Filter by timestamp + filtered := make([]*types.TraceUpdate, 0) + for _, update := range allUpdates { + if update.Timestamp >= since { + filtered = append(filtered, update) + } + } + + return filtered, nil +} + // Close closes the store driver func (d *Driver) Close() error { // Store connection is managed by gou, no cleanup needed diff --git a/trace/subscription.go b/trace/subscription.go index f77977e4..48a84959 100644 --- a/trace/subscription.go +++ b/trace/subscription.go @@ -1,122 +1,75 @@ package trace import ( + "time" + + gonanoid "github.com/matoous/go-nanoid/v2" "github.com/yaoapp/yao/trace/types" ) -// Subscription Operations - -// addUpdate adds an update to history and broadcasts to subscribers -func (m *manager) addUpdate(update *types.TraceUpdate) { - // Add to history - m.updatesMu.Lock() - m.updates = append(m.updates, update) - m.updatesMu.Unlock() - - // Only broadcast if there are subscribers - m.subMu.RLock() - hasSubscribers := len(m.subscribers) > 0 - m.subMu.RUnlock() - - if hasSubscribers { - // Broadcast to real-time subscribers (non-blocking, in goroutine) - go m.broadcast(update) - } -} - -// broadcast sends update to all active subscribers (non-blocking) -func (m *manager) broadcast(update *types.TraceUpdate) { - m.subMu.RLock() - defer m.subMu.RUnlock() - - for _, ch := range m.subscribers { - // Use recover to handle closed channels safely - func() { - defer func() { - if r := recover(); r != nil { - // Channel was closed, ignore (subscriber cleanup race condition) - } - }() - - select { - case ch <- update: - // Sent successfully - default: - // Channel full, skip (or could log warning) - } - }() - } -} - -// Subscribe subscribes to all trace updates (replay history + real-time) +// Subscribe creates a new subscription for trace updates (real-time from now) func (m *manager) Subscribe() (<-chan *types.TraceUpdate, error) { - return m.SubscribeFrom(0) + return m.subscribe(time.Now().Unix()) } -// SubscribeFrom subscribes from a specific timestamp (for resume) +// SubscribeFrom creates a subscription starting from a specific timestamp func (m *manager) SubscribeFrom(since int64) (<-chan *types.TraceUpdate, error) { - // Create subscriber channel with buffer - ch := make(chan *types.TraceUpdate, 100) - subID := genNodeID() + return m.subscribe(since) +} + +// subscribe is the internal implementation for subscriptions +func (m *manager) subscribe(since int64) (<-chan *types.TraceUpdate, error) { + // Generate unique subscriber ID + subID, _ := gonanoid.Generate("0123456789abcdefghijklmnopqrstuvwxyz", 12) + + // Create update channel + updateCh := make(chan *types.TraceUpdate, 100) // Register subscriber - m.subMu.Lock() - m.subscribers[subID] = ch - m.subMu.Unlock() + m.stateAddSubscriber(subID, updateCh) - // Start replay and streaming goroutine - go m.replayAndStream(ch, subID, since) + // Start replay and stream goroutine (will auto-cleanup on completion) + go m.replayAndStream(subID, updateCh, since) - return ch, nil + return updateCh, nil } -// replayAndStream replays history then streams real-time updates -func (m *manager) replayAndStream(ch chan *types.TraceUpdate, subID string, since int64) { +// replayAndStream replays historical updates and streams new ones +func (m *manager) replayAndStream(subID string, ch chan *types.TraceUpdate, since int64) { + // Auto-cleanup on exit - MUST remove from map before closing channel defer func() { - // Close channel and cleanup subscriber + // Remove from subscribers map first to prevent new broadcasts + m.stateRemoveSubscriber(subID) + // Close channel (any in-flight broadcasts will be caught by recover) close(ch) - m.subMu.Lock() - delete(m.subscribers, subID) - m.subMu.Unlock() }() - // Step 1: Replay history - m.updatesMu.RLock() - history := make([]*types.TraceUpdate, 0) - for _, update := range m.updates { - if update.Timestamp >= since { - history = append(history, update) - } - } - isCompleted := m.completed - m.updatesMu.RUnlock() + // Get historical updates + updates := m.stateGetUpdates(since) - // Send history in order - for _, update := range history { + // Replay historical updates + for _, update := range updates { select { case ch <- update: - // Sent successfully - // Optional: add small delay to control replay speed - // time.Sleep(10 * time.Millisecond) case <-m.ctx.Done(): - // Context cancelled, stop return } } - // Step 2: If already completed, exit - if isCompleted { - return + // Continue streaming new updates + // The channel will receive updates via broadcast from addUpdate + // Monitor completion to know when to exit + ticker := time.NewTicker(100 * time.Millisecond) + defer ticker.Stop() + + for { + select { + case <-ticker.C: + if m.stateIsCompleted() { + return + } + case <-m.ctx.Done(): + return + } } - - // Step 3: Wait for completion or context cancellation - // Real-time updates are sent by broadcast() method - <-m.ctx.Done() -} - -// IsComplete checks if the trace is completed -func (m *manager) IsComplete() bool { - m.updatesMu.RLock() - defer m.updatesMu.RUnlock() - return m.completed } diff --git a/trace/trace.go b/trace/trace.go index 45c1b741..83b1c7e8 100644 --- a/trace/trace.go +++ b/trace/trace.go @@ -147,6 +147,7 @@ func New(ctx context.Context, driver string, option *types.TraceOption, driverOp info := &types.TraceInfo{ ID: traceID, Driver: driver, + Status: types.TraceStatusPending, // Initial status is pending Options: driverOptions, Manager: manager, CreatedAt: now, diff --git a/trace/trace_basic_test.go b/trace/trace_basic_test.go index 1a544200..a5f3cee4 100644 --- a/trace/trace_basic_test.go +++ b/trace/trace_basic_test.go @@ -41,11 +41,22 @@ func TestTraceNew(t *testing.T) { // Verify trace is loaded assert.True(t, trace.IsLoaded(traceID)) - // Get root node + // Root node should be nil initially (lazy initialization) root, err := manager.GetRootNode() assert.NoError(t, err) + assert.Nil(t, root) + + // Add first node - this should become the root + node, err := manager.Add("test input", types.TraceNodeOption{Label: "First Node", Icon: "test"}) + assert.NoError(t, err) + assert.NotNil(t, node) + + // Now root node should exist + root, err = manager.GetRootNode() + assert.NoError(t, err) assert.NotNil(t, root) - assert.Equal(t, "Root", root.Label) + assert.Equal(t, "First Node", root.Label) + assert.Equal(t, "test", root.Icon) }) } } diff --git a/trace/trace_bench_test.go b/trace/trace_bench_test.go index 46afc44d..319fc8b3 100644 --- a/trace/trace_bench_test.go +++ b/trace/trace_bench_test.go @@ -378,6 +378,12 @@ func getTraceScenarios() []traceScenario { { name: "ParallelNodes", execute: func(m types.Manager) error { + // Add first node as root + _, err := m.Add("root", types.TraceNodeOption{Label: "Root"}) + if err != nil { + return err + } + nodes, err := m.Parallel([]types.TraceParallelInput{ {Input: "task1", Option: types.TraceNodeOption{Label: "Task 1"}}, {Input: "task2", Option: types.TraceNodeOption{Label: "Task 2"}}, diff --git a/trace/trace_concurrent_test.go b/trace/trace_concurrent_test.go index c98797ff..0cdb198f 100644 --- a/trace/trace_concurrent_test.go +++ b/trace/trace_concurrent_test.go @@ -18,13 +18,17 @@ func TestConcurrentNodeOperations(t *testing.T) { t.Run(d.Name, func(t *testing.T) { ctx := context.Background() - traceID, manager, err := trace.New(ctx, d.DriverType, nil, d.DriverOptions...) - assert.NoError(t, err) - defer trace.Release(traceID) - defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + traceID, manager, err := trace.New(ctx, d.DriverType, nil, d.DriverOptions...) + assert.NoError(t, err) + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) - // Create parallel nodes - nodes, err := manager.Parallel([]types.TraceParallelInput{ + // Add first node as root + _, err = manager.Add("root", types.TraceNodeOption{Label: "Root"}) + assert.NoError(t, err) + + // Create parallel nodes + nodes, err := manager.Parallel([]types.TraceParallelInput{ {Input: "task 1", Option: types.TraceNodeOption{Label: "Worker 1"}}, {Input: "task 2", Option: types.TraceNodeOption{Label: "Worker 2"}}, {Input: "task 3", Option: types.TraceNodeOption{Label: "Worker 3"}}, diff --git a/trace/trace_mem_test.go b/trace/trace_mem_test.go index 6b62be78..f4f227a5 100644 --- a/trace/trace_mem_test.go +++ b/trace/trace_mem_test.go @@ -208,6 +208,12 @@ func TestMemoryLeakComplexScenarios(t *testing.T) { { name: "ParallelNodes", execute: func(m types.Manager) error { + // Add first node as root + _, err := m.Add("root", types.TraceNodeOption{Label: "Root"}) + if err != nil { + return err + } + nodes, err := m.Parallel([]types.TraceParallelInput{ {Input: "task1", Option: types.TraceNodeOption{Label: "Task 1"}}, {Input: "task2", Option: types.TraceNodeOption{Label: "Task 2"}}, diff --git a/trace/trace_node_test.go b/trace/trace_node_test.go index 27d7de2e..9d668117 100644 --- a/trace/trace_node_test.go +++ b/trace/trace_node_test.go @@ -73,13 +73,17 @@ func TestParallelOperations(t *testing.T) { t.Run(d.Name, func(t *testing.T) { ctx := context.Background() - traceID, manager, err := trace.New(ctx, d.DriverType, nil, d.DriverOptions...) - assert.NoError(t, err) - defer trace.Release(traceID) - defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + traceID, manager, err := trace.New(ctx, d.DriverType, nil, d.DriverOptions...) + assert.NoError(t, err) + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) - // Create parallel nodes - nodes, err := manager.Parallel([]types.TraceParallelInput{ + // Add first node as root + _, err = manager.Add("root", types.TraceNodeOption{Label: "Root"}) + assert.NoError(t, err) + + // Create parallel nodes + nodes, err := manager.Parallel([]types.TraceParallelInput{ { Input: "task A", Option: types.TraceNodeOption{Label: "Worker A", Icon: "cpu"}, diff --git a/trace/trace_subscription_test.go b/trace/trace_subscription_test.go index a7a9c863..e1955b04 100644 --- a/trace/trace_subscription_test.go +++ b/trace/trace_subscription_test.go @@ -37,7 +37,13 @@ func TestSubscription(t *testing.T) { timeout := time.After(2 * time.Second) for { select { - case update := <-updates: + case update, ok := <-updates: + if !ok { + // Channel closed + done <- true + return + } + updatesMu.Lock() receivedUpdates = append(receivedUpdates, update) updatesMu.Unlock() diff --git a/trace/types/driver.go b/trace/types/driver.go index 2015cd7a..4b10e2a4 100644 --- a/trace/types/driver.go +++ b/trace/types/driver.go @@ -68,6 +68,12 @@ type Driver interface { // DeleteTrace removes entire trace and all its data DeleteTrace(ctx context.Context, traceID string) error + // SaveUpdate persists a trace update event to storage + SaveUpdate(ctx context.Context, traceID string, update *TraceUpdate) error + + // LoadUpdates loads trace update events from storage (filtering by timestamp) + LoadUpdates(ctx context.Context, traceID string, since int64) ([]*TraceUpdate, error) + // Close closes the driver and releases resources Close() error } diff --git a/trace/types/events.go b/trace/types/events.go index f34b507a..6aaef184 100644 --- a/trace/types/events.go +++ b/trace/types/events.go @@ -16,7 +16,7 @@ func NodesToStartData(nodes []*TraceNode) *NodeStartData { func (n *TraceNode) ToCompleteData() *NodeCompleteData { return &NodeCompleteData{ NodeID: n.ID, - Status: "success", + Status: CompleteStatusSuccess, EndTime: n.EndTime, Duration: (n.EndTime - n.StartTime) * 1000, // Convert to milliseconds Output: n.Output, @@ -27,7 +27,7 @@ func (n *TraceNode) ToCompleteData() *NodeCompleteData { func (n *TraceNode) ToFailedData(err error) *NodeFailedData { return &NodeFailedData{ NodeID: n.ID, - Status: "failed", + Status: CompleteStatusFailed, EndTime: n.EndTime, Duration: (n.EndTime - n.StartTime) * 1000, // Convert to milliseconds Error: err.Error(), @@ -68,7 +68,7 @@ func NewTraceInitData(traceID string, rootNode *TraceNode, agentName ...string) func NewTraceCompleteData(traceID string, totalDuration int64) *TraceCompleteData { return &TraceCompleteData{ TraceID: traceID, - Status: "completed", + Status: TraceStatusCompleted, TotalDuration: totalDuration, } } diff --git a/trace/types/types.go b/trace/types/types.go index e06366b9..282d0259 100644 --- a/trace/types/types.go +++ b/trace/types/types.go @@ -1,12 +1,38 @@ package types +// NodeStatus represents the status of a node +type NodeStatus string + // Node status constants const ( - StatusPending = "pending" // Node created but not started - StatusRunning = "running" // Node is currently executing - StatusCompleted = "completed" // Node finished successfully - StatusFailed = "failed" // Node failed with error - StatusSkipped = "skipped" // Node was skipped + StatusPending NodeStatus = "pending" // Node created but not started + StatusRunning NodeStatus = "running" // Node is currently executing + StatusCompleted NodeStatus = "completed" // Node finished successfully + StatusFailed NodeStatus = "failed" // Node failed with error + StatusSkipped NodeStatus = "skipped" // Node was skipped + StatusCancelled NodeStatus = "cancelled" // Node was cancelled +) + +// TraceStatus represents the status of a trace +type TraceStatus string + +// Trace status constants +const ( + TraceStatusPending TraceStatus = "pending" // Trace created but not started + TraceStatusRunning TraceStatus = "running" // Trace is running + TraceStatusCompleted TraceStatus = "completed" // Trace completed + TraceStatusFailed TraceStatus = "failed" // Trace failed + TraceStatusCancelled TraceStatus = "cancelled" // Trace was cancelled +) + +// CompleteStatus represents the completion status in events +type CompleteStatus string + +// Complete status constants (for event payloads) +const ( + CompleteStatusSuccess CompleteStatus = "success" // Operation succeeded + CompleteStatusFailed CompleteStatus = "failed" // Operation failed + CompleteStatusCancelled CompleteStatus = "cancelled" // Operation was cancelled ) // TraceNodeOption defines options for creating a node @@ -32,7 +58,7 @@ type TraceNode struct { ParentID string // Parent node ID Children []*TraceNode // Child nodes (for tree structure) TraceNodeOption // Embedded option fields (Label, Icon, Description, Metadata) - Status string // Node status (pending, running, completed, failed, skipped) + Status NodeStatus // Node status (pending, running, completed, failed, skipped) Input TraceInput // Node input data Output TraceOutput // Node output data CreatedAt int64 // Creation timestamp @@ -114,20 +140,20 @@ type NodeStartData struct { // NodeCompleteData payload for "node_complete" event type NodeCompleteData struct { - NodeID string `json:"nodeId"` - Status string `json:"status"` // "success" or "failed" - EndTime int64 `json:"endTime"` - Duration int64 `json:"duration"` // in milliseconds - Output TraceOutput `json:"output,omitempty"` + NodeID string `json:"nodeId"` + Status CompleteStatus `json:"status"` // "success" or "failed" + EndTime int64 `json:"endTime"` + Duration int64 `json:"duration"` // in milliseconds + Output TraceOutput `json:"output,omitempty"` } // NodeFailedData payload for "node_failed" event (same as NodeCompleteData but with error) type NodeFailedData struct { - NodeID string `json:"nodeId"` - Status string `json:"status"` // "failed" - EndTime int64 `json:"endTime"` - Duration int64 `json:"duration"` - Error string `json:"error"` + NodeID string `json:"nodeId"` + Status CompleteStatus `json:"status"` // "failed" + EndTime int64 `json:"endTime"` + Duration int64 `json:"duration"` + Error string `json:"error"` } // MemoryAddData payload for "memory_add" event @@ -148,9 +174,9 @@ type MemoryItem struct { // TraceCompleteData payload for "complete" event type TraceCompleteData struct { - TraceID string `json:"traceId"` - Status string `json:"status"` // "completed" - TotalDuration int64 `json:"totalDuration"` + TraceID string `json:"traceId"` + Status TraceStatus `json:"status"` // "completed" + TotalDuration int64 `json:"totalDuration"` } // SpaceDeletedData payload for "space_deleted" event @@ -169,6 +195,7 @@ type MemoryDeleteData struct { type TraceInfo struct { ID string `json:"id"` Driver string `json:"driver"` + Status TraceStatus `json:"status"` // Trace status Options []any `json:"options,omitempty"` Manager Manager `json:"-"` // Not persisted CreatedAt int64 `json:"created_at"`