From 71cd7b99937854c63381f7d4bc0ceb835bdcc2dd Mon Sep 17 00:00:00 2001 From: Max Date: Thu, 20 Nov 2025 18:15:48 +0800 Subject: [PATCH] Enhance trace management functionality and logging capabilities - Introduced new methods in the trace manager for retrieving events, trace info, nodes, logs, and spaces from storage, improving data access and management. - Updated the driver implementations to support loading and unarchiving traces, ensuring robust handling of trace data. - Enhanced error handling and logging in the OpenAI provider to facilitate better debugging and traceability of requests and responses. - Added a new TraceSpaceData struct to encapsulate space metadata along with key-value data for improved API responses. --- agent/llm/providers/openai/openai.go | 10 +- openapi/openapi.go | 4 + openapi/tests/trace/common_test.go | 164 ++++++++ openapi/tests/trace/events_test.go | 212 ++++++++++ openapi/tests/trace/info_test.go | 78 ++++ openapi/tests/trace/logs_test.go | 156 ++++++++ openapi/tests/trace/nodes_test.go | 144 +++++++ openapi/tests/trace/spaces_test.go | 168 ++++++++ openapi/trace/events.go | 179 +++++++++ openapi/trace/helpers.go | 210 ++++++++++ openapi/trace/info.go | 76 ++++ openapi/trace/logs.go | 96 +++++ openapi/trace/nodes.go | 155 +++++++ openapi/trace/spaces.go | 139 +++++++ openapi/trace/trace.go | 22 + trace/local/driver.go | 57 ++- trace/manager.go | 138 +++++++ trace/store/driver.go | 65 ++- trace/trace_resource_test.go | 576 +++++++++++++++++++++++++++ trace/types/interfaces.go | 20 + trace/types/types.go | 6 + 21 files changed, 2658 insertions(+), 17 deletions(-) create mode 100644 openapi/tests/trace/common_test.go create mode 100644 openapi/tests/trace/events_test.go create mode 100644 openapi/tests/trace/info_test.go create mode 100644 openapi/tests/trace/logs_test.go create mode 100644 openapi/tests/trace/nodes_test.go create mode 100644 openapi/tests/trace/spaces_test.go create mode 100644 openapi/trace/events.go create mode 100644 openapi/trace/helpers.go create mode 100644 openapi/trace/info.go create mode 100644 openapi/trace/logs.go create mode 100644 openapi/trace/nodes.go create mode 100644 openapi/trace/spaces.go create mode 100644 openapi/trace/trace.go create mode 100644 trace/trace_resource_test.go diff --git a/agent/llm/providers/openai/openai.go b/agent/llm/providers/openai/openai.go index 1e7b0e9c..7cb21243 100644 --- a/agent/llm/providers/openai/openai.go +++ b/agent/llm/providers/openai/openai.go @@ -647,7 +647,7 @@ func (p *Provider) streamWithRetry(ctx *context.Context, messages []context.Mess // Log request for debugging if trace != nil { - if requestBodyJSON, marshalErr := jsoniter.Marshal(requestBody); marshalErr == nil { + if requestBodyJSON, marshalErr := jsoniter.Marshal(requestBody); marshalErr == nil { trace.Debug("OpenAI Stream Request", map[string]any{ "url": url, "body": string(requestBodyJSON), @@ -766,12 +766,12 @@ func (p *Provider) streamWithRetry(ctx *context.Context, messages []context.Mess if trace != nil { trace.Warn("OpenAI stream completed but no data was received") - // Log request details for debugging - if requestBodyJSON, err := jsoniter.Marshal(requestBody); err == nil { + // Log request details for debugging + if requestBodyJSON, err := jsoniter.Marshal(requestBody); err == nil { trace.Error("Request body that caused empty response", map[string]any{ "body": string(requestBodyJSON), }) - } + } trace.Error("Request details", map[string]any{ "url": url, "model": accumulator.model, @@ -1057,7 +1057,7 @@ func (p *Provider) postWithRetry(ctx *context.Context, messages []context.Messag } // Log full response data for debugging if trace != nil { - if respJSON, err := jsoniter.Marshal(resp.Data); err == nil { + if respJSON, err := jsoniter.Marshal(resp.Data); err == nil { trace.Error("OpenAI API error response", map[string]any{ "response": string(respJSON), }) diff --git a/openapi/openapi.go b/openapi/openapi.go index d595b896..cab93dd6 100644 --- a/openapi/openapi.go +++ b/openapi/openapi.go @@ -21,6 +21,7 @@ import ( "github.com/yaoapp/yao/openapi/oauth/acl" "github.com/yaoapp/yao/openapi/oauth/types" "github.com/yaoapp/yao/openapi/team" + openapiTrace "github.com/yaoapp/yao/openapi/trace" "github.com/yaoapp/yao/openapi/user" ) @@ -146,6 +147,9 @@ func (openapi *OpenAPI) Attach(router *gin.Engine) { // MCP Server handlers mcp.Attach(group.Group("/mcp"), openapi.OAuth) + // Trace handlers + openapiTrace.Attach(group.Group("/trace"), openapi.OAuth) + // Custom handlers (Defined by developer) } diff --git a/openapi/tests/trace/common_test.go b/openapi/tests/trace/common_test.go new file mode 100644 index 00000000..10e41c4e --- /dev/null +++ b/openapi/tests/trace/common_test.go @@ -0,0 +1,164 @@ +package trace_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/openapi" + oauthtypes "github.com/yaoapp/yao/openapi/oauth/types" + "github.com/yaoapp/yao/openapi/tests/testutils" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +// testTraceData holds the prepared test trace and related information +type testTraceData struct { + TraceID string + Manager types.Manager + RootNodeID string + Node1ID string + Node2ID string + Node3ID string + TokenInfo *testutils.TokenInfo + TestClient *oauthtypes.ClientInfo + ServerURL string + BaseURL string + Ctx context.Context +} + +// prepareTestTrace creates a test trace with sample nodes, logs, and spaces +// This provides consistent test data for all trace API tests +func prepareTestTrace(t *testing.T) *testTraceData { + serverURL := testutils.Prepare(t) + + // Get base URL from server config + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + // Register test client and obtain token with trace permissions + testClient := testutils.RegisterTestClient(t, "Trace API Test Client", []string{"https://localhost/callback"}) + tokenInfo := testutils.ObtainAccessToken(t, serverURL, testClient.ClientID, testClient.ClientSecret, "https://localhost/callback", "openid profile trace:traces:read:all") + + // Create a test trace with proper user info + ctx := context.Background() + traceOption := &types.TraceOption{ + CreatedBy: tokenInfo.UserID, + Metadata: map[string]any{ + "test_type": "api_test", + "test_name": "common_trace_data", + }, + } + + traceID, manager, err := trace.New(ctx, trace.Local, traceOption) + assert.NoError(t, err) + assert.NotEmpty(t, traceID) + + // Add manager-level logs + manager.Info("Manager info log", map[string]any{"level": "manager", "action": "init"}) + manager.Debug("Manager debug log", map[string]any{"level": "manager", "action": "debug"}) + + // Create a memory space + space, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Test Space", + Description: "A test memory space", + }) + assert.NoError(t, err) + assert.NotNil(t, space) + spaceID := space.ID + + // Add some data to the space + err = manager.SetSpaceValue(spaceID, "key1", "value1") + assert.NoError(t, err) + err = manager.SetSpaceValue(spaceID, "key2", map[string]any{"nested": "data"}) + assert.NoError(t, err) + + // Add first node + node1, err := manager.Add("test input 1", types.TraceNodeOption{ + Label: "First Node", + Icon: "icon1", + Description: "First test node", + Metadata: map[string]any{"node_order": 1}, + }) + assert.NoError(t, err) + node1ID := node1.ID() + + node1.Info("Node 1 info log", map[string]any{"node": "1", "message": "info"}) + node1.Debug("Node 1 debug log", map[string]any{"node": "1", "message": "debug"}) + + err = node1.SetOutput(map[string]any{"result": "node1_output", "status": "processing"}) + assert.NoError(t, err) + + // Add second node + node2, err := manager.Add("test input 2", types.TraceNodeOption{ + Label: "Second Node", + Icon: "icon2", + Description: "Second test node", + Metadata: map[string]any{"node_order": 2}, + }) + assert.NoError(t, err) + node2ID := node2.ID() + + node2.Info("Node 2 info log", map[string]any{"node": "2", "message": "info"}) + node2.Warn("Node 2 warn log", map[string]any{"node": "2", "message": "warning"}) + + err = node2.Complete(map[string]any{"result": "node2_completed", "status": "success"}) + assert.NoError(t, err) + + // Add third node + node3, err := manager.Add("test input 3", types.TraceNodeOption{ + Label: "Third Node", + Icon: "icon3", + Description: "Third test node", + Metadata: map[string]any{"node_order": 3}, + }) + assert.NoError(t, err) + node3ID := node3.ID() + + node3.Debug("Node 3 debug log", map[string]any{"node": "3", "message": "debug"}) + node3.Error("Node 3 error log", map[string]any{"node": "3", "message": "error", "error_code": 500}) + + err = node3.Complete(map[string]any{"result": "node3_completed"}) + assert.NoError(t, err) + + // Complete the trace to flush all data to storage + err = manager.MarkComplete() + assert.NoError(t, err) + + // Get root node ID + rootNode, err := manager.GetRootNode() + assert.NoError(t, err) + rootNodeID := "" + if rootNode != nil { + rootNodeID = rootNode.ID + } + + return &testTraceData{ + TraceID: traceID, + Manager: manager, + RootNodeID: rootNodeID, + Node1ID: node1ID, + Node2ID: node2ID, + Node3ID: node3ID, + TokenInfo: tokenInfo, + TestClient: testClient, + ServerURL: serverURL, + BaseURL: baseURL, + Ctx: ctx, + } +} + +// cleanupTestTrace cleans up the test trace and related resources +func cleanupTestTrace(t *testing.T, data *testTraceData) { + if data.TraceID != "" { + trace.Release(data.TraceID) + trace.Remove(data.Ctx, trace.Local, data.TraceID) + } + if data.TestClient != nil { + testutils.CleanupTestClient(t, data.TestClient.ClientID) + } + testutils.Clean() +} + diff --git a/openapi/tests/trace/events_test.go b/openapi/tests/trace/events_test.go new file mode 100644 index 00000000..66ccd1da --- /dev/null +++ b/openapi/tests/trace/events_test.go @@ -0,0 +1,212 @@ +package trace_test + +import ( + "bufio" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +// TestGetEvents tests the events API endpoint +func TestGetEvents(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Test GET /traces/:traceID/events + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/events", data.ServerURL, data.BaseURL, data.TraceID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode, "Expected status code 200") + + // Parse response + body, err := io.ReadAll(resp.Body) + assert.NoError(t, err) + + var responseData map[string]interface{} + err = json.Unmarshal(body, &responseData) + assert.NoError(t, err) + + // Verify response structure + assert.Equal(t, data.TraceID, responseData["id"], "Trace ID should match") + assert.NotNil(t, responseData["events"], "Should have events field") + + events, ok := responseData["events"].([]interface{}) + assert.True(t, ok, "Events should be an array") + assert.NotEmpty(t, events, "Events array should not be empty") + + t.Logf("Retrieved %d events for trace %s", len(events), data.TraceID) + + // Verify event types + eventTypes := make(map[string]bool) + for _, e := range events { + event, ok := e.(map[string]interface{}) + if ok { + eventType, _ := event["Type"].(string) + eventTypes[eventType] = true + } + } + + assert.True(t, eventTypes["init"], "Should have init event") + assert.True(t, eventTypes["node_start"], "Should have node_start events") + assert.True(t, eventTypes["node_complete"], "Should have node_complete events") + assert.True(t, eventTypes["space_created"], "Should have space_created event") +} + +// TestGetEventsNotFound tests getting events for non-existent trace +func TestGetEventsNotFound(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Try to get events for non-existent trace + requestURL := fmt.Sprintf("%s%s/trace/traces/nonexistent/events", data.ServerURL, data.BaseURL) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusNotFound, resp.StatusCode, "Expected status code 404 for non-existent trace") +} + +// TestGetEventsUnauthorized tests getting events without authentication +func TestGetEventsUnauthorized(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Try to get events without token + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/events", data.ServerURL, data.BaseURL, data.TraceID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode, "Expected status code 401 without authentication") +} + +// TestGetEventsSSE tests the events API endpoint in SSE streaming mode +func TestGetEventsSSE(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Test GET /traces/:traceID/events?stream=true + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/events?stream=true", data.ServerURL, data.BaseURL, data.TraceID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + req.Header.Set("Accept", "text/event-stream") + + client := &http.Client{Timeout: 30 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + // Verify SSE response headers + assert.Equal(t, http.StatusOK, resp.StatusCode, "Expected status code 200") + assert.Equal(t, "text/event-stream", resp.Header.Get("Content-Type"), "Expected text/event-stream content type") + assert.Equal(t, "no-cache", resp.Header.Get("Cache-Control"), "Expected no-cache") + assert.Equal(t, "keep-alive", resp.Header.Get("Connection"), "Expected keep-alive connection") + + // Read SSE events + scanner := bufio.NewScanner(resp.Body) + events := make([]map[string]interface{}, 0) + var currentEvent map[string]interface{} + eventCount := 0 + maxEvents := 50 // Limit to prevent infinite loop + + for scanner.Scan() && eventCount < maxEvents { + line := scanner.Text() + + // SSE format: "data: {...}" + if strings.HasPrefix(line, "data: ") { + dataStr := strings.TrimPrefix(line, "data: ") + + // Check for [DONE] marker + if dataStr == "[DONE]" { + t.Log("Received [DONE] marker, stream completed") + break + } + + // Parse JSON event data + var eventData map[string]interface{} + if err := json.Unmarshal([]byte(dataStr), &eventData); err != nil { + t.Logf("Failed to parse event data: %s, error: %v", dataStr, err) + continue + } + + currentEvent = eventData + } else if line == "" && currentEvent != nil { + // Empty line marks end of an event + events = append(events, currentEvent) + eventCount++ + currentEvent = nil + } + } + + assert.NoError(t, scanner.Err(), "Should not have scanner errors") + assert.NotEmpty(t, events, "Should receive at least one SSE event") + + t.Logf("Received %d SSE events for trace %s", len(events), data.TraceID) + + // Verify event structure and types + eventTypes := make(map[string]int) + for i, event := range events { + // Verify required fields + assert.NotNil(t, event["Type"], "Event %d should have Type field", i) + assert.NotNil(t, event["TraceID"], "Event %d should have TraceID field", i) + assert.NotNil(t, event["Timestamp"], "Event %d should have Timestamp field", i) + + // Verify TraceID matches + if traceID, ok := event["TraceID"].(string); ok { + assert.Equal(t, data.TraceID, traceID, "Event %d TraceID should match", i) + } + + // Count event types + if eventType, ok := event["Type"].(string); ok { + eventTypes[eventType]++ + } + } + + // Verify expected event types + assert.Greater(t, eventTypes["init"], 0, "Should have at least one init event") + assert.Greater(t, eventTypes["node_start"], 0, "Should have at least one node_start event") + assert.Greater(t, eventTypes["node_complete"], 0, "Should have at least one node_complete event") + assert.Greater(t, eventTypes["complete"], 0, "Should have at least one complete event") + + // Log event type distribution + t.Logf("Event type distribution: %+v", eventTypes) + + // Verify event order: init should be first + if len(events) > 0 { + firstEventType, _ := events[0]["Type"].(string) + assert.Equal(t, "init", firstEventType, "First event should be init") + } + + // Verify complete event is last (before [DONE]) + if len(events) > 1 { + lastEventType, _ := events[len(events)-1]["Type"].(string) + assert.Equal(t, "complete", lastEventType, "Last event should be complete") + } +} diff --git a/openapi/tests/trace/info_test.go b/openapi/tests/trace/info_test.go new file mode 100644 index 00000000..ef1b1676 --- /dev/null +++ b/openapi/tests/trace/info_test.go @@ -0,0 +1,78 @@ +package trace_test + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +// TestGetInfo tests the trace info API endpoint +func TestGetInfo(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Test GET /traces/:traceID/info + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/info", data.ServerURL, data.BaseURL, data.TraceID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode, "Expected status code 200") + + // Parse response + body, err := io.ReadAll(resp.Body) + assert.NoError(t, err) + + var responseData map[string]interface{} + err = json.Unmarshal(body, &responseData) + assert.NoError(t, err) + + // Verify response structure + assert.Equal(t, data.TraceID, responseData["id"], "Trace ID should match") + assert.Equal(t, "local", responseData["driver"], "Driver should be local") + assert.NotNil(t, responseData["status"], "Should have status field") + assert.NotNil(t, responseData["created_at"], "Should have created_at field") + assert.NotNil(t, responseData["updated_at"], "Should have updated_at field") + + // Verify metadata + metadata, ok := responseData["metadata"].(map[string]interface{}) + assert.True(t, ok, "Should have metadata") + assert.Equal(t, "api_test", metadata["test_type"], "Metadata should match") + assert.Equal(t, "common_trace_data", metadata["test_name"], "Metadata should match") + + // Verify user info + assert.Equal(t, data.TokenInfo.UserID, responseData["created_by"], "Created by should match") + + t.Logf("Retrieved trace info for %s", data.TraceID) +} + +// TestGetInfoNotFound tests getting info for non-existent trace +func TestGetInfoNotFound(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Try to get info for non-existent trace + requestURL := fmt.Sprintf("%s%s/trace/traces/nonexistent/info", data.ServerURL, data.BaseURL) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusNotFound, resp.StatusCode, "Expected status code 404 for non-existent trace") +} diff --git a/openapi/tests/trace/logs_test.go b/openapi/tests/trace/logs_test.go new file mode 100644 index 00000000..b4bd90b9 --- /dev/null +++ b/openapi/tests/trace/logs_test.go @@ -0,0 +1,156 @@ +package trace_test + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +// TestGetLogs tests the get all logs API endpoint +func TestGetLogs(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Test GET /traces/:traceID/logs + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/logs", data.ServerURL, data.BaseURL, data.TraceID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode, "Expected status code 200") + + // Parse response + body, err := io.ReadAll(resp.Body) + assert.NoError(t, err) + + var responseData map[string]interface{} + err = json.Unmarshal(body, &responseData) + assert.NoError(t, err) + + // Verify response structure + assert.Equal(t, data.TraceID, responseData["trace_id"], "Trace ID should match") + assert.NotNil(t, responseData["logs"], "Should have logs field") + assert.NotNil(t, responseData["count"], "Should have count field") + + logs, ok := responseData["logs"].([]interface{}) + assert.True(t, ok, "Logs should be an array") + assert.NotEmpty(t, logs, "Logs array should not be empty") + + count := int(responseData["count"].(float64)) + assert.GreaterOrEqual(t, count, 6, "Should have at least 6 log entries (6 node logs)") + + // Verify log structure and collect log levels + logLevels := make(map[string]int) + for _, l := range logs { + log, ok := l.(map[string]interface{}) + assert.True(t, ok, "Each log should be an object") + assert.NotNil(t, log["timestamp"], "Log should have timestamp") + assert.NotEmpty(t, log["level"], "Log should have level") + assert.NotEmpty(t, log["message"], "Log should have message") + + level := log["level"].(string) + logLevels[level]++ + } + + assert.Greater(t, logLevels["info"], 0, "Should have info logs") + assert.Greater(t, logLevels["debug"], 0, "Should have debug logs") + assert.Greater(t, logLevels["warn"], 0, "Should have warn logs") + assert.Greater(t, logLevels["error"], 0, "Should have error logs") + + t.Logf("Retrieved %d logs for trace %s (info: %d, debug: %d, warn: %d, error: %d)", + count, data.TraceID, logLevels["info"], logLevels["debug"], logLevels["warn"], logLevels["error"]) +} + +// TestGetLogsByNode tests the get logs by node ID API endpoint +func TestGetLogsByNode(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Test GET /traces/:traceID/logs/:nodeID with Node1 + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/logs/%s", data.ServerURL, data.BaseURL, data.TraceID, data.Node1ID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode, "Expected status code 200") + + // Parse response + body, err := io.ReadAll(resp.Body) + assert.NoError(t, err) + + var responseData map[string]interface{} + err = json.Unmarshal(body, &responseData) + assert.NoError(t, err) + + // Verify response structure + assert.Equal(t, data.TraceID, responseData["trace_id"], "Trace ID should match") + assert.Equal(t, data.Node1ID, responseData["node_id"], "Node ID should match") + assert.NotNil(t, responseData["logs"], "Should have logs field") + + logs, ok := responseData["logs"].([]interface{}) + assert.True(t, ok, "Logs should be an array") + assert.NotEmpty(t, logs, "Logs array should not be empty") + + // Verify all logs belong to the specific node + for _, l := range logs { + log, ok := l.(map[string]interface{}) + assert.True(t, ok, "Each log should be an object") + assert.Equal(t, data.Node1ID, log["node_id"], "All logs should belong to the specified node") + assert.NotEmpty(t, log["message"], "Log should have message") + } + + // Should have at least 2 logs for Node1 (info + debug) + count := int(responseData["count"].(float64)) + assert.GreaterOrEqual(t, count, 2, "Node1 should have at least 2 log entries") + + t.Logf("Retrieved %d logs for node %s in trace %s", count, data.Node1ID, data.TraceID) +} + +// TestGetLogsByNodeNotFound tests getting logs for non-existent node +func TestGetLogsByNodeNotFound(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Try to get logs for non-existent node + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/logs/nonexistent", data.ServerURL, data.BaseURL, data.TraceID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + // Should return 200 with empty array (no logs for non-existent node) + assert.Equal(t, http.StatusOK, resp.StatusCode, "Expected status code 200") + + body, err := io.ReadAll(resp.Body) + assert.NoError(t, err) + + var responseData map[string]interface{} + err = json.Unmarshal(body, &responseData) + assert.NoError(t, err) + + logs, ok := responseData["logs"].([]interface{}) + assert.True(t, ok, "Logs should be an array") + assert.Empty(t, logs, "Should return empty array for non-existent node") +} diff --git a/openapi/tests/trace/nodes_test.go b/openapi/tests/trace/nodes_test.go new file mode 100644 index 00000000..877fc886 --- /dev/null +++ b/openapi/tests/trace/nodes_test.go @@ -0,0 +1,144 @@ +package trace_test + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +// TestGetNodes tests the get all nodes API endpoint +func TestGetNodes(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Test GET /traces/:traceID/nodes + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/nodes", data.ServerURL, data.BaseURL, data.TraceID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode, "Expected status code 200") + + // Parse response + body, err := io.ReadAll(resp.Body) + assert.NoError(t, err) + + var responseData map[string]interface{} + err = json.Unmarshal(body, &responseData) + assert.NoError(t, err) + + // Verify response structure + assert.Equal(t, data.TraceID, responseData["trace_id"], "Trace ID should match") + assert.NotNil(t, responseData["nodes"], "Should have nodes field") + assert.NotNil(t, responseData["count"], "Should have count field") + + nodes, ok := responseData["nodes"].([]interface{}) + assert.True(t, ok, "Nodes should be an array") + assert.NotEmpty(t, nodes, "Nodes array should not be empty") + + count := int(responseData["count"].(float64)) + assert.Equal(t, 3, count, "Should have 3 nodes (3 child nodes created)") + assert.Equal(t, count, len(nodes), "Count should match array length") + + // Verify node structure and metadata + metadataFound := 0 + for _, n := range nodes { + node, ok := n.(map[string]interface{}) + assert.True(t, ok, "Each node should be an object") + assert.NotEmpty(t, node["id"], "Node should have ID") + assert.NotNil(t, node["label"], "Node should have label") + assert.NotNil(t, node["status"], "Node should have status") + assert.NotNil(t, node["created_at"], "Node should have created_at") + + // Check if metadata is present (should be for all our test nodes) + if node["metadata"] != nil { + metadata, ok := node["metadata"].(map[string]interface{}) + assert.True(t, ok, "Metadata should be a map") + if nodeOrder, exists := metadata["node_order"]; exists { + assert.NotNil(t, nodeOrder, "node_order should exist in metadata") + metadataFound++ + } + } + } + + assert.Equal(t, 3, metadataFound, "All 3 nodes should have metadata with node_order") + + t.Logf("Retrieved %d nodes for trace %s (all with metadata)", count, data.TraceID) +} + +// TestGetNodeByID tests the get single node API endpoint +func TestGetNodeByID(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Test GET /traces/:traceID/nodes/:nodeID with Node1 + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/nodes/%s", data.ServerURL, data.BaseURL, data.TraceID, data.Node1ID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode, "Expected status code 200") + + // Parse response + body, err := io.ReadAll(resp.Body) + assert.NoError(t, err) + + var responseData map[string]interface{} + err = json.Unmarshal(body, &responseData) + assert.NoError(t, err) + + // Verify response structure + assert.Equal(t, data.Node1ID, responseData["id"], "Node ID should match") + assert.Equal(t, "First Node", responseData["label"], "Node label should match") + assert.Equal(t, "icon1", responseData["icon"], "Node icon should match") + assert.Equal(t, "First test node", responseData["description"], "Node description should match") + + // Verify metadata is present and correct + assert.NotNil(t, responseData["metadata"], "Metadata should be present") + metadata, ok := responseData["metadata"].(map[string]interface{}) + assert.True(t, ok, "Metadata should be a map") + assert.Equal(t, float64(1), metadata["node_order"], "Metadata node_order should be 1") + + // Verify input and output are present + assert.NotNil(t, responseData["input"], "Input should be present") + assert.NotNil(t, responseData["output"], "Output should be present") + + t.Logf("Retrieved node %s from trace %s with metadata: %+v", data.Node1ID, data.TraceID, metadata) +} + +// TestGetNodeByIDNotFound tests getting a non-existent node +func TestGetNodeByIDNotFound(t *testing.T) { + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Try to get non-existent node + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/nodes/nonexistent", data.ServerURL, data.BaseURL, data.TraceID) + + req, err := http.NewRequest("GET", requestURL, nil) + assert.NoError(t, err) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusNotFound, resp.StatusCode, "Expected status code 404 for non-existent node") +} diff --git a/openapi/tests/trace/spaces_test.go b/openapi/tests/trace/spaces_test.go new file mode 100644 index 00000000..b2e1258a --- /dev/null +++ b/openapi/tests/trace/spaces_test.go @@ -0,0 +1,168 @@ +package trace_test + +import ( + "encoding/json" + "fmt" + "net/http" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/trace/types" +) + +func TestGetSpaces(t *testing.T) { + // Prepare test trace with spaces + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Create additional spaces for this test + space1, err := data.Manager.CreateSpace(types.TraceSpaceOption{ + Label: "Memory Space", + Icon: "memory", + Description: "Test memory space", + }) + assert.NoError(t, err) + + space2, err := data.Manager.CreateSpace(types.TraceSpaceOption{ + Label: "Cache Space", + Icon: "cache", + Description: "Test cache space", + }) + assert.NoError(t, err) + + // Add some data to spaces + err = data.Manager.SetSpaceValue(space1.ID, "key1", "value1") + assert.NoError(t, err) + err = data.Manager.SetSpaceValue(space2.ID, "key2", "value2") + assert.NoError(t, err) + + // Make API request + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/spaces", data.ServerURL, data.BaseURL, data.TraceID) + req, _ := http.NewRequest("GET", requestURL, nil) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + // Verify response + assert.Equal(t, http.StatusOK, resp.StatusCode) + + var result map[string]any + err = json.NewDecoder(resp.Body).Decode(&result) + assert.NoError(t, err) + + // Verify structure + assert.Equal(t, data.TraceID, result["trace_id"]) + assert.NotNil(t, result["spaces"]) + assert.NotNil(t, result["count"]) + + spaces := result["spaces"].([]any) + assert.GreaterOrEqual(t, len(spaces), 2) // At least the 2 spaces we created, plus the one from prepareTestTrace + assert.Equal(t, float64(len(spaces)), result["count"]) + + // Verify space metadata (should not include data field) + spaceLabels := make(map[string]bool) + for _, s := range spaces { + space := s.(map[string]any) + assert.NotNil(t, space["id"]) + assert.NotNil(t, space["label"]) + assert.NotNil(t, space["created_at"]) + assert.NotNil(t, space["updated_at"]) + assert.Nil(t, space["data"]) // Should NOT include key-value data + spaceLabels[space["label"].(string)] = true + } + + assert.True(t, spaceLabels["Memory Space"]) + assert.True(t, spaceLabels["Cache Space"]) + + t.Logf("Retrieved %d spaces for trace %s", len(spaces), data.TraceID) +} + +func TestGetSpaceByID(t *testing.T) { + // Prepare test trace + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Create a space with specific data + space, err := data.Manager.CreateSpace(types.TraceSpaceOption{ + Label: "Detailed Space", + Icon: "memory", + Description: "Space with detailed data", + Metadata: map[string]any{"type": "cache"}, + }) + assert.NoError(t, err) + + // Add key-value data + err = data.Manager.SetSpaceValue(space.ID, "key1", "value1") + assert.NoError(t, err) + err = data.Manager.SetSpaceValue(space.ID, "key2", 123) + assert.NoError(t, err) + err = data.Manager.SetSpaceValue(space.ID, "key3", map[string]any{"nested": "data"}) + assert.NoError(t, err) + + // Make API request + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/spaces/%s", data.ServerURL, data.BaseURL, data.TraceID, space.ID) + req, _ := http.NewRequest("GET", requestURL, nil) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + // Verify response + assert.Equal(t, http.StatusOK, resp.StatusCode) + + var result map[string]any + err = json.NewDecoder(resp.Body).Decode(&result) + assert.NoError(t, err) + + // Verify space metadata + assert.Equal(t, space.ID, result["id"]) + assert.Equal(t, "Detailed Space", result["label"]) + assert.Equal(t, "memory", result["icon"]) + assert.Equal(t, "Space with detailed data", result["description"]) + assert.NotNil(t, result["created_at"]) + assert.NotNil(t, result["updated_at"]) + + // Verify metadata + metadata := result["metadata"].(map[string]any) + assert.Equal(t, "cache", metadata["type"]) + + // Verify key-value data + spaceData := result["data"].(map[string]any) + assert.Len(t, spaceData, 3) + assert.Equal(t, "value1", spaceData["key1"]) + assert.Equal(t, float64(123), spaceData["key2"]) // JSON numbers are float64 + nestedData := spaceData["key3"].(map[string]any) + assert.Equal(t, "data", nestedData["nested"]) + + t.Logf("Retrieved space %s with %d key-value pairs from trace %s", space.ID, len(spaceData), data.TraceID) +} + +func TestGetSpaceByIDNotFound(t *testing.T) { + // Prepare test trace + data := prepareTestTrace(t) + defer cleanupTestTrace(t, data) + + // Make API request with non-existent space ID + requestURL := fmt.Sprintf("%s%s/trace/traces/%s/spaces/non_existent_space", data.ServerURL, data.BaseURL, data.TraceID) + req, _ := http.NewRequest("GET", requestURL, nil) + req.Header.Set("Authorization", "Bearer "+data.TokenInfo.AccessToken) + + client := &http.Client{} + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + // Verify 404 response + assert.Equal(t, http.StatusNotFound, resp.StatusCode) + + var result map[string]any + err = json.NewDecoder(resp.Body).Decode(&result) + assert.NoError(t, err) + assert.NotNil(t, result["error"]) +} + diff --git a/openapi/trace/events.go b/openapi/trace/events.go new file mode 100644 index 00000000..a869e2aa --- /dev/null +++ b/openapi/trace/events.go @@ -0,0 +1,179 @@ +package trace + +import ( + "fmt" + "io" + + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/openapi/response" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +// GetEvents retrieves all trace events +// GET /api/__yao/openapi/v1/trace/traces/:traceID/events?stream=true +func GetEvents(c *gin.Context) { + // Get trace ID from URL parameter + traceID := c.Param("traceID") + if traceID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Trace ID is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Load trace manager and info with permission checking + manager, info, shouldRelease, err := loadTraceManager(c, traceID) + if err != nil { + respondWithLoadError(c, err) + return + } + + // Release after use if we loaded it temporarily + if shouldRelease { + defer trace.Release(traceID) + } + + // Check if stream mode is requested + streamMode := c.Query("stream") == "true" + + // Handle streaming mode + if streamMode { + handleStreamMode(c, manager, info) + return + } + + // Handle normal mode - return all events + handleNormalMode(c, manager, info) +} + +// handleStreamMode handles streaming mode for trace events (SSE) +func handleStreamMode(c *gin.Context, manager types.Manager, info *types.TraceInfo) { + // Set SSE headers + c.Header("Content-Type", "text/event-stream") + c.Header("Cache-Control", "no-cache") + c.Header("Connection", "keep-alive") + c.Header("X-Accel-Buffering", "no") + + // Subscribe to trace updates + updates, err := manager.Subscribe() + if err != nil { + // Send error as SSE event + fmt.Fprintf(c.Writer, "event: error\ndata: {\"error\":\"Failed to subscribe: %s\"}\n\n", err.Error()) + c.Writer.Flush() + return + } + + // Stream events + ctx := c.Request.Context() + clientGone := ctx.Done() + + for { + select { + case <-clientGone: + // Client disconnected + return + + case update, ok := <-updates: + if !ok { + // Channel closed + return + } + + // Format and send SSE event + err := sendSSEEvent(c.Writer, *update) + if err != nil { + return + } + + // Check if trace is complete + if update.Type == types.UpdateTypeComplete { + return + } + } + } +} + +// handleNormalMode handles normal mode for trace events (JSON array) +func handleNormalMode(c *gin.Context, manager types.Manager, info *types.TraceInfo) { + // Get all events from the beginning (timestamp 0 = all) + events, err := manager.GetEvents(0) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get events: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Determine trace status + var traceStatus types.TraceStatus + if manager.IsComplete() { + // Check last event for actual completion status + traceStatus = types.TraceStatusCompleted + for i := len(events) - 1; i >= 0; i-- { + if events[i].Type == types.UpdateTypeComplete { + if data, ok := events[i].Data.(*types.TraceCompleteData); ok { + traceStatus = data.Status + } + break + } + } + } else { + traceStatus = types.TraceStatusRunning + // Check if there are any events yet + if len(events) == 0 || (len(events) == 1 && events[0].Type == types.UpdateTypeInit) { + traceStatus = types.TraceStatusPending + } + } + + // Override with stored status if it indicates failure or cancellation + switch info.Status { + case types.TraceStatusFailed: + traceStatus = types.TraceStatusFailed + case types.TraceStatusCancelled: + traceStatus = types.TraceStatusCancelled + } + + // Prepare response data + eventsData := gin.H{ + "id": info.ID, + "status": traceStatus, + "created_at": info.CreatedAt, + "updated_at": info.UpdatedAt, + "archived": info.Archived, + "events": events, + } + + if info.ArchivedAt != nil { + eventsData["archived_at"] = *info.ArchivedAt + } + + response.RespondWithSuccess(c, response.StatusOK, eventsData) +} + +// sendSSEEvent sends a trace update as an SSE event +func sendSSEEvent(w io.Writer, update types.TraceUpdate) error { + // Write event type + _, err := fmt.Fprintf(w, "event: %s\n", update.Type) + if err != nil { + return err + } + + // Write data (JSON format) + dataJSON := formatUpdateData(update) + _, err = fmt.Fprintf(w, "data: %s\n\n", dataJSON) + if err != nil { + return err + } + + // Flush to client + if flusher, ok := w.(gin.ResponseWriter); ok { + flusher.Flush() + } + + return nil +} diff --git a/openapi/trace/helpers.go b/openapi/trace/helpers.go new file mode 100644 index 00000000..1a01651e --- /dev/null +++ b/openapi/trace/helpers.go @@ -0,0 +1,210 @@ +package trace + +import ( + "encoding/json" + "fmt" + + "github.com/gin-gonic/gin" + "github.com/yaoapp/gou/model" + "github.com/yaoapp/yao/config" + "github.com/yaoapp/yao/openapi/oauth/authorized" + oauthtypes "github.com/yaoapp/yao/openapi/oauth/types" + "github.com/yaoapp/yao/openapi/response" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +// loadTraceManager loads trace manager and info with permission checking +func loadTraceManager(c *gin.Context, traceID string) (manager types.Manager, info *types.TraceInfo, shouldRelease bool, err error) { + // Get authorized info for permission checking + authInfo := authorized.GetInfo(c) + + // Get trace info from application configuration + ctx := c.Request.Context() + + // Get configured driver + driverType, driverOptions, err := getTraceDriver() + if err != nil { + return nil, nil, false, err + } + + // Get trace info + info, err = trace.GetInfo(ctx, driverType, traceID, driverOptions...) + if err != nil { + return nil, nil, false, fmt.Errorf("trace not found: %w", err) + } + + // Check read permission + hasPermission, err := checkTracePermission(authInfo, info) + if err != nil { + return nil, nil, false, fmt.Errorf("permission check failed: %w", err) + } + + if !hasPermission { + return nil, nil, false, fmt.Errorf("no permission to access trace") + } + + // Load or get trace manager + if trace.IsLoaded(traceID) { + // Get from registry + manager, err = trace.Load(traceID) + if err != nil { + return nil, nil, false, fmt.Errorf("failed to load trace from registry: %w", err) + } + return manager, info, false, nil + } + + // Load from storage + _, manager, err = trace.LoadFromStorage(ctx, driverType, traceID, driverOptions...) + if err != nil { + return nil, nil, false, fmt.Errorf("failed to load trace from storage: %w", err) + } + + // Return true for shouldRelease since we loaded it temporarily + return manager, info, true, nil +} + +// respondWithLoadError responds with appropriate error based on load error +func respondWithLoadError(c *gin.Context, err error) { + var statusCode int + errMsg := err.Error() + + if errMsg == "trace not found" || containsString(errMsg, "trace not found:") { + statusCode = response.StatusNotFound + } else if errMsg == "no permission to access trace" || containsString(errMsg, "permission") { + statusCode = response.StatusForbidden + } else { + statusCode = response.StatusInternalServerError + } + + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: errMsg, + } + response.RespondWithError(c, statusCode, errorResp) +} + +// checkTracePermission checks if the user has permission to access the trace +func checkTracePermission(authInfo *oauthtypes.AuthorizedInfo, info *types.TraceInfo) (bool, error) { + // If no auth info, deny access + if authInfo == nil { + return false, nil + } + + // No constraints, allow access (root/admin) + if !authInfo.Constraints.TeamOnly && !authInfo.Constraints.OwnerOnly { + return true, nil + } + + // Combined Team and Owner permission validation + if authInfo.Constraints.TeamOnly && authInfo.Constraints.OwnerOnly { + if info.CreatedBy == authInfo.UserID && info.TeamID == authInfo.TeamID { + return true, nil + } + } + + // Owner only permission validation + if authInfo.Constraints.OwnerOnly && info.CreatedBy == authInfo.UserID { + return true, nil + } + + // Team only permission validation + if authInfo.Constraints.TeamOnly && info.TeamID == authInfo.TeamID { + return true, nil + } + + return false, fmt.Errorf("no permission to access trace: %s", info.ID) +} + +// getTraceDriver returns the configured trace driver type and options from global config +func getTraceDriver() (driverType string, driverOptions []any, err error) { + cfg := config.Conf + + switch cfg.Trace.Driver { + case "store": + if cfg.Trace.Store == "" { + return "", nil, fmt.Errorf("trace store ID not configured") + } + return trace.Store, []any{cfg.Trace.Store, cfg.Trace.Prefix}, nil + + case "local", "": + return trace.Local, []any{cfg.Trace.Path}, nil + + default: + return "", nil, fmt.Errorf("unsupported trace driver: %s", cfg.Trace.Driver) + } +} + +// formatUpdateData formats trace update data as JSON string +func formatUpdateData(update types.TraceUpdate) string { + // Use proper JSON marshaling + data, err := json.Marshal(update) + if err != nil { + // Fallback to basic JSON if marshaling fails + return fmt.Sprintf(`{"traceId":"%s","type":"%s","timestamp":%d,"error":"failed to marshal data"}`, + update.TraceID, update.Type, update.Timestamp) + } + return string(data) +} + +// containsString checks if a string contains a substring +func containsString(s, substr string) bool { + return len(s) >= len(substr) && (s == substr || len(s) > len(substr) && findSubstring(s, substr)) +} + +func findSubstring(s, substr string) bool { + for i := 0; i <= len(s)-len(substr); i++ { + if s[i:i+len(substr)] == substr { + return true + } + } + return false +} + +// AuthFilter returns query filters based on authorization info +// This can be used when listing traces with permission filtering +func AuthFilter(c *gin.Context, authInfo *oauthtypes.AuthorizedInfo) []model.QueryWhere { + var wheres []model.QueryWhere + + if authInfo == nil { + return wheres + } + + // No constraints, no filters needed + if !authInfo.Constraints.TeamOnly && !authInfo.Constraints.OwnerOnly { + return wheres + } + + // Combined Team and Owner constraint + if authInfo.Constraints.TeamOnly && authInfo.Constraints.OwnerOnly { + wheres = append(wheres, model.QueryWhere{ + Column: "__yao_created_by", + Value: authInfo.UserID, + }) + wheres = append(wheres, model.QueryWhere{ + Column: "__yao_team_id", + Value: authInfo.TeamID, + }) + return wheres + } + + // Owner only constraint + if authInfo.Constraints.OwnerOnly { + wheres = append(wheres, model.QueryWhere{ + Column: "__yao_created_by", + Value: authInfo.UserID, + }) + return wheres + } + + // Team only constraint + if authInfo.Constraints.TeamOnly { + wheres = append(wheres, model.QueryWhere{ + Column: "__yao_team_id", + Value: authInfo.TeamID, + }) + return wheres + } + + return wheres +} diff --git a/openapi/trace/info.go b/openapi/trace/info.go new file mode 100644 index 00000000..224a8951 --- /dev/null +++ b/openapi/trace/info.go @@ -0,0 +1,76 @@ +package trace + +import ( + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/openapi/response" + "github.com/yaoapp/yao/trace" +) + +// GetInfo retrieves trace information +// GET /api/__yao/openapi/v1/trace/traces/:traceID/info +func GetInfo(c *gin.Context) { + // Get trace ID from URL parameter + traceID := c.Param("traceID") + if traceID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Trace ID is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Load trace manager with permission checking + manager, _, shouldRelease, err := loadTraceManager(c, traceID) + if err != nil { + respondWithLoadError(c, err) + return + } + + // Release after use if we loaded it temporarily + if shouldRelease { + defer trace.Release(traceID) + } + + // Get trace info from manager (reads from storage) + info, err := manager.GetTraceInfo() + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get trace info: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Prepare response data + infoData := gin.H{ + "id": info.ID, + "driver": info.Driver, + "status": info.Status, + "created_at": info.CreatedAt, + "updated_at": info.UpdatedAt, + "archived": info.Archived, + } + + if info.ArchivedAt != nil { + infoData["archived_at"] = *info.ArchivedAt + } + + if info.Metadata != nil { + infoData["metadata"] = info.Metadata + } + + // Add user/team info if available + if info.CreatedBy != "" { + infoData["created_by"] = info.CreatedBy + } + if info.TeamID != "" { + infoData["team_id"] = info.TeamID + } + if info.TenantID != "" { + infoData["tenant_id"] = info.TenantID + } + + response.RespondWithSuccess(c, response.StatusOK, infoData) +} diff --git a/openapi/trace/logs.go b/openapi/trace/logs.go new file mode 100644 index 00000000..2315f3c5 --- /dev/null +++ b/openapi/trace/logs.go @@ -0,0 +1,96 @@ +package trace + +import ( + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/openapi/response" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +// GetLogs retrieves logs for a trace or specific node +// GET /api/__yao/openapi/v1/trace/traces/:traceID/logs?node_id=xxx +func GetLogs(c *gin.Context) { + // Get trace ID from URL parameter + traceID := c.Param("traceID") + if traceID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Trace ID is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Get optional node_id from URL parameter or query parameter + nodeID := c.Param("nodeID") + if nodeID == "" { + nodeID = c.Query("node_id") + } + + // Load trace manager with permission checking + manager, _, shouldRelease, err := loadTraceManager(c, traceID) + if err != nil { + respondWithLoadError(c, err) + return + } + + // Release after use if we loaded it temporarily + if shouldRelease { + defer trace.Release(traceID) + } + + // Get logs from manager (reads from storage) + var logs []*types.TraceLog + if nodeID != "" { + // Get logs for specific node + logs, err = manager.GetLogsByNode(nodeID) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get logs for node: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + } else { + // Get all logs + logs, err = manager.GetAllLogs() + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get logs: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + } + + // Prepare response + logList := make([]gin.H, 0, len(logs)) + for _, log := range logs { + logInfo := gin.H{ + "timestamp": log.Timestamp, + "level": log.Level, + "message": log.Message, + "node_id": log.NodeID, + } + + if len(log.Data) > 0 { + logInfo["data"] = log.Data + } + + logList = append(logList, logInfo) + } + + responseData := gin.H{ + "trace_id": traceID, + "logs": logList, + "count": len(logList), + } + + if nodeID != "" { + responseData["node_id"] = nodeID + } + + response.RespondWithSuccess(c, response.StatusOK, responseData) +} diff --git a/openapi/trace/nodes.go b/openapi/trace/nodes.go new file mode 100644 index 00000000..ab7fee1f --- /dev/null +++ b/openapi/trace/nodes.go @@ -0,0 +1,155 @@ +package trace + +import ( + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/openapi/response" + "github.com/yaoapp/yao/trace" +) + +// GetNodes retrieves all nodes in the trace +// GET /api/__yao/openapi/v1/trace/traces/:traceID/nodes +func GetNodes(c *gin.Context) { + // Get trace ID from URL parameter + traceID := c.Param("traceID") + if traceID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Trace ID is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Load trace manager with permission checking + manager, _, shouldRelease, err := loadTraceManager(c, traceID) + if err != nil { + respondWithLoadError(c, err) + return + } + + // Release after use if we loaded it temporarily + if shouldRelease { + defer trace.Release(traceID) + } + + // Get all nodes from manager (reads from storage) + nodes, err := manager.GetAllNodes() + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get nodes: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Prepare response - return flat list of nodes with basic info + nodeList := make([]gin.H, 0, len(nodes)) + for _, node := range nodes { + nodeInfo := gin.H{ + "id": node.ID, + "parent_id": node.ParentID, + "label": node.Label, + "icon": node.Icon, + "description": node.Description, + "status": node.Status, + "created_at": node.CreatedAt, + "start_time": node.StartTime, + "end_time": node.EndTime, + "updated_at": node.UpdatedAt, + } + + if node.Metadata != nil { + nodeInfo["metadata"] = node.Metadata + } + + nodeList = append(nodeList, nodeInfo) + } + + response.RespondWithSuccess(c, response.StatusOK, gin.H{ + "trace_id": traceID, + "nodes": nodeList, + "count": len(nodeList), + }) +} + +// GetNode retrieves a single node by ID +// GET /api/__yao/openapi/v1/trace/traces/:traceID/nodes/:nodeID +func GetNode(c *gin.Context) { + // Get trace ID and node ID from URL parameters + traceID := c.Param("traceID") + nodeID := c.Param("nodeID") + + if traceID == "" || nodeID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Trace ID and Node ID are required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Load trace manager with permission checking + manager, _, shouldRelease, err := loadTraceManager(c, traceID) + if err != nil { + respondWithLoadError(c, err) + return + } + + // Release after use if we loaded it temporarily + if shouldRelease { + defer trace.Release(traceID) + } + + // Get node by ID from manager (reads from storage) + node, err := manager.GetNodeByID(nodeID) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get node: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + if node == nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Node not found", + } + response.RespondWithError(c, response.StatusNotFound, errorResp) + return + } + + // Prepare detailed node response + nodeData := gin.H{ + "id": node.ID, + "parent_id": node.ParentID, + "label": node.Label, + "icon": node.Icon, + "description": node.Description, + "status": node.Status, + "input": node.Input, + "output": node.Output, + "created_at": node.CreatedAt, + "start_time": node.StartTime, + "end_time": node.EndTime, + "updated_at": node.UpdatedAt, + } + + if node.Metadata != nil { + nodeData["metadata"] = node.Metadata + } + + // Add children IDs (not full children objects to avoid deep nesting) + if len(node.Children) > 0 { + childrenIDs := make([]string, 0, len(node.Children)) + for _, child := range node.Children { + childrenIDs = append(childrenIDs, child.ID) + } + nodeData["children_ids"] = childrenIDs + } + + response.RespondWithSuccess(c, response.StatusOK, nodeData) +} + diff --git a/openapi/trace/spaces.go b/openapi/trace/spaces.go new file mode 100644 index 00000000..44bd336f --- /dev/null +++ b/openapi/trace/spaces.go @@ -0,0 +1,139 @@ +package trace + +import ( + "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/openapi/response" + "github.com/yaoapp/yao/trace" +) + +// GetSpaces retrieves all spaces in the trace (metadata only, without key-value data) +// GET /api/__yao/openapi/v1/trace/traces/:traceID/spaces +func GetSpaces(c *gin.Context) { + // Get trace ID from URL parameter + traceID := c.Param("traceID") + if traceID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Trace ID is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Load trace manager with permission checking + manager, _, shouldRelease, err := loadTraceManager(c, traceID) + if err != nil { + respondWithLoadError(c, err) + return + } + + // Release after use if we loaded it temporarily + if shouldRelease { + defer trace.Release(traceID) + } + + // Get all spaces from manager (reads from storage) + spaces, err := manager.GetAllSpaces() + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get spaces: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Prepare response - return flat list of spaces with metadata only + spaceList := make([]gin.H, 0, len(spaces)) + for _, space := range spaces { + spaceInfo := gin.H{ + "id": space.ID, + "label": space.Label, + "icon": space.Icon, + "description": space.Description, + "ttl": space.TTL, + "created_at": space.CreatedAt, + "updated_at": space.UpdatedAt, + } + + if space.Metadata != nil { + spaceInfo["metadata"] = space.Metadata + } + + spaceList = append(spaceList, spaceInfo) + } + + response.RespondWithSuccess(c, response.StatusOK, gin.H{ + "trace_id": traceID, + "spaces": spaceList, + "count": len(spaceList), + }) +} + +// GetSpace retrieves a single space by ID with all key-value data +// GET /api/__yao/openapi/v1/trace/traces/:traceID/spaces/:spaceID +func GetSpace(c *gin.Context) { + // Get trace ID and space ID from URL parameters + traceID := c.Param("traceID") + spaceID := c.Param("spaceID") + + if traceID == "" || spaceID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Trace ID and Space ID are required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Load trace manager with permission checking + manager, _, shouldRelease, err := loadTraceManager(c, traceID) + if err != nil { + respondWithLoadError(c, err) + return + } + + // Release after use if we loaded it temporarily + if shouldRelease { + defer trace.Release(traceID) + } + + // Get space by ID from manager (reads from storage with all data) + spaceData, err := manager.GetSpaceByID(spaceID) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to get space: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + if spaceData == nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Space not found", + } + response.RespondWithError(c, response.StatusNotFound, errorResp) + return + } + + // Prepare detailed space response with all key-value data + responseData := gin.H{ + "id": spaceData.ID, + "label": spaceData.Label, + "icon": spaceData.Icon, + "description": spaceData.Description, + "ttl": spaceData.TTL, + "created_at": spaceData.CreatedAt, + "updated_at": spaceData.UpdatedAt, + "data": spaceData.Data, // Include all key-value pairs + } + + if spaceData.Metadata != nil { + responseData["metadata"] = spaceData.Metadata + } + + response.RespondWithSuccess(c, response.StatusOK, responseData) +} + diff --git a/openapi/trace/trace.go b/openapi/trace/trace.go new file mode 100644 index 00000000..cad81a53 --- /dev/null +++ b/openapi/trace/trace.go @@ -0,0 +1,22 @@ +package trace + +import ( + "github.com/gin-gonic/gin" + oauthtypes "github.com/yaoapp/yao/openapi/oauth/types" +) + +// Attach attaches the trace API handlers to the router with OAuth protection +func Attach(group *gin.RouterGroup, oauth oauthtypes.OAuth) { + // Apply OAuth guard to all routes + group.Use(oauth.Guard) + + // Trace API endpoints + group.GET("/traces/:traceID/events", GetEvents) // GET /traces/:traceID/events?stream=true - Get trace events (support SSE streaming) + group.GET("/traces/:traceID/info", GetInfo) // GET /traces/:traceID/info - Get trace info + group.GET("/traces/:traceID/nodes", GetNodes) // GET /traces/:traceID/nodes - Get all nodes + group.GET("/traces/:traceID/nodes/:nodeID", GetNode) // GET /traces/:traceID/nodes/:nodeID - Get single node + group.GET("/traces/:traceID/logs", GetLogs) // GET /traces/:traceID/logs - Get all logs + group.GET("/traces/:traceID/logs/:nodeID", GetLogs) // GET /traces/:traceID/logs/:nodeID - Get logs for specific node + group.GET("/traces/:traceID/spaces", GetSpaces) // GET /traces/:traceID/spaces - Get all spaces (metadata only) + group.GET("/traces/:traceID/spaces/:spaceID", GetSpace) // GET /traces/:traceID/spaces/:spaceID - Get single space with all data +} diff --git a/trace/local/driver.go b/trace/local/driver.go index f6943049..0d7c23ab 100644 --- a/trace/local/driver.go +++ b/trace/local/driver.go @@ -225,17 +225,62 @@ func (d *Driver) LoadNode(ctx context.Context, traceID string, nodeID string) (* // LoadTrace loads the entire trace tree from disk func (d *Driver) LoadTrace(ctx context.Context, traceID string) (*types.TraceNode, error) { - // Load trace info to get root node ID - info, err := d.LoadTraceInfo(ctx, traceID) + // Check if archived and extract if needed + archived, err := d.IsArchived(ctx, traceID) if err != nil { - return nil, err + return nil, fmt.Errorf("failed to check archive status: %w", err) } - if info == nil { + if archived { + if err := d.unarchive(ctx, traceID); err != nil { + return nil, fmt.Errorf("failed to unarchive trace: %w", err) + } + } + + // Get all node files to find root + nodesDir := filepath.Join(d.getTracePath(traceID), "nodes") + entries, err := os.ReadDir(nodesDir) + if err != nil { + if os.IsNotExist(err) { + return nil, nil + } + return nil, fmt.Errorf("failed to read nodes directory: %w", err) + } + + if len(entries) == 0 { return nil, nil } - // For now, just return nil - full tree reconstruction can be implemented later - return nil, nil + // Find root node ID (node with empty ParentID) by checking each file + var rootNodeID string + for _, entry := range entries { + if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".json") { + continue + } + + nodeID := strings.TrimSuffix(entry.Name(), ".json") + filePath := filepath.Join(nodesDir, entry.Name()) + data, err := os.ReadFile(filePath) + if err != nil { + continue + } + + var pn persistNode + if err := json.Unmarshal(data, &pn); err != nil { + continue + } + + if pn.ParentID == "" { + rootNodeID = nodeID + break + } + } + + if rootNodeID == "" { + return nil, fmt.Errorf("no root node found in trace") + } + + // Load root node (this will recursively load all children) + return d.LoadNode(ctx, traceID, rootNodeID) } // SaveSpace persists a space to disk diff --git a/trace/manager.go b/trace/manager.go index ebe0d42c..7a5ef553 100644 --- a/trace/manager.go +++ b/trace/manager.go @@ -874,3 +874,141 @@ func (m *manager) ListSpaceKeys(spaceID string) []string { func (m *manager) IsComplete() bool { return m.stateIsCompleted() } + +// GetEvents retrieves all events since a specific timestamp +// since=0 returns all events from the beginning +func (m *manager) GetEvents(since int64) ([]*types.TraceUpdate, error) { + if err := m.checkContext(); err != nil { + return nil, err + } + return m.stateGetUpdates(since), nil +} + +// GetTraceInfo retrieves the trace info from storage +func (m *manager) GetTraceInfo() (*types.TraceInfo, error) { + if err := m.checkContext(); err != nil { + return nil, err + } + return m.driver.LoadTraceInfo(m.ctx, m.traceID) +} + +// GetAllNodes retrieves all nodes from storage +func (m *manager) GetAllNodes() ([]*types.TraceNode, error) { + if err := m.checkContext(); err != nil { + return nil, err + } + + // Load the root node tree from storage + rootNode, err := m.driver.LoadTrace(m.ctx, m.traceID) + if err != nil { + return nil, err + } + + if rootNode == nil { + return []*types.TraceNode{}, nil + } + + // Flatten the tree to get all nodes + var allNodes []*types.TraceNode + var collectNodes func(*types.TraceNode) + collectNodes = func(node *types.TraceNode) { + if node == nil { + return + } + allNodes = append(allNodes, node) + for _, child := range node.Children { + collectNodes(child) + } + } + collectNodes(rootNode) + + return allNodes, nil +} + +// GetNodeByID retrieves a specific node by ID from storage +func (m *manager) GetNodeByID(nodeID string) (*types.TraceNode, error) { + if err := m.checkContext(); err != nil { + return nil, err + } + return m.driver.LoadNode(m.ctx, m.traceID, nodeID) +} + +// GetAllLogs retrieves all logs from storage +func (m *manager) GetAllLogs() ([]*types.TraceLog, error) { + if err := m.checkContext(); err != nil { + return nil, err + } + return m.driver.LoadLogs(m.ctx, m.traceID, "") +} + +// GetLogsByNode retrieves logs for a specific node from storage +func (m *manager) GetLogsByNode(nodeID string) ([]*types.TraceLog, error) { + if err := m.checkContext(); err != nil { + return nil, err + } + return m.driver.LoadLogs(m.ctx, m.traceID, nodeID) +} + +// GetAllSpaces retrieves all spaces from storage +func (m *manager) GetAllSpaces() ([]*types.TraceSpace, error) { + if err := m.checkContext(); err != nil { + return nil, err + } + + // Get all space IDs from driver + spaceIDs, err := m.driver.ListSpaces(m.ctx, m.traceID) + if err != nil { + return nil, err + } + + // Load all spaces + spaces := make([]*types.TraceSpace, 0, len(spaceIDs)) + for _, spaceID := range spaceIDs { + space, err := m.driver.LoadSpace(m.ctx, m.traceID, spaceID) + if err != nil { + continue // Skip spaces that fail to load + } + if space != nil { + spaces = append(spaces, space) + } + } + + return spaces, nil +} + +// GetSpaceByID retrieves a specific space by ID from storage with all its key-value data +func (m *manager) GetSpaceByID(spaceID string) (*types.TraceSpaceData, error) { + if err := m.checkContext(); err != nil { + return nil, err + } + + // Load space metadata + space, err := m.driver.LoadSpace(m.ctx, m.traceID, spaceID) + if err != nil { + return nil, err + } + if space == nil { + return nil, nil + } + + // Load all keys in the space + keys, err := m.driver.ListSpaceKeys(m.ctx, m.traceID, spaceID) + if err != nil { + return nil, err + } + + // Load all key-value pairs + data := make(map[string]any) + for _, key := range keys { + value, err := m.driver.GetSpaceKey(m.ctx, m.traceID, spaceID, key) + if err != nil { + continue // Skip keys that fail to load + } + data[key] = value + } + + return &types.TraceSpaceData{ + TraceSpace: *space, + Data: data, + }, nil +} diff --git a/trace/store/driver.go b/trace/store/driver.go index e9e73976..e5a428a9 100644 --- a/trace/store/driver.go +++ b/trace/store/driver.go @@ -136,6 +136,10 @@ func (d *Driver) getKeyPrefix(traceID string) string { return d.prefix + ":" + traceID + ":" } +func (d *Driver) getNodeKeyPrefix(traceID string) string { + return d.getKey(traceID, "node") + ":" +} + // getTraceInfoKey returns the key for trace info func (d *Driver) getTraceInfoKey(traceID string) string { return d.getKey(traceID, "info") @@ -227,17 +231,66 @@ func (d *Driver) LoadNode(ctx context.Context, traceID string, nodeID string) (* // LoadTrace loads the entire trace tree from store func (d *Driver) LoadTrace(ctx context.Context, traceID string) (*types.TraceNode, error) { - // Load trace info to get root node ID - info, err := d.LoadTraceInfo(ctx, traceID) + // Check if archived and extract if needed + archived, err := d.IsArchived(ctx, traceID) if err != nil { - return nil, err + return nil, fmt.Errorf("failed to check archive status: %w", err) } - if info == nil { + if archived { + if err := d.unarchive(ctx, traceID); err != nil { + return nil, fmt.Errorf("failed to unarchive trace: %w", err) + } + } + + // List all node keys + nodePrefix := d.getNodeKeyPrefix(traceID) + nodeKeys, err := d.listKeysByPrefix(ctx, nodePrefix) + if err != nil { + return nil, fmt.Errorf("failed to list node keys: %w", err) + } + + if len(nodeKeys) == 0 { return nil, nil } - // For now, just return nil - full tree reconstruction can be implemented later - return nil, nil + // Find root node ID (node with empty ParentID) by checking each node + var rootNodeID string + for _, key := range nodeKeys { + // Extract node ID from key (format: prefix:traceID:nodes:nodeID) + parts := strings.Split(key, ":") + if len(parts) < 4 { + continue + } + nodeID := parts[len(parts)-1] + + // Read node data to check if it's root + data, exists := d.store.Get(key) + if !exists { + continue + } + + dataStr, ok := data.(string) + if !ok { + continue + } + + var pn persistNode + if err := json.Unmarshal([]byte(dataStr), &pn); err != nil { + continue + } + + if pn.ParentID == "" { + rootNodeID = nodeID + break + } + } + + if rootNodeID == "" { + return nil, fmt.Errorf("no root node found in trace") + } + + // Load root node (this will recursively load all children) + return d.LoadNode(ctx, traceID, rootNodeID) } // SaveSpace persists a space to store diff --git a/trace/trace_resource_test.go b/trace/trace_resource_test.go new file mode 100644 index 00000000..5ffd4a0a --- /dev/null +++ b/trace/trace_resource_test.go @@ -0,0 +1,576 @@ +package trace_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/yaoapp/yao/trace" + "github.com/yaoapp/yao/trace/types" +) + +// Note: TestMain is defined in trace_basic_test.go and applies to all tests in this package + +func TestManagerGetTraceInfo(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + t.Run(d.Name, func(t *testing.T) { + ctx := context.Background() + + // Create trace with custom metadata + option := &types.TraceOption{ + CreatedBy: "test@example.com", + TeamID: "team-001", + TenantID: "tenant-001", + Metadata: map[string]any{"test_key": "test_value"}, + } + + traceID, manager, err := trace.New(ctx, d.DriverType, option, d.DriverOptions...) + assert.NoError(t, err) + assert.NotNil(t, manager) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Get trace info through manager + info, err := manager.GetTraceInfo() + assert.NoError(t, err) + assert.NotNil(t, info) + assert.Equal(t, traceID, info.ID) + assert.Equal(t, "test@example.com", info.CreatedBy) + assert.Equal(t, "team-001", info.TeamID) + assert.Equal(t, "tenant-001", info.TenantID) + assert.Equal(t, "test_value", info.Metadata["test_key"]) + assert.Equal(t, d.DriverType, info.Driver) + }) + } +} + +func TestManagerGetAllNodes(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) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Initially no nodes + nodes, err := manager.GetAllNodes() + assert.NoError(t, err) + assert.Empty(t, nodes) + + // Add root node + node1, err := manager.Add("input1", types.TraceNodeOption{Label: "Node 1", Icon: "icon1"}) + assert.NoError(t, err) + + // Add child node + node2, err := manager.Add("input2", types.TraceNodeOption{Label: "Node 2", Icon: "icon2"}) + assert.NoError(t, err) + + // Add another child + node3, err := manager.Add("input3", types.TraceNodeOption{Label: "Node 3", Icon: "icon3"}) + assert.NoError(t, err) + + // Complete nodes to ensure they are fully persisted + err = node3.Complete() + assert.NoError(t, err) + err = node2.Complete() + assert.NoError(t, err) + err = node1.Complete() + assert.NoError(t, err) + + // Get all nodes + nodes, err = manager.GetAllNodes() + assert.NoError(t, err) + assert.Len(t, nodes, 3) + + // Verify node IDs are present + nodeIDs := make(map[string]bool) + for _, node := range nodes { + nodeIDs[node.ID] = true + } + assert.True(t, nodeIDs[node1.ID()]) + assert.True(t, nodeIDs[node2.ID()]) + assert.True(t, nodeIDs[node3.ID()]) + + // Verify node labels + nodeLabels := make(map[string]string) + for _, node := range nodes { + nodeLabels[node.ID] = node.Label + } + assert.Equal(t, "Node 1", nodeLabels[node1.ID()]) + assert.Equal(t, "Node 2", nodeLabels[node2.ID()]) + assert.Equal(t, "Node 3", nodeLabels[node3.ID()]) + }) + } +} + +func TestManagerGetNodeByID(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) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Add a node + node, err := manager.Add("test input", types.TraceNodeOption{ + Label: "Test Node", + Icon: "test", + Description: "Test Description", + }) + assert.NoError(t, err) + nodeID := node.ID() + + // Get node by ID + retrievedNode, err := manager.GetNodeByID(nodeID) + assert.NoError(t, err) + assert.NotNil(t, retrievedNode) + assert.Equal(t, nodeID, retrievedNode.ID) + assert.Equal(t, "Test Node", retrievedNode.Label) + assert.Equal(t, "test", retrievedNode.Icon) + assert.Equal(t, "Test Description", retrievedNode.Description) + assert.Equal(t, "test input", retrievedNode.Input) + + // Try to get non-existent node (should return error or nil) + nonExistentNode, err := manager.GetNodeByID("non_existent_id") + if err == nil { + // If no error, node should be nil + assert.Nil(t, nonExistentNode) + } else { + // If error, that's also acceptable + assert.Nil(t, nonExistentNode) + } + }) + } +} + +func TestManagerGetAllLogs(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) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Initially no logs + logs, err := manager.GetAllLogs() + assert.NoError(t, err) + assert.Empty(t, logs) + + // Add a node and log some messages + node, err := manager.Add("test", types.TraceNodeOption{Label: "Test Node"}) + assert.NoError(t, err) + + // Log different levels + node.Info("Info message", map[string]any{"key1": "value1"}) + node.Debug("Debug message", map[string]any{"key2": "value2"}) + node.Warn("Warning message", map[string]any{"key3": "value3"}) + node.Error("Error message", map[string]any{"key4": "value4"}) + + // Get all logs + logs, err = manager.GetAllLogs() + assert.NoError(t, err) + assert.GreaterOrEqual(t, len(logs), 4) + + // Verify log levels + levels := make(map[string]int) + for _, log := range logs { + levels[log.Level]++ + } + assert.GreaterOrEqual(t, levels["info"], 1) + assert.GreaterOrEqual(t, levels["debug"], 1) + assert.GreaterOrEqual(t, levels["warn"], 1) + assert.GreaterOrEqual(t, levels["error"], 1) + + // Verify log messages + messages := make([]string, 0) + for _, log := range logs { + messages = append(messages, log.Message) + } + assert.Contains(t, messages, "Info message") + assert.Contains(t, messages, "Debug message") + assert.Contains(t, messages, "Warning message") + assert.Contains(t, messages, "Error message") + }) + } +} + +func TestManagerGetLogsByNode(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) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Add two nodes + node1, err := manager.Add("test1", types.TraceNodeOption{Label: "Node 1"}) + assert.NoError(t, err) + + node2, err := manager.Add("test2", types.TraceNodeOption{Label: "Node 2"}) + assert.NoError(t, err) + + // Log to node1 + node1.Info("Node 1 message 1") + node1.Debug("Node 1 message 2") + + // Log to node2 + node2.Info("Node 2 message 1") + node2.Warn("Node 2 message 2") + node2.Error("Node 2 message 3") + + // Get logs for node1 + logs1, err := manager.GetLogsByNode(node1.ID()) + assert.NoError(t, err) + assert.GreaterOrEqual(t, len(logs1), 2) + + // Verify all logs belong to node1 + for _, log := range logs1 { + assert.Equal(t, node1.ID(), log.NodeID) + } + + // Get logs for node2 + logs2, err := manager.GetLogsByNode(node2.ID()) + assert.NoError(t, err) + assert.GreaterOrEqual(t, len(logs2), 3) + + // Verify all logs belong to node2 + for _, log := range logs2 { + assert.Equal(t, node2.ID(), log.NodeID) + } + + // Verify node1 and node2 logs are different + assert.NotEqual(t, len(logs1), len(logs2)) + }) + } +} + +func TestManagerResourceAccessAfterLoadFromStorage(t *testing.T) { + drivers := trace.GetTestDrivers() + + for _, d := range drivers { + t.Run(d.Name, func(t *testing.T) { + ctx := context.Background() + + // Create trace with metadata + traceID, manager, err := trace.New(ctx, d.DriverType, &types.TraceOption{ + CreatedBy: "test@example.com", + TeamID: "team-001", + TenantID: "tenant-001", + Metadata: map[string]any{"test_key": "test_value"}, + }, d.DriverOptions...) + assert.NoError(t, err) + + // Add root node + node1, err := manager.Add("input1", types.TraceNodeOption{ + Label: "Root Node", + Icon: "root", + Description: "Root node description", + }) + assert.NoError(t, err) + node1.Info("Root node info log", map[string]any{"data": "info1"}) + node1.Debug("Root node debug log", map[string]any{"data": "debug1"}) + + // Add child node + node2, err := manager.Add("input2", types.TraceNodeOption{ + Label: "Child Node", + Icon: "child", + Description: "Child node description", + }) + assert.NoError(t, err) + node2.Info("Child node info log", map[string]any{"data": "info2"}) + node2.Warn("Child node warning log", map[string]any{"data": "warn2"}) + + // Add another child node + node3, err := manager.Add("input3", types.TraceNodeOption{ + Label: "Second Child Node", + Icon: "child2", + Description: "Second child description", + }) + assert.NoError(t, err) + node3.Error("Child node error log", map[string]any{"data": "error3"}) + + // Complete nodes to ensure data is persisted + err = node3.Complete(map[string]any{"result": "success3"}) + assert.NoError(t, err) + err = node2.Complete(map[string]any{"result": "success2"}) + assert.NoError(t, err) + err = node1.Complete(map[string]any{"result": "success1"}) + assert.NoError(t, err) + + // Release from registry + err = trace.Release(traceID) + assert.NoError(t, err) + + // Load from storage + _, loadedManager, err := trace.LoadFromStorage(ctx, d.DriverType, traceID, d.DriverOptions...) + assert.NoError(t, err) + assert.NotNil(t, loadedManager) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Test GetTraceInfo + info, err := loadedManager.GetTraceInfo() + assert.NoError(t, err) + assert.Equal(t, traceID, info.ID) + assert.Equal(t, "test@example.com", info.CreatedBy) + assert.Equal(t, "team-001", info.TeamID) + assert.Equal(t, "tenant-001", info.TenantID) + assert.Equal(t, "test_value", info.Metadata["test_key"]) + + // Test GetAllNodes + nodes, err := loadedManager.GetAllNodes() + assert.NoError(t, err) + assert.Len(t, nodes, 3, "Should have 3 nodes") + + // Verify node labels + nodeLabels := make(map[string]bool) + for _, node := range nodes { + nodeLabels[node.Label] = true + } + assert.True(t, nodeLabels["Root Node"]) + assert.True(t, nodeLabels["Child Node"]) + assert.True(t, nodeLabels["Second Child Node"]) + + // Test GetNodeByID + retrievedNode, err := loadedManager.GetNodeByID(nodes[0].ID) + assert.NoError(t, err) + assert.NotNil(t, retrievedNode) + assert.Equal(t, nodes[0].Label, retrievedNode.Label) + + // Test GetAllLogs (should have at least 5 logs) + logs, err := loadedManager.GetAllLogs() + assert.NoError(t, err) + assert.GreaterOrEqual(t, len(logs), 5, "Should have at least 5 logs") + + // Verify different log levels exist + logLevels := make(map[string]bool) + for _, log := range logs { + logLevels[log.Level] = true + } + assert.True(t, logLevels["info"], "Should have info logs") + assert.True(t, logLevels["debug"], "Should have debug logs") + assert.True(t, logLevels["warn"], "Should have warn logs") + assert.True(t, logLevels["error"], "Should have error logs") + + // Test GetLogsByNode (get logs for first node) + nodeLogs, err := loadedManager.GetLogsByNode(nodes[0].ID) + assert.NoError(t, err) + assert.NotEmpty(t, nodeLogs) + // Verify all logs belong to the same node + for _, log := range nodeLogs { + assert.Equal(t, nodes[0].ID, log.NodeID) + } + }) + } +} + +func TestManagerGetEventsWithResourceAccess(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 nodes + node1, err := manager.Add("test1", types.TraceNodeOption{Label: "Node 1"}) + assert.NoError(t, err) + + node1.Info("Test message") + err = node1.Complete(map[string]any{"result": "success"}) + assert.NoError(t, err) + + // Get events + events, err := manager.GetEvents(0) + assert.NoError(t, err) + assert.NotEmpty(t, events) + + // Verify event types + eventTypes := make(map[string]bool) + for _, event := range events { + eventTypes[event.Type] = true + } + assert.True(t, eventTypes[types.UpdateTypeInit]) + assert.True(t, eventTypes[types.UpdateTypeNodeStart]) + assert.True(t, eventTypes[types.UpdateTypeLogAdded]) + assert.True(t, eventTypes[types.UpdateTypeNodeComplete]) + + // Get all nodes - should match nodes in events + nodes, err := manager.GetAllNodes() + assert.NoError(t, err) + assert.Len(t, nodes, 1) + assert.Equal(t, node1.ID(), nodes[0].ID) + + // Get logs - should match log events + logs, err := manager.GetAllLogs() + assert.NoError(t, err) + assert.NotEmpty(t, logs) + }) + } +} + +func TestManagerGetAllSpaces(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) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Initially no spaces + spaces, err := manager.GetAllSpaces() + assert.NoError(t, err) + assert.Empty(t, spaces) + + // Create spaces + space1, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Space 1", + Icon: "memory", + Description: "First test space", + }) + assert.NoError(t, err) + + space2, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Space 2", + Icon: "cache", + Description: "Second test space", + }) + assert.NoError(t, err) + + space3, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Space 3", + Icon: "store", + Description: "Third test space", + }) + assert.NoError(t, err) + + // Get all spaces + spaces, err = manager.GetAllSpaces() + assert.NoError(t, err) + assert.Len(t, spaces, 3) + + // Verify space IDs + spaceIDs := make(map[string]bool) + for _, space := range spaces { + spaceIDs[space.ID] = true + } + assert.True(t, spaceIDs[space1.ID]) + assert.True(t, spaceIDs[space2.ID]) + assert.True(t, spaceIDs[space3.ID]) + + // Verify space labels + spaceLabels := make(map[string]string) + for _, space := range spaces { + spaceLabels[space.ID] = space.Label + } + assert.Equal(t, "Space 1", spaceLabels[space1.ID]) + assert.Equal(t, "Space 2", spaceLabels[space2.ID]) + assert.Equal(t, "Space 3", spaceLabels[space3.ID]) + }) + } +} + +func TestManagerGetSpaceByID(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) + + defer trace.Release(traceID) + defer trace.Remove(ctx, d.DriverType, traceID, d.DriverOptions...) + + // Create a space + space, err := manager.CreateSpace(types.TraceSpaceOption{ + Label: "Test Space", + Icon: "memory", + Description: "Test space with data", + Metadata: map[string]any{"type": "cache"}, + }) + assert.NoError(t, err) + + // Set some key-value pairs + err = manager.SetSpaceValue(space.ID, "key1", "value1") + assert.NoError(t, err) + err = manager.SetSpaceValue(space.ID, "key2", 123) + assert.NoError(t, err) + err = manager.SetSpaceValue(space.ID, "key3", map[string]any{"nested": "data"}) + assert.NoError(t, err) + + // Get space by ID with all data + spaceData, err := manager.GetSpaceByID(space.ID) + assert.NoError(t, err) + assert.NotNil(t, spaceData) + assert.Equal(t, space.ID, spaceData.ID) + assert.Equal(t, "Test Space", spaceData.Label) + assert.Equal(t, "memory", spaceData.Icon) + assert.Equal(t, "Test space with data", spaceData.Description) + assert.Equal(t, "cache", spaceData.Metadata["type"]) + + // Verify key-value data + assert.Len(t, spaceData.Data, 3) + assert.Equal(t, "value1", spaceData.Data["key1"]) + // Note: Store driver may serialize numbers as float64 through JSON + key2Value := spaceData.Data["key2"] + if floatVal, ok := key2Value.(float64); ok { + assert.Equal(t, float64(123), floatVal) + } else { + assert.Equal(t, 123, key2Value) + } + nestedData, ok := spaceData.Data["key3"].(map[string]any) + assert.True(t, ok) + assert.Equal(t, "data", nestedData["nested"]) + + // Try to get non-existent space + nonExistentSpace, err := manager.GetSpaceByID("non_existent_id") + if err == nil { + assert.Nil(t, nonExistentSpace) + } else { + assert.Nil(t, nonExistentSpace) + } + }) + } +} diff --git a/trace/types/interfaces.go b/trace/types/interfaces.go index 18e704d6..125c0a16 100644 --- a/trace/types/interfaces.go +++ b/trace/types/interfaces.go @@ -54,6 +54,26 @@ type Manager interface { SubscribeFrom(since int64) (<-chan *TraceUpdate, error) // IsComplete checks if the trace is completed IsComplete() bool + + // Query Operations for Events + // GetEvents retrieves all events since a specific timestamp (0 = all events) + GetEvents(since int64) ([]*TraceUpdate, error) + + // Resource Access Operations - read directly from storage + // GetTraceInfo retrieves the trace info from storage + GetTraceInfo() (*TraceInfo, error) + // GetAllNodes retrieves all nodes from storage + GetAllNodes() ([]*TraceNode, error) + // GetNodeByID retrieves a specific node by ID from storage + GetNodeByID(nodeID string) (*TraceNode, error) + // GetAllLogs retrieves all logs from storage + GetAllLogs() ([]*TraceLog, error) + // GetLogsByNode retrieves logs for a specific node from storage + GetLogsByNode(nodeID string) ([]*TraceLog, error) + // GetAllSpaces retrieves all spaces metadata from storage (without key-value data) + GetAllSpaces() ([]*TraceSpace, error) + // GetSpaceByID retrieves a specific space by ID from storage (includes all key-value data) + GetSpaceByID(spaceID string) (*TraceSpaceData, error) } // Node represents a trace node with operations for tree building and logging diff --git a/trace/types/types.go b/trace/types/types.go index 9f2ed068..579c09e3 100644 --- a/trace/types/types.go +++ b/trace/types/types.go @@ -77,6 +77,12 @@ type TraceSpace struct { // Internal data storage will be managed by implementation } +// TraceSpaceData represents a space with all its key-value data (for API responses) +type TraceSpaceData struct { + TraceSpace // Embedded space metadata + Data map[string]any `json:"data"` // All key-value pairs in the space +} + // TraceParallelInput defines input and options for a parallel node type TraceParallelInput struct { Input TraceInput // Input data for the node