From be098081b4d81e500de4dd89e50bfec281a076b1 Mon Sep 17 00:00:00 2001 From: Algis Dumbris Date: Fri, 13 Mar 2026 11:57:05 +0200 Subject: [PATCH] feat: implement trace logging, REST API, and observability Add comprehensive trace logging and observability features: - Enhanced trace store with owner-scoped queries, filtering (agent, action, time range), pagination, streaming export, and retention cleanup - REST API endpoints: GET /api/traces (list with filters), GET /api/traces/export (streaming JSON/CSV), GET /api/traces/stats (action counts) - Owner isolation enforced at every layer (store, API, tests) - Hand-rolled Prometheus metrics (pure Go, zero CGO): traces_total, traces_by_action, errors_total, active_agents at GET /metrics - Configurable slog JSON handler with --log-level flag - Request ID middleware for cross-referencing logs and traces - Batch trace writing (64 entries or 100ms flush interval) - Background retention cleanup via --trace-retention flag - SQL migration 002 adds owner_id column and composite indexes Co-Authored-By: Claude Opus 4.6 --- cmd/synapbus/main.go | 98 +++- internal/api/middleware.go | 101 ++++ internal/api/router.go | 40 ++ internal/api/traces_export_handler.go | 120 +++++ internal/api/traces_handler.go | 105 ++++ internal/api/traces_handler_test.go | 502 ++++++++++++++++++ internal/storage/schema/002_trace_logging.sql | 9 + internal/trace/metrics.go | 103 ++++ internal/trace/metrics_test.go | 122 +++++ internal/trace/model.go | 54 ++ internal/trace/retention.go | 85 +++ internal/trace/sqlite_store_test.go | 344 ++++++++++++ internal/trace/store.go | 197 ++++++- internal/trace/tracer.go | 150 ++++-- internal/trace/tracer_test.go | 130 ++++- schema/002_trace_logging.sql | 9 + 16 files changed, 2119 insertions(+), 50 deletions(-) create mode 100644 internal/api/middleware.go create mode 100644 internal/api/router.go create mode 100644 internal/api/traces_export_handler.go create mode 100644 internal/api/traces_handler.go create mode 100644 internal/api/traces_handler_test.go create mode 100644 internal/storage/schema/002_trace_logging.sql create mode 100644 internal/trace/metrics.go create mode 100644 internal/trace/metrics_test.go create mode 100644 internal/trace/model.go create mode 100644 internal/trace/retention.go create mode 100644 internal/trace/sqlite_store_test.go create mode 100644 schema/002_trace_logging.sql diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index f16cccd..346b366 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -7,6 +7,8 @@ import ( "net/http" "os" "os/signal" + "strconv" + "strings" "syscall" "time" @@ -14,6 +16,7 @@ import ( "github.com/spf13/cobra" "github.com/smart-mcp-proxy/synapbus/internal/agents" + "github.com/smart-mcp-proxy/synapbus/internal/api" mcpserver "github.com/smart-mcp-proxy/synapbus/internal/mcp" "github.com/smart-mcp-proxy/synapbus/internal/messaging" "github.com/smart-mcp-proxy/synapbus/internal/storage" @@ -21,8 +24,11 @@ import ( ) var ( - port int - dataDir string + port int + dataDir string + logLevel string + metricsEnabled bool + traceRetention string ) func main() { @@ -40,15 +46,52 @@ func main() { serveCmd.Flags().IntVar(&port, "port", 8080, "HTTP server port") serveCmd.Flags().StringVar(&dataDir, "data", "./data", "Data directory for storage") + serveCmd.Flags().StringVar(&logLevel, "log-level", "info", "Log level: debug, info, warn, error") + serveCmd.Flags().BoolVar(&metricsEnabled, "metrics", false, "Enable Prometheus metrics endpoint at /metrics") + serveCmd.Flags().StringVar(&traceRetention, "trace-retention", "0", "Trace retention period (e.g. 30d, 90d, 0 for unlimited)") rootCmd.AddCommand(serveCmd) if err := rootCmd.Execute(); err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) + slog.Error("command failed", "error", err) os.Exit(1) } } +// parseLogLevel converts a string log level to slog.Level. +func parseLogLevel(level string) slog.Level { + switch strings.ToLower(level) { + case "debug": + return slog.LevelDebug + case "warn": + return slog.LevelWarn + case "error": + return slog.LevelError + default: + return slog.LevelInfo + } +} + +// parseRetentionDuration parses a retention string like "30d", "90d", or "0". +func parseRetentionDuration(s string) time.Duration { + if s == "" || s == "0" { + return 0 + } + // Try parsing as "Nd" format (days) + if strings.HasSuffix(s, "d") { + days, err := strconv.Atoi(strings.TrimSuffix(s, "d")) + if err == nil && days > 0 { + return time.Duration(days) * 24 * time.Hour + } + } + // Try standard duration parsing + d, err := time.ParseDuration(s) + if err != nil { + return 0 + } + return d +} + func runServe(cmd *cobra.Command, args []string) error { // Check for environment variable overrides if p := os.Getenv("SYNAPBUS_PORT"); p != "" { @@ -57,6 +100,21 @@ func runServe(cmd *cobra.Command, args []string) error { if d := os.Getenv("SYNAPBUS_DATA_DIR"); d != "" { dataDir = d } + if ll := os.Getenv("SYNAPBUS_LOG_LEVEL"); ll != "" { + logLevel = ll + } + if me := os.Getenv("SYNAPBUS_METRICS"); me != "" { + metricsEnabled = me == "true" || me == "1" + } + if tr := os.Getenv("SYNAPBUS_TRACE_RETENTION"); tr != "" { + traceRetention = tr + } + + // Configure slog with JSON handler + level := parseLogLevel(logLevel) + handler := slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: level}) + logger := slog.New(handler) + slog.SetDefault(logger) ctx, cancel := context.WithCancel(context.Background()) defer cancel() @@ -64,6 +122,9 @@ func runServe(cmd *cobra.Command, args []string) error { slog.Info("starting SynapBus", "port", port, "data_dir", dataDir, + "log_level", logLevel, + "metrics_enabled", metricsEnabled, + "trace_retention", traceRetention, ) // Initialize SQLite database @@ -83,6 +144,28 @@ func runServe(cmd *cobra.Command, args []string) error { tracer := trace.NewTracer(db.DB) defer tracer.Close() + // Set up metrics + var metrics *trace.Metrics + if metricsEnabled { + metrics = trace.NewMetrics() + tracer.SetMetrics(metrics) + slog.Info("prometheus metrics enabled") + } + + // Create trace store + traceStore := trace.NewSQLiteTraceStore(db.DB) + + // Set up trace retention cleanup + retentionDuration := parseRetentionDuration(traceRetention) + var retentionCleaner *trace.RetentionCleaner + if retentionDuration > 0 { + retentionCleaner = trace.NewRetentionCleaner(traceStore, retentionDuration, 1*time.Hour) + retentionCleaner.Start() + slog.Info("trace retention cleanup enabled", + "retention", retentionDuration.String(), + ) + } + // Create services msgStore := messaging.NewSQLiteMessageStore(db.DB) msgService := messaging.NewMessagingService(msgStore, tracer) @@ -103,6 +186,10 @@ func runServe(cmd *cobra.Command, args []string) error { // MCP SSE endpoint r.Mount("/mcp", mcpSrv.SSEHandler()) + // Mount API routes (traces, export, stats, metrics) + apiRouter := api.NewRouter(traceStore, metrics) + r.Mount("/", apiRouter) + // Start HTTP server addr := fmt.Sprintf(":%d", port) srv := &http.Server{ @@ -133,6 +220,11 @@ func runServe(cmd *cobra.Command, args []string) error { shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 30*time.Second) defer shutdownCancel() + // Stop retention cleaner + if retentionCleaner != nil { + retentionCleaner.Stop() + } + if err := mcpSrv.Shutdown(shutdownCtx); err != nil { slog.Error("MCP server shutdown error", "error", err) } diff --git a/internal/api/middleware.go b/internal/api/middleware.go new file mode 100644 index 0000000..a4ef0a4 --- /dev/null +++ b/internal/api/middleware.go @@ -0,0 +1,101 @@ +// Package api provides REST API handlers for the SynapBus Web UI. +package api + +import ( + "context" + "fmt" + "log/slog" + "net/http" + "time" + + "github.com/google/uuid" +) + +type contextKey string + +const ( + ownerIDKey contextKey = "owner_id" + requestIDKey contextKey = "request_id" +) + +// OwnerIDFromContext extracts the owner ID from the context. +func OwnerIDFromContext(ctx context.Context) (int64, bool) { + id, ok := ctx.Value(ownerIDKey).(int64) + return id, ok +} + +// ContextWithOwnerID stores the owner ID in the context. +func ContextWithOwnerID(ctx context.Context, id int64) context.Context { + return context.WithValue(ctx, ownerIDKey, id) +} + +// RequestIDFromContext extracts the request ID from the context. +func RequestIDFromContext(ctx context.Context) string { + id, _ := ctx.Value(requestIDKey).(string) + return id +} + +// RequestIDMiddleware generates a unique request ID per request and adds it to the context and slog. +func RequestIDMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + reqID := uuid.New().String() + ctx := context.WithValue(r.Context(), requestIDKey, reqID) + w.Header().Set("X-Request-ID", reqID) + next.ServeHTTP(w, r.WithContext(ctx)) + }) +} + +// LoggingMiddleware logs every HTTP request with structured fields. +func LoggingMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + start := time.Now() + ww := &responseWriter{ResponseWriter: w, status: http.StatusOK} + + next.ServeHTTP(ww, r) + + reqID := RequestIDFromContext(r.Context()) + slog.Info("http request", + "method", r.Method, + "path", r.URL.Path, + "status", ww.status, + "duration_ms", time.Since(start).Milliseconds(), + "request_id", reqID, + ) + }) +} + +// responseWriter wraps http.ResponseWriter to capture status code. +type responseWriter struct { + http.ResponseWriter + status int +} + +func (w *responseWriter) WriteHeader(code int) { + w.status = code + w.ResponseWriter.WriteHeader(code) +} + +// OwnerAuthMiddleware is a simple middleware that extracts owner_id from an authenticated session. +// In the full system this would validate session tokens. For now it extracts from +// a header or query param for testing purposes. In production, this integrates with +// the auth/session system. +func OwnerAuthMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Check X-Owner-ID header (set by session middleware in production) + ownerIDStr := r.Header.Get("X-Owner-ID") + if ownerIDStr == "" { + http.Error(w, `{"error":"unauthorized","message":"Authentication required"}`, http.StatusUnauthorized) + return + } + + var ownerID int64 + _, err := fmt.Sscan(ownerIDStr, &ownerID) + if err != nil || ownerID <= 0 { + http.Error(w, `{"error":"unauthorized","message":"Invalid owner ID"}`, http.StatusUnauthorized) + return + } + + ctx := ContextWithOwnerID(r.Context(), ownerID) + next.ServeHTTP(w, r.WithContext(ctx)) + }) +} diff --git a/internal/api/router.go b/internal/api/router.go new file mode 100644 index 0000000..27c0968 --- /dev/null +++ b/internal/api/router.go @@ -0,0 +1,40 @@ +package api + +import ( + "net/http" + + "github.com/go-chi/chi/v5" + + "github.com/smart-mcp-proxy/synapbus/internal/trace" +) + +// NewRouter creates a chi router with all API routes configured. +// metricsInstance may be nil if metrics are disabled. +func NewRouter(traceStore trace.TraceStore, metricsInstance *trace.Metrics) chi.Router { + r := chi.NewRouter() + + // Global middleware + r.Use(RequestIDMiddleware) + r.Use(LoggingMiddleware) + + tracesHandler := NewTracesHandler(traceStore) + + // Authenticated API routes + r.Group(func(r chi.Router) { + r.Use(OwnerAuthMiddleware) + + r.Get("/api/traces", tracesHandler.ListTraces) + r.Get("/api/traces/export", tracesHandler.ExportTraces) + r.Get("/api/traces/stats", tracesHandler.TraceStats) + }) + + // Metrics endpoint (unauthenticated, only registered when enabled) + if metricsInstance != nil { + r.Get("/metrics", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8") + metricsInstance.WritePrometheus(w) + }) + } + + return r +} diff --git a/internal/api/traces_export_handler.go b/internal/api/traces_export_handler.go new file mode 100644 index 0000000..9a72bda --- /dev/null +++ b/internal/api/traces_export_handler.go @@ -0,0 +1,120 @@ +package api + +import ( + "encoding/csv" + "encoding/json" + "fmt" + "net/http" + "strconv" + "time" + + "github.com/smart-mcp-proxy/synapbus/internal/trace" +) + +// ExportTraces handles GET /api/traces/export. +func (h *TracesHandler) ExportTraces(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized) + return + } + + filter := trace.TraceFilter{ + OwnerID: fmt.Sprintf("%d", ownerID), + AgentName: r.URL.Query().Get("agent_name"), + Action: r.URL.Query().Get("action"), + } + + if since := r.URL.Query().Get("since"); since != "" { + t, err := time.Parse(time.RFC3339, since) + if err != nil { + http.Error(w, `{"error":"invalid 'since' parameter"}`, http.StatusBadRequest) + return + } + filter.Since = &t + } + + if until := r.URL.Query().Get("until"); until != "" { + t, err := time.Parse(time.RFC3339, until) + if err != nil { + http.Error(w, `{"error":"invalid 'until' parameter"}`, http.StatusBadRequest) + return + } + filter.Until = &t + } + + // Determine format + format := r.URL.Query().Get("format") + if format == "" { + accept := r.Header.Get("Accept") + switch accept { + case "text/csv": + format = "csv" + default: + format = "json" + } + } + + dateStr := time.Now().Format("2006-01-02") + + switch format { + case "csv": + h.exportCSV(w, r, filter, dateStr) + default: + h.exportJSON(w, r, filter, dateStr) + } +} + +func (h *TracesHandler) exportJSON(w http.ResponseWriter, r *http.Request, filter trace.TraceFilter, dateStr string) { + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Content-Disposition", fmt.Sprintf(`attachment; filename="traces-%s.json"`, dateStr)) + w.Header().Set("Transfer-Encoding", "chunked") + + // Write opening bracket + w.Write([]byte("[")) + + first := true + err := h.store.QueryStream(r.Context(), filter, func(t trace.Trace) error { + if !first { + w.Write([]byte(",")) + } + first = false + + b, err := json.Marshal(t) + if err != nil { + return err + } + _, err = w.Write(b) + return err + }) + + if err != nil { + // Best effort — headers already sent + w.Write([]byte(fmt.Sprintf(`]`))) + return + } + + w.Write([]byte("]")) +} + +func (h *TracesHandler) exportCSV(w http.ResponseWriter, r *http.Request, filter trace.TraceFilter, dateStr string) { + w.Header().Set("Content-Type", "text/csv") + w.Header().Set("Content-Disposition", fmt.Sprintf(`attachment; filename="traces-%s.csv"`, dateStr)) + w.Header().Set("Transfer-Encoding", "chunked") + + cw := csv.NewWriter(w) + defer cw.Flush() + + // Write header + cw.Write([]string{"id", "agent_name", "action", "details", "timestamp"}) + + h.store.QueryStream(r.Context(), filter, func(t trace.Trace) error { + return cw.Write([]string{ + strconv.FormatInt(t.ID, 10), + t.AgentName, + t.Action, + string(t.Details), + t.Timestamp.Format(time.RFC3339), + }) + }) +} diff --git a/internal/api/traces_handler.go b/internal/api/traces_handler.go new file mode 100644 index 0000000..ed003df --- /dev/null +++ b/internal/api/traces_handler.go @@ -0,0 +1,105 @@ +package api + +import ( + "encoding/json" + "fmt" + "net/http" + "strconv" + "time" + + "github.com/smart-mcp-proxy/synapbus/internal/trace" +) + +// TracesHandler handles REST API requests for trace data. +type TracesHandler struct { + store trace.TraceStore +} + +// NewTracesHandler creates a new TracesHandler. +func NewTracesHandler(store trace.TraceStore) *TracesHandler { + return &TracesHandler{store: store} +} + +// ListTraces handles GET /api/traces. +func (h *TracesHandler) ListTraces(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized) + return + } + + filter := trace.TraceFilter{ + OwnerID: fmt.Sprintf("%d", ownerID), + AgentName: r.URL.Query().Get("agent_name"), + Action: r.URL.Query().Get("action"), + } + + if since := r.URL.Query().Get("since"); since != "" { + t, err := time.Parse(time.RFC3339, since) + if err != nil { + http.Error(w, `{"error":"invalid 'since' parameter, expected ISO 8601 format"}`, http.StatusBadRequest) + return + } + filter.Since = &t + } + + if until := r.URL.Query().Get("until"); until != "" { + t, err := time.Parse(time.RFC3339, until) + if err != nil { + http.Error(w, `{"error":"invalid 'until' parameter, expected ISO 8601 format"}`, http.StatusBadRequest) + return + } + filter.Until = &t + } + + if page := r.URL.Query().Get("page"); page != "" { + p, err := strconv.Atoi(page) + if err == nil { + filter.Page = p + } + } + + if pageSize := r.URL.Query().Get("page_size"); pageSize != "" { + ps, err := strconv.Atoi(pageSize) + if err == nil { + filter.PageSize = ps + } + } + + traces, total, err := h.store.Query(r.Context(), filter) + if err != nil { + http.Error(w, fmt.Sprintf(`{"error":"query failed: %s"}`, err), http.StatusInternalServerError) + return + } + + filter.Normalize() + resp := map[string]any{ + "traces": traces, + "total": total, + "page": filter.Page, + "page_size": filter.PageSize, + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(resp) +} + +// TraceStats handles GET /api/traces/stats. +func (h *TracesHandler) TraceStats(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized) + return + } + + counts, err := h.store.CountByAction(r.Context(), fmt.Sprintf("%d", ownerID)) + if err != nil { + http.Error(w, fmt.Sprintf(`{"error":"stats query failed: %s"}`, err), http.StatusInternalServerError) + return + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "stats": counts, + }) +} diff --git a/internal/api/traces_handler_test.go b/internal/api/traces_handler_test.go new file mode 100644 index 0000000..9432b65 --- /dev/null +++ b/internal/api/traces_handler_test.go @@ -0,0 +1,502 @@ +package api + +import ( + "context" + "database/sql" + "encoding/csv" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/smart-mcp-proxy/synapbus/internal/trace" + + _ "modernc.org/sqlite" +) + +func newTestDB(t *testing.T) *sql.DB { + t.Helper() + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + t.Fatalf("open database: %v", err) + } + t.Cleanup(func() { db.Close() }) + + _, err = db.Exec(` + CREATE TABLE IF NOT EXISTS traces ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + agent_name TEXT NOT NULL, + action TEXT NOT NULL, + details TEXT NOT NULL DEFAULT '{}', + error TEXT, + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + owner_id TEXT NOT NULL DEFAULT '' + ) + `) + if err != nil { + t.Fatalf("create traces table: %v", err) + } + + return db +} + +func seedTraces(t *testing.T, store *trace.SQLiteTraceStore, ownerID string, agentName string, actions []string) { + t.Helper() + ctx := context.Background() + now := time.Now().UTC() + for i, action := range actions { + details, _ := json.Marshal(map[string]int{"i": i}) + tr := &trace.Trace{ + OwnerID: ownerID, + AgentName: agentName, + Action: action, + Details: json.RawMessage(details), + Timestamp: now.Add(time.Duration(-i) * time.Minute), + } + if err := store.Insert(ctx, tr); err != nil { + t.Fatalf("seed trace: %v", err) + } + } +} + +func makeRequest(t *testing.T, handler http.Handler, method, path string, ownerID string) *httptest.ResponseRecorder { + t.Helper() + req := httptest.NewRequest(method, path, nil) + if ownerID != "" { + req.Header.Set("X-Owner-ID", ownerID) + } + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + return rr +} + +func TestListTraces_OwnerIsolation(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + // Seed traces for two owners + seedTraces(t, store, "1", "alice-bot", []string{"send_message", "read_inbox", "send_message"}) + seedTraces(t, store, "2", "bob-bot", []string{"send_message", "error"}) + + router := NewRouter(store, nil) + + t.Run("owner 1 sees only own traces", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK) + } + + var resp map[string]any + if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + total := int(resp["total"].(float64)) + if total != 3 { + t.Errorf("total = %d, want 3", total) + } + + traces := resp["traces"].([]any) + for _, tr := range traces { + trMap := tr.(map[string]any) + if trMap["owner_id"] != "1" { + t.Errorf("leaked trace from owner %v", trMap["owner_id"]) + } + } + }) + + t.Run("owner 2 sees only own traces", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces", "2") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", rr.Code, http.StatusOK) + } + + var resp map[string]any + if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + total := int(resp["total"].(float64)) + if total != 2 { + t.Errorf("total = %d, want 2", total) + } + }) + + t.Run("unauthenticated request rejected", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces", "") + if rr.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want %d", rr.Code, http.StatusUnauthorized) + } + }) +} + +func TestListTraces_Filters(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + // Seed various traces + seedTraces(t, store, "1", "agent-a", []string{"send_message", "read_inbox", "send_message", "error"}) + seedTraces(t, store, "1", "agent-b", []string{"send_message", "join_channel"}) + + router := NewRouter(store, nil) + + t.Run("filter by agent_name", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces?agent_name=agent-a", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + if int(resp["total"].(float64)) != 4 { + t.Errorf("total = %v, want 4", resp["total"]) + } + }) + + t.Run("filter by action", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces?action=send_message", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + if int(resp["total"].(float64)) != 3 { + t.Errorf("total = %v, want 3", resp["total"]) + } + }) + + t.Run("filter combined", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces?agent_name=agent-a&action=send_message", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + if int(resp["total"].(float64)) != 2 { + t.Errorf("total = %v, want 2", resp["total"]) + } + }) + + t.Run("pagination", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces?page=1&page_size=2", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + traces := resp["traces"].([]any) + if len(traces) != 2 { + t.Errorf("got %d traces, want 2", len(traces)) + } + if int(resp["total"].(float64)) != 6 { + t.Errorf("total = %v, want 6", resp["total"]) + } + }) + + t.Run("empty result", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces?action=nonexistent", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + if int(resp["total"].(float64)) != 0 { + t.Errorf("total = %v, want 0", resp["total"]) + } + traces := resp["traces"].([]any) + if len(traces) != 0 { + t.Errorf("got %d traces, want 0", len(traces)) + } + }) +} + +func TestListTraces_ResponseFormat(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + seedTraces(t, store, "1", "agent", []string{"send_message"}) + + router := NewRouter(store, nil) + rr := makeRequest(t, router, "GET", "/api/traces", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + // Verify JSON structure + var resp struct { + Traces []trace.Trace `json:"traces"` + Total int `json:"total"` + Page int `json:"page"` + PageSize int `json:"page_size"` + } + if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if resp.Total != 1 { + t.Errorf("total = %d, want 1", resp.Total) + } + if resp.Page != 1 { + t.Errorf("page = %d, want 1", resp.Page) + } + if resp.PageSize != 50 { + t.Errorf("page_size = %d, want 50", resp.PageSize) + } + if len(resp.Traces) != 1 { + t.Errorf("traces count = %d, want 1", len(resp.Traces)) + } + + tr := resp.Traces[0] + if tr.AgentName != "agent" { + t.Errorf("agent_name = %q, want %q", tr.AgentName, "agent") + } + if tr.Action != "send_message" { + t.Errorf("action = %q, want %q", tr.Action, "send_message") + } +} + +func TestTraceStats(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + seedTraces(t, store, "1", "agent", []string{"send_message", "send_message", "read_inbox", "error"}) + + router := NewRouter(store, nil) + rr := makeRequest(t, router, "GET", "/api/traces/stats", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + + stats := resp["stats"].(map[string]any) + if stats["send_message"].(float64) != 2 { + t.Errorf("send_message = %v, want 2", stats["send_message"]) + } +} + +func TestExportTraces_JSON(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + seedTraces(t, store, "1", "agent", []string{"send_message", "read_inbox"}) + + router := NewRouter(store, nil) + rr := makeRequest(t, router, "GET", "/api/traces/export?format=json", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + if ct := rr.Header().Get("Content-Type"); ct != "application/json" { + t.Errorf("Content-Type = %q, want application/json", ct) + } + + cd := rr.Header().Get("Content-Disposition") + if !strings.Contains(cd, "traces-") || !strings.Contains(cd, ".json") { + t.Errorf("unexpected Content-Disposition: %q", cd) + } + + var traces []trace.Trace + if err := json.Unmarshal(rr.Body.Bytes(), &traces); err != nil { + t.Fatalf("unmarshal JSON export: %v", err) + } + if len(traces) != 2 { + t.Errorf("got %d traces, want 2", len(traces)) + } +} + +func TestExportTraces_CSV(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + seedTraces(t, store, "1", "agent", []string{"send_message", "read_inbox"}) + + router := NewRouter(store, nil) + rr := makeRequest(t, router, "GET", "/api/traces/export?format=csv", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + if ct := rr.Header().Get("Content-Type"); ct != "text/csv" { + t.Errorf("Content-Type = %q, want text/csv", ct) + } + + reader := csv.NewReader(rr.Body) + records, err := reader.ReadAll() + if err != nil { + t.Fatalf("parse CSV: %v", err) + } + + // Header + 2 data rows + if len(records) != 3 { + t.Errorf("got %d CSV rows, want 3 (header + 2 data)", len(records)) + } + + // Verify header + header := records[0] + expectedHeaders := []string{"id", "agent_name", "action", "details", "timestamp"} + for i, h := range expectedHeaders { + if header[i] != h { + t.Errorf("header[%d] = %q, want %q", i, header[i], h) + } + } +} + +func TestExportTraces_Empty(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + router := NewRouter(store, nil) + + t.Run("empty JSON export", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces/export?format=json", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var traces []trace.Trace + if err := json.Unmarshal(rr.Body.Bytes(), &traces); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if len(traces) != 0 { + t.Errorf("got %d traces, want 0", len(traces)) + } + }) + + t.Run("empty CSV export", func(t *testing.T) { + rr := makeRequest(t, router, "GET", "/api/traces/export?format=csv", "1") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + reader := csv.NewReader(rr.Body) + records, err := reader.ReadAll() + if err != nil && err != io.EOF { + t.Fatalf("parse CSV: %v", err) + } + // Should have only header + if len(records) != 1 { + t.Errorf("got %d CSV rows, want 1 (header only)", len(records)) + } + }) +} + +func TestMetricsEndpoint(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + t.Run("metrics enabled", func(t *testing.T) { + metrics := trace.NewMetrics() + metrics.IncTrace("send_message") + metrics.IncTrace("send_message") + metrics.IncTrace("read_inbox") + metrics.IncError() + metrics.SetActiveAgents(3) + + router := NewRouter(store, metrics) + rr := makeRequest(t, router, "GET", "/metrics", "") + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + body := rr.Body.String() + if !strings.Contains(body, "synapbus_traces_total 3") { + t.Error("missing synapbus_traces_total metric") + } + if !strings.Contains(body, `synapbus_traces_by_action{action="send_message"} 2`) { + t.Error("missing synapbus_traces_by_action send_message metric") + } + if !strings.Contains(body, "synapbus_errors_total 1") { + t.Error("missing synapbus_errors_total metric") + } + if !strings.Contains(body, "synapbus_active_agents 3") { + t.Error("missing synapbus_active_agents metric") + } + + if ct := rr.Header().Get("Content-Type"); !strings.Contains(ct, "text/plain") { + t.Errorf("Content-Type = %q, want text/plain", ct) + } + }) + + t.Run("metrics disabled returns 404", func(t *testing.T) { + router := NewRouter(store, nil) + rr := makeRequest(t, router, "GET", "/metrics", "") + if rr.Code != http.StatusNotFound { + t.Errorf("status = %d, want %d", rr.Code, http.StatusNotFound) + } + }) +} + +func TestMultiOwnerIsolation(t *testing.T) { + db := newTestDB(t) + store := trace.NewSQLiteTraceStore(db) + + // Three owners with different agents and traces + seedTraces(t, store, "1", "alice-bot", []string{"send_message", "read_inbox"}) + seedTraces(t, store, "2", "bob-bot", []string{"send_message", "error", "join_channel"}) + seedTraces(t, store, "3", "charlie-bot", []string{"send_message"}) + + router := NewRouter(store, nil) + + for _, tc := range []struct { + ownerID string + expectedTotal int + }{ + {"1", 2}, + {"2", 3}, + {"3", 1}, + } { + t.Run("owner_"+tc.ownerID, func(t *testing.T) { + // Test list endpoint + rr := makeRequest(t, router, "GET", "/api/traces", tc.ownerID) + if rr.Code != http.StatusOK { + t.Fatalf("status = %d", rr.Code) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + total := int(resp["total"].(float64)) + if total != tc.expectedTotal { + t.Errorf("owner %s: total = %d, want %d", tc.ownerID, total, tc.expectedTotal) + } + + traces := resp["traces"].([]any) + for _, tr := range traces { + trMap := tr.(map[string]any) + if trMap["owner_id"] != tc.ownerID { + t.Errorf("owner %s: leaked trace from owner %v", tc.ownerID, trMap["owner_id"]) + } + } + + // Test export endpoint + rr = makeRequest(t, router, "GET", "/api/traces/export?format=json", tc.ownerID) + if rr.Code != http.StatusOK { + t.Fatalf("export status = %d", rr.Code) + } + + var exported []trace.Trace + json.Unmarshal(rr.Body.Bytes(), &exported) + if len(exported) != tc.expectedTotal { + t.Errorf("owner %s: exported %d, want %d", tc.ownerID, len(exported), tc.expectedTotal) + } + for _, tr := range exported { + if tr.OwnerID != tc.ownerID { + t.Errorf("owner %s: exported trace from owner %s", tc.ownerID, tr.OwnerID) + } + } + + // Test stats endpoint + rr = makeRequest(t, router, "GET", "/api/traces/stats", tc.ownerID) + if rr.Code != http.StatusOK { + t.Fatalf("stats status = %d", rr.Code) + } + }) + } +} diff --git a/internal/storage/schema/002_trace_logging.sql b/internal/storage/schema/002_trace_logging.sql new file mode 100644 index 0000000..5d51e9c --- /dev/null +++ b/internal/storage/schema/002_trace_logging.sql @@ -0,0 +1,9 @@ +-- Add owner_id to traces table for owner-scoped queries +ALTER TABLE traces ADD COLUMN owner_id TEXT NOT NULL DEFAULT ''; + +-- Composite indexes for efficient owner-scoped trace queries +CREATE INDEX idx_traces_owner_timestamp ON traces(owner_id, created_at); +CREATE INDEX idx_traces_owner_agent_timestamp ON traces(owner_id, agent_name, created_at); +CREATE INDEX idx_traces_owner_action_timestamp ON traces(owner_id, action, created_at); + +INSERT INTO schema_migrations (version) VALUES (2); diff --git a/internal/trace/metrics.go b/internal/trace/metrics.go new file mode 100644 index 0000000..4d5412f --- /dev/null +++ b/internal/trace/metrics.go @@ -0,0 +1,103 @@ +package trace + +import ( + "fmt" + "io" + "sort" + "sync" + "sync/atomic" +) + +// Metrics provides Prometheus-compatible metrics for SynapBus. +// It uses a hand-rolled implementation to avoid CGO dependencies +// from prometheus/client_golang. +type Metrics struct { + tracesTotal atomic.Int64 + errorsTotal atomic.Int64 + activeAgents atomic.Int64 + + mu sync.RWMutex + tracesByAction map[string]*atomic.Int64 +} + +// NewMetrics creates a new Metrics instance. +func NewMetrics() *Metrics { + return &Metrics{ + tracesByAction: make(map[string]*atomic.Int64), + } +} + +// IncTrace increments the total trace counter and the per-action counter. +func (m *Metrics) IncTrace(action string) { + m.tracesTotal.Add(1) + m.getOrCreateActionCounter(action).Add(1) +} + +// IncError increments the error counter. +func (m *Metrics) IncError() { + m.errorsTotal.Add(1) +} + +// SetActiveAgents sets the active agents gauge. +func (m *Metrics) SetActiveAgents(n int) { + m.activeAgents.Store(int64(n)) +} + +func (m *Metrics) getOrCreateActionCounter(action string) *atomic.Int64 { + m.mu.RLock() + counter, ok := m.tracesByAction[action] + m.mu.RUnlock() + if ok { + return counter + } + + m.mu.Lock() + defer m.mu.Unlock() + // Double-check + if counter, ok := m.tracesByAction[action]; ok { + return counter + } + counter = &atomic.Int64{} + m.tracesByAction[action] = counter + return counter +} + +// WritePrometheus writes all metrics in Prometheus exposition format to the writer. +func (m *Metrics) WritePrometheus(w io.Writer) { + fmt.Fprintf(w, "# HELP synapbus_traces_total Total number of trace entries recorded.\n") + fmt.Fprintf(w, "# TYPE synapbus_traces_total counter\n") + fmt.Fprintf(w, "synapbus_traces_total %d\n", m.tracesTotal.Load()) + fmt.Fprintf(w, "\n") + + fmt.Fprintf(w, "# HELP synapbus_traces_by_action Total traces by action type.\n") + fmt.Fprintf(w, "# TYPE synapbus_traces_by_action counter\n") + m.mu.RLock() + // Sort actions for deterministic output + actions := make([]string, 0, len(m.tracesByAction)) + for action := range m.tracesByAction { + actions = append(actions, action) + } + sort.Strings(actions) + for _, action := range actions { + counter := m.tracesByAction[action] + fmt.Fprintf(w, "synapbus_traces_by_action{action=%q} %d\n", action, counter.Load()) + } + m.mu.RUnlock() + fmt.Fprintf(w, "\n") + + fmt.Fprintf(w, "# HELP synapbus_errors_total Total number of errors recorded.\n") + fmt.Fprintf(w, "# TYPE synapbus_errors_total counter\n") + fmt.Fprintf(w, "synapbus_errors_total %d\n", m.errorsTotal.Load()) + fmt.Fprintf(w, "\n") + + fmt.Fprintf(w, "# HELP synapbus_active_agents Number of currently active agents.\n") + fmt.Fprintf(w, "# TYPE synapbus_active_agents gauge\n") + fmt.Fprintf(w, "synapbus_active_agents %d\n", m.activeAgents.Load()) +} + +// NullMetrics is a no-op metrics implementation for when metrics are disabled. +type NullMetrics struct{} + +func (NullMetrics) IncTrace(action string) {} +func (NullMetrics) IncError() {} +func (NullMetrics) SetActiveAgents(n int) {} diff --git a/internal/trace/metrics_test.go b/internal/trace/metrics_test.go new file mode 100644 index 0000000..9b4aad4 --- /dev/null +++ b/internal/trace/metrics_test.go @@ -0,0 +1,122 @@ +package trace + +import ( + "bytes" + "strings" + "testing" +) + +func TestMetrics_IncTrace(t *testing.T) { + m := NewMetrics() + + m.IncTrace("send_message") + m.IncTrace("send_message") + m.IncTrace("read_inbox") + + if got := m.tracesTotal.Load(); got != 3 { + t.Errorf("tracesTotal = %d, want 3", got) + } + + m.mu.RLock() + if got := m.tracesByAction["send_message"].Load(); got != 2 { + t.Errorf("tracesByAction[send_message] = %d, want 2", got) + } + if got := m.tracesByAction["read_inbox"].Load(); got != 1 { + t.Errorf("tracesByAction[read_inbox] = %d, want 1", got) + } + m.mu.RUnlock() +} + +func TestMetrics_IncError(t *testing.T) { + m := NewMetrics() + + m.IncError() + m.IncError() + + if got := m.errorsTotal.Load(); got != 2 { + t.Errorf("errorsTotal = %d, want 2", got) + } +} + +func TestMetrics_SetActiveAgents(t *testing.T) { + m := NewMetrics() + + m.SetActiveAgents(5) + if got := m.activeAgents.Load(); got != 5 { + t.Errorf("activeAgents = %d, want 5", got) + } + + m.SetActiveAgents(3) + if got := m.activeAgents.Load(); got != 3 { + t.Errorf("activeAgents = %d, want 3", got) + } +} + +func TestMetrics_WritePrometheus(t *testing.T) { + m := NewMetrics() + + m.IncTrace("send_message") + m.IncTrace("send_message") + m.IncTrace("read_inbox") + m.IncError() + m.SetActiveAgents(7) + + var buf bytes.Buffer + m.WritePrometheus(&buf) + output := buf.String() + + expected := []string{ + "# HELP synapbus_traces_total", + "# TYPE synapbus_traces_total counter", + "synapbus_traces_total 3", + "# HELP synapbus_traces_by_action", + "# TYPE synapbus_traces_by_action counter", + `synapbus_traces_by_action{action="read_inbox"} 1`, + `synapbus_traces_by_action{action="send_message"} 2`, + "# HELP synapbus_errors_total", + "# TYPE synapbus_errors_total counter", + "synapbus_errors_total 1", + "# HELP synapbus_active_agents", + "# TYPE synapbus_active_agents gauge", + "synapbus_active_agents 7", + } + + for _, exp := range expected { + if !strings.Contains(output, exp) { + t.Errorf("missing expected output: %q", exp) + } + } +} + +func TestNullMetrics_DoesNotPanic(t *testing.T) { + var m NullMetrics + // These should not panic + m.IncTrace("action") + m.IncError() + m.SetActiveAgents(5) +} + +func TestMetrics_ConcurrentAccess(t *testing.T) { + m := NewMetrics() + + // Concurrent writes should not race + done := make(chan struct{}) + for i := 0; i < 10; i++ { + go func() { + for j := 0; j < 100; j++ { + m.IncTrace("action") + m.IncError() + m.SetActiveAgents(j) + } + done <- struct{}{} + }() + } + + for i := 0; i < 10; i++ { + <-done + } + + if got := m.tracesTotal.Load(); got != 1000 { + t.Errorf("tracesTotal = %d, want 1000", got) + } +} diff --git a/internal/trace/model.go b/internal/trace/model.go new file mode 100644 index 0000000..4f2f783 --- /dev/null +++ b/internal/trace/model.go @@ -0,0 +1,54 @@ +package trace + +import ( + "encoding/json" + "time" +) + +// Trace represents a single recorded agent action. +type Trace struct { + ID int64 `json:"id"` + OwnerID string `json:"owner_id"` + AgentName string `json:"agent_name"` + Action string `json:"action"` + Details json.RawMessage `json:"details"` + Error string `json:"error,omitempty"` + Timestamp time.Time `json:"timestamp"` +} + +// TraceFilter defines query parameters for filtering traces. +type TraceFilter struct { + OwnerID string `json:"owner_id"` + AgentName string `json:"agent_name,omitempty"` + Action string `json:"action,omitempty"` + Since *time.Time `json:"since,omitempty"` + Until *time.Time `json:"until,omitempty"` + Page int `json:"page"` + PageSize int `json:"page_size"` +} + +// Normalize sets defaults for missing filter values. +func (f *TraceFilter) Normalize() { + if f.Page < 1 { + f.Page = 1 + } + if f.PageSize <= 0 { + f.PageSize = 50 + } + if f.PageSize > 200 { + f.PageSize = 200 + } +} + +// Offset returns the SQL offset for the current page. +func (f *TraceFilter) Offset() int { + return (f.Page - 1) * f.PageSize +} + +// TraceExportFormat defines the output format for trace exports. +type TraceExportFormat string + +const ( + ExportFormatJSON TraceExportFormat = "json" + ExportFormatCSV TraceExportFormat = "csv" +) diff --git a/internal/trace/retention.go b/internal/trace/retention.go new file mode 100644 index 0000000..c09a469 --- /dev/null +++ b/internal/trace/retention.go @@ -0,0 +1,85 @@ +package trace + +import ( + "context" + "log/slog" + "time" +) + +// RetentionCleaner periodically deletes traces older than the retention period. +type RetentionCleaner struct { + store TraceStore + retention time.Duration + interval time.Duration + logger *slog.Logger + cancel context.CancelFunc + done chan struct{} +} + +// NewRetentionCleaner creates a new retention cleaner. +// retention is the maximum age of traces to keep. +// interval is how often to run the cleanup (defaults to 1 hour). +func NewRetentionCleaner(store TraceStore, retention time.Duration, interval time.Duration) *RetentionCleaner { + if interval == 0 { + interval = 1 * time.Hour + } + return &RetentionCleaner{ + store: store, + retention: retention, + interval: interval, + logger: slog.Default().With("component", "trace-retention"), + done: make(chan struct{}), + } +} + +// Start begins the background retention cleanup goroutine. +func (rc *RetentionCleaner) Start() { + ctx, cancel := context.WithCancel(context.Background()) + rc.cancel = cancel + go rc.loop(ctx) +} + +// Stop stops the retention cleaner and waits for the goroutine to finish. +func (rc *RetentionCleaner) Stop() { + if rc.cancel != nil { + rc.cancel() + <-rc.done + } +} + +func (rc *RetentionCleaner) loop(ctx context.Context) { + defer close(rc.done) + + // Run once immediately + rc.cleanup(ctx) + + ticker := time.NewTicker(rc.interval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + rc.cleanup(ctx) + } + } +} + +func (rc *RetentionCleaner) cleanup(ctx context.Context) { + cutoff := time.Now().Add(-rc.retention) + deleted, err := rc.store.DeleteOlderThan(ctx, cutoff) + if err != nil { + rc.logger.Error("trace retention cleanup failed", + "error", err, + "cutoff", cutoff, + ) + return + } + if deleted > 0 { + rc.logger.Info("trace retention cleanup completed", + "deleted", deleted, + "cutoff", cutoff, + ) + } +} diff --git a/internal/trace/sqlite_store_test.go b/internal/trace/sqlite_store_test.go new file mode 100644 index 0000000..c2a0f5a --- /dev/null +++ b/internal/trace/sqlite_store_test.go @@ -0,0 +1,344 @@ +package trace + +import ( + "context" + "encoding/json" + "testing" + "time" +) + +func TestSQLiteTraceStore_Insert(t *testing.T) { + db := newTestDB(t) + store := NewSQLiteTraceStore(db) + ctx := context.Background() + + tr := &Trace{ + OwnerID: "1", + AgentName: "test-agent", + Action: "send_message", + Details: json.RawMessage(`{"to":"other-agent"}`), + Timestamp: time.Now().UTC(), + } + + if err := store.Insert(ctx, tr); err != nil { + t.Fatalf("Insert: %v", err) + } + + if tr.ID == 0 { + t.Error("expected non-zero ID after insert") + } +} + +func TestSQLiteTraceStore_Query(t *testing.T) { + db := newTestDB(t) + store := NewSQLiteTraceStore(db) + ctx := context.Background() + + now := time.Now().UTC() + + // Insert traces for two owners + traces := []Trace{ + {OwnerID: "1", AgentName: "agent-a", Action: "send_message", Details: json.RawMessage(`{}`), Timestamp: now.Add(-5 * time.Minute)}, + {OwnerID: "1", AgentName: "agent-a", Action: "read_inbox", Details: json.RawMessage(`{}`), Timestamp: now.Add(-4 * time.Minute)}, + {OwnerID: "1", AgentName: "agent-b", Action: "send_message", Details: json.RawMessage(`{}`), Timestamp: now.Add(-3 * time.Minute)}, + {OwnerID: "2", AgentName: "agent-c", Action: "send_message", Details: json.RawMessage(`{}`), Timestamp: now.Add(-2 * time.Minute)}, + {OwnerID: "2", AgentName: "agent-c", Action: "error", Details: json.RawMessage(`{"err":"boom"}`), Timestamp: now.Add(-1 * time.Minute)}, + } + + for i := range traces { + if err := store.Insert(ctx, &traces[i]); err != nil { + t.Fatalf("Insert trace %d: %v", i, err) + } + } + + t.Run("owner isolation", func(t *testing.T) { + result, total, err := store.Query(ctx, TraceFilter{OwnerID: "1"}) + if err != nil { + t.Fatalf("Query: %v", err) + } + if total != 3 { + t.Errorf("total = %d, want 3", total) + } + if len(result) != 3 { + t.Errorf("got %d traces, want 3", len(result)) + } + // Verify no owner 2 traces leak through + for _, tr := range result { + if tr.OwnerID != "1" { + t.Errorf("got trace with owner_id = %q, want %q", tr.OwnerID, "1") + } + } + }) + + t.Run("filter by agent_name", func(t *testing.T) { + result, total, err := store.Query(ctx, TraceFilter{OwnerID: "1", AgentName: "agent-a"}) + if err != nil { + t.Fatalf("Query: %v", err) + } + if total != 2 { + t.Errorf("total = %d, want 2", total) + } + if len(result) != 2 { + t.Errorf("got %d traces, want 2", len(result)) + } + }) + + t.Run("filter by action", func(t *testing.T) { + result, total, err := store.Query(ctx, TraceFilter{OwnerID: "1", Action: "send_message"}) + if err != nil { + t.Fatalf("Query: %v", err) + } + if total != 2 { + t.Errorf("total = %d, want 2", total) + } + if len(result) != 2 { + t.Errorf("got %d traces, want 2", len(result)) + } + }) + + t.Run("filter by time range", func(t *testing.T) { + since := now.Add(-4 * time.Minute) + until := now.Add(-2 * time.Minute) + result, total, err := store.Query(ctx, TraceFilter{ + OwnerID: "1", + Since: &since, + Until: &until, + }) + if err != nil { + t.Fatalf("Query: %v", err) + } + if total != 2 { + t.Errorf("total = %d, want 2", total) + } + if len(result) != 2 { + t.Errorf("got %d traces, want 2", len(result)) + } + }) + + t.Run("combined filters", func(t *testing.T) { + result, total, err := store.Query(ctx, TraceFilter{ + OwnerID: "1", + AgentName: "agent-a", + Action: "send_message", + }) + if err != nil { + t.Fatalf("Query: %v", err) + } + if total != 1 { + t.Errorf("total = %d, want 1", total) + } + if len(result) != 1 { + t.Errorf("got %d traces, want 1", len(result)) + } + }) + + t.Run("pagination", func(t *testing.T) { + result, total, err := store.Query(ctx, TraceFilter{OwnerID: "1", PageSize: 2, Page: 1}) + if err != nil { + t.Fatalf("Query: %v", err) + } + if total != 3 { + t.Errorf("total = %d, want 3", total) + } + if len(result) != 2 { + t.Errorf("got %d traces on page 1, want 2", len(result)) + } + + // Page 2 + result2, total2, err := store.Query(ctx, TraceFilter{OwnerID: "1", PageSize: 2, Page: 2}) + if err != nil { + t.Fatalf("Query page 2: %v", err) + } + if total2 != 3 { + t.Errorf("total = %d, want 3", total2) + } + if len(result2) != 1 { + t.Errorf("got %d traces on page 2, want 1", len(result2)) + } + }) + + t.Run("page_size capped at 200", func(t *testing.T) { + _, _, err := store.Query(ctx, TraceFilter{OwnerID: "1", PageSize: 500}) + if err != nil { + t.Fatalf("Query: %v", err) + } + // Should not error — just cap silently + }) + + t.Run("reverse chronological order", func(t *testing.T) { + result, _, err := store.Query(ctx, TraceFilter{OwnerID: "1"}) + if err != nil { + t.Fatalf("Query: %v", err) + } + for i := 1; i < len(result); i++ { + if result[i].Timestamp.After(result[i-1].Timestamp) { + t.Error("traces not in reverse chronological order") + } + } + }) +} + +func TestSQLiteTraceStore_QueryStream(t *testing.T) { + db := newTestDB(t) + store := NewSQLiteTraceStore(db) + ctx := context.Background() + + now := time.Now().UTC() + for i := 0; i < 5; i++ { + tr := &Trace{ + OwnerID: "1", + AgentName: "stream-agent", + Action: "action", + Details: json.RawMessage(`{}`), + Timestamp: now.Add(time.Duration(i) * time.Second), + } + if err := store.Insert(ctx, tr); err != nil { + t.Fatalf("Insert: %v", err) + } + } + + var count int + err := store.QueryStream(ctx, TraceFilter{OwnerID: "1"}, func(tr Trace) error { + count++ + if tr.AgentName != "stream-agent" { + t.Errorf("unexpected agent: %s", tr.AgentName) + } + return nil + }) + if err != nil { + t.Fatalf("QueryStream: %v", err) + } + if count != 5 { + t.Errorf("streamed %d traces, want 5", count) + } +} + +func TestSQLiteTraceStore_CountByAction(t *testing.T) { + db := newTestDB(t) + store := NewSQLiteTraceStore(db) + ctx := context.Background() + + now := time.Now().UTC() + actions := []string{"send_message", "send_message", "read_inbox", "error"} + for _, action := range actions { + tr := &Trace{ + OwnerID: "1", + AgentName: "agent", + Action: action, + Details: json.RawMessage(`{}`), + Timestamp: now, + } + if err := store.Insert(ctx, tr); err != nil { + t.Fatalf("Insert: %v", err) + } + } + + counts, err := store.CountByAction(ctx, "1") + if err != nil { + t.Fatalf("CountByAction: %v", err) + } + if counts["send_message"] != 2 { + t.Errorf("send_message count = %d, want 2", counts["send_message"]) + } + if counts["read_inbox"] != 1 { + t.Errorf("read_inbox count = %d, want 1", counts["read_inbox"]) + } + if counts["error"] != 1 { + t.Errorf("error count = %d, want 1", counts["error"]) + } +} + +func TestSQLiteTraceStore_DeleteOlderThan(t *testing.T) { + db := newTestDB(t) + store := NewSQLiteTraceStore(db) + ctx := context.Background() + + now := time.Now().UTC() + + // Insert old and new traces + old := &Trace{ + OwnerID: "1", + AgentName: "agent", + Action: "old_action", + Details: json.RawMessage(`{}`), + Timestamp: now.Add(-48 * time.Hour), + } + recent := &Trace{ + OwnerID: "1", + AgentName: "agent", + Action: "new_action", + Details: json.RawMessage(`{}`), + Timestamp: now, + } + + if err := store.Insert(ctx, old); err != nil { + t.Fatalf("Insert old: %v", err) + } + if err := store.Insert(ctx, recent); err != nil { + t.Fatalf("Insert recent: %v", err) + } + + // Delete traces older than 24 hours + cutoff := now.Add(-24 * time.Hour) + deleted, err := store.DeleteOlderThan(ctx, cutoff) + if err != nil { + t.Fatalf("DeleteOlderThan: %v", err) + } + if deleted != 1 { + t.Errorf("deleted = %d, want 1", deleted) + } + + // Verify the recent one still exists + traces, total, err := store.Query(ctx, TraceFilter{OwnerID: "1"}) + if err != nil { + t.Fatalf("Query: %v", err) + } + if total != 1 { + t.Errorf("remaining = %d, want 1", total) + } + if len(traces) != 1 { + t.Fatalf("expected 1 trace, got %d", len(traces)) + } + if traces[0].Action != "new_action" { + t.Errorf("remaining trace action = %q, want %q", traces[0].Action, "new_action") + } +} + +func TestSQLiteTraceStore_OwnerIsolation(t *testing.T) { + db := newTestDB(t) + store := NewSQLiteTraceStore(db) + ctx := context.Background() + + now := time.Now().UTC() + + // Three owners + for _, ownerID := range []string{"1", "2", "3"} { + for i := 0; i < 5; i++ { + tr := &Trace{ + OwnerID: ownerID, + AgentName: "agent-" + ownerID, + Action: "action", + Details: json.RawMessage(`{}`), + Timestamp: now, + } + if err := store.Insert(ctx, tr); err != nil { + t.Fatalf("Insert: %v", err) + } + } + } + + for _, ownerID := range []string{"1", "2", "3"} { + traces, total, err := store.Query(ctx, TraceFilter{OwnerID: ownerID}) + if err != nil { + t.Fatalf("Query owner %s: %v", ownerID, err) + } + if total != 5 { + t.Errorf("owner %s: total = %d, want 5", ownerID, total) + } + for _, tr := range traces { + if tr.OwnerID != ownerID { + t.Errorf("owner %s got trace with owner_id = %q", ownerID, tr.OwnerID) + } + } + } +} diff --git a/internal/trace/store.go b/internal/trace/store.go index ebf596c..6a673f8 100644 --- a/internal/trace/store.go +++ b/internal/trace/store.go @@ -3,11 +3,13 @@ package trace import ( "context" "database/sql" + "encoding/json" "fmt" + "strings" "time" ) -// StoredTrace represents a trace entry read from the database. +// StoredTrace represents a trace entry read from the database (legacy type). type StoredTrace struct { ID int64 AgentName string @@ -17,10 +19,18 @@ type StoredTrace struct { CreatedAt time.Time } -// TraceStore defines the interface for reading trace entries. +// TraceStore defines the interface for reading and managing trace entries. type TraceStore interface { + // Legacy methods (kept for backward compatibility) GetTraces(ctx context.Context, agentName string, limit int) ([]*StoredTrace, error) GetTracesByAction(ctx context.Context, action string, limit int) ([]*StoredTrace, error) + + // Enhanced query methods + Insert(ctx context.Context, t *Trace) error + Query(ctx context.Context, f TraceFilter) ([]Trace, int, error) + QueryStream(ctx context.Context, f TraceFilter, fn func(Trace) error) error + CountByAction(ctx context.Context, ownerID string) (map[string]int64, error) + DeleteOlderThan(ctx context.Context, before time.Time) (int64, error) } // SQLiteTraceStore implements TraceStore using SQLite. @@ -49,7 +59,7 @@ func (s *SQLiteTraceStore) GetTraces(ctx context.Context, agentName string, limi } defer rows.Close() - return scanTraces(rows) + return scanStoredTraces(rows) } func (s *SQLiteTraceStore) GetTracesByAction(ctx context.Context, action string, limit int) ([]*StoredTrace, error) { @@ -68,10 +78,166 @@ func (s *SQLiteTraceStore) GetTracesByAction(ctx context.Context, action string, } defer rows.Close() - return scanTraces(rows) + return scanStoredTraces(rows) } -func scanTraces(rows *sql.Rows) ([]*StoredTrace, error) { +// Insert adds a single trace entry to the database. +func (s *SQLiteTraceStore) Insert(ctx context.Context, t *Trace) error { + result, err := s.db.ExecContext(ctx, + `INSERT INTO traces (owner_id, agent_name, action, details, error, created_at) + VALUES (?, ?, ?, ?, ?, ?)`, + t.OwnerID, t.AgentName, t.Action, string(t.Details), nullString(t.Error), t.Timestamp, + ) + if err != nil { + return fmt.Errorf("insert trace: %w", err) + } + id, err := result.LastInsertId() + if err != nil { + return fmt.Errorf("get trace id: %w", err) + } + t.ID = id + return nil +} + +// Query returns traces matching the filter, along with the total count of matching records. +func (s *SQLiteTraceStore) Query(ctx context.Context, f TraceFilter) ([]Trace, int, error) { + f.Normalize() + + where, args := buildWhereClause(f) + + // Get total count + countQuery := "SELECT COUNT(*) FROM traces" + where + var total int + if err := s.db.QueryRowContext(ctx, countQuery, args...).Scan(&total); err != nil { + return nil, 0, fmt.Errorf("count traces: %w", err) + } + + // Get paginated results + query := "SELECT id, owner_id, agent_name, action, details, error, created_at FROM traces" + + where + " ORDER BY created_at DESC LIMIT ? OFFSET ?" + args = append(args, f.PageSize, f.Offset()) + + rows, err := s.db.QueryContext(ctx, query, args...) + if err != nil { + return nil, 0, fmt.Errorf("query traces: %w", err) + } + defer rows.Close() + + traces, err := scanNewTraces(rows) + if err != nil { + return nil, 0, err + } + + return traces, total, nil +} + +// QueryStream iterates over matching traces and calls fn for each one, without loading all into memory. +func (s *SQLiteTraceStore) QueryStream(ctx context.Context, f TraceFilter, fn func(Trace) error) error { + // For streaming, remove pagination — stream all matching rows + f.Page = 0 + f.PageSize = 0 + + where, args := buildWhereClause(f) + + query := "SELECT id, owner_id, agent_name, action, details, error, created_at FROM traces" + + where + " ORDER BY created_at DESC" + + rows, err := s.db.QueryContext(ctx, query, args...) + if err != nil { + return fmt.Errorf("stream traces: %w", err) + } + defer rows.Close() + + for rows.Next() { + var t Trace + var details string + var errStr sql.NullString + if err := rows.Scan(&t.ID, &t.OwnerID, &t.AgentName, &t.Action, &details, &errStr, &t.Timestamp); err != nil { + return fmt.Errorf("scan trace: %w", err) + } + t.Details = json.RawMessage(details) + if errStr.Valid { + t.Error = errStr.String + } + if err := fn(t); err != nil { + return err + } + } + return rows.Err() +} + +// CountByAction returns the count of traces grouped by action type for a given owner. +func (s *SQLiteTraceStore) CountByAction(ctx context.Context, ownerID string) (map[string]int64, error) { + query := "SELECT action, COUNT(*) FROM traces WHERE owner_id = ? GROUP BY action ORDER BY COUNT(*) DESC" + rows, err := s.db.QueryContext(ctx, query, ownerID) + if err != nil { + return nil, fmt.Errorf("count by action: %w", err) + } + defer rows.Close() + + counts := make(map[string]int64) + for rows.Next() { + var action string + var count int64 + if err := rows.Scan(&action, &count); err != nil { + return nil, err + } + counts[action] = count + } + return counts, rows.Err() +} + +// DeleteOlderThan removes traces older than the given time and returns the count of deleted rows. +func (s *SQLiteTraceStore) DeleteOlderThan(ctx context.Context, before time.Time) (int64, error) { + result, err := s.db.ExecContext(ctx, + "DELETE FROM traces WHERE created_at < ?", before, + ) + if err != nil { + return 0, fmt.Errorf("delete old traces: %w", err) + } + return result.RowsAffected() +} + +// buildWhereClause constructs the WHERE clause and args from a TraceFilter. +func buildWhereClause(f TraceFilter) (string, []any) { + var conditions []string + var args []any + + if f.OwnerID != "" { + conditions = append(conditions, "owner_id = ?") + args = append(args, f.OwnerID) + } + if f.AgentName != "" { + conditions = append(conditions, "agent_name = ?") + args = append(args, f.AgentName) + } + if f.Action != "" { + conditions = append(conditions, "action = ?") + args = append(args, f.Action) + } + if f.Since != nil { + conditions = append(conditions, "created_at >= ?") + args = append(args, *f.Since) + } + if f.Until != nil { + conditions = append(conditions, "created_at <= ?") + args = append(args, *f.Until) + } + + if len(conditions) == 0 { + return "", args + } + return " WHERE " + strings.Join(conditions, " AND "), args +} + +func nullString(s string) sql.NullString { + if s == "" { + return sql.NullString{} + } + return sql.NullString{String: s, Valid: true} +} + +func scanStoredTraces(rows *sql.Rows) ([]*StoredTrace, error) { var traces []*StoredTrace for rows.Next() { var t StoredTrace @@ -85,3 +251,24 @@ func scanTraces(rows *sql.Rows) ([]*StoredTrace, error) { } return traces, rows.Err() } + +func scanNewTraces(rows *sql.Rows) ([]Trace, error) { + var traces []Trace + for rows.Next() { + var t Trace + var details string + var errStr sql.NullString + if err := rows.Scan(&t.ID, &t.OwnerID, &t.AgentName, &t.Action, &details, &errStr, &t.Timestamp); err != nil { + return nil, fmt.Errorf("scan trace: %w", err) + } + t.Details = json.RawMessage(details) + if errStr.Valid { + t.Error = errStr.String + } + traces = append(traces, t) + } + if traces == nil { + traces = []Trace{} + } + return traces, rows.Err() +} diff --git a/internal/trace/tracer.go b/internal/trace/tracer.go index a4c9371..1e784a9 100644 --- a/internal/trace/tracer.go +++ b/internal/trace/tracer.go @@ -6,26 +6,35 @@ import ( "database/sql" "encoding/json" "log/slog" + "time" ) -// TraceEntry represents a single trace record. +// TraceEntry represents a single trace record to be written. type TraceEntry struct { + OwnerID string AgentName string Action string Details any Error string } -// Tracer records agent actions to the traces table. -// It uses a buffered channel for async recording. -type Tracer struct { - db *sql.DB - logger *slog.Logger - ch chan TraceEntry - done chan struct{} +// MetricsRecorder is an optional interface for recording trace metrics. +type MetricsRecorder interface { + IncTrace(action string) + IncError() } -// NewTracer creates a new Tracer with a buffered channel. +// Tracer records agent actions to the traces table. +// It uses a buffered channel for async recording and batches writes. +type Tracer struct { + db *sql.DB + logger *slog.Logger + ch chan TraceEntry + done chan struct{} + metrics MetricsRecorder +} + +// NewTracer creates a new Tracer with a buffered channel and batch writing. func NewTracer(db *sql.DB) *Tracer { t := &Tracer{ db: db, @@ -37,9 +46,20 @@ func NewTracer(db *sql.DB) *Tracer { return t } -// Record enqueues a trace entry for async storage. +// SetMetrics sets the metrics recorder for the tracer. +func (t *Tracer) SetMetrics(m MetricsRecorder) { + t.metrics = m +} + +// Record enqueues a trace entry for async storage (no owner ID — legacy API). func (t *Tracer) Record(ctx context.Context, agentName, action string, details any) { + t.RecordWithOwner(ctx, "", agentName, action, details) +} + +// RecordWithOwner enqueues a trace entry with an explicit owner ID. +func (t *Tracer) RecordWithOwner(ctx context.Context, ownerID, agentName, action string, details any) { entry := TraceEntry{ + OwnerID: ownerID, AgentName: agentName, Action: action, Details: details, @@ -60,9 +80,15 @@ func (t *Tracer) Record(ctx context.Context, agentName, action string, details a ) } -// RecordError enqueues a trace entry with an error. +// RecordError enqueues a trace entry with an error (no owner ID — legacy API). func (t *Tracer) RecordError(ctx context.Context, agentName, action string, details any, traceErr error) { + t.RecordErrorWithOwner(ctx, "", agentName, action, details, traceErr) +} + +// RecordErrorWithOwner enqueues a trace entry with an error and explicit owner ID. +func (t *Tracer) RecordErrorWithOwner(ctx context.Context, ownerID, agentName, action string, details any, traceErr error) { entry := TraceEntry{ + OwnerID: ownerID, AgentName: agentName, Action: action, Details: details, @@ -89,37 +115,99 @@ func (t *Tracer) Close() { func (t *Tracer) processLoop() { defer close(t.done) - for entry := range t.ch { - t.writeEntry(entry) + + batch := make([]TraceEntry, 0, 64) + ticker := time.NewTicker(100 * time.Millisecond) + defer ticker.Stop() + + for { + select { + case entry, ok := <-t.ch: + if !ok { + // Channel closed, flush remaining + if len(batch) > 0 { + t.writeBatch(batch) + } + return + } + batch = append(batch, entry) + if len(batch) >= 64 { + t.writeBatch(batch) + batch = batch[:0] + } + case <-ticker.C: + if len(batch) > 0 { + t.writeBatch(batch) + batch = batch[:0] + } + } } } -func (t *Tracer) writeEntry(entry TraceEntry) { - detailsJSON, err := json.Marshal(entry.Details) +func (t *Tracer) writeBatch(entries []TraceEntry) { + tx, err := t.db.Begin() if err != nil { - t.logger.Error("failed to marshal trace details", + t.logger.Error("failed to begin trace batch transaction", "error", err, - "agent", entry.AgentName, - "action", entry.Action, + "batch_size", len(entries), ) - detailsJSON = []byte("{}") + return } - var traceErr sql.NullString - if entry.Error != "" { - traceErr = sql.NullString{String: entry.Error, Valid: true} - } - - _, err = t.db.Exec( - `INSERT INTO traces (agent_name, action, details, error, created_at) - VALUES (?, ?, ?, ?, CURRENT_TIMESTAMP)`, - entry.AgentName, entry.Action, string(detailsJSON), traceErr, + stmt, err := tx.Prepare( + `INSERT INTO traces (owner_id, agent_name, action, details, error, created_at) + VALUES (?, ?, ?, ?, ?, ?)`, ) if err != nil { - t.logger.Error("failed to write trace entry", + t.logger.Error("failed to prepare trace insert", "error", err, - "agent", entry.AgentName, - "action", entry.Action, + ) + tx.Rollback() + return + } + defer stmt.Close() + + now := time.Now().UTC() + for _, entry := range entries { + detailsJSON, marshalErr := json.Marshal(entry.Details) + if marshalErr != nil { + t.logger.Error("failed to marshal trace details", + "error", marshalErr, + "agent", entry.AgentName, + "action", entry.Action, + ) + detailsJSON = []byte("{}") + } + + var traceErr sql.NullString + if entry.Error != "" { + traceErr = sql.NullString{String: entry.Error, Valid: true} + } + + _, execErr := stmt.Exec( + entry.OwnerID, entry.AgentName, entry.Action, string(detailsJSON), traceErr, now, + ) + if execErr != nil { + t.logger.Error("failed to write trace entry", + "error", execErr, + "agent", entry.AgentName, + "action", entry.Action, + ) + } + + // Record metrics + if t.metrics != nil { + t.metrics.IncTrace(entry.Action) + if entry.Error != "" { + t.metrics.IncError() + } + } + } + + if err := tx.Commit(); err != nil { + t.logger.Error("failed to commit trace batch", + "error", err, + "batch_size", len(entries), ) } } diff --git a/internal/trace/tracer_test.go b/internal/trace/tracer_test.go index 3f872fd..269c63a 100644 --- a/internal/trace/tracer_test.go +++ b/internal/trace/tracer_test.go @@ -18,7 +18,7 @@ func newTestDB(t *testing.T) *sql.DB { } t.Cleanup(func() { db.Close() }) - // Create traces table + // Create traces table with owner_id (matches migration 001 + 002) _, err = db.Exec(` CREATE TABLE IF NOT EXISTS traces ( id INTEGER PRIMARY KEY AUTOINCREMENT, @@ -26,7 +26,8 @@ func newTestDB(t *testing.T) *sql.DB { action TEXT NOT NULL, details TEXT NOT NULL DEFAULT '{}', error TEXT, - created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + owner_id TEXT NOT NULL DEFAULT '' ) `) if err != nil { @@ -38,10 +39,10 @@ func newTestDB(t *testing.T) *sql.DB { func TestTracer_Record(t *testing.T) { tests := []struct { - name string - agent string - action string - details any + name string + agent string + action string + details any }{ { name: "simple trace", @@ -59,9 +60,9 @@ func TestTracer_Record(t *testing.T) { details: nil, }, { - name: "trace with string details", - agent: "agent-b", - action: "search", + name: "trace with string details", + agent: "agent-b", + action: "search", details: "query string", }, } @@ -76,7 +77,7 @@ func TestTracer_Record(t *testing.T) { tracer.Record(ctx, tt.agent, tt.action, tt.details) // Give async writer time to process - time.Sleep(50 * time.Millisecond) + time.Sleep(200 * time.Millisecond) // Verify trace was written var count int @@ -94,6 +95,28 @@ func TestTracer_Record(t *testing.T) { } } +func TestTracer_RecordWithOwner(t *testing.T) { + db := newTestDB(t) + tracer := NewTracer(db) + defer tracer.Close() + + ctx := context.Background() + tracer.RecordWithOwner(ctx, "42", "test-agent", "send_message", map[string]any{"to": "other"}) + + time.Sleep(200 * time.Millisecond) + + var ownerID string + err := db.QueryRow( + "SELECT owner_id FROM traces WHERE agent_name = 'test-agent'", + ).Scan(&ownerID) + if err != nil { + t.Fatalf("query trace: %v", err) + } + if ownerID != "42" { + t.Errorf("owner_id = %q, want %q", ownerID, "42") + } +} + func TestTracer_RecordError(t *testing.T) { db := newTestDB(t) tracer := NewTracer(db) @@ -105,7 +128,7 @@ func TestTracer_RecordError(t *testing.T) { fmt.Errorf("something went wrong"), ) - time.Sleep(50 * time.Millisecond) + time.Sleep(200 * time.Millisecond) var errorText sql.NullString err := db.QueryRow( @@ -119,6 +142,91 @@ func TestTracer_RecordError(t *testing.T) { } } +func TestTracer_BatchFlush(t *testing.T) { + db := newTestDB(t) + tracer := NewTracer(db) + + ctx := context.Background() + + // Record multiple entries to trigger batch + for i := 0; i < 10; i++ { + tracer.Record(ctx, "batch-agent", "action", map[string]any{"i": i}) + } + + // Close flushes remaining entries + tracer.Close() + + var count int + err := db.QueryRow("SELECT COUNT(*) FROM traces WHERE agent_name = 'batch-agent'").Scan(&count) + if err != nil { + t.Fatalf("query: %v", err) + } + if count != 10 { + t.Errorf("got %d traces, want 10", count) + } +} + +func TestTracer_GracefulShutdownFlushes(t *testing.T) { + db := newTestDB(t) + tracer := NewTracer(db) + + ctx := context.Background() + + // Record entries + tracer.Record(ctx, "shutdown-agent", "action1", nil) + tracer.Record(ctx, "shutdown-agent", "action2", nil) + + // Close should flush + tracer.Close() + + var count int + err := db.QueryRow("SELECT COUNT(*) FROM traces WHERE agent_name = 'shutdown-agent'").Scan(&count) + if err != nil { + t.Fatalf("query: %v", err) + } + if count != 2 { + t.Errorf("got %d traces after shutdown, want 2", count) + } +} + +func TestTracer_RecordReturnsImmediately(t *testing.T) { + db := newTestDB(t) + tracer := NewTracer(db) + defer tracer.Close() + + ctx := context.Background() + + start := time.Now() + tracer.Record(ctx, "fast-agent", "action", nil) + elapsed := time.Since(start) + + // Record should return in under 5ms (it's async) + if elapsed > 5*time.Millisecond { + t.Errorf("Record took %v, expected < 5ms (should be non-blocking)", elapsed) + } +} + +func TestTracer_Metrics(t *testing.T) { + db := newTestDB(t) + tracer := NewTracer(db) + + metrics := NewMetrics() + tracer.SetMetrics(metrics) + + ctx := context.Background() + tracer.RecordWithOwner(ctx, "1", "agent", "send_message", nil) + tracer.RecordErrorWithOwner(ctx, "1", "agent", "error", nil, fmt.Errorf("boom")) + + tracer.Close() + + if got := metrics.tracesTotal.Load(); got != 2 { + t.Errorf("tracesTotal = %d, want 2", got) + } + if got := metrics.errorsTotal.Load(); got != 1 { + t.Errorf("errorsTotal = %d, want 1", got) + } +} + func TestTraceStore(t *testing.T) { db := newTestDB(t) tracer := NewTracer(db) diff --git a/schema/002_trace_logging.sql b/schema/002_trace_logging.sql new file mode 100644 index 0000000..5d51e9c --- /dev/null +++ b/schema/002_trace_logging.sql @@ -0,0 +1,9 @@ +-- Add owner_id to traces table for owner-scoped queries +ALTER TABLE traces ADD COLUMN owner_id TEXT NOT NULL DEFAULT ''; + +-- Composite indexes for efficient owner-scoped trace queries +CREATE INDEX idx_traces_owner_timestamp ON traces(owner_id, created_at); +CREATE INDEX idx_traces_owner_agent_timestamp ON traces(owner_id, agent_name, created_at); +CREATE INDEX idx_traces_owner_action_timestamp ON traces(owner_id, action, created_at); + +INSERT INTO schema_migrations (version) VALUES (2);