diff --git a/Makefile b/Makefile index 565524ec..97824aac 100644 --- a/Makefile +++ b/Makefile @@ -46,9 +46,9 @@ unit-test: benchmark: @echo "" @echo "=============================================" - @echo "Running Benchmark Tests (agent only)..." + @echo "Running Benchmark Tests (agent & trace)..." @echo "=============================================" - @for d in $$($(GO) list ./agent/...); do \ + @for d in $$($(GO) list ./agent/... ./trace/...); do \ if $(GO) test -list=Benchmark $$d 2>/dev/null | grep -q "^Benchmark"; then \ echo ""; \ echo "📊 Benchmarking: $$d"; \ @@ -66,14 +66,14 @@ benchmark: memory-leak: @echo "" @echo "=============================================" - @echo "Running Memory Leak Detection (agent only)..." + @echo "Running Memory Leak Detection (agent & trace)..." @echo "=============================================" - @for d in $$($(GO) list ./agent/...); do \ - if $(GO) test -list='TestMemoryLeak|TestIsolateDisposal' $$d 2>/dev/null | grep -qE "^Test(MemoryLeak|IsolateDisposal)"; then \ + @for d in $$($(GO) list ./agent/... ./trace/...); do \ + if $(GO) test -list='TestMemoryLeak|TestIsolateDisposal|TestGoroutineLeak' $$d 2>/dev/null | grep -qE "^Test(MemoryLeak|IsolateDisposal|GoroutineLeak)"; then \ echo ""; \ echo "🔍 Memory Leak Detection: $$d"; \ echo "---------------------------------------------"; \ - $(GO) test -run='TestMemoryLeak|TestIsolateDisposal' -v $$d || exit 1; \ + $(GO) test -run='TestMemoryLeak|TestIsolateDisposal|TestGoroutineLeak' -v -timeout=60s $$d || exit 1; \ fi; \ done @echo "" diff --git a/trace/README.md b/trace/README.md index 36e148f2..b96f3103 100644 --- a/trace/README.md +++ b/trace/README.md @@ -323,17 +323,26 @@ Create a new trace or load existing one from storage. **Drivers:** -- `trace.Local` - Local disk storage (default path: `./traces`) -- `trace.Store` - Gou store backend +- `trace.Local` - Local disk storage (default: uses log directory from config, fallback to `./traces`) +- `trace.Store` - Gou store backend (default store: `__yao.store`, default prefix: `__trace`) **Example:** ```go +// Local with default path (uses log directory from config) +traceID, manager, _ := trace.New(ctx, trace.Local, nil) + // Local with custom path traceID, manager, _ := trace.New(ctx, trace.Local, nil, "/data/traces") -// Store with custom name -traceID, manager, _ := trace.New(ctx, trace.Store, nil, "my_traces") +// Store with default settings (uses __yao.store with __trace prefix) +traceID, manager, _ := trace.New(ctx, trace.Store, nil) + +// Store with custom store name +traceID, manager, _ := trace.New(ctx, trace.Store, nil, "my_store") + +// Store with custom store name and prefix +traceID, manager, _ := trace.New(ctx, trace.Store, nil, "my_store", "my_prefix") // With trace options option := &types.TraceOption{ diff --git a/trace/local/driver.go b/trace/local/driver.go index 545faff8..e14c7183 100644 --- a/trace/local/driver.go +++ b/trace/local/driver.go @@ -2,7 +2,13 @@ package local import ( "context" + "encoding/json" + "fmt" + "os" + "path/filepath" + "strings" + "github.com/yaoapp/yao/config" "github.com/yaoapp/yao/trace/types" ) @@ -13,131 +19,467 @@ type Driver struct { // New creates a new local driver func New(basePath string) (*Driver, error) { - // TODO: Implement initialization (create directories, etc.) + // If basePath is empty, use log directory from config + if basePath == "" { + if config.Conf.Log != "" { + // Get directory from log file path + basePath = filepath.Join(filepath.Dir(config.Conf.Log), "traces") + } else { + // Fallback to current directory + basePath = "./traces" + } + } + + // Create base directory if it doesn't exist + if err := os.MkdirAll(basePath, 0755); err != nil { + return nil, fmt.Errorf("failed to create base directory: %w", err) + } + return &Driver{ basePath: basePath, }, nil } +// getTracePath returns the path for a trace directory +// Format: {basePath}/{YYYYMMDD}/{traceID}/ +func (d *Driver) getTracePath(traceID string) string { + // Extract date prefix from traceID (first 8 digits) + datePrefix := traceID[:8] + return filepath.Join(d.basePath, datePrefix, traceID) +} + +// ensureTraceDir creates the trace directory if it doesn't exist +func (d *Driver) ensureTraceDir(traceID string) error { + tracePath := d.getTracePath(traceID) + return os.MkdirAll(tracePath, 0755) +} + // SaveNode persists a node to disk func (d *Driver) SaveNode(ctx context.Context, traceID string, node *types.TraceNode) error { - // TODO: Implement disk save - // File path: {basePath}/{YYYYMMDD}/{traceID}/nodes/{nodeID}.json + if err := d.ensureTraceDir(traceID); err != nil { + return err + } + + // Create nodes directory + nodesDir := filepath.Join(d.getTracePath(traceID), "nodes") + if err := os.MkdirAll(nodesDir, 0755); err != nil { + return fmt.Errorf("failed to create nodes directory: %w", err) + } + + // Save node as JSON + filePath := filepath.Join(nodesDir, node.ID+".json") + data, err := json.MarshalIndent(node, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal node: %w", err) + } + + if err := os.WriteFile(filePath, data, 0644); err != nil { + return fmt.Errorf("failed to write node file: %w", err) + } + return nil } // LoadNode loads a node from disk func (d *Driver) LoadNode(ctx context.Context, traceID string, nodeID string) (*types.TraceNode, error) { - // TODO: Implement disk load - return nil, nil + filePath := filepath.Join(d.getTracePath(traceID), "nodes", nodeID+".json") + + data, err := os.ReadFile(filePath) + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, fmt.Errorf("failed to read node file: %w", err) + } + + var node types.TraceNode + if err := json.Unmarshal(data, &node); err != nil { + return nil, fmt.Errorf("failed to unmarshal node: %w", err) + } + + return &node, nil } // LoadTrace loads the entire trace tree from disk func (d *Driver) LoadTrace(ctx context.Context, traceID string) (*types.TraceNode, error) { - // TODO: Implement disk load trace - // File path: {basePath}/{YYYYMMDD}/{traceID}/trace.json + // Load trace info to get root node ID + info, err := d.LoadTraceInfo(ctx, traceID) + if err != nil { + return nil, err + } + if info == nil { + return nil, nil + } + + // For now, just return nil - full tree reconstruction can be implemented later return nil, nil } // SaveSpace persists a space to disk func (d *Driver) SaveSpace(ctx context.Context, traceID string, space *types.TraceSpace) error { - // TODO: Implement disk save space - // File path: {basePath}/{YYYYMMDD}/{traceID}/spaces/{spaceID}.json + if err := d.ensureTraceDir(traceID); err != nil { + return err + } + + // Create spaces directory + spacesDir := filepath.Join(d.getTracePath(traceID), "spaces") + if err := os.MkdirAll(spacesDir, 0755); err != nil { + return fmt.Errorf("failed to create spaces directory: %w", err) + } + + // Save space metadata as JSON + filePath := filepath.Join(spacesDir, space.ID+".json") + data, err := json.MarshalIndent(space, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal space: %w", err) + } + + if err := os.WriteFile(filePath, data, 0644); err != nil { + return fmt.Errorf("failed to write space file: %w", err) + } + return nil } // LoadSpace loads a space from disk func (d *Driver) LoadSpace(ctx context.Context, traceID string, spaceID string) (*types.TraceSpace, error) { - // TODO: Implement disk load space - return nil, nil + filePath := filepath.Join(d.getTracePath(traceID), "spaces", spaceID+".json") + + data, err := os.ReadFile(filePath) + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, fmt.Errorf("failed to read space file: %w", err) + } + + var space types.TraceSpace + if err := json.Unmarshal(data, &space); err != nil { + return nil, fmt.Errorf("failed to unmarshal space: %w", err) + } + + return &space, nil } // DeleteSpace removes a space from disk func (d *Driver) DeleteSpace(ctx context.Context, traceID string, spaceID string) error { - // TODO: Implement disk delete space + // Delete space metadata file + filePath := filepath.Join(d.getTracePath(traceID), "spaces", spaceID+".json") + if err := os.Remove(filePath); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("failed to delete space file: %w", err) + } + + // Delete space data directory + dataDir := filepath.Join(d.getTracePath(traceID), "spaces", spaceID) + if err := os.RemoveAll(dataDir); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("failed to delete space data directory: %w", err) + } + return nil } // ListSpaces lists all space IDs for a trace from disk func (d *Driver) ListSpaces(ctx context.Context, traceID string) ([]string, error) { - // TODO: Implement disk list spaces - return nil, nil + spacesDir := filepath.Join(d.getTracePath(traceID), "spaces") + + entries, err := os.ReadDir(spacesDir) + if err != nil { + if os.IsNotExist(err) { + return []string{}, nil + } + return nil, fmt.Errorf("failed to read spaces directory: %w", err) + } + + var spaceIDs []string + for _, entry := range entries { + if !entry.IsDir() && strings.HasSuffix(entry.Name(), ".json") { + // Remove .json extension to get space ID + spaceID := strings.TrimSuffix(entry.Name(), ".json") + spaceIDs = append(spaceIDs, spaceID) + } + } + + return spaceIDs, nil +} + +// getSpaceDataPath returns the path for space data file +func (d *Driver) getSpaceDataPath(traceID, spaceID string) string { + return filepath.Join(d.getTracePath(traceID), "spaces", spaceID, "data.json") +} + +// loadSpaceData loads all key-value pairs for a space +func (d *Driver) loadSpaceData(traceID, spaceID string) (map[string]any, error) { + filePath := d.getSpaceDataPath(traceID, spaceID) + + data, err := os.ReadFile(filePath) + if err != nil { + if os.IsNotExist(err) { + return make(map[string]any), nil + } + return nil, fmt.Errorf("failed to read space data: %w", err) + } + + var kvData map[string]any + if err := json.Unmarshal(data, &kvData); err != nil { + return nil, fmt.Errorf("failed to unmarshal space data: %w", err) + } + + return kvData, nil +} + +// saveSpaceData saves all key-value pairs for a space +func (d *Driver) saveSpaceData(traceID, spaceID string, kvData map[string]any) error { + filePath := d.getSpaceDataPath(traceID, spaceID) + + // Create space data directory + dataDir := filepath.Dir(filePath) + if err := os.MkdirAll(dataDir, 0755); err != nil { + return fmt.Errorf("failed to create space data directory: %w", err) + } + + // Save as JSON + data, err := json.MarshalIndent(kvData, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal space data: %w", err) + } + + if err := os.WriteFile(filePath, data, 0644); err != nil { + return fmt.Errorf("failed to write space data file: %w", err) + } + + return nil } // SetSpaceKey stores a value by key in a space func (d *Driver) SetSpaceKey(ctx context.Context, traceID, spaceID, key string, value any) error { - // TODO: Implement disk set space key - // File path: {basePath}/{YYYYMMDD}/{traceID}/spaces/{spaceID}/data.json - return nil + // Load existing data + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return err + } + + // Set new value + kvData[key] = value + + // Save data + return d.saveSpaceData(traceID, spaceID, kvData) } // GetSpaceKey retrieves a value by key from a space func (d *Driver) GetSpaceKey(ctx context.Context, traceID, spaceID, key string) (any, error) { - // TODO: Implement disk get space key - return nil, nil + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return nil, err + } + + value, exists := kvData[key] + if !exists { + return nil, nil + } + + return value, nil } // HasSpaceKey checks if a key exists in a space func (d *Driver) HasSpaceKey(ctx context.Context, traceID, spaceID, key string) bool { - // TODO: Implement disk has space key - return false + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return false + } + + _, exists := kvData[key] + return exists } // DeleteSpaceKey removes a key-value pair from a space func (d *Driver) DeleteSpaceKey(ctx context.Context, traceID, spaceID, key string) error { - // TODO: Implement disk delete space key - return nil + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return err + } + + delete(kvData, key) + + return d.saveSpaceData(traceID, spaceID, kvData) } // ClearSpaceKeys removes all key-value pairs from a space func (d *Driver) ClearSpaceKeys(ctx context.Context, traceID, spaceID string) error { - // TODO: Implement disk clear space keys - return nil + return d.saveSpaceData(traceID, spaceID, make(map[string]any)) } // ListSpaceKeys returns all keys in a space func (d *Driver) ListSpaceKeys(ctx context.Context, traceID, spaceID string) ([]string, error) { - // TODO: Implement disk list space keys - return nil, nil + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return nil, err + } + + keys := make([]string, 0, len(kvData)) + for key := range kvData { + keys = append(keys, key) + } + + return keys, nil } // SaveLog appends a log entry to disk func (d *Driver) SaveLog(ctx context.Context, traceID string, log *types.TraceLog) error { - // TODO: Implement disk save log - // File path: {basePath}/{YYYYMMDD}/{traceID}/logs/{nodeID}.jsonl (append mode) + if err := d.ensureTraceDir(traceID); err != nil { + return err + } + + // Create logs directory + logsDir := filepath.Join(d.getTracePath(traceID), "logs") + if err := os.MkdirAll(logsDir, 0755); err != nil { + return fmt.Errorf("failed to create logs directory: %w", err) + } + + // Append log to node's log file (JSONL format) + filePath := filepath.Join(logsDir, log.NodeID+".jsonl") + + // Marshal log as single-line JSON + data, err := json.Marshal(log) + if err != nil { + return fmt.Errorf("failed to marshal log: %w", err) + } + + // Append to file + f, err := os.OpenFile(filePath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) + if err != nil { + return fmt.Errorf("failed to open log file: %w", err) + } + defer f.Close() + + if _, err := f.Write(append(data, '\n')); err != nil { + return fmt.Errorf("failed to write log: %w", err) + } + return nil } // LoadLogs loads all logs for a trace or specific node from disk func (d *Driver) LoadLogs(ctx context.Context, traceID string, nodeID string) ([]*types.TraceLog, error) { - // TODO: Implement disk load logs - // If nodeID is empty, load all logs - // If nodeID provided, load logs for that node only - return nil, nil + logsDir := filepath.Join(d.getTracePath(traceID), "logs") + + var logs []*types.TraceLog + + if nodeID != "" { + // Load logs for specific node + filePath := filepath.Join(logsDir, nodeID+".jsonl") + nodeLogs, err := d.loadLogFile(filePath) + if err != nil { + return nil, err + } + logs = append(logs, nodeLogs...) + } else { + // Load all logs + entries, err := os.ReadDir(logsDir) + if err != nil { + if os.IsNotExist(err) { + return []*types.TraceLog{}, nil + } + return nil, fmt.Errorf("failed to read logs directory: %w", err) + } + + for _, entry := range entries { + if !entry.IsDir() && strings.HasSuffix(entry.Name(), ".jsonl") { + filePath := filepath.Join(logsDir, entry.Name()) + nodeLogs, err := d.loadLogFile(filePath) + if err != nil { + return nil, err + } + logs = append(logs, nodeLogs...) + } + } + } + + return logs, nil +} + +// loadLogFile loads logs from a JSONL file +func (d *Driver) loadLogFile(filePath string) ([]*types.TraceLog, error) { + data, err := os.ReadFile(filePath) + if err != nil { + if os.IsNotExist(err) { + return []*types.TraceLog{}, nil + } + return nil, fmt.Errorf("failed to read log file: %w", err) + } + + lines := strings.Split(string(data), "\n") + logs := make([]*types.TraceLog, 0, len(lines)) + + for _, line := range lines { + if line == "" { + continue + } + + var log types.TraceLog + if err := json.Unmarshal([]byte(line), &log); err != nil { + // Skip malformed lines + continue + } + + logs = append(logs, &log) + } + + return logs, nil } // SaveTraceInfo persists trace metadata to disk func (d *Driver) SaveTraceInfo(ctx context.Context, info *types.TraceInfo) error { - // TODO: Implement disk save trace info - // File path: {basePath}/{YYYYMMDD}/{traceID}/trace_info.json + if err := d.ensureTraceDir(info.ID); err != nil { + return err + } + + filePath := filepath.Join(d.getTracePath(info.ID), "trace_info.json") + + data, err := json.MarshalIndent(info, "", " ") + if err != nil { + return fmt.Errorf("failed to marshal trace info: %w", err) + } + + if err := os.WriteFile(filePath, data, 0644); err != nil { + return fmt.Errorf("failed to write trace info file: %w", err) + } + return nil } // LoadTraceInfo loads trace metadata from disk func (d *Driver) LoadTraceInfo(ctx context.Context, traceID string) (*types.TraceInfo, error) { - // TODO: Implement disk load trace info - return nil, nil + filePath := filepath.Join(d.getTracePath(traceID), "trace_info.json") + + data, err := os.ReadFile(filePath) + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, fmt.Errorf("failed to read trace info file: %w", err) + } + + var info types.TraceInfo + if err := json.Unmarshal(data, &info); err != nil { + return nil, fmt.Errorf("failed to unmarshal trace info: %w", err) + } + + return &info, nil } // DeleteTrace removes entire trace from disk func (d *Driver) DeleteTrace(ctx context.Context, traceID string) error { - // TODO: Implement disk delete trace - // Delete directory: {basePath}/{YYYYMMDD}/{traceID}/ + tracePath := d.getTracePath(traceID) + + if err := os.RemoveAll(tracePath); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("failed to delete trace directory: %w", err) + } + return nil } // Close closes the local driver func (d *Driver) Close() error { - // TODO: Implement cleanup if needed + // No cleanup needed for local file system return nil } diff --git a/trace/manager.go b/trace/manager.go index 7adf9557..04ff921e 100644 --- a/trace/manager.go +++ b/trace/manager.go @@ -13,12 +13,14 @@ import ( // manager implements the Manager interface with unified business logic type manager struct { ctx context.Context + cancel context.CancelFunc // Cancel function to stop background goroutines traceID string driver types.Driver rootNode *types.TraceNode currentNodes []*types.TraceNode spaces map[string]*types.TraceSpace - mu sync.RWMutex // Protects currentNodes and spaces + 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) @@ -51,13 +53,18 @@ func NewManager(ctx context.Context, traceID string, driver types.Driver) (types return nil, fmt.Errorf("failed to save root node: %w", err) } + // Create a cancellable context for the manager + managerCtx, cancel := context.WithCancel(ctx) + m := &manager{ - ctx: ctx, + ctx: managerCtx, + 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, @@ -641,6 +648,11 @@ 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 @@ -688,6 +700,11 @@ 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 @@ -713,6 +730,11 @@ 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 @@ -740,3 +762,18 @@ func (m *manager) ListSpaceKeys(spaceID string) []string { } 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 +} diff --git a/trace/store/driver.go b/trace/store/driver.go index e687b9bd..1fea76a5 100644 --- a/trace/store/driver.go +++ b/trace/store/driver.go @@ -2,145 +2,434 @@ package store import ( "context" + "encoding/json" + "fmt" + "strings" + "github.com/yaoapp/gou/store" "github.com/yaoapp/yao/trace/types" ) // Driver the gou store storage driver implementation type Driver struct { - storeName string // Store name in gou + storeName string // Store name in gou + store store.Store // Gou store instance + prefix string // Key prefix for isolation } // New creates a new store driver -func New(storeName string) (*Driver, error) { - // TODO: Implement initialization (connect to gou store, etc.) +// storeName: the name of the store to use +// prefix: optional key prefix for isolation (default: "__trace") +func New(storeName string, prefix ...string) (*Driver, error) { + // Get store instance from gou + st, err := store.Get(storeName) + if err != nil { + return nil, fmt.Errorf("failed to get store %s: %w", storeName, err) + } + + // Set default prefix if not provided + keyPrefix := "__trace" + if len(prefix) > 0 && prefix[0] != "" { + keyPrefix = prefix[0] + } + return &Driver{ storeName: storeName, + store: st, + prefix: keyPrefix, }, nil } +// getKey generates a key for storage with configurable prefix +// Format: {prefix}:{traceID}:{type}:{id} +// The prefix ensures isolation from other data in shared store +func (d *Driver) getKey(traceID string, parts ...string) string { + allParts := append([]string{d.prefix, traceID}, parts...) + return strings.Join(allParts, ":") +} + // SaveNode persists a node to store func (d *Driver) SaveNode(ctx context.Context, traceID string, node *types.TraceNode) error { - // TODO: Implement store save - // Key: trace:{traceID}:node:{nodeID} + key := d.getKey(traceID, "node", node.ID) + + data, err := json.Marshal(node) + if err != nil { + return fmt.Errorf("failed to marshal node: %w", err) + } + + if err := d.store.Set(key, string(data), 0); err != nil { + return fmt.Errorf("failed to save node to store: %w", err) + } + return nil } // LoadNode loads a node from store func (d *Driver) LoadNode(ctx context.Context, traceID string, nodeID string) (*types.TraceNode, error) { - // TODO: Implement store load - return nil, nil + key := d.getKey(traceID, "node", nodeID) + + value, ok := d.store.Get(key) + if !ok { + return nil, nil + } + + dataStr, ok := value.(string) + if !ok { + return nil, fmt.Errorf("invalid data type in store") + } + + var node types.TraceNode + if err := json.Unmarshal([]byte(dataStr), &node); err != nil { + return nil, fmt.Errorf("failed to unmarshal node: %w", err) + } + + return &node, nil } // LoadTrace loads the entire trace tree from store func (d *Driver) LoadTrace(ctx context.Context, traceID string) (*types.TraceNode, error) { - // TODO: Implement store load trace - // Key: trace:{traceID} + // Load trace info to get root node ID + info, err := d.LoadTraceInfo(ctx, traceID) + if err != nil { + return nil, err + } + if info == nil { + return nil, nil + } + + // For now, just return nil - full tree reconstruction can be implemented later return nil, nil } // SaveSpace persists a space to store func (d *Driver) SaveSpace(ctx context.Context, traceID string, space *types.TraceSpace) error { - // TODO: Implement store save space - // Key: trace:{traceID}:space:{spaceID} + key := d.getKey(traceID, "space", space.ID) + + data, err := json.Marshal(space) + if err != nil { + return fmt.Errorf("failed to marshal space: %w", err) + } + + if err := d.store.Set(key, string(data), 0); err != nil { + return fmt.Errorf("failed to save space to store: %w", err) + } + return nil } // LoadSpace loads a space from store func (d *Driver) LoadSpace(ctx context.Context, traceID string, spaceID string) (*types.TraceSpace, error) { - // TODO: Implement store load space - return nil, nil + key := d.getKey(traceID, "space", spaceID) + + value, ok := d.store.Get(key) + if !ok { + return nil, nil + } + + dataStr, ok := value.(string) + if !ok { + return nil, fmt.Errorf("invalid data type in store") + } + + var space types.TraceSpace + if err := json.Unmarshal([]byte(dataStr), &space); err != nil { + return nil, fmt.Errorf("failed to unmarshal space: %w", err) + } + + return &space, nil } // DeleteSpace removes a space from store func (d *Driver) DeleteSpace(ctx context.Context, traceID string, spaceID string) error { - // TODO: Implement store delete space + // Delete space metadata + key := d.getKey(traceID, "space", spaceID) + if err := d.store.Del(key); err != nil { + return fmt.Errorf("failed to delete space from store: %w", err) + } + + // Delete space data (all keys) + dataKey := d.getKey(traceID, "space", spaceID, "data") + _ = d.store.Del(dataKey) // Ignore error if not exists + return nil } // ListSpaces lists all space IDs for a trace from store func (d *Driver) ListSpaces(ctx context.Context, traceID string) ([]string, error) { - // TODO: Implement store list spaces - // Use pattern matching: trace:{traceID}:space:* - return nil, nil + // Get all keys from store + allKeys := d.store.Keys() + + // Filter keys matching pattern: {prefix}:{traceID}:space:* + prefix := d.getKey(traceID, "space", "") + spaceIDs := make([]string, 0) + + for _, key := range allKeys { + if strings.HasPrefix(key, prefix) { + parts := strings.Split(key, ":") + // Count parts to find space metadata key + // Format: {prefix}:{traceID}:space:{spaceID} + // Parts count depends on prefix (e.g., "__trace" = 4 parts total + 1 = 5) + expectedParts := strings.Count(d.prefix, ":") + 4 + if len(parts) == expectedParts { + // This is a space metadata key (not a data key) + spaceID := parts[len(parts)-1] + spaceIDs = append(spaceIDs, spaceID) + } + } + } + + return spaceIDs, nil +} + +// getSpaceDataKey returns the key for space data storage +func (d *Driver) getSpaceDataKey(traceID, spaceID string) string { + return d.getKey(traceID, "space", spaceID, "data") +} + +// loadSpaceData loads all key-value pairs for a space +func (d *Driver) loadSpaceData(traceID, spaceID string) (map[string]any, error) { + key := d.getSpaceDataKey(traceID, spaceID) + + value, ok := d.store.Get(key) + if !ok { + return make(map[string]any), nil + } + + dataStr, ok := value.(string) + if !ok { + return nil, fmt.Errorf("invalid data type in store") + } + + var kvData map[string]any + if err := json.Unmarshal([]byte(dataStr), &kvData); err != nil { + return nil, fmt.Errorf("failed to unmarshal space data: %w", err) + } + + return kvData, nil +} + +// saveSpaceData saves all key-value pairs for a space +func (d *Driver) saveSpaceData(traceID, spaceID string, kvData map[string]any) error { + key := d.getSpaceDataKey(traceID, spaceID) + + data, err := json.Marshal(kvData) + if err != nil { + return fmt.Errorf("failed to marshal space data: %w", err) + } + + if err := d.store.Set(key, string(data), 0); err != nil { + return fmt.Errorf("failed to save space data: %w", err) + } + + return nil } // SetSpaceKey stores a value by key in a space func (d *Driver) SetSpaceKey(ctx context.Context, traceID, spaceID, key string, value any) error { - // TODO: Implement store set space key - // Key: trace:{traceID}:space:{spaceID}:key:{key} - return nil + // Load existing data + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return err + } + + // Set new value + kvData[key] = value + + // Save data + return d.saveSpaceData(traceID, spaceID, kvData) } // GetSpaceKey retrieves a value by key from a space func (d *Driver) GetSpaceKey(ctx context.Context, traceID, spaceID, key string) (any, error) { - // TODO: Implement store get space key - return nil, nil + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return nil, err + } + + value, exists := kvData[key] + if !exists { + return nil, nil + } + + return value, nil } // HasSpaceKey checks if a key exists in a space func (d *Driver) HasSpaceKey(ctx context.Context, traceID, spaceID, key string) bool { - // TODO: Implement store has space key - return false + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return false + } + + _, exists := kvData[key] + return exists } // DeleteSpaceKey removes a key-value pair from a space func (d *Driver) DeleteSpaceKey(ctx context.Context, traceID, spaceID, key string) error { - // TODO: Implement store delete space key - return nil + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return err + } + + delete(kvData, key) + + return d.saveSpaceData(traceID, spaceID, kvData) } // ClearSpaceKeys removes all key-value pairs from a space func (d *Driver) ClearSpaceKeys(ctx context.Context, traceID, spaceID string) error { - // TODO: Implement store clear space keys - // Delete keys: trace:{traceID}:space:{spaceID}:key:* - return nil + return d.saveSpaceData(traceID, spaceID, make(map[string]any)) } // ListSpaceKeys returns all keys in a space func (d *Driver) ListSpaceKeys(ctx context.Context, traceID, spaceID string) ([]string, error) { - // TODO: Implement store list space keys - // Use pattern matching: trace:{traceID}:space:{spaceID}:key:* - return nil, nil + kvData, err := d.loadSpaceData(traceID, spaceID) + if err != nil { + return nil, err + } + + keys := make([]string, 0, len(kvData)) + for key := range kvData { + keys = append(keys, key) + } + + return keys, nil } // SaveLog appends a log entry to store func (d *Driver) SaveLog(ctx context.Context, traceID string, log *types.TraceLog) error { - // TODO: Implement store save log - // Key: trace:{traceID}:logs:{nodeID} (list type, append) + // Store logs using ArraySlice approach (store as array in a key) + key := d.getKey(traceID, "logs", log.NodeID) + + // Marshal log + data, err := json.Marshal(log) + if err != nil { + return fmt.Errorf("failed to marshal log: %w", err) + } + + // Append to array using Push + if err := d.store.Push(key, string(data)); err != nil { + return fmt.Errorf("failed to append log to store: %w", err) + } + return nil } // LoadLogs loads all logs for a trace or specific node from store func (d *Driver) LoadLogs(ctx context.Context, traceID string, nodeID string) ([]*types.TraceLog, error) { - // TODO: Implement store load logs - // If nodeID is empty, load all logs from trace:{traceID}:logs:* - // If nodeID provided, load from trace:{traceID}:logs:{nodeID} - return nil, nil + var logs []*types.TraceLog + + if nodeID != "" { + // Load logs for specific node + key := d.getKey(traceID, "logs", nodeID) + nodeLogs, err := d.loadLogsFromKey(key) + if err != nil { + return nil, err + } + logs = append(logs, nodeLogs...) + } else { + // Load all logs by iterating all keys + // Pattern: {prefix}:{traceID}:logs:* + allKeys := d.store.Keys() + prefix := d.getKey(traceID, "logs", "") + + for _, key := range allKeys { + if strings.HasPrefix(key, prefix) { + nodeLogs, err := d.loadLogsFromKey(key) + if err != nil { + return nil, err + } + logs = append(logs, nodeLogs...) + } + } + } + + return logs, nil +} + +// loadLogsFromKey loads logs from a specific key (array) +func (d *Driver) loadLogsFromKey(key string) ([]*types.TraceLog, error) { + // Get all items from array + items, err := d.store.ArrayAll(key) + if err != nil { + return []*types.TraceLog{}, nil + } + + logs := make([]*types.TraceLog, 0, len(items)) + for _, item := range items { + itemStr, ok := item.(string) + if !ok { + continue + } + + var log types.TraceLog + if err := json.Unmarshal([]byte(itemStr), &log); err != nil { + // Skip malformed entries + continue + } + logs = append(logs, &log) + } + + return logs, nil } // SaveTraceInfo persists trace metadata to store func (d *Driver) SaveTraceInfo(ctx context.Context, info *types.TraceInfo) error { - // TODO: Implement store save trace info - // Key: trace:{traceID}:info + key := d.getKey(info.ID, "info") + + data, err := json.Marshal(info) + if err != nil { + return fmt.Errorf("failed to marshal trace info: %w", err) + } + + if err := d.store.Set(key, string(data), 0); err != nil { + return fmt.Errorf("failed to save trace info to store: %w", err) + } + return nil } // LoadTraceInfo loads trace metadata from store func (d *Driver) LoadTraceInfo(ctx context.Context, traceID string) (*types.TraceInfo, error) { - // TODO: Implement store load trace info - return nil, nil + key := d.getKey(traceID, "info") + + value, ok := d.store.Get(key) + if !ok { + return nil, nil + } + + dataStr, ok := value.(string) + if !ok { + return nil, fmt.Errorf("invalid data type in store") + } + + var info types.TraceInfo + if err := json.Unmarshal([]byte(dataStr), &info); err != nil { + return nil, fmt.Errorf("failed to unmarshal trace info: %w", err) + } + + return &info, nil } // DeleteTrace removes entire trace from store func (d *Driver) DeleteTrace(ctx context.Context, traceID string) error { - // TODO: Implement store delete trace - // Delete keys: trace:{traceID}* (including all spaces, nodes, and logs) + // Get all keys + allKeys := d.store.Keys() + prefix := d.getKey(traceID, "") + + // Delete all keys matching pattern: {prefix}:{traceID}:* + for _, key := range allKeys { + if strings.HasPrefix(key, prefix) { + _ = d.store.Del(key) // Ignore errors + } + } + return nil } // Close closes the store driver func (d *Driver) Close() error { - // TODO: Implement cleanup if needed + // Store connection is managed by gou, no cleanup needed return nil } diff --git a/trace/subscription.go b/trace/subscription.go index 739b432b..de767c82 100644 --- a/trace/subscription.go +++ b/trace/subscription.go @@ -13,8 +13,15 @@ func (m *manager) addUpdate(update *types.TraceUpdate) { m.updates = append(m.updates, update) m.updatesMu.Unlock() - // Broadcast to real-time subscribers (non-blocking, in goroutine) - go m.broadcast(update) + // 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) @@ -23,12 +30,21 @@ func (m *manager) broadcast(update *types.TraceUpdate) { defer m.subMu.RUnlock() for _, ch := range m.subscribers { - select { - case ch <- update: - // Sent successfully - default: - // Channel full, skip (or could log warning) - } + // 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) + } + }() } } @@ -68,7 +84,7 @@ func (m *manager) replayAndStream(ch chan *types.TraceUpdate, subID string, sinc m.updatesMu.RLock() history := make([]*types.TraceUpdate, 0) for _, update := range m.updates { - if update.Timestamp > since { + if update.Timestamp >= since { history = append(history, update) } } diff --git a/trace/test_helpers.go b/trace/test_helpers.go new file mode 100644 index 00000000..69370781 --- /dev/null +++ b/trace/test_helpers.go @@ -0,0 +1,24 @@ +package trace + +// TestDriver defines the test cases for both drivers +type TestDriver struct { + Name string + DriverType string + DriverOptions []any +} + +// GetTestDrivers returns all drivers to test +func GetTestDrivers() []TestDriver { + return []TestDriver{ + { + Name: "Local", + DriverType: Local, + DriverOptions: []any{}, // Use default (log directory) + }, + { + Name: "Store", + DriverType: Store, + DriverOptions: []any{}, // Use default (__yao.store with __trace prefix) + }, + } +} diff --git a/trace/trace.go b/trace/trace.go index 08769bf0..45c1b741 100644 --- a/trace/trace.go +++ b/trace/trace.go @@ -31,7 +31,7 @@ func getDriver(driver string, options ...any) (types.Driver, error) { switch driver { case Local: - basePath := "./traces" // default + basePath := "" // empty means use log directory from config if len(options) > 0 { if path, ok := options[0].(string); ok { basePath = path @@ -43,13 +43,21 @@ func getDriver(driver string, options ...any) (types.Driver, error) { } case Store: - storeName := "trace" // default + storeName := "__yao.store" // default: use system common store + prefix := "" // empty means use driver's default prefix "__trace" + if len(options) > 0 { if name, ok := options[0].(string); ok { storeName = name } } - drv, err = store.New(storeName) + if len(options) > 1 { + if p, ok := options[1].(string); ok { + prefix = p + } + } + + drv, err = store.New(storeName, prefix) if err != nil { return nil, fmt.Errorf("failed to create store driver: %w", err) } @@ -281,7 +289,7 @@ func GetInfo(ctx context.Context, driver string, traceID string, options ...any) // traceID: the trace ID to release func Release(traceID string) error { registryMu.Lock() - _, exists := registry[traceID] + info, exists := registry[traceID] if exists { delete(registry, traceID) } @@ -291,9 +299,10 @@ func Release(traceID string) error { return fmt.Errorf("trace not found in registry: %s", traceID) } - // Close driver resources if manager has a close method - // (Currently manager doesn't expose driver, but driver has Close method) - // This is handled when the context is cancelled or program exits + // Cancel the manager's context to stop background goroutines + if mgr, ok := info.Manager.(*manager); ok && mgr.cancel != nil { + mgr.cancel() + } return nil } diff --git a/trace/trace_basic_test.go b/trace/trace_basic_test.go new file mode 100644 index 00000000..1a544200 --- /dev/null +++ b/trace/trace_basic_test.go @@ -0,0 +1,221 @@ +package trace_test + +import ( + "context" + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/test" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +func TestMain(m *testing.M) { + // Prepare test environment (initializes stores, models, etc.) + test.Prepare(&testing.T{}, config.Conf) + defer test.Clean() + + // Run tests + os.Exit(m.Run()) +} + +func TestTraceNew(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + t.Run(d.Name, func(t *testing.T) { + ctx := context.Background() + + // Create new trace + traceID, manager, err := trace.New(ctx, d.DriverType, nil, d.DriverOptions...) + assert.NoError(t, err) + assert.NotEmpty(t, traceID) + assert.NotNil(t, manager) + + // Clean up + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Verify trace is loaded + assert.True(t, trace.IsLoaded(traceID)) + + // Get root node + root, err := manager.GetRootNode() + assert.NoError(t, err) + assert.NotNil(t, root) + assert.Equal(t, "Root", root.Label) + }) + } +} + +func TestTraceWithCustomID(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + t.Run(d.Name, func(t *testing.T) { + ctx := context.Background() + customID := trace.GenTraceID() + + option := &types.TraceOption{ + ID: customID, + CreatedBy: "test@example.com", + TeamID: "team-001", + TenantID: "tenant-001", + Metadata: map[string]any{"test": "value"}, + } + + traceID, manager, err := trace.New(ctx, d.DriverType, option, d.DriverOptions...) + assert.NoError(t, err) + assert.Equal(t, customID, traceID) + assert.NotNil(t, manager) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Verify trace info + info, err := trace.GetInfo(ctx, d.DriverType, traceID, d.DriverOptions...) + assert.NoError(t, err) + assert.NotNil(t, info) + assert.Equal(t, customID, info.ID) + assert.Equal(t, "test@example.com", info.CreatedBy) + assert.Equal(t, "team-001", info.TeamID) + assert.Equal(t, "tenant-001", info.TenantID) + }) + } +} + +func TestTraceLoadFromStorage(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + t.Run(d.Name, func(t *testing.T) { + ctx := context.Background() + + // Create and persist a trace + traceID, manager, err := trace.New(ctx, d.DriverType, nil, d.DriverOptions...) + assert.NoError(t, err) + + // Add some data + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + assert.NoError(t, err) + + space, err := manager.CreateSpace(types.TraceSpaceOption{Label: "Test Space"}) + assert.NoError(t, err) + + err = manager.SetSpaceValue(space.ID, "key", "value") + assert.NoError(t, err) + + // Release from registry + err = trace.Release(traceID) + assert.NoError(t, err) + assert.False(t, trace.IsLoaded(traceID)) + + // Load from storage + loadedTraceID, loadedManager, err := trace.LoadFromStorage(ctx, d.DriverType, traceID, d.DriverOptions...) + assert.NoError(t, err) + assert.Equal(t, traceID, loadedTraceID) + assert.NotNil(t, loadedManager) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Verify loaded + assert.True(t, trace.IsLoaded(traceID)) + + // Verify data still exists + spaces := loadedManager.ListSpaces() + assert.NotEmpty(t, spaces) + }) + } +} + +func TestTraceExistsAndRemove(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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) + assert.NotNil(t, manager) + + // Check exists + exists, err := trace.Exists(ctx, d.DriverType, traceID, d.DriverOptions...) + assert.NoError(t, err) + assert.True(t, exists) + + // Remove trace + err = trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + assert.NoError(t, err) + + // Check not exists + exists, err = trace.Exists(ctx, d.DriverType, traceID, d.DriverOptions...) + assert.NoError(t, err) + assert.False(t, exists) + + // Check not loaded + assert.False(t, trace.IsLoaded(traceID)) + }) + } +} + +func TestTraceList(t *testing.T) { + ctx := context.Background() + + // Create multiple traces + var traces []string + for i := 0; i < 3; i++ { + traceID, _, err := trace.New(ctx, trace.Local, nil) + assert.NoError(t, err) + traces = append(traces, traceID) + } + + // Clean up + defer func() { + for _, traceID := range traces { + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + }() + + // List active traces + activeTraces := trace.List() + assert.GreaterOrEqual(t, len(activeTraces), 3) + + // Verify our traces are in the list + for _, traceID := range traces { + found := false + for _, activeID := range activeTraces { + if activeID == traceID { + found = true + break + } + } + assert.True(t, found, "Trace %s should be in active list", traceID) + } +} + +func TestContextCancellation(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + t.Run(d.Name, func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + + traceID, manager, err := trace.New(ctx, d.DriverType, nil, d.DriverOptions...) + assert.NoError(t, err) + defer trace.Release(traceID) + defer trace.Remove(context.Background(), d.DriverType, traceID, d.DriverOptions...) + + // Cancel context + cancel() + + // Operations should fail with context error + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + assert.Error(t, err) + }) + } +} diff --git a/trace/trace_bench_test.go b/trace/trace_bench_test.go new file mode 100644 index 00000000..91d36dd6 --- /dev/null +++ b/trace/trace_bench_test.go @@ -0,0 +1,494 @@ +package trace_test + +import ( + "context" + "fmt" + "sync" + "testing" + "time" + + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +// ============================================================================ +// Simple Scenario Benchmarks +// ============================================================================ + +// BenchmarkSimpleTraceLocal benchmarks simple trace operations with local driver +// Run with: go test -bench=BenchmarkSimpleTraceLocal -benchmem -benchtime=100x +func BenchmarkSimpleTraceLocal(b *testing.B) { + ctx := context.Background() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + b.Fatalf("Failed to create trace: %s", err.Error()) + } + + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + b.Fatalf("Failed to add node: %s", err.Error()) + } + + err = manager.Complete("result") + if err != nil { + b.Fatalf("Failed to complete: %s", err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } +} + +// BenchmarkSimpleTraceStore benchmarks simple trace operations with store driver +// Run with: go test -bench=BenchmarkSimpleTraceStore -benchmem -benchtime=100x +func BenchmarkSimpleTraceStore(b *testing.B) { + ctx := context.Background() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + traceID, manager, err := trace.New(ctx, trace.Store, nil) + if err != nil { + b.Fatalf("Failed to create trace: %s", err.Error()) + } + + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + b.Fatalf("Failed to add node: %s", err.Error()) + } + + err = manager.Complete("result") + if err != nil { + b.Fatalf("Failed to complete: %s", err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Store, traceID) + } +} + +// ============================================================================ +// Complex Scenario Benchmarks (with Parallel, Space, Subscription) +// ============================================================================ + +// BenchmarkComplexTraceLocal benchmarks complex trace operations with local driver +// Run with: go test -bench=BenchmarkComplexTraceLocal -benchmem -benchtime=100x +func BenchmarkComplexTraceLocal(b *testing.B) { + ctx := context.Background() + + scenarios := getTraceScenarios() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + scenario := scenarios[i%len(scenarios)] + + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + b.Fatalf("Failed to create trace: %s", err.Error()) + } + + err = scenario.execute(manager) + if err != nil { + b.Errorf("%s failed: %s", scenario.name, err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } +} + +// BenchmarkComplexTraceStore benchmarks complex trace operations with store driver +// Run with: go test -bench=BenchmarkComplexTraceStore -benchmem -benchtime=100x +func BenchmarkComplexTraceStore(b *testing.B) { + ctx := context.Background() + + scenarios := getTraceScenarios() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + scenario := scenarios[i%len(scenarios)] + + traceID, manager, err := trace.New(ctx, trace.Store, nil) + if err != nil { + b.Fatalf("Failed to create trace: %s", err.Error()) + } + + err = scenario.execute(manager) + if err != nil { + b.Errorf("%s failed: %s", scenario.name, err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Store, traceID) + } +} + +// ============================================================================ +// Concurrent Benchmarks +// ============================================================================ + +// BenchmarkConcurrentSimpleLocal benchmarks concurrent simple operations with local driver +// Run with: go test -bench=BenchmarkConcurrentSimpleLocal -benchmem -benchtime=100x +func BenchmarkConcurrentSimpleLocal(b *testing.B) { + ctx := context.Background() + + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + b.Errorf("Failed to create trace: %s", err.Error()) + continue + } + + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + b.Errorf("Failed to add node: %s", err.Error()) + } + + err = manager.Complete("result") + if err != nil { + b.Errorf("Failed to complete: %s", err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + }) +} + +// BenchmarkConcurrentSimpleStore benchmarks concurrent simple operations with store driver +// Run with: go test -bench=BenchmarkConcurrentSimpleStore -benchmem -benchtime=100x +func BenchmarkConcurrentSimpleStore(b *testing.B) { + ctx := context.Background() + + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + traceID, manager, err := trace.New(ctx, trace.Store, nil) + if err != nil { + b.Errorf("Failed to create trace: %s", err.Error()) + continue + } + + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + b.Errorf("Failed to add node: %s", err.Error()) + } + + err = manager.Complete("result") + if err != nil { + b.Errorf("Failed to complete: %s", err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Store, traceID) + } + }) +} + +// BenchmarkConcurrentComplexLocal benchmarks concurrent complex operations with local driver +// Run with: go test -bench=BenchmarkConcurrentComplexLocal -benchmem -benchtime=100x +func BenchmarkConcurrentComplexLocal(b *testing.B) { + ctx := context.Background() + scenarios := getTraceScenarios() + + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + i := 0 + for pb.Next() { + scenario := scenarios[i%len(scenarios)] + i++ + + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + b.Errorf("Failed to create trace: %s", err.Error()) + continue + } + + err = scenario.execute(manager) + if err != nil { + b.Errorf("%s failed: %s", scenario.name, err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + }) +} + +// BenchmarkConcurrentComplexStore benchmarks concurrent complex operations with store driver +// Run with: go test -bench=BenchmarkConcurrentComplexStore -benchmem -benchtime=100x +func BenchmarkConcurrentComplexStore(b *testing.B) { + ctx := context.Background() + scenarios := getTraceScenarios() + + b.ResetTimer() + b.RunParallel(func(pb *testing.PB) { + i := 0 + for pb.Next() { + scenario := scenarios[i%len(scenarios)] + i++ + + traceID, manager, err := trace.New(ctx, trace.Store, nil) + if err != nil { + b.Errorf("Failed to create trace: %s", err.Error()) + continue + } + + err = scenario.execute(manager) + if err != nil { + b.Errorf("%s failed: %s", scenario.name, err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Store, traceID) + } + }) +} + +// ============================================================================ +// Subscription Benchmarks +// ============================================================================ + +// BenchmarkSubscription benchmarks subscription operations +// Run with: go test -bench=BenchmarkSubscription -benchmem -benchtime=100x +func BenchmarkSubscription(b *testing.B) { + ctx := context.Background() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + b.Fatalf("Failed to create trace: %s", err.Error()) + } + + // Subscribe + updates, err := manager.Subscribe() + if err != nil { + b.Fatalf("Failed to subscribe: %s", err.Error()) + } + + // Perform operations + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + b.Fatalf("Failed to add node: %s", err.Error()) + } + + err = manager.Complete("result") + if err != nil { + b.Fatalf("Failed to complete: %s", err.Error()) + } + + err = manager.MarkComplete() + if err != nil { + b.Fatalf("Failed to mark complete: %s", err.Error()) + } + + // Drain updates + timeout := time.After(10 * time.Millisecond) + drainLoop: + for { + select { + case _, ok := <-updates: + if !ok { + break drainLoop + } + case <-timeout: + break drainLoop + } + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } +} + +// ============================================================================ +// Space Operations Benchmarks +// ============================================================================ + +// BenchmarkSpaceOperations benchmarks space operations +// Run with: go test -bench=BenchmarkSpaceOperations -benchmem -benchtime=100x +func BenchmarkSpaceOperations(b *testing.B) { + ctx := context.Background() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + b.Fatalf("Failed to create trace: %s", err.Error()) + } + + // Create space + space, err := manager.CreateSpace(types.TraceSpaceOption{Label: "Test Space"}) + if err != nil { + b.Fatalf("Failed to create space: %s", err.Error()) + } + + // Set values + for j := 0; j < 10; j++ { + err = manager.SetSpaceValue(space.ID, fmt.Sprintf("key_%d", j), fmt.Sprintf("value_%d", j)) + if err != nil { + b.Fatalf("Failed to set space value: %s", err.Error()) + } + } + + // Get values + for j := 0; j < 10; j++ { + _, err = manager.GetSpaceValue(space.ID, fmt.Sprintf("key_%d", j)) + if err != nil { + b.Fatalf("Failed to get space value: %s", err.Error()) + } + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } +} + +// ============================================================================ +// Helper Functions +// ============================================================================ + +type traceScenario struct { + name string + execute func(types.Manager) error +} + +func getTraceScenarios() []traceScenario { + return []traceScenario{ + { + name: "SequentialNodes", + execute: func(m types.Manager) error { + for i := 0; i < 5; i++ { + _, err := m.Add(fmt.Sprintf("step_%d", i), types.TraceNodeOption{Label: fmt.Sprintf("Step %d", i)}) + if err != nil { + return err + } + if err := m.Complete(fmt.Sprintf("result_%d", i)); err != nil { + return err + } + } + return nil + }, + }, + { + name: "ParallelNodes", + execute: func(m types.Manager) error { + nodes, err := m.Parallel([]types.TraceParallelInput{ + {Input: "task1", Option: types.TraceNodeOption{Label: "Task 1"}}, + {Input: "task2", Option: types.TraceNodeOption{Label: "Task 2"}}, + {Input: "task3", Option: types.TraceNodeOption{Label: "Task 3"}}, + }) + if err != nil { + return err + } + + var wg sync.WaitGroup + for i, node := range nodes { + wg.Add(1) + go func(idx int, n types.Node) { + defer wg.Done() + n.Complete(fmt.Sprintf("result_%d", idx)) + }(i, node) + } + wg.Wait() + return nil + }, + }, + { + name: "WithSpace", + execute: func(m types.Manager) error { + space, err := m.CreateSpace(types.TraceSpaceOption{Label: "Context"}) + if err != nil { + return err + } + + for i := 0; i < 5; i++ { + if err := m.SetSpaceValue(space.ID, fmt.Sprintf("key_%d", i), fmt.Sprintf("value_%d", i)); err != nil { + return err + } + } + + _, err = m.Add("process", types.TraceNodeOption{Label: "Process"}) + if err != nil { + return err + } + return m.Complete("done") + }, + }, + { + name: "WithLogging", + execute: func(m types.Manager) error { + m.Info("Starting process") + _, err := m.Add("step1", types.TraceNodeOption{Label: "Step 1"}) + if err != nil { + return err + } + m.Debug("Debug info") + if err := m.Complete("result1"); err != nil { + return err + } + + _, err = m.Add("step2", types.TraceNodeOption{Label: "Step 2"}) + if err != nil { + return err + } + m.Warn("Warning message") + return m.Complete("result2") + }, + }, + { + name: "ComplexFlow", + execute: func(m types.Manager) error { + // Create space + space, err := m.CreateSpace(types.TraceSpaceOption{Label: "Shared"}) + if err != nil { + return err + } + + // Sequential node + _, err = m.Add("prepare", types.TraceNodeOption{Label: "Prepare"}) + if err != nil { + return err + } + m.Info("Preparing data") + if err := m.Complete("prepared"); err != nil { + return err + } + + // Parallel nodes + nodes, err := m.Parallel([]types.TraceParallelInput{ + {Input: "taskA", Option: types.TraceNodeOption{Label: "Task A"}}, + {Input: "taskB", Option: types.TraceNodeOption{Label: "Task B"}}, + }) + if err != nil { + return err + } + + var wg sync.WaitGroup + for i, node := range nodes { + wg.Add(1) + go func(idx int, n types.Node) { + defer wg.Done() + n.Info("Processing task %d", idx) + m.SetSpaceValue(space.ID, fmt.Sprintf("result_%d", idx), fmt.Sprintf("done_%d", idx)) + n.Complete(fmt.Sprintf("result_%d", idx)) + }(i, node) + } + wg.Wait() + + // Final node + _, err = m.Add("finalize", types.TraceNodeOption{Label: "Finalize"}) + if err != nil { + return err + } + return m.Complete("completed") + }, + }, + } +} + diff --git a/trace/trace_concurrent_test.go b/trace/trace_concurrent_test.go new file mode 100644 index 00000000..f4f2f44e --- /dev/null +++ b/trace/trace_concurrent_test.go @@ -0,0 +1,285 @@ +package trace_test + +import ( + "context" + "fmt" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +func TestConcurrentNodeOperations(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // 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"}}, + {Input: "task 4", Option: types.TraceNodeOption{Label: "Worker 4"}}, + {Input: "task 5", Option: types.TraceNodeOption{Label: "Worker 5"}}, + }) + assert.NoError(t, err) + assert.Len(t, nodes, 5) + + // Concurrent operations on each node + var wg sync.WaitGroup + for i, node := range nodes { + wg.Add(1) + go func(idx int, n types.Node) { + defer wg.Done() + + // Concurrent logging + n.Info("Starting worker %d", idx+1) + n.Debug("Debug info %d", idx+1) + + // Set metadata + err := n.SetMetadata("worker_id", idx+1) + assert.NoError(t, err) + + // Complete + err = n.Complete(map[string]any{"worker": idx + 1}) + assert.NoError(t, err) + }(i, node) + } + wg.Wait() + }) + } +} + +func TestConcurrentSpaceOperations(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Create shared space + space, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Shared Space", + }) + assert.NoError(t, err) + + // Concurrent writes to the SAME space (now thread-safe with per-space locks) + var wg sync.WaitGroup + numWorkers := 10 + + for i := 0; i < numWorkers; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + + key := fmt.Sprintf("key_%d", idx) + value := fmt.Sprintf("value_%d", idx) + + err := manager.SetSpaceValue(space.ID, key, value) + assert.NoError(t, err) + }(i) + } + wg.Wait() + + // Verify all keys were set + keys := manager.ListSpaceKeys(space.ID) + assert.Len(t, keys, numWorkers) + + // Concurrent reads + wg = sync.WaitGroup{} + for i := 0; i < numWorkers; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + + key := fmt.Sprintf("key_%d", idx) + val, err := manager.GetSpaceValue(space.ID, key) + assert.NoError(t, err) + assert.Equal(t, fmt.Sprintf("value_%d", idx), val) + }(i) + } + wg.Wait() + }) + } +} + +func TestConcurrentSpaceCreation(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Create multiple spaces concurrently + var wg sync.WaitGroup + numSpaces := 10 + spaces := make([]*types.TraceSpace, numSpaces) + var mu sync.Mutex + + for i := 0; i < numSpaces; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + + space, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: fmt.Sprintf("Space %d", idx), + }) + assert.NoError(t, err) + + mu.Lock() + spaces[idx] = space + mu.Unlock() + }(i) + } + wg.Wait() + + // Verify all spaces were created + allSpaces := manager.ListSpaces() + assert.Len(t, allSpaces, numSpaces) + }) + } +} + +func TestConcurrentSubscribers(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Create multiple subscribers concurrently + var wg sync.WaitGroup + numSubscribers := 5 + subscribers := make([]<-chan *types.TraceUpdate, numSubscribers) + var mu sync.Mutex + + for i := 0; i < numSubscribers; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + + sub, err := manager.Subscribe() + assert.NoError(t, err) + + mu.Lock() + subscribers[idx] = sub + mu.Unlock() + }(i) + } + wg.Wait() + + // Verify all subscriptions were created + for i, sub := range subscribers { + assert.NotNil(t, sub, "Subscriber %d should not be nil", i) + } + + // Perform operations and verify all subscribers receive updates + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + assert.NoError(t, err) + }) + } +} + +func TestConcurrentTraceCreation(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + t.Run(d.Name, func(t *testing.T) { + ctx := context.Background() + + // Create multiple traces concurrently + var wg sync.WaitGroup + numTraces := 10 + traceIDs := make([]string, numTraces) + var mu sync.Mutex + + for i := 0; i < numTraces; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + + traceID, manager, err := trace.New(ctx, d.DriverType, nil, d.DriverOptions...) + assert.NoError(t, err) + assert.NotNil(t, manager) + + mu.Lock() + traceIDs[idx] = traceID + mu.Unlock() + }(i) + } + wg.Wait() + + // Clean up all traces + defer func() { + for _, traceID := range traceIDs { + trace.Release(traceID) + trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + } + }() + + // Verify all traces were created and loaded + for i, traceID := range traceIDs { + assert.NotEmpty(t, traceID, "Trace %d should have ID", i) + assert.True(t, trace.IsLoaded(traceID), "Trace %d should be loaded", i) + } + }) + } +} + +func TestConcurrentLogging(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Concurrent logging + var wg sync.WaitGroup + numLogs := 50 + for i := 0; i < numLogs; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + + manager.Info("Log message %d", idx) + manager.Debug("Debug message %d", idx) + manager.Warn("Warning message %d", idx) + }(i) + } + wg.Wait() + + // Note: We can't easily verify log count without exposing LoadLogs, + // but we verify no errors occurred during concurrent logging + }) + } +} + diff --git a/trace/trace_mem_test.go b/trace/trace_mem_test.go new file mode 100644 index 00000000..3aef6ae9 --- /dev/null +++ b/trace/trace_mem_test.go @@ -0,0 +1,621 @@ +package trace_test + +import ( + "context" + "fmt" + "runtime" + "sync" + "testing" + "time" + + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +// ============================================================================ +// Memory Leak Detection Tests +// ============================================================================ + +// TestMemoryLeakLocal checks for memory leaks with local driver +// Run with: go test -run=TestMemoryLeakLocal -v +func TestMemoryLeakLocal(t *testing.T) { + ctx := context.Background() + + // Warm up - execute a few times to stabilize memory + for i := 0; i < 10; i++ { + traceID, manager, _ := trace.New(ctx, trace.Local, nil) + manager.Add("test", types.TraceNodeOption{Label: "Test"}) + manager.Complete("result") + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + + // Force GC and get baseline memory + runtime.GC() + time.Sleep(100 * time.Millisecond) + var baseline runtime.MemStats + runtime.ReadMemStats(&baseline) + + // Execute many iterations + iterations := 1000 + for i := 0; i < iterations; i++ { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + t.Errorf("Create failed at iteration %d: %s", i, err.Error()) + continue + } + + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + t.Errorf("Add failed at iteration %d: %s", i, err.Error()) + } + + err = manager.Complete("result") + if err != nil { + t.Errorf("Complete failed at iteration %d: %s", i, err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + + // Periodic GC to help detect leaks faster + if i%100 == 0 { + runtime.GC() + } + } + + // Force GC and check final memory + runtime.GC() + time.Sleep(100 * time.Millisecond) + var final runtime.MemStats + runtime.ReadMemStats(&final) + + // Calculate memory growth + baselineHeap := baseline.HeapAlloc + finalHeap := final.HeapAlloc + growth := int64(finalHeap) - int64(baselineHeap) + growthPerIteration := float64(growth) / float64(iterations) + + t.Logf("Memory Statistics (Local Driver):") + t.Logf(" Iterations: %d", iterations) + t.Logf(" Baseline HeapAlloc: %d bytes (%.2f MB)", baselineHeap, float64(baselineHeap)/1024/1024) + t.Logf(" Final HeapAlloc: %d bytes (%.2f MB)", finalHeap, float64(finalHeap)/1024/1024) + t.Logf(" Total Growth: %d bytes (%.2f MB)", growth, float64(growth)/1024/1024) + t.Logf(" Growth per iteration: %.2f bytes", growthPerIteration) + t.Logf(" Total Alloc: %d bytes (%.2f MB)", final.TotalAlloc, float64(final.TotalAlloc)/1024/1024) + t.Logf(" Mallocs: %d", final.Mallocs) + t.Logf(" Frees: %d", final.Frees) + t.Logf(" Live Objects: %d", final.Mallocs-final.Frees) + t.Logf(" GC Runs: %d", final.NumGC-baseline.NumGC) + + // Check for memory leak + // Local driver involves file I/O, allow up to 10KB growth per iteration + maxGrowthPerIteration := 10240.0 + if growthPerIteration > maxGrowthPerIteration { + t.Errorf("Possible memory leak detected: %.2f bytes/iteration (threshold: %.2f bytes/iteration)", + growthPerIteration, maxGrowthPerIteration) + } else { + t.Logf("✓ Memory growth is within acceptable range") + } +} + +// TestMemoryLeakStore checks for memory leaks with store driver +// Run with: go test -run=TestMemoryLeakStore -v +func TestMemoryLeakStore(t *testing.T) { + ctx := context.Background() + + // Warm up + for i := 0; i < 10; i++ { + traceID, manager, _ := trace.New(ctx, trace.Store, nil) + manager.Add("test", types.TraceNodeOption{Label: "Test"}) + manager.Complete("result") + trace.Release(traceID) + trace.Remove(ctx, trace.Store, traceID) + } + + // Force GC and get baseline memory + runtime.GC() + time.Sleep(100 * time.Millisecond) + var baseline runtime.MemStats + runtime.ReadMemStats(&baseline) + + // Execute many iterations + iterations := 1000 + for i := 0; i < iterations; i++ { + traceID, manager, err := trace.New(ctx, trace.Store, nil) + if err != nil { + t.Errorf("Create failed at iteration %d: %s", i, err.Error()) + continue + } + + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + t.Errorf("Add failed at iteration %d: %s", i, err.Error()) + } + + err = manager.Complete("result") + if err != nil { + t.Errorf("Complete failed at iteration %d: %s", i, err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Store, traceID) + + // Periodic GC + if i%100 == 0 { + runtime.GC() + } + } + + // Force GC and check final memory + runtime.GC() + time.Sleep(100 * time.Millisecond) + var final runtime.MemStats + runtime.ReadMemStats(&final) + + // Calculate memory growth + baselineHeap := baseline.HeapAlloc + finalHeap := final.HeapAlloc + growth := int64(finalHeap) - int64(baselineHeap) + growthPerIteration := float64(growth) / float64(iterations) + + t.Logf("Memory Statistics (Store Driver):") + t.Logf(" Iterations: %d", iterations) + t.Logf(" Baseline HeapAlloc: %d bytes (%.2f MB)", baselineHeap, float64(baselineHeap)/1024/1024) + t.Logf(" Final HeapAlloc: %d bytes (%.2f MB)", finalHeap, float64(finalHeap)/1024/1024) + t.Logf(" Total Growth: %d bytes (%.2f MB)", growth, float64(growth)/1024/1024) + t.Logf(" Growth per iteration: %.2f bytes", growthPerIteration) + t.Logf(" Total Alloc: %d bytes (%.2f MB)", final.TotalAlloc, float64(final.TotalAlloc)/1024/1024) + t.Logf(" Mallocs: %d", final.Mallocs) + t.Logf(" Frees: %d", final.Frees) + t.Logf(" Live Objects: %d", final.Mallocs-final.Frees) + t.Logf(" GC Runs: %d", final.NumGC-baseline.NumGC) + + // Store driver should have similar or better performance than local + maxGrowthPerIteration := 10240.0 + if growthPerIteration > maxGrowthPerIteration { + t.Errorf("Possible memory leak detected: %.2f bytes/iteration (threshold: %.2f bytes/iteration)", + growthPerIteration, maxGrowthPerIteration) + } else { + t.Logf("✓ Memory growth is within acceptable range") + } +} + +// TestMemoryLeakComplexScenarios checks for memory leaks with complex operations +// Run with: go test -run=TestMemoryLeakComplexScenarios -v +func TestMemoryLeakComplexScenarios(t *testing.T) { + ctx := context.Background() + + scenarios := []struct { + name string + execute func(types.Manager) error + }{ + { + name: "SequentialNodes", + execute: func(m types.Manager) error { + for i := 0; i < 5; i++ { + _, err := m.Add(fmt.Sprintf("step_%d", i), types.TraceNodeOption{Label: fmt.Sprintf("Step %d", i)}) + if err != nil { + return err + } + if err := m.Complete(fmt.Sprintf("result_%d", i)); err != nil { + return err + } + } + return nil + }, + }, + { + name: "ParallelNodes", + execute: func(m types.Manager) error { + nodes, err := m.Parallel([]types.TraceParallelInput{ + {Input: "task1", Option: types.TraceNodeOption{Label: "Task 1"}}, + {Input: "task2", Option: types.TraceNodeOption{Label: "Task 2"}}, + {Input: "task3", Option: types.TraceNodeOption{Label: "Task 3"}}, + }) + if err != nil { + return err + } + + var wg sync.WaitGroup + for i, node := range nodes { + wg.Add(1) + go func(idx int, n types.Node) { + defer wg.Done() + n.Complete(fmt.Sprintf("result_%d", idx)) + }(i, node) + } + wg.Wait() + return nil + }, + }, + { + name: "WithSpace", + execute: func(m types.Manager) error { + space, err := m.CreateSpace(types.TraceSpaceOption{Label: "Context"}) + if err != nil { + return err + } + + for i := 0; i < 10; i++ { + if err := m.SetSpaceValue(space.ID, fmt.Sprintf("key_%d", i), fmt.Sprintf("value_%d", i)); err != nil { + return err + } + } + + _, err = m.Add("process", types.TraceNodeOption{Label: "Process"}) + if err != nil { + return err + } + return m.Complete("done") + }, + }, + { + name: "WithSubscription", + execute: func(m types.Manager) error { + updates, err := m.Subscribe() + if err != nil { + return err + } + + // Drain updates in background with timeout + done := make(chan bool) + go func() { + timeout := time.After(100 * time.Millisecond) + for { + select { + case _, ok := <-updates: + if !ok { + done <- true + return + } + case <-timeout: + done <- true + return + } + } + }() + + _, err = m.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + return err + } + if err := m.Complete("result"); err != nil { + return err + } + if err := m.MarkComplete(); err != nil { + return err + } + + // Wait for subscription to drain (with timeout) + <-done + return nil + }, + }, + } + + // Warm up + for i := 0; i < 10; i++ { + traceID, manager, _ := trace.New(ctx, trace.Local, nil) + manager.Add("warmup", types.TraceNodeOption{Label: "Warmup"}) + manager.Complete("done") + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + + // Test each scenario + for _, scenario := range scenarios { + t.Run(scenario.name, func(t *testing.T) { + // Get baseline + runtime.GC() + time.Sleep(50 * time.Millisecond) + var baseline runtime.MemStats + runtime.ReadMemStats(&baseline) + + // Execute iterations + iterations := 200 + for i := 0; i < iterations; i++ { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + t.Errorf("Create failed at iteration %d: %s", i, err.Error()) + continue + } + + err = scenario.execute(manager) + if err != nil { + t.Errorf("Scenario failed at iteration %d: %s", i, err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + + if i%50 == 0 { + runtime.GC() + } + } + + // Check final memory + runtime.GC() + time.Sleep(50 * time.Millisecond) + var final runtime.MemStats + runtime.ReadMemStats(&final) + + growth := int64(final.HeapAlloc) - int64(baseline.HeapAlloc) + growthPerIteration := float64(growth) / float64(iterations) + + t.Logf(" Baseline HeapAlloc: %d bytes (%.2f MB)", baseline.HeapAlloc, float64(baseline.HeapAlloc)/1024/1024) + t.Logf(" Final HeapAlloc: %d bytes (%.2f MB)", final.HeapAlloc, float64(final.HeapAlloc)/1024/1024) + t.Logf(" Growth: %d bytes (%.2f MB)", growth, float64(growth)/1024/1024) + t.Logf(" Growth/iteration: %.2f bytes", growthPerIteration) + + // Complex scenarios may have more memory usage + maxGrowthPerIteration := 15360.0 + if growthPerIteration > maxGrowthPerIteration { + t.Errorf("Possible memory leak: %.2f bytes/iteration (threshold: %.2f)", + growthPerIteration, maxGrowthPerIteration) + } else { + t.Logf(" ✓ Memory growth is within acceptable range") + } + }) + } +} + +// TestMemoryLeakConcurrent checks for memory leaks under concurrent load +// Run with: go test -run=TestMemoryLeakConcurrent -v +func TestMemoryLeakConcurrent(t *testing.T) { + ctx := context.Background() + + // Warm up + for i := 0; i < 20; i++ { + traceID, manager, _ := trace.New(ctx, trace.Local, nil) + manager.Add("warmup", types.TraceNodeOption{Label: "Warmup"}) + manager.Complete("done") + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + + // Get baseline + runtime.GC() + time.Sleep(100 * time.Millisecond) + var baseline runtime.MemStats + runtime.ReadMemStats(&baseline) + + // Run concurrent load + iterations := 1000 + concurrency := 10 + iterPerGoroutine := iterations / concurrency + + done := make(chan bool, concurrency) + for g := 0; g < concurrency; g++ { + go func(id int) { + defer func() { done <- true }() + for i := 0; i < iterPerGoroutine; i++ { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + t.Errorf("Goroutine %d: Create failed at iteration %d: %s", id, i, err.Error()) + continue + } + + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + t.Errorf("Goroutine %d: Add failed at iteration %d: %s", id, i, err.Error()) + } + + err = manager.Complete("result") + if err != nil { + t.Errorf("Goroutine %d: Complete failed at iteration %d: %s", id, i, err.Error()) + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + }(g) + } + + // Wait for all goroutines + for g := 0; g < concurrency; g++ { + <-done + } + + // Check final memory + runtime.GC() + time.Sleep(100 * time.Millisecond) + var final runtime.MemStats + runtime.ReadMemStats(&final) + + growth := int64(final.HeapAlloc) - int64(baseline.HeapAlloc) + growthPerIteration := float64(growth) / float64(iterations) + + t.Logf("Memory Statistics (Concurrent Load):") + t.Logf(" Iterations: %d", iterations) + t.Logf(" Concurrency: %d", concurrency) + t.Logf(" Baseline HeapAlloc: %d bytes (%.2f MB)", baseline.HeapAlloc, float64(baseline.HeapAlloc)/1024/1024) + t.Logf(" Final HeapAlloc: %d bytes (%.2f MB)", final.HeapAlloc, float64(final.HeapAlloc)/1024/1024) + t.Logf(" Growth: %d bytes (%.2f MB)", growth, float64(growth)/1024/1024) + t.Logf(" Growth/iteration: %.2f bytes", growthPerIteration) + t.Logf(" GC Runs: %d", final.NumGC-baseline.NumGC) + + // Concurrent scenarios may have slightly more overhead + maxGrowthPerIteration := 15360.0 + if growthPerIteration > maxGrowthPerIteration { + t.Errorf("Possible memory leak: %.2f bytes/iteration (threshold: %.2f)", + growthPerIteration, maxGrowthPerIteration) + } else { + t.Logf("✓ Memory growth is within acceptable range") + } +} + +// TestMemoryLeakSpaceOperations checks for memory leaks with space operations +// Run with: go test -run=TestMemoryLeakSpaceOperations -v +func TestMemoryLeakSpaceOperations(t *testing.T) { + ctx := context.Background() + + // Warm up + for i := 0; i < 10; i++ { + traceID, manager, _ := trace.New(ctx, trace.Local, nil) + space, _ := manager.CreateSpace(types.TraceSpaceOption{Label: "Test"}) + manager.SetSpaceValue(space.ID, "key", "value") + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + + // Get baseline + runtime.GC() + time.Sleep(100 * time.Millisecond) + var baseline runtime.MemStats + runtime.ReadMemStats(&baseline) + + // Execute iterations with space operations + iterations := 500 + for i := 0; i < iterations; i++ { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + t.Errorf("Create failed at iteration %d: %s", i, err.Error()) + continue + } + + // Create space and perform operations + space, err := manager.CreateSpace(types.TraceSpaceOption{Label: "Test Space"}) + if err != nil { + t.Errorf("CreateSpace failed at iteration %d: %s", i, err.Error()) + } + + // Set multiple values + for j := 0; j < 20; j++ { + err = manager.SetSpaceValue(space.ID, fmt.Sprintf("key_%d", j), fmt.Sprintf("value_%d", j)) + if err != nil { + t.Errorf("SetSpaceValue failed at iteration %d: %s", i, err.Error()) + } + } + + // Get values + for j := 0; j < 20; j++ { + _, err = manager.GetSpaceValue(space.ID, fmt.Sprintf("key_%d", j)) + if err != nil { + t.Errorf("GetSpaceValue failed at iteration %d: %s", i, err.Error()) + } + } + + // Delete some values + for j := 0; j < 10; j++ { + err = manager.DeleteSpaceValue(space.ID, fmt.Sprintf("key_%d", j)) + if err != nil { + t.Errorf("DeleteSpaceValue failed at iteration %d: %s", i, err.Error()) + } + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + + if i%100 == 0 { + runtime.GC() + } + } + + // Check final memory + runtime.GC() + time.Sleep(100 * time.Millisecond) + var final runtime.MemStats + runtime.ReadMemStats(&final) + + growth := int64(final.HeapAlloc) - int64(baseline.HeapAlloc) + growthPerIteration := float64(growth) / float64(iterations) + + t.Logf("Memory Statistics (Space Operations):") + t.Logf(" Iterations: %d", iterations) + t.Logf(" Baseline HeapAlloc: %d bytes (%.2f MB)", baseline.HeapAlloc, float64(baseline.HeapAlloc)/1024/1024) + t.Logf(" Final HeapAlloc: %d bytes (%.2f MB)", final.HeapAlloc, float64(final.HeapAlloc)/1024/1024) + t.Logf(" Growth: %d bytes (%.2f MB)", growth, float64(growth)/1024/1024) + t.Logf(" Growth/iteration: %.2f bytes", growthPerIteration) + t.Logf(" GC Runs: %d", final.NumGC-baseline.NumGC) + + // Space operations involve maps and persistence + maxGrowthPerIteration := 20480.0 + if growthPerIteration > maxGrowthPerIteration { + t.Errorf("Possible memory leak: %.2f bytes/iteration (threshold: %.2f)", + growthPerIteration, maxGrowthPerIteration) + } else { + t.Logf("✓ Memory growth is within acceptable range") + } +} + +// TestGoroutineLeak verifies that no goroutines are leaked +// Run with: go test -run=TestGoroutineLeak -v +func TestGoroutineLeak(t *testing.T) { + ctx := context.Background() + + // Track goroutine count to detect goroutine leaks + initialGoroutines := runtime.NumGoroutine() + + // Execute multiple iterations + iterations := 100 + for i := 0; i < iterations; i++ { + traceID, manager, err := trace.New(ctx, trace.Local, nil) + if err != nil { + t.Errorf("Create failed at iteration %d: %s", i, err.Error()) + continue + } + + // Subscribe (creates goroutines) + updates, err := manager.Subscribe() + if err != nil { + t.Errorf("Subscribe failed at iteration %d: %s", i, err.Error()) + } + + // Perform operations + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + if err != nil { + t.Errorf("Add failed at iteration %d: %s", i, err.Error()) + } + + err = manager.Complete("result") + if err != nil { + t.Errorf("Complete failed at iteration %d: %s", i, err.Error()) + } + + err = manager.MarkComplete() + if err != nil { + t.Errorf("MarkComplete failed at iteration %d: %s", i, err.Error()) + } + + // Drain subscription + timeout := time.After(10 * time.Millisecond) + drainLoop: + for { + select { + case _, ok := <-updates: + if !ok { + break drainLoop + } + case <-timeout: + break drainLoop + } + } + + trace.Release(traceID) + trace.Remove(ctx, trace.Local, traceID) + } + + // Give time for cleanup + time.Sleep(200 * time.Millisecond) + runtime.GC() + time.Sleep(200 * time.Millisecond) + + finalGoroutines := runtime.NumGoroutine() + goroutineGrowth := finalGoroutines - initialGoroutines + + t.Logf("Goroutine Statistics:") + t.Logf(" Initial: %d", initialGoroutines) + t.Logf(" Final: %d", finalGoroutines) + t.Logf(" Growth: %d", goroutineGrowth) + + // Allow some goroutine growth for runtime internals, but not proportional to iterations + maxGoroutineGrowth := 20 + if goroutineGrowth > maxGoroutineGrowth { + t.Errorf("Possible goroutine leak: %d new goroutines (threshold: %d)", + goroutineGrowth, maxGoroutineGrowth) + } else { + t.Logf("✓ No goroutine leak detected") + } +} + diff --git a/trace/trace_node_test.go b/trace/trace_node_test.go new file mode 100644 index 00000000..58d45d49 --- /dev/null +++ b/trace/trace_node_test.go @@ -0,0 +1,214 @@ +package trace_test + +import ( + "context" + "fmt" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +func TestNodeOperations(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Add sequential node + node1, err := manager.Add("input data", types.TraceNodeOption{ + Label: "Input Processing", + Icon: "processor", + Description: "Process input data", + }) + assert.NoError(t, err) + assert.NotNil(t, node1) + + // Log messages (chainable) + manager.Info("Processing started"). + Debug("Debug info"). + Warn("Warning message") + + // Set output and complete + err = manager.Complete(map[string]any{"result": "success"}) + assert.NoError(t, err) + + // Add another node + node2, err := manager.Add("processing", types.TraceNodeOption{ + Label: "Processing", + Icon: "cpu", + }) + assert.NoError(t, err) + assert.NotNil(t, node2) + + // Set metadata + err = manager.SetMetadata("key1", "value1") + assert.NoError(t, err) + + err = manager.Complete(map[string]any{"status": "done"}) + assert.NoError(t, err) + + // Get current nodes + currentNodes, err := manager.GetCurrentNodes() + assert.NoError(t, err) + assert.NotEmpty(t, currentNodes) + assert.Equal(t, types.StatusCompleted, currentNodes[0].Status) + }) + } +} + +func TestParallelOperations(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Create parallel nodes + nodes, err := manager.Parallel([]types.TraceParallelInput{ + { + Input: "task A", + Option: types.TraceNodeOption{Label: "Worker A", Icon: "cpu"}, + }, + { + Input: "task B", + Option: types.TraceNodeOption{Label: "Worker B", Icon: "cpu"}, + }, + { + Input: "task C", + Option: types.TraceNodeOption{Label: "Worker C", Icon: "cpu"}, + }, + }) + assert.NoError(t, err) + assert.Len(t, nodes, 3) + + // Each node completes itself + var wg sync.WaitGroup + for i, node := range nodes { + wg.Add(1) + go func(idx int, n types.Node) { + defer wg.Done() + + n.Info("Worker %d processing", idx+1) + time.Sleep(10 * time.Millisecond) + err := n.Complete(map[string]any{"worker": idx + 1, "status": "done"}) + assert.NoError(t, err) + }(i, node) + } + wg.Wait() + + // Add node after parallel (auto-join) + node, err := manager.Add("merge", types.TraceNodeOption{ + Label: "Merge", + Icon: "merge", + }) + assert.NoError(t, err) + assert.NotNil(t, node) + + err = manager.Complete(map[string]any{"merged": true}) + assert.NoError(t, err) + }) + } +} + +func TestNodeFailOperation(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Add node + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + assert.NoError(t, err) + + // Fail node + testErr := fmt.Errorf("test error") + err = manager.Fail(testErr) + assert.NoError(t, err) + + // Verify node status + currentNodes, err := manager.GetCurrentNodes() + assert.NoError(t, err) + assert.NotEmpty(t, currentNodes) + assert.Equal(t, types.StatusFailed, currentNodes[0].Status) + }) + } +} + +func TestNodeChaining(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Test chainable logging + result := manager.Info("Step 1"). + Debug("Debug step 1"). + Warn("Warning step 1") + + // Should return Manager interface + assert.NotNil(t, result) + + // Should still be able to call Manager methods + _, err = result.Add("next", types.TraceNodeOption{Label: "Next"}) + assert.NoError(t, err) + }) + } +} + +func TestCompleteWithOutput(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Add node + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + assert.NoError(t, err) + + // Complete with output directly + output := map[string]any{"result": "success", "count": 42} + err = manager.Complete(output) + assert.NoError(t, err) + + // Verify output was set + currentNodes, err := manager.GetCurrentNodes() + assert.NoError(t, err) + assert.NotEmpty(t, currentNodes) + assert.Equal(t, output, currentNodes[0].Output) + }) + } +} + diff --git a/trace/trace_space_test.go b/trace/trace_space_test.go new file mode 100644 index 00000000..3caa3aec --- /dev/null +++ b/trace/trace_space_test.go @@ -0,0 +1,186 @@ +package trace_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +func TestSpaceOperations(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Create space + space, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Test Space", + Icon: "database", + Description: "Test space for unit tests", + TTL: 3600, + }) + assert.NoError(t, err) + assert.NotNil(t, space) + assert.NotEmpty(t, space.ID) + + // Set values + err = manager.SetSpaceValue(space.ID, "key1", "value1") + assert.NoError(t, err) + + err = manager.SetSpaceValue(space.ID, "key2", map[string]any{"nested": "data"}) + assert.NoError(t, err) + + err = manager.SetSpaceValue(space.ID, "key3", 12345) + assert.NoError(t, err) + + // Get values + val1, err := manager.GetSpaceValue(space.ID, "key1") + assert.NoError(t, err) + assert.Equal(t, "value1", val1) + + val2, err := manager.GetSpaceValue(space.ID, "key2") + assert.NoError(t, err) + assert.NotNil(t, val2) + + // Has value + exists := manager.HasSpaceValue(space.ID, "key1") + assert.True(t, exists) + + exists = manager.HasSpaceValue(space.ID, "nonexistent") + assert.False(t, exists) + + // List keys + keys := manager.ListSpaceKeys(space.ID) + assert.Len(t, keys, 3) + + // Delete value + err = manager.DeleteSpaceValue(space.ID, "key1") + assert.NoError(t, err) + + exists = manager.HasSpaceValue(space.ID, "key1") + assert.False(t, exists) + + // Clear all values + err = manager.ClearSpaceValues(space.ID) + assert.NoError(t, err) + + keys = manager.ListSpaceKeys(space.ID) + assert.Empty(t, keys) + + // List spaces + spaces := manager.ListSpaces() + assert.NotEmpty(t, spaces) + assert.True(t, manager.HasSpace(space.ID)) + + // Delete space + err = manager.DeleteSpace(space.ID) + assert.NoError(t, err) + + assert.False(t, manager.HasSpace(space.ID)) + }) + } +} + +func TestMultipleSpaces(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Create multiple spaces + space1, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Context", + Icon: "context", + }) + assert.NoError(t, err) + + space2, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Memory", + Icon: "memory", + }) + assert.NoError(t, err) + + space3, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Cache", + Icon: "cache", + }) + assert.NoError(t, err) + + // Set values in different spaces + err = manager.SetSpaceValue(space1.ID, "context_key", "context_value") + assert.NoError(t, err) + + err = manager.SetSpaceValue(space2.ID, "memory_key", "memory_value") + assert.NoError(t, err) + + err = manager.SetSpaceValue(space3.ID, "cache_key", "cache_value") + assert.NoError(t, err) + + // Verify isolation + val1, err := manager.GetSpaceValue(space1.ID, "context_key") + assert.NoError(t, err) + assert.Equal(t, "context_value", val1) + + // Key from space1 should not exist in space2 + exists := manager.HasSpaceValue(space2.ID, "context_key") + assert.False(t, exists) + + // List all spaces + spaces := manager.ListSpaces() + assert.Len(t, spaces, 3) + }) + } +} + +func TestSpaceGetSpace(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Create space + space, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Test Space", + Description: "Test description", + TTL: 7200, + }) + assert.NoError(t, err) + + // Get space by ID + retrieved, err := manager.GetSpace(space.ID) + assert.NoError(t, err) + assert.NotNil(t, retrieved) + assert.Equal(t, space.ID, retrieved.ID) + assert.Equal(t, "Test Space", retrieved.Label) + assert.Equal(t, "Test description", retrieved.Description) + assert.Equal(t, int64(7200), retrieved.TTL) + + // Get non-existent space (returns nil, nil) + nonExistent, err := manager.GetSpace("nonexistent") + assert.NoError(t, err) + assert.Nil(t, nonExistent) + }) + } +} + diff --git a/trace/trace_subscription_test.go b/trace/trace_subscription_test.go new file mode 100644 index 00000000..77ddf086 --- /dev/null +++ b/trace/trace_subscription_test.go @@ -0,0 +1,283 @@ +package trace_test + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +func TestSubscription(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Subscribe to updates + updates, err := manager.Subscribe() + assert.NoError(t, err) + assert.NotNil(t, updates) + + // Collect updates in background + var receivedUpdates []*types.TraceUpdate + var updatesMu sync.Mutex + done := make(chan bool) + + go func() { + timeout := time.After(2 * time.Second) + for { + select { + case update := <-updates: + updatesMu.Lock() + receivedUpdates = append(receivedUpdates, update) + updatesMu.Unlock() + + // Check for trace completion + if update.Type == types.UpdateTypeComplete { + done <- true + return + } + case <-timeout: + done <- true + return + } + } + }() + + // Perform operations + manager.Info("Test operation") + _, err = manager.Add("test node", types.TraceNodeOption{Label: "Test"}) + assert.NoError(t, err) + + err = manager.Complete(map[string]any{"test": "data"}) + assert.NoError(t, err) + + // Create space and set value + space, err := manager.CreateSpace(types.TraceSpaceOption{Label: "Test Space"}) + assert.NoError(t, err) + + err = manager.SetSpaceValue(space.ID, "key", "value") + assert.NoError(t, err) + + // Mark trace complete + err = manager.MarkComplete() + assert.NoError(t, err) + + // Wait for completion or timeout + <-done + + // Verify we received updates + updatesMu.Lock() + defer updatesMu.Unlock() + + assert.NotEmpty(t, receivedUpdates) + + // Check for specific event types + eventTypes := make(map[string]bool) + for _, update := range receivedUpdates { + eventTypes[update.Type] = true + } + + assert.True(t, eventTypes[types.UpdateTypeInit], "Should receive init event") + assert.True(t, eventTypes[types.UpdateTypeNodeStart], "Should receive node_start event") + assert.True(t, eventTypes[types.UpdateTypeComplete], "Should receive complete event") + }) + } +} + +func TestSubscribeFrom(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Real scenario: User starts a trace, performs some operations + _, err = manager.Add("Step 1", types.TraceNodeOption{Label: "Processing"}) + assert.NoError(t, err) + manager.Info("Processing step 1") + err = manager.Complete("step1 result") + assert.NoError(t, err) + + // Wait to ensure different timestamp (simulate time passing) + time.Sleep(1100 * time.Millisecond) + + // Record timestamp (simulate user noting current time before refresh) + resumeTimestamp := time.Now().Unix() + + // Wait again to ensure next operations are after resumeTimestamp + time.Sleep(100 * time.Millisecond) + + // Continue with more operations + _, err = manager.Add("Step 2", types.TraceNodeOption{Label: "Finalizing"}) + assert.NoError(t, err) + manager.Info("Processing step 2") + err = manager.Complete("step2 result") + assert.NoError(t, err) + + // Mark trace complete + err = manager.MarkComplete() + assert.NoError(t, err) + + // Real scenario: User refreshes page and resumes from last known timestamp + // This should replay events from resumeTimestamp onwards + updates, err := manager.SubscribeFrom(resumeTimestamp) + assert.NoError(t, err) + assert.NotNil(t, updates) + + // Collect updates + var receivedUpdates []*types.TraceUpdate + timeout := time.After(1 * time.Second) + foundStep2 := false + + collectLoop: + for { + select { + case update, ok := <-updates: + if !ok { + // Channel closed + break collectLoop + } + receivedUpdates = append(receivedUpdates, update) + // Check if we received step 2 events + if update.Type == types.UpdateTypeNodeStart { + if data, ok := update.Data.(*types.NodeStartData); ok { + if data.Node != nil && data.Node.Label == "Finalizing" { + foundStep2 = true + } + } + } + // Stop after receiving trace_complete + if update.Type == types.UpdateTypeComplete { + break collectLoop + } + case <-timeout: + break collectLoop + } + } + + // Verify we received events from step 2 onwards + assert.NotEmpty(t, receivedUpdates, "Should receive events from resume point") + assert.True(t, foundStep2, "Should receive Step 2 events") + + // All events should be at or after the resume timestamp + for _, update := range receivedUpdates { + assert.GreaterOrEqual(t, update.Timestamp, resumeTimestamp, + "Event timestamp %d should be >= resume timestamp %d (event type: %s)", + update.Timestamp, resumeTimestamp, update.Type) + } + }) + } +} + +func TestIsComplete(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Initially not complete + assert.False(t, manager.IsComplete()) + + // Mark complete + err = manager.MarkComplete() + assert.NoError(t, err) + + // Now should be complete + assert.True(t, manager.IsComplete()) + }) + } +} + +func TestMultipleSubscribers(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + 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...) + + // Create multiple subscribers + sub1, err := manager.Subscribe() + assert.NoError(t, err) + + sub2, err := manager.Subscribe() + assert.NoError(t, err) + + sub3, err := manager.Subscribe() + assert.NoError(t, err) + + // Collect updates from all subscribers + var wg sync.WaitGroup + counts := make([]int, 3) + var mu sync.Mutex + + for i, sub := range []<-chan *types.TraceUpdate{sub1, sub2, sub3} { + wg.Add(1) + go func(idx int, ch <-chan *types.TraceUpdate) { + defer wg.Done() + timeout := time.After(1 * time.Second) + for { + select { + case update := <-ch: + if update != nil { + mu.Lock() + counts[idx]++ + mu.Unlock() + if update.Type == types.UpdateTypeComplete { + return + } + } + case <-timeout: + return + } + } + }(i, sub) + } + + // Perform operations + _, err = manager.Add("test", types.TraceNodeOption{Label: "Test"}) + assert.NoError(t, err) + err = manager.Complete(nil) + assert.NoError(t, err) + err = manager.MarkComplete() + assert.NoError(t, err) + + // Wait for all subscribers + wg.Wait() + + // All subscribers should receive updates + mu.Lock() + defer mu.Unlock() + for i, count := range counts { + assert.Greater(t, count, 0, "Subscriber %d should receive updates", i+1) + } + }) + } +} +