diff --git a/.golangci.yml b/.golangci.yml index 4d93ae7..22225c1 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -1,76 +1,70 @@ -# golangci-lint configuration for dev-cli -# Focused on Docker API safety and common Go bugs -# Run: golangci-lint run ./... - run: timeout: 5m - modules-download-mode: readonly + tests: true + relative-path-mode: cfg + issues-exit-code: 1 + +output: + formats: + - format: colored-line-number # Best for CLI usage + print-issued-lines: true + print-linter-name: true linters: + disable-all: true enable: - # Critical for Docker SDK - catch unhandled errors - - errcheck - # Static analysis for common bugs - - staticcheck - # Lightweight linter, catches common mistakes - - revive - # Detect ineffectual assignments - - ineffassign - # Check for unchecked type assertions - - unconvert - # Detect unused code - - unused - # Check for goroutine leaks (important for streaming logs) - - govet - # Check printf-style functions - - goprintffuncname + - errcheck # Checks for unchecked errors + - govet # Official Go tool + - ineffassign # Detects unused assignments + - typecheck # Parses and type-checks Go code + - unused # Checks for unused constants, variables, functions - disable: - # Too noisy for dev tools - - gocritic - - gocyclo - - funlen - - gocognit + - staticcheck # Massive set of best practice & style rules -linters-settings: - errcheck: - # Check for ignored errors in defer statements - check-blank: true - # Don't report on explicitly ignored errors with _ - exclude-functions: - - io.Copy - - (io.Closer).Close + - revive # Faster, configurable replacement for deprecated 'golint' + - gosec # Security scanner (SQL injection, hardcoded creds) + - godoclint # (New 2025) Validates comments against Go standards [web:101] + + - bodyclose # Checks if HTTP response bodies are closed + - noctx # Ensures you send context.Context to functions that need it + - dogsled # Checks for too many blank identifiers (e.g. _, _, _, err) + - unconvert # Remove unnecessary type conversions + - goconst # Finds repeated strings that could be constants + - exportloopref # Checks for pointers to enclosing loop variables +linters-settings: revive: rules: - - name: blank-imports - - name: context-as-argument - - name: context-keys-type - - name: error-return - - name: error-strings - name: exported - - name: increment-decrement - - name: var-declaration + severity: warning + disabled: false + arguments: ["disableStuttering"] # Check for types like "UserUser" - name: package-comments - disabled: true # Not needed for internal packages + severity: warning + disabled: false - staticcheck: - checks: - - all - - -SA1019 # Ignore deprecation warnings for now + # gosec: + # excludes: + # - G101 # "Potential hardcoded credential" (often false positives in tests) + # config: + # G306: "0600" # Allow only strict file permissions + # + # staticcheck: + # checks: ["all", "-ST1000"] # Enable all, but disable package comment check -issues: - # Maximum issues count per one linter - max-issues-per-linter: 50 - max-same-issues: 10 + # Goconst + goconst: + min-len: 3 + min-occurrences: 3 +issues: + max-issues-per-linter: 0 + max-same-issues: 0 exclude-rules: - # Exclude some linters from running on test files - path: _test\.go linters: + - gosec - errcheck - - # Exclude lll issues for long lines in go.mod - - path: go\.mod + - path: tools/ linters: - - lll + - gochecknoglobals diff --git a/cmd/ask.go b/cmd/ask.go index 5f624dc..b3a81e9 100644 --- a/cmd/ask.go +++ b/cmd/ask.go @@ -2,6 +2,8 @@ package cmd import ( "bytes" + "dev-cli/internal/ai" + "dev-cli/internal/core" "encoding/json" "fmt" "io" @@ -10,9 +12,6 @@ import ( "strings" "time" - "dev-cli/internal/config" - "dev-cli/internal/llm" - "github.com/briandowns/spinner" "github.com/spf13/cobra" ) @@ -47,7 +46,7 @@ Two modes: os.Setenv("DEV_CLI_FORCE_LOCAL", "1") } - if err := llm.EnsureOllamaRunning(); err != nil { + if err := ai.EnsureOllamaRunning(); err != nil { fmt.Fprintf(os.Stderr, "\033[33m⚠\033[0m Ollama not available: %v\n", err) } @@ -105,7 +104,7 @@ func looksLikeToolName(args []string) bool { } func fetchSolutions(query string) { - client := llm.NewHybridClient() + client := ai.NewHybridClient() backend := "Ollama" if client.HasPerplexity() { @@ -166,7 +165,7 @@ func fetchSolutions(query string) { } func fetchCommands(toolName, topic string, count int) { - cfg := config.Load() + cfg := core.LoadConfig() baseURL := cfg.OllamaURL model := cfg.OllamaModel diff --git a/cmd/doctor.go b/cmd/doctor.go index 7928ef2..5d369f7 100644 --- a/cmd/doctor.go +++ b/cmd/doctor.go @@ -2,6 +2,7 @@ package cmd import ( "context" + "encoding/json" "fmt" "net/http" "os" @@ -16,6 +17,7 @@ import ( var ( doctorFix bool doctorQuiet bool + doctorJSON bool ) type CheckResult struct { @@ -26,6 +28,29 @@ type CheckResult struct { FixFunc func() error } +// DoctorReport is the JSON output format for agent consumption. +type DoctorReport struct { + Timestamp string `json:"timestamp"` + Checks []CheckResultJSON `json:"checks"` + Summary DoctorSummary `json:"summary"` +} + +// CheckResultJSON is the JSON-serializable version of CheckResult. +type CheckResultJSON struct { + Name string `json:"name"` + Status string `json:"status"` + Message string `json:"message"` + FixCmd string `json:"fix_cmd,omitempty"` +} + +// DoctorSummary contains aggregate check results. +type DoctorSummary struct { + Passed int `json:"passed"` + Warned int `json:"warned"` + Failed int `json:"failed"` + Total int `json:"total"` +} + var doctorCmd = &cobra.Command{ Use: "doctor", Short: "Check system health and fix issues", @@ -49,12 +74,10 @@ func init() { rootCmd.AddCommand(doctorCmd) doctorCmd.Flags().BoolVar(&doctorFix, "fix", false, "Attempt to auto-fix issues") doctorCmd.Flags().BoolVar(&doctorQuiet, "quiet", false, "Only show failures") + doctorCmd.Flags().BoolVar(&doctorJSON, "json", false, "Output results as JSON for agent consumption") } func runDoctor(cmd *cobra.Command, args []string) { - fmt.Println("\033[1m🔍 dev-cli doctor\033[0m") - fmt.Println() - checks := []func() CheckResult{ checkDocker, checkDockerCompose, @@ -66,19 +89,20 @@ func runDoctor(cmd *cobra.Command, args []string) { } var failed, warned, passed int + var jsonResults []CheckResultJSON for _, check := range checks { result := check() - if doctorQuiet && result.Status == "ok" { - passed++ - continue + if doctorJSON { + jsonResults = append(jsonResults, CheckResultJSON{ + Name: result.Name, + Status: result.Status, + Message: result.Message, + FixCmd: result.FixCmd, + }) } - icon := getStatusIcon(result.Status) - fmt.Printf("%s \033[1m%s\033[0m\n", icon, result.Name) - fmt.Printf(" %s\n", result.Message) - switch result.Status { case "ok": passed++ @@ -86,6 +110,21 @@ func runDoctor(cmd *cobra.Command, args []string) { warned++ case "fail": failed++ + } + + if doctorJSON { + continue + } + + if doctorQuiet && result.Status == "ok" { + continue + } + + icon := getStatusIcon(result.Status) + fmt.Printf("%s \033[1m%s\033[0m\n", icon, result.Name) + fmt.Printf(" %s\n", result.Message) + + if result.Status == "fail" { if doctorFix && (result.FixCmd != "" || result.FixFunc != nil) { fmt.Printf(" \033[33m➜ Attempting fix...\033[0m\n") if err := attemptFix(result); err != nil { @@ -102,7 +141,27 @@ func runDoctor(cmd *cobra.Command, args []string) { fmt.Println() } - // Summary + if doctorJSON { + report := DoctorReport{ + Timestamp: time.Now().Format(time.RFC3339), + Checks: jsonResults, + Summary: DoctorSummary{ + Passed: passed, + Warned: warned, + Failed: failed, + Total: len(checks), + }, + } + enc := json.NewEncoder(os.Stdout) + enc.SetIndent("", " ") + _ = enc.Encode(report) + if failed > 0 { + os.Exit(1) + } + return + } + + fmt.Println("\033[1m🔍 dev-cli doctor\033[0m") fmt.Println("\033[90m────────────────────────────────\033[0m") fmt.Printf("✓ %d passed ", passed) if warned > 0 { @@ -178,7 +237,6 @@ func checkDocker() CheckResult { } } - // Get version versionCmd := exec.Command("docker", "--version") versionOutput, _ := versionCmd.Output() version := strings.TrimSpace(string(versionOutput)) @@ -191,7 +249,6 @@ func checkDocker() CheckResult { func checkDockerCompose() CheckResult { result := CheckResult{Name: "Docker Compose"} - // Check docker compose (plugin) if err := exec.Command("docker", "compose", "version").Run(); err == nil { cmd := exec.Command("docker", "compose", "version", "--short") output, _ := cmd.Output() @@ -200,7 +257,6 @@ func checkDockerCompose() CheckResult { return result } - // Check docker-compose (standalone) if _, err := exec.LookPath("docker-compose"); err == nil { cmd := exec.Command("docker-compose", "--version") output, _ := cmd.Output() @@ -209,14 +265,13 @@ func checkDockerCompose() CheckResult { return result } - // Neither available return CheckResult{ Name: "Docker Compose", Status: "fail", Message: "Docker Compose not installed", FixCmd: "sudo pacman -S docker-compose", FixFunc: func() error { - // Try to install docker-compose via pacman (Arch) + cmd := exec.Command("sudo", "pacman", "-S", "--noconfirm", "docker-compose") cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr @@ -240,13 +295,12 @@ func checkOllama() CheckResult { } } - // Ollama not responding - check if Docker container exists dockerCheck := exec.Command("docker", "ps", "-a", "--filter", "name=ollama", "--format", "{{.Names}}") output, _ := dockerCheck.Output() containerExists := strings.TrimSpace(string(output)) != "" if containerExists { - // Container exists but not responding - try starting it + return CheckResult{ Name: "Ollama", Status: "fail", @@ -261,7 +315,6 @@ func checkOllama() CheckResult { } } - // Check if infra/ollama/docker-compose.yml exists projectRoot := getProjectRoot() composeFile := filepath.Join(projectRoot, "infra", "ollama", "docker-compose.yml") if _, err := os.Stat(composeFile); err == nil { @@ -276,7 +329,6 @@ func checkOllama() CheckResult { } } - // Fallback: check if native ollama is installed if _, err := exec.LookPath("ollama"); err == nil { return CheckResult{ Name: "Ollama", @@ -286,14 +338,13 @@ func checkOllama() CheckResult { } } - // Nothing available - suggest Docker setup return CheckResult{ Name: "Ollama", Status: "fail", Message: "Ollama not installed", FixCmd: "cd infra/ollama && docker compose up -d", FixFunc: func() error { - // Create infra/ollama if needed and start + if projectRoot != "" { composeFile := filepath.Join(projectRoot, "infra", "ollama", "docker-compose.yml") if _, err := os.Stat(composeFile); err == nil { @@ -320,7 +371,6 @@ func checkOllamaModel() CheckResult { } defer resp.Body.Close() - // Check if any model is available cmd := exec.Command("sh", "-c", "curl -s http://localhost:11434/api/tags | grep -o '\"name\":\"[^\"]*\"' | head -1") output, _ := cmd.Output() @@ -415,7 +465,7 @@ func checkNetwork() CheckResult { } func getProjectRoot() string { - // Try to find project root by looking for go.mod + dir, _ := os.Getwd() for { if _, err := os.Stat(filepath.Join(dir, "go.mod")); err == nil { @@ -433,15 +483,15 @@ func getProjectRoot() string { // getDockerComposeCmd returns the correct docker compose command // Returns ("docker-compose", []string{}) or ("docker", []string{"compose"}) func getDockerComposeCmd() (string, []string) { - // Try docker compose (plugin) first + if err := exec.Command("docker", "compose", "version").Run(); err == nil { return "docker", []string{"compose"} } - // Fall back to docker-compose (standalone) + if _, err := exec.LookPath("docker-compose"); err == nil { return "docker-compose", []string{} } - // Default to docker compose + return "docker", []string{"compose"} } diff --git a/cmd/explain.go b/cmd/explain.go index 8491a91..9bc87c2 100644 --- a/cmd/explain.go +++ b/cmd/explain.go @@ -2,6 +2,8 @@ package cmd import ( "bufio" + "dev-cli/internal/ai" + "dev-cli/internal/core" "encoding/json" "fmt" "os" @@ -9,10 +11,6 @@ import ( "strings" "time" - "dev-cli/internal/config" - "dev-cli/internal/llm" - "dev-cli/internal/storage" - "github.com/briandowns/spinner" "github.com/spf13/cobra" @@ -59,7 +57,7 @@ Reads from your command history (requires shell integration via 'dev-cli init zs if explainExitCode == 130 { return } - analyzeEntry(storage.LogEntry{ + analyzeEntry(core.LogEntry{ Command: explainCommand, ExitCode: explainExitCode, Output: explainOutput, @@ -81,7 +79,7 @@ func init() { } func analyzeFromLog(limit int, filterStr, sinceStr string, interactive bool) { - db, err := storage.InitDB() + db, err := core.InitDB() if err != nil { fmt.Fprintf(os.Stderr, "⚠️ Failed to open db: %v\n", err) return @@ -101,7 +99,7 @@ func analyzeFromLog(limit int, filterStr, sinceStr string, interactive bool) { limit = 1 } - items, err := storage.GetFailures(db, storage.QueryOpts{ + items, err := core.GetFailures(db, core.QueryOpts{ Limit: limit, Filter: filterStr, Since: sinceDur, @@ -127,7 +125,7 @@ func analyzeFromLog(limit int, filterStr, sinceStr string, interactive bool) { } } - analyzeEntry(storage.LogEntry{ + analyzeEntry(core.LogEntry{ Command: item.Command, ExitCode: item.ExitCode, Output: output, @@ -135,10 +133,10 @@ func analyzeFromLog(limit int, filterStr, sinceStr string, interactive bool) { } } -func analyzeEntry(entry storage.LogEntry, interactive bool) { +func analyzeEntry(entry core.LogEntry, interactive bool) { fmt.Printf("\n\033[31m×\033[0m %s \033[90m(exit %d)\033[0m\n", entry.Command, entry.ExitCode) - if err := llm.EnsureOllamaRunning(); err != nil { + if err := ai.EnsureOllamaRunning(); err != nil { fmt.Fprintf(os.Stderr, "\033[33m⚠\033[0m Ollama not available: %v\n", err) return } @@ -147,7 +145,7 @@ func analyzeEntry(entry storage.LogEntry, interactive bool) { s.Suffix = " 🧠 Analyzing failure..." s.Start() - client := llm.NewClient(config.Load()) + client := ai.NewOllamaClient(core.LoadConfig()) result, err := client.Explain(entry.Command, entry.ExitCode, entry.Output) s.Stop() diff --git a/cmd/export.go b/cmd/export.go index b2ab9d2..8a562c8 100644 --- a/cmd/export.go +++ b/cmd/export.go @@ -68,7 +68,6 @@ func runExport(cmd *cobra.Command, args []string) { os.Exit(1) } - // Format output for OpenCode output := formatForOpenCode(source, logs) if exportSave { diff --git a/cmd/fix.go b/cmd/fix.go index e32c414..2fe8a0a 100644 --- a/cmd/fix.go +++ b/cmd/fix.go @@ -1,7 +1,7 @@ package cmd import ( - "dev-cli/internal/agent" + "dev-cli/internal/ai" "fmt" "github.com/spf13/cobra" @@ -22,7 +22,7 @@ The agent will: dev-cli fix "kubectl can't connect to cluster"`, Args: cobra.MinimumNArgs(1), Run: func(cmd *cobra.Command, args []string) { - ag := agent.New() + ag := ai.NewAgent() err := ag.Resolve(args[0], func(proposal string) bool { fmt.Printf("> Proposal: %s\n", proposal) diff --git a/cmd/mark_resolved.go b/cmd/mark_resolved.go new file mode 100644 index 0000000..833dadd --- /dev/null +++ b/cmd/mark_resolved.go @@ -0,0 +1,84 @@ +package cmd + +import ( + "fmt" + "os" + + "dev-cli/internal/storage" + + "github.com/spf13/cobra" +) + +var ( + resolveID int64 + resolveResolution string +) + +var markResolvedCmd = &cobra.Command{ + Use: "mark-resolved", + Short: "Mark a failed command as resolved", + Hidden: true, + Run: func(cmd *cobra.Command, args []string) { + if resolveID <= 0 { + fmt.Fprintln(os.Stderr, "error: --id is required") + os.Exit(1) + } + + validResolutions := map[string]bool{ + "solution": true, + "unrelated": true, + "skipped": true, + } + if !validResolutions[resolveResolution] { + fmt.Fprintf(os.Stderr, "error: --resolution must be one of: solution, unrelated, skipped\n") + os.Exit(1) + } + + db, err := storage.InitDB() + if err != nil { + fmt.Fprintf(os.Stderr, "error opening db: %v\n", err) + os.Exit(1) + } + defer db.Close() + + if err := storage.MarkResolution(db, resolveID, resolveResolution); err != nil { + fmt.Fprintf(os.Stderr, "error marking resolution: %v\n", err) + os.Exit(1) + } + }, +} + +var checkLastFailureCmd = &cobra.Command{ + Use: "check-last-failure", + Short: "Check if there's an unresolved failure", + Hidden: true, + Run: func(cmd *cobra.Command, args []string) { + db, err := storage.InitDB() + if err != nil { + os.Exit(1) + } + defer db.Close() + + failure, err := storage.GetLastUnresolvedFailure(db) + if err != nil { + os.Exit(1) + } + if failure == nil { + os.Exit(1) + } + + cmdStr := failure.Command + if len(cmdStr) > 50 { + cmdStr = cmdStr[:47] + "..." + } + fmt.Printf("%d|%s\n", failure.ID, cmdStr) + }, +} + +func init() { + rootCmd.AddCommand(markResolvedCmd) + markResolvedCmd.Flags().Int64Var(&resolveID, "id", 0, "History entry ID") + markResolvedCmd.Flags().StringVar(&resolveResolution, "resolution", "", "Resolution type: solution, unrelated, skipped") + + rootCmd.AddCommand(checkLastFailureCmd) +} diff --git a/cmd/watch.go b/cmd/watch.go index 460806e..16ea79a 100644 --- a/cmd/watch.go +++ b/cmd/watch.go @@ -132,7 +132,7 @@ func runWatch(cmd *cobra.Command, args []string) { logContent := strings.Join(buffer, "\n") if watchOpenCode { - // OpenCode handoff mode + fmt.Println("\n\033[33m[!] Error detected! Saving for OpenCode...\033[0m") savePath, err := saveErrorForOpenCode(source, logContent) if err != nil { @@ -142,7 +142,7 @@ func runWatch(cmd *cobra.Command, args []string) { fmt.Printf("\033[36mRun 'opencode' and use: @%s\033[0m\n", savePath) } } else { - // Traditional AI analysis mode + fmt.Println("\n\033[33m[!] Error detected! Analyzing...\033[0m") result, err := client.AnalyzeLog(logContent, watchAI) if err != nil { diff --git a/cmd/workflow.go b/cmd/workflow.go new file mode 100644 index 0000000..a51c90c --- /dev/null +++ b/cmd/workflow.go @@ -0,0 +1,413 @@ +package cmd + +import ( + "context" + "database/sql" + "fmt" + "os" + "os/signal" + "path/filepath" + "strings" + "syscall" + "text/tabwriter" + "time" + + "dev-cli/internal/pipeline" + "dev-cli/internal/storage" + "dev-cli/internal/workflow" + + "github.com/spf13/cobra" +) + +var ( + workflowVerbose bool +) + +var workflowCmd = &cobra.Command{ + Use: "workflow", + Short: "Manage and execute multi-step workflows", + Long: `Execute, resume, and manage multi-step workflow automations. + +Workflows are defined in YAML files and support: + - Sequential step execution + - Conditional branching + - Automatic retry on failure + - Rollback capabilities + - Checkpoint/resume for long operations`, +} + +var workflowRunCmd = &cobra.Command{ + Use: "run ", + Short: "Execute a workflow from a YAML file", + Example: ` dev-cli workflow run deploy.yaml + dev-cli workflow run ~/.devlogs/workflows/cleanup.yaml --verbose`, + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + filePath := args[0] + + wf, err := workflow.ParseFile(filePath) + if err != nil { + return fmt.Errorf("failed to parse workflow: %w", err) + } + + fmt.Printf("🚀 Starting workflow: %s\n", wf.Name) + if wf.Description != "" { + fmt.Printf(" %s\n", wf.Description) + } + fmt.Printf(" Steps: %d\n\n", len(wf.Steps)) + + db, err := storage.InitDB() + if err != nil { + return fmt.Errorf("failed to initialize database: %w", err) + } + defer db.Close() + + store := workflow.NewCheckpointStore(db) + if err := store.InitSchema(); err != nil { + return fmt.Errorf("failed to initialize workflow schema: %w", err) + } + + bus := pipeline.NewEventBus() + engine := workflow.NewEngine(store, bus) + engine.SetVerbose(workflowVerbose) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + go func() { + <-sigCh + fmt.Println("\n⏸ Received interrupt, saving checkpoint...") + cancel() + }() + + result, err := engine.Run(ctx, wf) + if err != nil && result == nil { + return fmt.Errorf("workflow execution failed: %w", err) + } + + fmt.Println() + printRunResult(result) + + return nil + }, +} + +var workflowResumeCmd = &cobra.Command{ + Use: "resume ", + Short: "Resume a paused or failed workflow", + Example: ` dev-cli workflow resume run_1703548800000000000`, + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + runID := args[0] + + db, err := storage.InitDB() + if err != nil { + return fmt.Errorf("failed to initialize database: %w", err) + } + defer db.Close() + + store := workflow.NewCheckpointStore(db) + + state, err := store.LoadRun(runID) + if err != nil { + return fmt.Errorf("failed to load run: %w", err) + } + + workflowFile, err := findWorkflowFile(state.WorkflowID, state.WorkflowName) + if err != nil { + return fmt.Errorf("workflow file not found: %w\n\nPlease provide the workflow file path with: dev-cli workflow resume-file %s ", err, runID) + } + + wf, err := workflow.ParseFile(workflowFile) + if err != nil { + return fmt.Errorf("failed to parse workflow: %w", err) + } + + fmt.Printf("▶ Resuming workflow: %s (run: %s)\n", wf.Name, runID) + fmt.Printf(" Current step: %d/%d\n\n", state.CurrentStepIdx+1, len(wf.Steps)) + + bus := pipeline.NewEventBus() + engine := workflow.NewEngine(store, bus) + engine.SetVerbose(workflowVerbose) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM) + go func() { + <-sigCh + fmt.Println("\n⏸ Received interrupt, saving checkpoint...") + cancel() + }() + + result, err := engine.Resume(ctx, wf, runID) + if err != nil && result == nil { + return fmt.Errorf("resume failed: %w", err) + } + + fmt.Println() + printRunResult(result) + + return nil + }, +} + +var workflowListCmd = &cobra.Command{ + Use: "list", + Short: "List recent workflow runs", + RunE: func(cmd *cobra.Command, args []string) error { + db, err := storage.InitDB() + if err != nil { + return fmt.Errorf("failed to initialize database: %w", err) + } + defer db.Close() + + store := workflow.NewCheckpointStore(db) + if err := store.InitSchema(); err != nil { + return err + } + + runs, err := store.ListRuns(20) + if err != nil { + return fmt.Errorf("failed to list runs: %w", err) + } + + if len(runs) == 0 { + fmt.Println("No workflow runs found.") + return nil + } + + w := tabwriter.NewWriter(os.Stdout, 0, 0, 2, ' ', 0) + fmt.Fprintln(w, "RUN ID\tWORKFLOW\tSTATUS\tSTARTED\tDURATION") + fmt.Fprintln(w, "------\t--------\t------\t-------\t--------") + + for _, run := range runs { + duration := "" + if !run.CompletedAt.IsZero() { + duration = run.CompletedAt.Sub(run.StartedAt).Truncate(time.Second).String() + } else if run.Status == workflow.StatusRunning { + duration = time.Since(run.StartedAt).Truncate(time.Second).String() + " (running)" + } + + fmt.Fprintf(w, "%s\t%s\t%s\t%s\t%s\n", + run.RunID, + run.WorkflowName, + formatStatus(run.Status), + run.StartedAt.Format("2006-01-02 15:04"), + duration, + ) + } + + return w.Flush() + }, +} + +var workflowStatusCmd = &cobra.Command{ + Use: "status ", + Short: "Show detailed status of a workflow run", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + runID := args[0] + + db, err := storage.InitDB() + if err != nil { + return fmt.Errorf("failed to initialize database: %w", err) + } + defer db.Close() + + store := workflow.NewCheckpointStore(db) + state, err := store.LoadRun(runID) + if err != nil { + return fmt.Errorf("failed to load run: %w", err) + } + + fmt.Printf("Workflow: %s\n", state.WorkflowName) + fmt.Printf("Run ID: %s\n", state.RunID) + fmt.Printf("Status: %s\n", formatStatus(state.Status)) + fmt.Printf("Started: %s\n", state.StartedAt.Format(time.RFC3339)) + + if !state.CompletedAt.IsZero() { + fmt.Printf("Finished: %s\n", state.CompletedAt.Format(time.RFC3339)) + fmt.Printf("Duration: %s\n", state.CompletedAt.Sub(state.StartedAt).Truncate(time.Second)) + } + + if state.Error != "" { + fmt.Printf("Error: %s\n", state.Error) + } + + fmt.Printf("\nSteps:\n") + w := tabwriter.NewWriter(os.Stdout, 0, 0, 2, ' ', 0) + fmt.Fprintln(w, " STEP\tSTATUS\tEXIT\tDURATION") + fmt.Fprintln(w, " ----\t------\t----\t--------") + + for stepID, result := range state.StepResults { + fmt.Fprintf(w, " %s\t%s\t%d\t%s\n", + stepID, + formatStepStatus(result.Status), + result.ExitCode, + result.Duration.Truncate(time.Millisecond), + ) + } + + return w.Flush() + }, +} + +var workflowRollbackCmd = &cobra.Command{ + Use: "rollback ", + Short: "Manually trigger rollback for a workflow run", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + runID := args[0] + + db, err := storage.InitDB() + if err != nil { + return fmt.Errorf("failed to initialize database: %w", err) + } + defer db.Close() + + store := workflow.NewCheckpointStore(db) + state, err := store.LoadRun(runID) + if err != nil { + return fmt.Errorf("failed to load run: %w", err) + } + + workflowFile, err := findWorkflowFile(state.WorkflowID, state.WorkflowName) + if err != nil { + return fmt.Errorf("workflow file not found: %w", err) + } + + wf, err := workflow.ParseFile(workflowFile) + if err != nil { + return fmt.Errorf("failed to parse workflow: %w", err) + } + + fmt.Printf("↺ Rolling back workflow: %s\n", wf.Name) + + bus := pipeline.NewEventBus() + engine := workflow.NewEngine(store, bus) + engine.SetVerbose(true) + + ctx := context.Background() + if err := engine.Rollback(ctx, wf, runID); err != nil { + return fmt.Errorf("rollback failed: %w", err) + } + + fmt.Println("\n✓ Rollback completed") + return nil + }, +} + +func init() { + rootCmd.AddCommand(workflowCmd) + + workflowCmd.PersistentFlags().BoolVarP(&workflowVerbose, "verbose", "v", false, "Enable verbose output") + + workflowCmd.AddCommand(workflowRunCmd) + workflowCmd.AddCommand(workflowResumeCmd) + workflowCmd.AddCommand(workflowListCmd) + workflowCmd.AddCommand(workflowStatusCmd) + workflowCmd.AddCommand(workflowRollbackCmd) +} + +func printRunResult(result *workflow.RunResult) { + if result == nil { + return + } + + switch result.Status { + case workflow.StatusCompleted: + fmt.Printf("✓ Workflow completed successfully in %s\n", result.Duration.Truncate(time.Second)) + case workflow.StatusPaused: + fmt.Printf("⏸ Workflow paused. Resume with:\n dev-cli workflow resume %s\n", result.RunID) + case workflow.StatusFailed: + fmt.Printf("✗ Workflow failed: %s\n", result.Error) + fmt.Printf(" Resume with: dev-cli workflow resume %s\n", result.RunID) + case workflow.StatusRolledBack: + fmt.Printf("↺ Workflow rolled back after failure: %s\n", result.Error) + default: + fmt.Printf("? Workflow ended with status: %s\n", result.Status) + } +} + +func formatStatus(status workflow.RunStatus) string { + switch status { + case workflow.StatusCompleted: + return "✓ completed" + case workflow.StatusRunning: + return "▶ running" + case workflow.StatusPaused: + return "⏸ paused" + case workflow.StatusFailed: + return "✗ failed" + case workflow.StatusRolledBack: + return "↺ rolledback" + default: + return string(status) + } +} + +func formatStepStatus(status workflow.StepStatus) string { + switch status { + case workflow.StepSuccess: + return "✓" + case workflow.StepFailed: + return "✗" + case workflow.StepSkipped: + return "⏭" + case workflow.StepRolledBack: + return "↺" + case workflow.StepRunning: + return "▶" + default: + return string(status) + } +} + +func findWorkflowFile(workflowID, workflowName string) (string, error) { + + home, _ := os.UserHomeDir() + searchPaths := []string{ + filepath.Join(home, ".devlogs", "workflows"), + ".", + filepath.Join(home, ".config", "dev-cli", "workflows"), + } + + for _, dir := range searchPaths { + files, err := os.ReadDir(dir) + if err != nil { + continue + } + + for _, f := range files { + if f.IsDir() { + continue + } + + name := f.Name() + if !strings.HasSuffix(name, ".yaml") && !strings.HasSuffix(name, ".yml") { + continue + } + + fullPath := filepath.Join(dir, name) + wf, err := workflow.ParseFile(fullPath) + if err != nil { + continue + } + + if wf.ID == workflowID || wf.Name == workflowName { + return fullPath, nil + } + } + } + + return "", fmt.Errorf("could not find workflow %q", workflowName) +} + +// GetDB returns a database connection (for use by external callers) +func GetDB() (*sql.DB, error) { + return storage.InitDB() +} diff --git a/dev-cli b/dev-cli index 8204442..32fc57d 100755 Binary files a/dev-cli and b/dev-cli differ diff --git a/internal/agent/agent.go b/internal/ai/agent.go similarity index 69% rename from internal/agent/agent.go rename to internal/ai/agent.go index 4a9c29d..b1f43ff 100644 --- a/internal/agent/agent.go +++ b/internal/ai/agent.go @@ -1,4 +1,4 @@ -package agent +package ai import ( "bytes" @@ -10,19 +10,34 @@ import ( "time" "github.com/briandowns/spinner" - - "dev-cli/internal/llm" ) const maxRetries = 3 +type Solver interface { + Solve(goal string) (string, error) +} + +type Executor interface { + Execute(command string) (success bool, errOutput string) +} + type Agent struct { - client *llm.HybridClient + solver Solver + executor Executor } -func New() *Agent { +func NewAgent() *Agent { return &Agent{ - client: llm.NewHybridClient(), + solver: NewHybridClient(), + executor: &shellExecutor{}, + } +} + +func NewAgentWithDeps(solver Solver, executor Executor) *Agent { + return &Agent{ + solver: solver, + executor: executor, } } @@ -44,7 +59,7 @@ func (a *Agent) Resolve(issue string, approval func(string) bool) error { prompt = fmt.Sprintf("Previous command failed with:\n%s\n\nOriginal task: %s\n\nPlease provide a corrected command.", lastError, issue) } - proposal, err := a.client.Solve(prompt) + proposal, err := a.solver.Solve(prompt) s.Stop() if err != nil { @@ -62,14 +77,14 @@ func (a *Agent) Resolve(issue string, approval func(string) bool) error { } fmt.Printf("\n > Running: %s\n", proposal) - success, errOutput := a.executeWithCapture(proposal) + success, errOutput := a.executor.Execute(proposal) if success { fmt.Println(" + Done") return nil } - lastError = truncate(errOutput, 500) + lastError = truncateAgent(errOutput, 500) fmt.Printf(" x Failed. Retrying...\n") } @@ -77,7 +92,9 @@ func (a *Agent) Resolve(issue string, approval func(string) bool) error { return fmt.Errorf("max retries exceeded") } -func (a *Agent) executeWithCapture(command string) (bool, string) { +type shellExecutor struct{} + +func (e *shellExecutor) Execute(command string) (bool, string) { cmd := exec.Command("sh", "-c", command) var stderrBuf bytes.Buffer @@ -89,7 +106,7 @@ func (a *Agent) executeWithCapture(command string) (bool, string) { return err == nil, stderrBuf.String() } -func truncate(s string, maxLen int) string { +func truncateAgent(s string, maxLen int) string { s = strings.TrimSpace(s) if len(s) <= maxLen { return s diff --git a/internal/ai/cache.go b/internal/ai/cache.go new file mode 100644 index 0000000..941360c --- /dev/null +++ b/internal/ai/cache.go @@ -0,0 +1,184 @@ +package ai + +import ( + "sync" + "time" +) + +type CachedAnalysis struct { + Signature string + RootCauseNodes []string + RemediationSteps []string + Explanation string + Fix string + Confidence float64 + HitCount int + LastHit time.Time + CreatedAt time.Time +} + +type ErrorCache struct { + cache map[string]*CachedAnalysis + order []string + maxSize int + mu sync.RWMutex + hits int64 + misses int64 +} + +func NewErrorCache(maxSize int) *ErrorCache { + if maxSize <= 0 { + maxSize = 100 + } + return &ErrorCache{ + cache: make(map[string]*CachedAnalysis), + order: make([]string, 0, maxSize), + maxSize: maxSize, + } +} + +func (c *ErrorCache) Get(signature string) *CachedAnalysis { + c.mu.Lock() + defer c.mu.Unlock() + + analysis, ok := c.cache[signature] + if !ok { + c.misses++ + return nil + } + + c.hits++ + analysis.HitCount++ + analysis.LastHit = time.Now() + c.moveToFront(signature) + + return analysis +} + +func (c *ErrorCache) Put(signature string, analysis *CachedAnalysis) { + c.mu.Lock() + defer c.mu.Unlock() + + if _, exists := c.cache[signature]; exists { + c.cache[signature] = analysis + c.moveToFront(signature) + return + } + + if len(c.cache) >= c.maxSize { + c.evictOldest() + } + + analysis.CreatedAt = time.Now() + analysis.LastHit = time.Now() + analysis.HitCount = 1 + analysis.Signature = signature + c.cache[signature] = analysis + c.order = append([]string{signature}, c.order...) +} + +func (c *ErrorCache) Contains(signature string) bool { + c.mu.RLock() + defer c.mu.RUnlock() + _, ok := c.cache[signature] + return ok +} + +func (c *ErrorCache) Delete(signature string) { + c.mu.Lock() + defer c.mu.Unlock() + + delete(c.cache, signature) + c.removeFromOrder(signature) +} + +func (c *ErrorCache) Clear() { + c.mu.Lock() + defer c.mu.Unlock() + + c.cache = make(map[string]*CachedAnalysis) + c.order = make([]string, 0, c.maxSize) + c.hits = 0 + c.misses = 0 +} + +func (c *ErrorCache) Size() int { + c.mu.RLock() + defer c.mu.RUnlock() + return len(c.cache) +} + +type CacheStats struct { + Hits int64 + Misses int64 + Size int + MaxSize int + HitRate float64 +} + +func (c *ErrorCache) Stats() CacheStats { + c.mu.RLock() + defer c.mu.RUnlock() + + hitRate := float64(0) + total := c.hits + c.misses + if total > 0 { + hitRate = float64(c.hits) / float64(total) + } + + return CacheStats{ + Hits: c.hits, + Misses: c.misses, + Size: len(c.cache), + MaxSize: c.maxSize, + HitRate: hitRate, + } +} + +func (c *ErrorCache) GetTopHits(limit int) []*CachedAnalysis { + c.mu.RLock() + defer c.mu.RUnlock() + + entries := make([]*CachedAnalysis, 0, len(c.cache)) + for _, v := range c.cache { + entries = append(entries, v) + } + + for i := 0; i < len(entries)-1; i++ { + for j := i + 1; j < len(entries); j++ { + if entries[j].HitCount > entries[i].HitCount { + entries[i], entries[j] = entries[j], entries[i] + } + } + } + + if limit > len(entries) { + limit = len(entries) + } + + return entries[:limit] +} + +func (c *ErrorCache) moveToFront(signature string) { + c.removeFromOrder(signature) + c.order = append([]string{signature}, c.order...) +} + +func (c *ErrorCache) removeFromOrder(signature string) { + for i, s := range c.order { + if s == signature { + c.order = append(c.order[:i], c.order[i+1:]...) + break + } + } +} + +func (c *ErrorCache) evictOldest() { + if len(c.order) == 0 { + return + } + + oldest := c.order[len(c.order)-1] + delete(c.cache, oldest) + c.order = c.order[:len(c.order)-1] +} diff --git a/internal/ai/client.go b/internal/ai/client.go new file mode 100644 index 0000000..2bcbe41 --- /dev/null +++ b/internal/ai/client.go @@ -0,0 +1,823 @@ +package ai + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net/http" + "os" + "os/exec" + "strings" + "sync" + "time" + + "dev-cli/internal/core" +) + +const ( + DefaultOllamaURL = "http://localhost:11434" + DefaultModel = "qwen2.5-coder:3b-instruct" + FallbackModel = "qwen2.5-coder:3b-instruct-q8_0" + RequestTimeout = 30 * time.Second + PerplexityAPIURL = "https://api.perplexity.ai/chat/completions" +) + +type ExplainResult struct { + Explanation string `json:"explanation"` + Fix string `json:"fix"` +} + +type Step struct { + Type string `json:"type"` + Content string `json:"content"` + File string `json:"file,omitempty"` + Note string `json:"note,omitempty"` +} + +type Solution struct { + ID int `json:"id"` + Title string `json:"title"` + Description string `json:"description"` + Steps []Step `json:"steps"` + Source string `json:"source,omitempty"` +} + +type ResearchResult struct { + Query string `json:"query"` + Solutions []Solution `json:"solutions"` +} + +type LogAnalysisResult struct { + Explanation string `json:"explanation"` + Fix string `json:"fix"` +} + +type ToolCallResult struct { + ToolName string `json:"tool_name"` + Parameters map[string]any `json:"parameters"` + Reasoning string `json:"reasoning,omitempty"` +} + +type generateRequest struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + Stream bool `json:"stream"` + Format string `json:"format,omitempty"` + KeepAlive string `json:"keep_alive,omitempty"` +} + +type generateResponse struct { + Response string `json:"response"` + Done bool `json:"done"` +} + +type OllamaClient struct { + baseURL string + model string + httpClient *http.Client +} + +func NewOllamaClient(cfg *core.Config) *OllamaClient { + baseURL := DefaultOllamaURL + if cfg.OllamaURL != "" { + baseURL = cfg.OllamaURL + } + + model := DefaultModel + if cfg.OllamaModel != "" { + model = cfg.OllamaModel + } + + return &OllamaClient{ + baseURL: baseURL, + model: model, + httpClient: &http.Client{ + Timeout: RequestTimeout, + }, + } +} + +func EnsureOllamaRunning() error { + client := &http.Client{Timeout: 2 * time.Second} + resp, err := client.Get(DefaultOllamaURL + "/api/tags") + if err == nil { + resp.Body.Close() + return nil + } + + fmt.Println("\033[33m⚡ Ollama not running, starting...\033[0m") + + startCmd := exec.Command("docker", "start", "ollama") + if err := startCmd.Run(); err == nil { + return waitForOllama(client, 30*time.Second) + } + + fmt.Println("\033[90m Creating Ollama container...\033[0m") + createCmd := exec.Command("docker", "run", "-d", + "--name", "ollama", + "-p", "11434:11434", + "-v", "ollama:/root/.ollama", + "--restart", "unless-stopped", + "ollama/ollama") + + output, err := createCmd.CombinedOutput() + if err != nil { + return fmt.Errorf("failed to create Ollama container: %w\n%s", err, string(output)) + } + + return waitForOllama(client, 60*time.Second) +} + +func waitForOllama(client *http.Client, timeout time.Duration) error { + start := time.Now() + for { + if time.Since(start) > timeout { + return fmt.Errorf("timeout waiting for Ollama to start") + } + + resp, err := client.Get(DefaultOllamaURL + "/api/tags") + if err == nil { + resp.Body.Close() + fmt.Println("\033[32m✓ Ollama is ready\033[0m") + return nil + } + + time.Sleep(500 * time.Millisecond) + } +} + +func (c *OllamaClient) Explain(cmd string, exitCode int, output string) (*ExplainResult, error) { + if len(output) > 2000 { + output = output[len(output)-2000:] + } + + prompt := fmt.Sprintf(`You are a CLI error analyzer. Analyze this failed command and respond with JSON only. + +RULES: +1. "explanation" = Brief 1-sentence error cause can attend for more precision only if needed. +2. "fix" = EXACT shell command to run (NOT advice, NOT instructions - just the command) + - Good fix: "npm init -y new line and more command if needed to run in sequence" + - Bad fix: "Make sure package.json exists" + - If no fix possible, refer to sources more authentic to that problem to precise documentation etc "" + +EXAMPLES: +- package.json missing → {"explanation": "Missing package.json", "fix": "npm init -y"} +- permission denied → {"explanation": "Permission denied", "fix": "sudo !!"} +- command not found → {"explanation": "Command not installed", "fix": ""} + +Command: %s +Exit Code: %d +Output: %s + +JSON response:`, cmd, exitCode, output) + + return c.generateExplain(prompt) +} + +func (c *OllamaClient) generateExplain(prompt string) (*ExplainResult, error) { + req := generateRequest{ + Model: c.model, + Prompt: prompt, + Stream: false, + Format: "json", + } + + if os.Getenv("DEV_CLI_OLLAMA_UNLOAD") == "true" { + req.KeepAlive = "0m" + } + + reqBody, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + resp, err := c.httpClient.Post(c.baseURL+"/api/generate", "application/json", bytes.NewReader(reqBody)) + if err != nil { + return nil, fmt.Errorf("call Ollama: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("ollama status %d: %s", resp.StatusCode, string(body)) + } + + var genResp generateResponse + if err := json.NewDecoder(resp.Body).Decode(&genResp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + var result ExplainResult + responseText := strings.TrimSpace(genResp.Response) + if err := json.Unmarshal([]byte(responseText), &result); err != nil { + return &ExplainResult{Explanation: responseText, Fix: ""}, nil + } + + return &result, nil +} + +func (c *OllamaClient) Research(query string) (*ResearchResult, error) { + prompt := fmt.Sprintf(`You are a Senior Developer Assistant. The user needs to: "%s". +Provide the TOP 3 distinct ways to achieve this. + +RULES: +1. Option 1 = "Best Practice" / Modern way +2. Option 2 = "Quickest/Easiest" way +3. Option 3 = "Alternative" (edge case or manual approach) +4. Each solution can have multiple steps +5. Step type is "command" for shell commands, "file" for code snippets +6. For "file" type, include the target filename in "file" field + +OUTPUT JSON ONLY: +{ + "solutions": [ + { + "id": 1, + "title": "Using Docker (Recommended)", + "description": "Isolated environment", + "steps": [ + {"type": "command", "content": "docker run -d postgres", "note": "Start container"} + ], + "source": "" + } + ] +}`, query) + + req := generateRequest{ + Model: c.model, + Prompt: prompt, + Stream: false, + Format: "json", + } + + if os.Getenv("DEV_CLI_OLLAMA_UNLOAD") == "true" { + req.KeepAlive = "0m" + } + + reqBody, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + resp, err := c.httpClient.Post(c.baseURL+"/api/generate", "application/json", bytes.NewReader(reqBody)) + if err != nil { + return nil, fmt.Errorf("call Ollama: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("ollama status %d: %s", resp.StatusCode, string(body)) + } + + var genResp generateResponse + if err := json.NewDecoder(resp.Body).Decode(&genResp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + responseText := strings.TrimSpace(genResp.Response) + + var result ResearchResult + if err := json.Unmarshal([]byte(responseText), &result); err != nil { + return nil, fmt.Errorf("parse solutions: %w", err) + } + + result.Query = query + return &result, nil +} + +func (c *OllamaClient) AnalyzeLog(logLines string) (*LogAnalysisResult, error) { + prompt := fmt.Sprintf(`You are a Log Analyzer. Identify the error in these log lines. + +OUTPUT JSON ONLY: +{ + "explanation": "Brief description of the error (1 sentence)", + "fix": "Suggested command or action to fix it (or empty if unknown)" +} + +LOGS: +%s`, logLines) + + req := generateRequest{ + Model: c.model, + Prompt: prompt, + Stream: false, + Format: "json", + } + + if os.Getenv("DEV_CLI_OLLAMA_UNLOAD") == "true" { + req.KeepAlive = "0m" + } + + reqBody, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + resp, err := c.httpClient.Post(c.baseURL+"/api/generate", "application/json", bytes.NewReader(reqBody)) + if err != nil { + return nil, fmt.Errorf("call Ollama: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("ollama status %d: %s", resp.StatusCode, string(body)) + } + + var genResp generateResponse + if err := json.NewDecoder(resp.Body).Decode(&genResp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + var result LogAnalysisResult + responseText := strings.TrimSpace(genResp.Response) + if err := json.Unmarshal([]byte(responseText), &result); err != nil { + return &LogAnalysisResult{Explanation: responseText}, nil + } + + return &result, nil +} + +func (c *OllamaClient) Solve(goal string) (string, error) { + prompt := fmt.Sprintf(`You are an Autonomous CLI Agent. The user wants to: "%s". +Provide a SINGLE shell command to achieve this. + +RULES: +1. Output ONLY the command. No markdown, no explanations. +2. If multiple steps are needed, chain them with && or ; +3. Assume a standard Linux environment. +4. BE SAFE. Do not return commands that delete data without confirmation unless explicitly asked. + +GOAL: %s +COMMAND:`, goal, goal) + + req := generateRequest{ + Model: c.model, + Prompt: prompt, + Stream: false, + } + + if os.Getenv("DEV_CLI_OLLAMA_UNLOAD") == "true" { + req.KeepAlive = "0m" + } + + reqBody, err := json.Marshal(req) + if err != nil { + return "", fmt.Errorf("marshal request: %w", err) + } + + resp, err := c.httpClient.Post(c.baseURL+"/api/generate", "application/json", bytes.NewReader(reqBody)) + if err != nil { + return "", fmt.Errorf("call Ollama: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return "", fmt.Errorf("ollama status %d: %s", resp.StatusCode, string(body)) + } + + var genResp generateResponse + if err := json.NewDecoder(resp.Body).Decode(&genResp); err != nil { + return "", fmt.Errorf("decode response: %w", err) + } + + return strings.TrimSpace(genResp.Response), nil +} + +func (c *OllamaClient) GenerateWithTools(prompt string, toolSchemas string) (*ToolCallResult, error) { + systemPrompt := fmt.Sprintf(`You are an AI assistant with access to tools. Based on the user's request, determine which tool to use and with what parameters. + +AVAILABLE TOOLS: +%s + +RULES: +1. Analyze the user's request carefully +2. Select the most appropriate tool +3. Determine the correct parameters +4. Respond with ONLY valid JSON in this exact format: +{ + "tool_name": "name_of_tool", + "parameters": { + "param1": "value1", + "param2": "value2" + }, + "reasoning": "brief explanation of why this tool was chosen" +} + +Do NOT include any text outside the JSON object.`, toolSchemas) + + fullPrompt := fmt.Sprintf(`%s + +USER REQUEST: %s + +JSON RESPONSE:`, systemPrompt, prompt) + + req := generateRequest{ + Model: c.model, + Prompt: fullPrompt, + Stream: false, + Format: "json", + } + + if os.Getenv("DEV_CLI_OLLAMA_UNLOAD") == "true" { + req.KeepAlive = "0m" + } + + reqBody, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + resp, err := c.httpClient.Post(c.baseURL+"/api/generate", "application/json", bytes.NewReader(reqBody)) + if err != nil { + return nil, fmt.Errorf("call Ollama: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("ollama status %d: %s", resp.StatusCode, string(body)) + } + + var genResp generateResponse + if err := json.NewDecoder(resp.Body).Decode(&genResp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + var result ToolCallResult + responseText := strings.TrimSpace(genResp.Response) + if err := json.Unmarshal([]byte(responseText), &result); err != nil { + return nil, fmt.Errorf("parse tool call: %w (response: %s)", err, responseText) + } + + return &result, nil +} + +type perplexityMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +type perplexityRequest struct { + Model string `json:"model"` + Messages []perplexityMessage `json:"messages"` +} + +type perplexityChoice struct { + Message struct { + Content string `json:"content"` + } `json:"message"` +} + +type perplexityResponse struct { + Choices []perplexityChoice `json:"choices"` +} + +type PerplexityClient struct { + apiKey string + model string + httpClient *http.Client +} + +func NewPerplexityClient(cfg *core.Config) *PerplexityClient { + if cfg.PerplexityKey == "" { + return nil + } + + return &PerplexityClient{ + apiKey: cfg.PerplexityKey, + model: cfg.PerplexityModel, + httpClient: &http.Client{ + Timeout: 60 * time.Second, + }, + } +} + +func (c *PerplexityClient) Research(ctx context.Context, query string) (*ResearchResult, error) { + prompt := fmt.Sprintf(`You are a Senior Developer Assistant. The user needs to: "%s". +Provide the TOP 3 distinct ways to achieve this. + +OUTPUT JSON ONLY (No markdown, no code fences): +{ + "solutions": [ + { + "id": 1, + "title": "Using npm (Recommended)", + "description": "Modern package manager with better caching", + "steps": [ + {"type": "command", "content": "npm install tailwindcss", "note": "Install package"} + ], + "source": "https://tailwindcss.com/docs" + } + ] +}`, query) + + reqBody, err := json.Marshal(perplexityRequest{ + Model: c.model, + Messages: []perplexityMessage{ + {Role: "system", Content: "You are a helpful developer assistant. Always respond with valid JSON only, no markdown formatting."}, + {Role: "user", Content: prompt}, + }, + }) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + req, err := http.NewRequestWithContext(ctx, "POST", PerplexityAPIURL, bytes.NewReader(reqBody)) + if err != nil { + return nil, fmt.Errorf("create request: %w", err) + } + + req.Header.Set("Authorization", "Bearer "+c.apiKey) + req.Header.Set("Content-Type", "application/json") + + resp, err := c.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("call Perplexity: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("perplexity status %d: %s", resp.StatusCode, string(body)) + } + + var pResp perplexityResponse + if err := json.NewDecoder(resp.Body).Decode(&pResp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + if len(pResp.Choices) == 0 { + return nil, fmt.Errorf("no response from Perplexity") + } + + content := strings.TrimSpace(pResp.Choices[0].Message.Content) + content = stripMarkdownFences(content) + + var result ResearchResult + if err := json.Unmarshal([]byte(content), &result); err != nil { + return nil, fmt.Errorf("parse solutions: %w", err) + } + + result.Query = query + return &result, nil +} + +func (c *PerplexityClient) AnalyzeLog(ctx context.Context, logLines string) (*LogAnalysisResult, error) { + prompt := fmt.Sprintf(`You are a Log Analyzer. Identify the error in these log lines. + +OUTPUT JSON ONLY (No markdown): +{ + "explanation": "Brief description of the error (1 sentence)", + "fix": "Suggested command or action to fix it (or empty if unknown)" +} + +LOGS: +%s`, logLines) + + reqBody, err := json.Marshal(perplexityRequest{ + Model: c.model, + Messages: []perplexityMessage{ + {Role: "system", Content: "You are a helpful developer assistant. Always respond with valid JSON only, no markdown formatting."}, + {Role: "user", Content: prompt}, + }, + }) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + req, err := http.NewRequestWithContext(ctx, "POST", PerplexityAPIURL, bytes.NewReader(reqBody)) + if err != nil { + return nil, fmt.Errorf("create request: %w", err) + } + + req.Header.Set("Authorization", "Bearer "+c.apiKey) + req.Header.Set("Content-Type", "application/json") + + resp, err := c.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("call Perplexity: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("perplexity status %d: %s", resp.StatusCode, string(body)) + } + + var pResp perplexityResponse + if err := json.NewDecoder(resp.Body).Decode(&pResp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + if len(pResp.Choices) == 0 { + return nil, fmt.Errorf("no response from Perplexity") + } + + content := strings.TrimSpace(pResp.Choices[0].Message.Content) + content = stripMarkdownFences(content) + + var result LogAnalysisResult + if err := json.Unmarshal([]byte(content), &result); err != nil { + return &LogAnalysisResult{Explanation: content}, nil + } + + return &result, nil +} + +func stripMarkdownFences(s string) string { + s = strings.TrimSpace(s) + if strings.HasPrefix(s, "```json") { + s = strings.TrimPrefix(s, "```json") + } else if strings.HasPrefix(s, "```") { + s = strings.TrimPrefix(s, "```") + } + s = strings.TrimSuffix(s, "```") + return strings.TrimSpace(s) +} + +var webKeywords = []string{ + "install", "latest", "version", "how to", "compare", + "why", "best", "setup", "configure", "deploy", "update", "upgrade", +} + +type cacheEntry struct { + result *ResearchResult + timestamp time.Time +} + +type ResponseCache struct { + mu sync.RWMutex + entries map[string]cacheEntry + keys []string + maxSize int + ttl time.Duration +} + +func NewResponseCache(maxSize int, ttl time.Duration) *ResponseCache { + return &ResponseCache{ + entries: make(map[string]cacheEntry), + keys: make([]string, 0), + maxSize: maxSize, + ttl: ttl, + } +} + +func hashQuery(query string) string { + h := sha256.New() + h.Write([]byte(strings.ToLower(strings.TrimSpace(query)))) + return hex.EncodeToString(h.Sum(nil))[:16] +} + +func (c *ResponseCache) Get(query string) (*ResearchResult, bool) { + c.mu.RLock() + defer c.mu.RUnlock() + + key := hashQuery(query) + entry, ok := c.entries[key] + if !ok { + return nil, false + } + + if time.Since(entry.timestamp) > c.ttl { + return nil, false + } + + return entry.result, true +} + +func (c *ResponseCache) Set(query string, result *ResearchResult) { + c.mu.Lock() + defer c.mu.Unlock() + + key := hashQuery(query) + + if _, exists := c.entries[key]; exists { + c.entries[key] = cacheEntry{result: result, timestamp: time.Now()} + c.moveToEnd(key) + return + } + + if len(c.keys) >= c.maxSize { + oldest := c.keys[0] + delete(c.entries, oldest) + c.keys = c.keys[1:] + } + + c.entries[key] = cacheEntry{result: result, timestamp: time.Now()} + c.keys = append(c.keys, key) +} + +func (c *ResponseCache) moveToEnd(key string) { + for i, k := range c.keys { + if k == key { + c.keys = append(c.keys[:i], c.keys[i+1:]...) + c.keys = append(c.keys, key) + return + } + } +} + +func (c *ResponseCache) Stats() (size int, capacity int) { + c.mu.RLock() + defer c.mu.RUnlock() + return len(c.entries), c.maxSize +} + +func (c *ResponseCache) Clear() { + c.mu.Lock() + defer c.mu.Unlock() + c.entries = make(map[string]cacheEntry) + c.keys = make([]string, 0) +} + +type HybridClient struct { + perplexity *PerplexityClient + ollama *OllamaClient + cache *ResponseCache +} + +var defaultCache = NewResponseCache(50, 10*time.Minute) + +func NewHybridClient() *HybridClient { + cfg := core.LoadConfig() + return &HybridClient{ + perplexity: NewPerplexityClient(cfg), + ollama: NewOllamaClient(cfg), + cache: defaultCache, + } +} + +func (h *HybridClient) Research(query string) (*ResearchResult, error) { + if cached, ok := h.cache.Get(query); ok { + return cached, nil + } + + var result *ResearchResult + var err error + + if h.perplexity != nil && needsWebSearch(query) { + result, err = h.perplexity.Research(context.Background(), query) + if err == nil { + h.cache.Set(query, result) + return result, nil + } + } + + result, err = h.ollama.Research(query) + if err == nil { + h.cache.Set(query, result) + } + return result, err +} + +func (h *HybridClient) HasPerplexity() bool { + return h.perplexity != nil +} + +func (h *HybridClient) CacheStats() (size int, capacity int) { + return h.cache.Stats() +} + +func (h *HybridClient) ClearCache() { + h.cache.Clear() +} + +func (h *HybridClient) AnalyzeLog(logLines string, aiMode string) (*LogAnalysisResult, error) { + if os.Getenv("DEV_CLI_FORCE_LOCAL") != "" || aiMode == "local" { + return h.ollama.AnalyzeLog(logLines) + } + + if aiMode == "cloud" { + if h.perplexity != nil { + return h.perplexity.AnalyzeLog(context.Background(), logLines) + } + return nil, fmt.Errorf("cloud AI requested but PERPLEXITY_API_KEY is not set") + } + + return h.ollama.AnalyzeLog(logLines) +} + +func (h *HybridClient) Solve(goal string) (string, error) { + return h.ollama.Solve(goal) +} + +func needsWebSearch(query string) bool { + if os.Getenv("DEV_CLI_FORCE_LOCAL") != "" { + return false + } + + lower := strings.ToLower(query) + for _, kw := range webKeywords { + if strings.Contains(lower, kw) { + return true + } + } + return false +} diff --git a/internal/ai/sanitizer.go b/internal/ai/sanitizer.go new file mode 100644 index 0000000..9bb2cc5 --- /dev/null +++ b/internal/ai/sanitizer.go @@ -0,0 +1,151 @@ +package ai + +import ( + "regexp" + "strings" +) + +type Sanitizer struct { + patterns []*secretPattern +} + +type secretPattern struct { + regex *regexp.Regexp + replacement string + name string +} + +func DefaultSanitizer() *Sanitizer { + return &Sanitizer{ + patterns: []*secretPattern{ + { + regex: regexp.MustCompile(`(?i)(api[_-]?key|apikey)[=:]["']?([a-zA-Z0-9_\-]{20,})["']?`), + replacement: `$1=[REDACTED_API_KEY]`, + name: "API Key", + }, + { + regex: regexp.MustCompile(`(?i)bearer\s+([a-zA-Z0-9_\-\.]+)`), + replacement: `Bearer [REDACTED_TOKEN]`, + name: "Bearer Token", + }, + { + regex: regexp.MustCompile(`(?i)(AKIA|ABIA|ACCA|ASIA)[A-Z0-9]{16}`), + replacement: `[REDACTED_AWS_KEY]`, + name: "AWS Access Key", + }, + { + regex: regexp.MustCompile(`(?i)(aws[_-]?secret[_-]?access[_-]?key)[=:]["']?([a-zA-Z0-9/+=]{40})["']?`), + replacement: `$1=[REDACTED_AWS_SECRET]`, + name: "AWS Secret Key", + }, + { + regex: regexp.MustCompile(`-----BEGIN\s+(RSA\s+)?PRIVATE KEY-----[\s\S]*?-----END\s+(RSA\s+)?PRIVATE KEY-----`), + replacement: `[REDACTED_PRIVATE_KEY]`, + name: "Private Key", + }, + { + regex: regexp.MustCompile(`ghp_[a-zA-Z0-9]{36}`), + replacement: `[REDACTED_GITHUB_TOKEN]`, + name: "GitHub PAT", + }, + { + regex: regexp.MustCompile(`(?i)(password|passwd|pwd|secret|token)[=:]["']?([^\s"']{8,})["']?`), + replacement: `$1=[REDACTED]`, + name: "Password/Secret", + }, + { + regex: regexp.MustCompile(`(?i)(mongodb|postgresql|mysql|redis)://[^:]+:([^@]+)@`), + replacement: `$1://[user]:[REDACTED]@`, + name: "Database Password", + }, + { + regex: regexp.MustCompile(`eyJ[a-zA-Z0-9_-]*\.eyJ[a-zA-Z0-9_-]*\.[a-zA-Z0-9_-]*`), + replacement: `[REDACTED_JWT]`, + name: "JWT Token", + }, + { + regex: regexp.MustCompile(`xox[baprs]-[a-zA-Z0-9-]+`), + replacement: `[REDACTED_SLACK_TOKEN]`, + name: "Slack Token", + }, + }, + } +} + +func (s *Sanitizer) Sanitize(input string) string { + result := input + for _, pattern := range s.patterns { + result = pattern.regex.ReplaceAllString(result, pattern.replacement) + } + return result +} + +func (s *Sanitizer) SanitizeWithReport(input string) (sanitized string, found []string) { + result := input + for _, pattern := range s.patterns { + if pattern.regex.MatchString(result) { + found = append(found, pattern.name) + result = pattern.regex.ReplaceAllString(result, pattern.replacement) + } + } + return result, found +} + +func (s *Sanitizer) ContainsSecrets(input string) bool { + for _, pattern := range s.patterns { + if pattern.regex.MatchString(input) { + return true + } + } + return false +} + +var globalSanitizer = DefaultSanitizer() + +func SanitizeForLLM(input string) string { + return globalSanitizer.Sanitize(input) +} + +func MaskEnvVars(input string) string { + sensitiveVars := regexp.MustCompile(`(?i)(export\s+)?(API_KEY|SECRET|PASSWORD|TOKEN|PRIVATE_KEY|AWS_SECRET)[=]["']?([^\s"'\n]+)["']?`) + return sensitiveVars.ReplaceAllString(input, `$1$2=[REDACTED]`) +} + +func SanitizeOutput(output string) string { + result := globalSanitizer.Sanitize(output) + result = MaskEnvVars(result) + return result +} + +func (s *Sanitizer) AddPattern(name, pattern, replacement string) error { + re, err := regexp.Compile(pattern) + if err != nil { + return err + } + s.patterns = append(s.patterns, &secretPattern{ + regex: re, + replacement: replacement, + name: name, + }) + return nil +} + +func (s *Sanitizer) PatternCount() int { + return len(s.patterns) +} + +func TruncateForLLM(input string, maxLen int) string { + if len(input) <= maxLen { + return input + } + half := (maxLen - 20) / 2 + return input[:half] + "\n...[truncated]...\n" + input[len(input)-half:] +} + +func PrepareForLLM(input string, maxLen int) string { + sanitized := SanitizeForLLM(input) + if maxLen > 0 { + sanitized = TruncateForLLM(sanitized, maxLen) + } + return strings.TrimSpace(sanitized) +} diff --git a/internal/textutil/textutil.go b/internal/core/config.go similarity index 52% rename from internal/textutil/textutil.go rename to internal/core/config.go index a231acf..4a0e093 100644 --- a/internal/textutil/textutil.go +++ b/internal/core/config.go @@ -1,12 +1,64 @@ -package textutil +package core import ( + "os" + "path/filepath" "strings" "github.com/charmbracelet/x/ansi" "github.com/muesli/reflow/wordwrap" ) +type Config struct { + OllamaURL string + OllamaModel string + PerplexityKey string + PerplexityModel string + ForceLocalLLM bool + LogDir string +} + +func LoadConfig() *Config { + cfg := &Config{ + OllamaURL: "http://localhost:11434", + OllamaModel: "qwen2.5-coder:3b-instruct", + PerplexityModel: "sonar-pro", + ForceLocalLLM: false, + } + + if val := os.Getenv("DEV_CLI_OLLAMA_URL"); val != "" { + cfg.OllamaURL = val + } + if val := os.Getenv("DEV_CLI_OLLAMA_MODEL"); val != "" { + cfg.OllamaModel = val + } + if val := os.Getenv("DEV_CLI_PERPLEXITY_KEY"); val != "" { + cfg.PerplexityKey = val + } else if val := os.Getenv("PERPLEXITY_API_KEY"); val != "" { + cfg.PerplexityKey = val + } + if val := os.Getenv("DEV_CLI_PERPLEXITY_MODEL"); val != "" { + cfg.PerplexityModel = val + } + if os.Getenv("DEV_CLI_FORCE_LOCAL") != "" { + cfg.ForceLocalLLM = true + } + if val := os.Getenv("DEV_CLI_LOG_DIR"); val != "" { + cfg.LogDir = val + } else { + home, _ := os.UserHomeDir() + cfg.LogDir = filepath.Join(home, ".devlogs") + } + + return cfg +} + +func (c *Config) IsWebSearchEnabled() bool { + return !c.ForceLocalLLM && c.PerplexityKey != "" +} + +var CurrentConfig = LoadConfig() + func CutLine(line string, start, end int) string { if start >= end { return "" @@ -76,3 +128,10 @@ func MaxLineWidth(lines []string) int { } return maxWidth } + +func TruncateOutput(output string, maxLen int) string { + if len(output) <= maxLen { + return output + } + return output[:maxLen-20] + "\n...[truncated]..." +} diff --git a/internal/core/db.go b/internal/core/db.go new file mode 100644 index 0000000..5942245 --- /dev/null +++ b/internal/core/db.go @@ -0,0 +1,354 @@ +package core + +import ( + "database/sql" + "encoding/json" + "fmt" + "os" + "path/filepath" + "strings" + "time" + + _ "modernc.org/sqlite" +) + +func InitDB() (*sql.DB, error) { + var dbPath string + if envDir := os.Getenv("DEV_CLI_LOG_DIR"); envDir != "" { + if err := os.MkdirAll(envDir, 0755); err != nil { + return nil, fmt.Errorf("create log dir: %w", err) + } + dbPath = filepath.Join(envDir, "history.db") + } else { + home, err := os.UserHomeDir() + if err != nil { + return nil, fmt.Errorf("get user home dir: %w", err) + } + dir := filepath.Join(home, ".devlogs") + if err := os.MkdirAll(dir, 0755); err != nil { + return nil, fmt.Errorf("create data dir: %w", err) + } + dbPath = filepath.Join(dir, "history.db") + } + return OpenDB(dbPath) +} + +func OpenDB(path string) (*sql.DB, error) { + db, err := sql.Open("sqlite", path) + if err != nil { + return nil, fmt.Errorf("open db: %w", err) + } + + if err := db.Ping(); err != nil { + return nil, fmt.Errorf("ping db: %w", err) + } + + if err := migrate(db); err != nil { + db.Close() + return nil, fmt.Errorf("migrate: %w", err) + } + + return db, nil +} + +func migrate(db *sql.DB) error { + schema := ` + CREATE TABLE IF NOT EXISTS history ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + timestamp INTEGER NOT NULL, + command TEXT NOT NULL, + exit_code INTEGER, + duration_ms INTEGER, + directory TEXT, + session_id TEXT, + details TEXT, + resolution TEXT + ); + + CREATE INDEX IF NOT EXISTS idx_history_timestamp ON history(timestamp); + CREATE INDEX IF NOT EXISTS idx_history_exit_code ON history(exit_code); + CREATE INDEX IF NOT EXISTS idx_history_session ON history(session_id); + + CREATE TABLE IF NOT EXISTS workflow_runs ( + id TEXT PRIMARY KEY, + workflow_id TEXT NOT NULL, + workflow_name TEXT, + status TEXT NOT NULL, + current_step INTEGER DEFAULT 0, + started_at DATETIME, + updated_at DATETIME, + completed_at DATETIME, + error TEXT + ); + + CREATE TABLE IF NOT EXISTS workflow_step_results ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + run_id TEXT NOT NULL, + step_id TEXT NOT NULL, + status TEXT NOT NULL, + exit_code INTEGER, + output TEXT, + error TEXT, + retries INTEGER DEFAULT 0, + started_at DATETIME, + completed_at DATETIME, + duration_ms INTEGER, + FOREIGN KEY (run_id) REFERENCES workflow_runs(id) + ); + + CREATE INDEX IF NOT EXISTS idx_workflow_runs_status ON workflow_runs(status); + CREATE INDEX IF NOT EXISTS idx_step_results_run_id ON workflow_step_results(run_id); + + CREATE TABLE IF NOT EXISTS root_causes ( + id TEXT PRIMARY KEY, + error_signature TEXT NOT NULL, + timestamp INTEGER NOT NULL, + root_cause_nodes TEXT, + remediation_steps TEXT, + confidence REAL DEFAULT 0.0, + history_item_id INTEGER, + FOREIGN KEY (history_item_id) REFERENCES history(id) + ); + + CREATE TABLE IF NOT EXISTS runbooks ( + id TEXT PRIMARY KEY, + project_id TEXT, + name TEXT NOT NULL, + description TEXT, + steps TEXT NOT NULL, + success_rate REAL DEFAULT 0.0, + last_used INTEGER, + usage_count INTEGER DEFAULT 0, + tags TEXT + ); + + CREATE TABLE IF NOT EXISTS project_fingerprints ( + id TEXT PRIMARY KEY, + project_type TEXT NOT NULL, + package_manager TEXT, + common_issues TEXT, + associated_runbooks TEXT, + detected_at TEXT NOT NULL, + detected_time INTEGER + ); + + CREATE INDEX IF NOT EXISTS idx_root_cause_signature ON root_causes(error_signature); + CREATE INDEX IF NOT EXISTS idx_root_cause_history ON root_causes(history_item_id); + CREATE INDEX IF NOT EXISTS idx_runbook_project ON runbooks(project_id); + CREATE INDEX IF NOT EXISTS idx_fingerprint_type ON project_fingerprints(project_type); + CREATE INDEX IF NOT EXISTS idx_fingerprint_path ON project_fingerprints(detected_at); + ` + + _, err := db.Exec(schema) + if err != nil { + return err + } + + _, _ = db.Exec("ALTER TABLE history ADD COLUMN resolution TEXT") + + return nil +} + +type LogEntry struct { + Command string `json:"command"` + ExitCode int `json:"exit_code"` + Output string `json:"output,omitempty"` + Cwd string `json:"cwd"` + DurationMs int64 `json:"duration_ms"` + Timestamp string `json:"timestamp"` + SessionID string `json:"session_id,omitempty"` + Details string `json:"details,omitempty"` +} + +type HistoryItem struct { + ID int64 + Timestamp time.Time + Command string + ExitCode int + DurationMs int64 + Directory string + SessionID string + Details string + Resolution string +} + +func SaveCommand(db *sql.DB, entry LogEntry) error { + ts, err := time.Parse(time.RFC3339, entry.Timestamp) + if err != nil { + ts = time.Now() + } + + detailsMap := map[string]interface{}{ + "output": entry.Output, + } + + detailsJSON, err := json.Marshal(detailsMap) + if err != nil { + return fmt.Errorf("marshal details: %w", err) + } + + query := `INSERT INTO history (timestamp, command, exit_code, duration_ms, directory, session_id, details) + VALUES (?, ?, ?, ?, ?, ?, ?)` + + _, err = db.Exec(query, ts.Unix(), entry.Command, entry.ExitCode, entry.DurationMs, entry.Cwd, entry.SessionID, string(detailsJSON)) + return err +} + +func GetRecentHistory(db *sql.DB, limit int) ([]HistoryItem, error) { + query := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details, COALESCE(resolution, '') + FROM history ORDER BY id DESC LIMIT ?` + + rows, err := db.Query(query, limit) + if err != nil { + return nil, err + } + defer rows.Close() + + var items []HistoryItem + for rows.Next() { + var item HistoryItem + var ts int64 + if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution); err != nil { + return nil, err + } + item.Timestamp = time.Unix(ts, 0) + items = append(items, item) + } + return items, nil +} + +func SearchHistory(db *sql.DB, query string) ([]HistoryItem, error) { + sqlQuery := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details, COALESCE(resolution, '') + FROM history + WHERE command LIKE ? OR details LIKE ? + ORDER BY id DESC + LIMIT 50` + + wildcard := "%" + query + "%" + rows, err := db.Query(sqlQuery, wildcard, wildcard) + if err != nil { + return nil, err + } + defer rows.Close() + + var items []HistoryItem + for rows.Next() { + var item HistoryItem + var ts int64 + if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution); err != nil { + return nil, err + } + item.Timestamp = time.Unix(ts, 0) + items = append(items, item) + } + return items, nil +} + +type QueryOpts struct { + Limit int + Filter string + Since time.Duration +} + +func GetFailures(db *sql.DB, opts QueryOpts) ([]HistoryItem, error) { + queryBuilder := `SELECT h.id, h.timestamp, h.command, h.exit_code, h.duration_ms, h.directory, h.session_id, h.details, COALESCE(h.resolution, '') + FROM history h` + var args []interface{} + var whereClauses []string + + whereClauses = append(whereClauses, "h.exit_code != 0") + + if opts.Filter != "" { + whereClauses = append(whereClauses, "h.command LIKE ?") + args = append(args, "%"+opts.Filter+"%") + } + + if opts.Since > 0 { + cutoff := time.Now().Add(-opts.Since).Unix() + whereClauses = append(whereClauses, "h.timestamp >= ?") + args = append(args, cutoff) + } + + if len(whereClauses) > 0 { + queryBuilder += " WHERE " + strings.Join(whereClauses, " AND ") + } + + queryBuilder += " ORDER BY h.id DESC" + + if opts.Limit > 0 { + queryBuilder += " LIMIT ?" + args = append(args, opts.Limit) + } + + rows, err := db.Query(queryBuilder, args...) + if err != nil { + return nil, err + } + defer rows.Close() + + var items []HistoryItem + for rows.Next() { + var item HistoryItem + var ts int64 + if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution); err != nil { + return nil, err + } + item.Timestamp = time.Unix(ts, 0) + items = append(items, item) + } + return items, nil +} + +func GetLastUnresolvedFailure(db *sql.DB) (*HistoryItem, error) { + query := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details, COALESCE(resolution, '') + FROM history + WHERE exit_code != 0 AND exit_code != 130 AND (resolution IS NULL OR resolution = '') + ORDER BY id DESC LIMIT 1` + + row := db.QueryRow(query) + var item HistoryItem + var ts int64 + err := row.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + item.Timestamp = time.Unix(ts, 0) + return &item, nil +} + +func GetHistoryByID(db *sql.DB, id int64) (*HistoryItem, error) { + query := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details, COALESCE(resolution, '') + FROM history WHERE id = ?` + + row := db.QueryRow(query, id) + var item HistoryItem + var ts int64 + err := row.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + item.Timestamp = time.Unix(ts, 0) + return &item, nil +} + +func MarkResolution(db *sql.DB, id int64, resolution string) error { + query := `UPDATE history SET resolution = ? WHERE id = ?` + result, err := db.Exec(query, resolution, id) + if err != nil { + return err + } + rowsAffected, err := result.RowsAffected() + if err != nil { + return err + } + if rowsAffected == 0 { + return fmt.Errorf("history entry not found: %d", id) + } + return nil +} diff --git a/internal/core/events.go b/internal/core/events.go new file mode 100644 index 0000000..1d3a1d5 --- /dev/null +++ b/internal/core/events.go @@ -0,0 +1,325 @@ +package core + +import ( + "sync" + "time" + + "dev-cli/internal/infra" +) + +type EventType string + +const ( + EventCommandStart EventType = "command.start" + EventCommandOutput EventType = "command.output" + EventCommandComplete EventType = "command.complete" + EventCommandError EventType = "command.error" + + EventContainerLog EventType = "container.log" + EventContainerStatus EventType = "container.status" + EventContainerAlert EventType = "container.alert" + + EventGitStatus EventType = "git.status" + EventGitChanged EventType = "git.changed" + + EventAISuggestion EventType = "ai.suggestion" + EventAIAnalysis EventType = "ai.analysis" + + EventSystemAlert EventType = "system.alert" + EventSystemStats EventType = "system.stats" + + EventWorkflowStart EventType = "workflow.start" + EventWorkflowStep EventType = "workflow.step" + EventWorkflowCheckpoint EventType = "workflow.checkpoint" + EventWorkflowComplete EventType = "workflow.complete" + EventWorkflowRollback EventType = "workflow.rollback" + + EventRCAStart EventType = "rca.start" + EventRCANodeFound EventType = "rca.node_found" + EventRCAComplete EventType = "rca.complete" + EventRCACacheHit EventType = "rca.cache_hit" + + EventRemediationPending EventType = "remediation.pending" + EventRemediationApproved EventType = "remediation.approved" + EventRemediationExecuted EventType = "remediation.executed" + EventRemediationRolledBack EventType = "remediation.rollback" + EventRemediationSkipped EventType = "remediation.skipped" +) + +type Event struct { + Type EventType + Timestamp time.Time + Source string + Data interface{} + BlockID string +} + +type EventHandler func(Event) + +type EventBus struct { + mu sync.RWMutex + subscribers map[EventType][]EventHandler + history []Event + maxHistory int +} + +func NewEventBus() *EventBus { + return &EventBus{ + subscribers: make(map[EventType][]EventHandler), + history: make([]Event, 0), + maxHistory: 100, + } +} + +func (e *EventBus) Subscribe(eventType EventType, handler EventHandler) { + e.mu.Lock() + defer e.mu.Unlock() + e.subscribers[eventType] = append(e.subscribers[eventType], handler) +} + +func (e *EventBus) SubscribeAll(handler EventHandler) { + e.mu.Lock() + defer e.mu.Unlock() + e.subscribers["*"] = append(e.subscribers["*"], handler) +} + +func (e *EventBus) Publish(event Event) { + e.mu.Lock() + e.history = append(e.history, event) + if len(e.history) > e.maxHistory { + e.history = e.history[1:] + } + + handlers := make([]EventHandler, 0) + handlers = append(handlers, e.subscribers[event.Type]...) + handlers = append(handlers, e.subscribers["*"]...) + e.mu.Unlock() + + for _, handler := range handlers { + handler(event) + } +} + +func (e *EventBus) RecentEvents(n int) []Event { + e.mu.RLock() + defer e.mu.RUnlock() + + if n > len(e.history) { + n = len(e.history) + } + return e.history[len(e.history)-n:] +} + +func (e *EventBus) RecentByType(eventType EventType, n int) []Event { + e.mu.RLock() + defer e.mu.RUnlock() + + var result []Event + for i := len(e.history) - 1; i >= 0 && len(result) < n; i-- { + if e.history[i].Type == eventType { + result = append(result, e.history[i]) + } + } + return result +} + +type BlockType string + +const ( + BlockTypeCommand BlockType = "command" + BlockTypeAI BlockType = "ai" + BlockTypeOutput BlockType = "output" + BlockTypeError BlockType = "error" + BlockTypeSuggestion BlockType = "suggestion" +) + +type Block struct { + ID string + Type BlockType + Timestamp time.Time + Command string + Output string + ExitCode int + Duration time.Duration + Folded bool + + AISuggestion string + AIAnalyzed bool + + WorkingDir string +} + +type Suggestion struct { + ForBlockID string + Type string + Title string + Command string + Explanation string + Confidence float64 +} + +type StateStore struct { + mu sync.RWMutex + + Blocks []Block + blockIndex map[string]int + SelectedIdx int + MaxBlocks int + + DockerHealth infra.DockerHealth + GPUStats infra.GPUStats + StarshipLine string + + Suggestions []Suggestion + LastError *Block + ErrorPatterns map[string]string + + Cwd string + Shell string + IsLoading bool +} + +func NewStateStore() *StateStore { + return &StateStore{ + Blocks: make([]Block, 0), + blockIndex: make(map[string]int), + SelectedIdx: -1, + MaxBlocks: 100, + Suggestions: make([]Suggestion, 0), + ErrorPatterns: make(map[string]string), + } +} + +func (s *StateStore) AddBlock(block Block) { + s.mu.Lock() + defer s.mu.Unlock() + + if len(s.Blocks) >= s.MaxBlocks { + oldest := s.Blocks[0] + delete(s.blockIndex, oldest.ID) + s.Blocks = s.Blocks[1:] + s.rebuildIndex() + } + + s.Blocks = append(s.Blocks, block) + s.blockIndex[block.ID] = len(s.Blocks) - 1 + s.SelectedIdx = len(s.Blocks) - 1 + + if block.ExitCode != 0 { + s.LastError = &block + } +} + +func (s *StateStore) GetBlock(id string) *Block { + s.mu.RLock() + defer s.mu.RUnlock() + + if idx, ok := s.blockIndex[id]; ok && idx < len(s.Blocks) { + return &s.Blocks[idx] + } + return nil +} + +func (s *StateStore) GetRecentBlocks(n int) []Block { + s.mu.RLock() + defer s.mu.RUnlock() + + if n > len(s.Blocks) { + n = len(s.Blocks) + } + result := make([]Block, n) + copy(result, s.Blocks[len(s.Blocks)-n:]) + return result +} + +func (s *StateStore) GetBlocks() []Block { + s.mu.RLock() + defer s.mu.RUnlock() + + result := make([]Block, len(s.Blocks)) + copy(result, s.Blocks) + return result +} + +func (s *StateStore) UpdateBlock(id string, fn func(*Block)) { + s.mu.Lock() + defer s.mu.Unlock() + + if idx, ok := s.blockIndex[id]; ok && idx < len(s.Blocks) { + fn(&s.Blocks[idx]) + } +} + +func (s *StateStore) AddSuggestion(suggestion Suggestion) { + s.mu.Lock() + defer s.mu.Unlock() + + s.Suggestions = append(s.Suggestions, suggestion) + if len(s.Suggestions) > 10 { + s.Suggestions = s.Suggestions[1:] + } +} + +func (s *StateStore) GetSuggestionsForBlock(blockID string) []Suggestion { + s.mu.RLock() + defer s.mu.RUnlock() + + var result []Suggestion + for _, sug := range s.Suggestions { + if sug.ForBlockID == blockID { + result = append(result, sug) + } + } + return result +} + +func (s *StateStore) ClearBlocks() { + s.mu.Lock() + defer s.mu.Unlock() + s.Blocks = make([]Block, 0) + s.blockIndex = make(map[string]int) + s.SelectedIdx = -1 +} + +func (s *StateStore) rebuildIndex() { + s.blockIndex = make(map[string]int, len(s.Blocks)) + for i, block := range s.Blocks { + s.blockIndex[block.ID] = i + } +} + +func (s *StateStore) SetDockerHealth(h infra.DockerHealth) { + s.mu.Lock() + defer s.mu.Unlock() + s.DockerHealth = h +} + +func (s *StateStore) SetGPUStats(g infra.GPUStats) { + s.mu.Lock() + defer s.mu.Unlock() + s.GPUStats = g +} + +func (s *StateStore) SetStarshipLine(line string) { + s.mu.Lock() + defer s.mu.Unlock() + s.StarshipLine = line +} + +func (s *StateStore) SetCwd(cwd string) { + s.mu.Lock() + defer s.mu.Unlock() + s.Cwd = cwd +} + +func (s *StateStore) GetContext() map[string]interface{} { + s.mu.RLock() + defer s.mu.RUnlock() + + return map[string]interface{}{ + "cwd": s.Cwd, + "container_count": len(s.DockerHealth.Containers), + "has_last_error": s.LastError != nil, + "recent_commands": len(s.Blocks), + } +} diff --git a/internal/core/executor.go b/internal/core/executor.go new file mode 100644 index 0000000..4ac5769 --- /dev/null +++ b/internal/core/executor.go @@ -0,0 +1,399 @@ +package core + +import ( + "bytes" + "context" + "database/sql" + "fmt" + "io" + "os" + "os/exec" + "strings" + "time" + + "github.com/creack/pty" +) + +type ExecResult struct { + Command string + Output string + ExitCode int + Duration time.Duration + Timestamp time.Time + Shell string + Cwd string +} + +type EventPublisher interface { + PublishCommandEvent(command string, exitCode int, duration time.Duration, output string) +} + +var globalDB *sql.DB + +func SetDatabase(db *sql.DB) { + globalDB = db +} + +func ExecuteAndLog(command string, publisher EventPublisher) ExecResult { + return ExecuteAndLogWithTimeout(command, 60*time.Second, publisher) +} + +func ExecuteAndLogWithTimeout(command string, timeout time.Duration, publisher EventPublisher) ExecResult { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + result := ExecuteWithContext(ctx, command) + + if globalDB != nil { + logEntry := LogEntry{ + Command: result.Command, + ExitCode: result.ExitCode, + Output: TruncateOutput(result.Output, 10240), + Cwd: result.Cwd, + DurationMs: result.Duration.Milliseconds(), + Timestamp: result.Timestamp.Format(time.RFC3339), + } + if err := SaveCommand(globalDB, logEntry); err != nil { + fmt.Fprintf(os.Stderr, "Warning: failed to log command: %v\n", err) + } + } + + if publisher != nil { + publisher.PublishCommandEvent(result.Command, result.ExitCode, result.Duration, result.Output) + } + + return result +} + +func getShell() string { + if shell := os.Getenv("SHELL"); shell != "" { + return shell + } + for _, shell := range []string{"/bin/zsh", "/usr/bin/zsh", "/bin/bash", "/bin/sh"} { + if _, err := os.Stat(shell); err == nil { + return shell + } + } + return "/bin/sh" +} + +func Execute(command string) ExecResult { + return ExecuteWithTimeout(command, 60*time.Second) +} + +func ExecuteWithTimeout(command string, timeout time.Duration) ExecResult { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + return ExecuteWithContext(ctx, command) +} + +func ExecuteWithContext(ctx context.Context, command string) ExecResult { + start := time.Now() + shell := getShell() + cwd, _ := os.Getwd() + + var cmd *exec.Cmd + var wrappedCmd string + if strings.HasSuffix(shell, "zsh") { + wrappedCmd = fmt.Sprintf("source ~/.zshrc 2>/dev/null; %s", command) + cmd = exec.CommandContext(ctx, shell, "-c", wrappedCmd) + } else if strings.HasSuffix(shell, "bash") { + wrappedCmd = fmt.Sprintf("source ~/.bashrc 2>/dev/null; %s", command) + cmd = exec.CommandContext(ctx, shell, "-c", wrappedCmd) + } else { + wrappedCmd = command + cmd = exec.CommandContext(ctx, shell, "-c", command) + } + + var stdout, stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + cmd.Dir = cwd + cmd.Env = os.Environ() + + hasTermEnv := false + for _, env := range cmd.Env { + if strings.HasPrefix(env, "TERM=") { + hasTermEnv = true + break + } + } + if !hasTermEnv { + cmd.Env = append(cmd.Env, "TERM=xterm-256color") + } + + err := cmd.Run() + + duration := time.Since(start) + + output := stdout.String() + stderrStr := stderr.String() + + stderrStr = filterShellNoise(stderrStr) + + if stderrStr != "" { + if output != "" && !strings.HasSuffix(output, "\n") { + output += "\n" + } + output += stderrStr + } + + output = strings.TrimSuffix(output, "\n") + + exitCode := 0 + if err != nil { + if exitError, ok := err.(*exec.ExitError); ok { + exitCode = exitError.ExitCode() + } else { + exitCode = 1 + if output == "" { + output = err.Error() + } + } + } + + return ExecResult{ + Command: command, + Output: output, + ExitCode: exitCode, + Duration: duration, + Timestamp: start, + Shell: shell, + Cwd: cwd, + } +} + +func filterShellNoise(stderr string) string { + lines := strings.Split(stderr, "\n") + var filtered []string + + for _, line := range lines { + if strings.Contains(line, "compinit") || + strings.Contains(line, "compdef") || + strings.Contains(line, "zinit") || + strings.Contains(line, "Loading") || + strings.Contains(line, "Loaded") || + strings.TrimSpace(line) == "" { + continue + } + filtered = append(filtered, line) + } + + return strings.Join(filtered, "\n") +} + +func ExecuteSimple(command string) ExecResult { + start := time.Now() + + cmd := exec.Command("sh", "-c", command) + + var stdout, stderr bytes.Buffer + cmd.Stdout = &stdout + cmd.Stderr = &stderr + + cwd, _ := os.Getwd() + cmd.Dir = cwd + cmd.Env = os.Environ() + + err := cmd.Run() + + duration := time.Since(start) + + output := stdout.String() + if stderr.Len() > 0 { + if output != "" && !strings.HasSuffix(output, "\n") { + output += "\n" + } + output += stderr.String() + } + + output = strings.TrimSuffix(output, "\n") + + exitCode := 0 + if err != nil { + if exitError, ok := err.(*exec.ExitError); ok { + exitCode = exitError.ExitCode() + } else { + exitCode = 1 + if output == "" { + output = err.Error() + } + } + } + + return ExecResult{ + Command: command, + Output: output, + ExitCode: exitCode, + Duration: duration, + Timestamp: start, + Shell: "sh", + } +} + +func IsAIQuery(input string) bool { + input = strings.TrimSpace(input) + return strings.HasPrefix(input, "?") || strings.HasPrefix(input, "@") +} + +func ParseAIQuery(input string) (queryType string, query string) { + input = strings.TrimSpace(input) + + if strings.HasPrefix(input, "?") { + return "question", strings.TrimPrefix(input, "?") + } + + if strings.HasPrefix(input, "@fix") { + return "fix", strings.TrimSpace(strings.TrimPrefix(input, "@fix")) + } + + if strings.HasPrefix(input, "@explain") { + return "explain", strings.TrimSpace(strings.TrimPrefix(input, "@explain")) + } + + if strings.HasPrefix(input, "@") { + parts := strings.SplitN(strings.TrimPrefix(input, "@"), " ", 2) + if len(parts) == 2 { + return parts[0], parts[1] + } + return parts[0], "" + } + + return "", input +} + +func ExecutePTY(command string) ExecResult { + return ExecutePTYWithTimeout(command, 60*time.Second) +} + +func ExecutePTYWithTimeout(command string, timeout time.Duration) ExecResult { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + return ExecutePTYWithContext(ctx, command) +} + +func ExecutePTYWithContext(ctx context.Context, command string) ExecResult { + start := time.Now() + shell := getShell() + + var cmd *exec.Cmd + if strings.HasSuffix(shell, "zsh") { + cmd = exec.CommandContext(ctx, shell, "-i", "-c", command) + } else if strings.HasSuffix(shell, "bash") { + cmd = exec.CommandContext(ctx, shell, "-i", "-c", command) + } else { + cmd = exec.CommandContext(ctx, shell, "-c", command) + } + + cwd, _ := os.Getwd() + cmd.Dir = cwd + + cmd.Env = os.Environ() + cmd.Env = append(cmd.Env, "TERM=xterm-256color") + + ptmx, err := pty.Start(cmd) + if err != nil { + return ExecResult{ + Command: command, + Output: "Failed to start PTY: " + err.Error(), + ExitCode: 1, + Duration: time.Since(start), + Timestamp: start, + Shell: shell, + } + } + defer ptmx.Close() + + var output bytes.Buffer + done := make(chan error, 1) + + go func() { + io.Copy(&output, ptmx) + done <- cmd.Wait() + }() + + select { + case err = <-done: + case <-ctx.Done(): + cmd.Process.Kill() + err = ctx.Err() + } + + duration := time.Since(start) + outputStr := cleanPTYOutput(output.String()) + + exitCode := 0 + if err != nil { + if exitError, ok := err.(*exec.ExitError); ok { + exitCode = exitError.ExitCode() + } else if err == context.DeadlineExceeded { + exitCode = 124 + outputStr = "Command timed out" + } else { + exitCode = 1 + if outputStr == "" { + outputStr = err.Error() + } + } + } + + return ExecResult{ + Command: command, + Output: outputStr, + ExitCode: exitCode, + Duration: duration, + Timestamp: start, + Shell: shell, + } +} + +func cleanPTYOutput(output string) string { + output = stripANSI(output) + + lines := strings.Split(output, "\n") + var cleaned []string + + for _, line := range lines { + if len(cleaned) == 0 && strings.TrimSpace(line) == "" { + continue + } + if strings.Contains(line, "compinit") || + strings.Contains(line, "zinit") || + strings.Contains(line, "Loading") || + strings.HasPrefix(strings.TrimSpace(line), "❯") || + strings.HasPrefix(strings.TrimSpace(line), "$") { + continue + } + cleaned = append(cleaned, line) + } + + result := strings.Join(cleaned, "\n") + return strings.TrimSpace(result) +} + +func stripANSI(str string) string { + var result strings.Builder + inEscape := false + + for i := 0; i < len(str); i++ { + if str[i] == '\x1b' { + inEscape = true + continue + } + if inEscape { + if (str[i] >= 'a' && str[i] <= 'z') || (str[i] >= 'A' && str[i] <= 'Z') { + inEscape = false + } + continue + } + if str[i] < 32 && str[i] != '\n' && str[i] != '\t' && str[i] != '\r' { + continue + } + result.WriteByte(str[i]) + } + + return result.String() +} diff --git a/internal/core/rca.go b/internal/core/rca.go new file mode 100644 index 0000000..6bade9c --- /dev/null +++ b/internal/core/rca.go @@ -0,0 +1,410 @@ +package core + +import ( + "database/sql" + "encoding/json" + "fmt" + "hash/fnv" + "time" +) + +type RootCause struct { + ID string `json:"id"` + ErrorSignature string `json:"error_signature"` + Timestamp time.Time `json:"timestamp"` + RootCauseNodes []string `json:"root_cause_nodes"` + RemediationSteps []string `json:"remediation_steps"` + Confidence float64 `json:"confidence"` + HistoryItemID int64 `json:"history_item_id"` +} + +type Runbook struct { + ID string `json:"id"` + ProjectID string `json:"project_id"` + Name string `json:"name"` + Description string `json:"description"` + Steps []RunbookStep `json:"steps"` + SuccessRate float64 `json:"success_rate"` + LastUsed time.Time `json:"last_used"` + UsageCount int `json:"usage_count"` + Tags []string `json:"tags"` +} + +type RunbookStep struct { + ID string `json:"id"` + Name string `json:"name"` + Command string `json:"command"` + Description string `json:"description"` + Rollback string `json:"rollback,omitempty"` + Condition string `json:"condition,omitempty"` +} + +type ProjectFingerprint struct { + ID string `json:"id"` + ProjectType string `json:"project_type"` + PackageManager string `json:"package_manager"` + CommonIssues []string `json:"common_issues"` + AssociatedRunbooks []string `json:"associated_runbooks"` + DetectedAt string `json:"detected_at"` + DetectedTime time.Time `json:"detected_time"` +} + +func GenerateErrorSignature(command string, exitCode int, output string) string { + firstLine := output + if idx := indexOfRune(output, '\n'); idx > 0 { + firstLine = output[:idx] + } + if len(firstLine) > 100 { + firstLine = firstLine[:100] + } + + combined := fmt.Sprintf("%s|%d|%s", command, exitCode, firstLine) + return hashString(combined) +} + +func indexOfRune(s string, r rune) int { + for i, c := range s { + if c == r { + return i + } + } + return -1 +} + +func hashString(s string) string { + h := fnv.New64a() + h.Write([]byte(s)) + return fmt.Sprintf("%016x", h.Sum64()) +} + +func SaveRootCause(db *sql.DB, rc RootCause) error { + nodesJSON, err := json.Marshal(rc.RootCauseNodes) + if err != nil { + return fmt.Errorf("marshal root_cause_nodes: %w", err) + } + stepsJSON, err := json.Marshal(rc.RemediationSteps) + if err != nil { + return fmt.Errorf("marshal remediation_steps: %w", err) + } + + query := `INSERT OR REPLACE INTO root_causes + (id, error_signature, timestamp, root_cause_nodes, remediation_steps, confidence, history_item_id) + VALUES (?, ?, ?, ?, ?, ?, ?)` + + _, err = db.Exec(query, + rc.ID, + rc.ErrorSignature, + rc.Timestamp.Unix(), + string(nodesJSON), + string(stepsJSON), + rc.Confidence, + rc.HistoryItemID, + ) + return err +} + +func GetRootCauseBySignature(db *sql.DB, signature string) (*RootCause, error) { + query := `SELECT id, error_signature, timestamp, root_cause_nodes, remediation_steps, confidence, history_item_id + FROM root_causes WHERE error_signature = ? ORDER BY timestamp DESC LIMIT 1` + + row := db.QueryRow(query, signature) + return scanRootCause(row) +} + +func GetRootCauseByID(db *sql.DB, id string) (*RootCause, error) { + query := `SELECT id, error_signature, timestamp, root_cause_nodes, remediation_steps, confidence, history_item_id + FROM root_causes WHERE id = ?` + + row := db.QueryRow(query, id) + return scanRootCause(row) +} + +func GetRecentRootCauses(db *sql.DB, limit int) ([]RootCause, error) { + query := `SELECT id, error_signature, timestamp, root_cause_nodes, remediation_steps, confidence, history_item_id + FROM root_causes ORDER BY timestamp DESC LIMIT ?` + + rows, err := db.Query(query, limit) + if err != nil { + return nil, err + } + defer rows.Close() + + var results []RootCause + for rows.Next() { + rc, err := scanRootCauseRow(rows) + if err != nil { + return nil, err + } + results = append(results, *rc) + } + return results, nil +} + +func scanRootCause(row *sql.Row) (*RootCause, error) { + var rc RootCause + var ts int64 + var nodesJSON, stepsJSON string + + err := row.Scan(&rc.ID, &rc.ErrorSignature, &ts, &nodesJSON, &stepsJSON, &rc.Confidence, &rc.HistoryItemID) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + + rc.Timestamp = time.Unix(ts, 0) + if err := json.Unmarshal([]byte(nodesJSON), &rc.RootCauseNodes); err != nil { + rc.RootCauseNodes = []string{} + } + if err := json.Unmarshal([]byte(stepsJSON), &rc.RemediationSteps); err != nil { + rc.RemediationSteps = []string{} + } + + return &rc, nil +} + +func scanRootCauseRow(rows *sql.Rows) (*RootCause, error) { + var rc RootCause + var ts int64 + var nodesJSON, stepsJSON string + + err := rows.Scan(&rc.ID, &rc.ErrorSignature, &ts, &nodesJSON, &stepsJSON, &rc.Confidence, &rc.HistoryItemID) + if err != nil { + return nil, err + } + + rc.Timestamp = time.Unix(ts, 0) + if err := json.Unmarshal([]byte(nodesJSON), &rc.RootCauseNodes); err != nil { + rc.RootCauseNodes = []string{} + } + if err := json.Unmarshal([]byte(stepsJSON), &rc.RemediationSteps); err != nil { + rc.RemediationSteps = []string{} + } + + return &rc, nil +} + +func SaveRunbook(db *sql.DB, rb Runbook) error { + stepsJSON, err := json.Marshal(rb.Steps) + if err != nil { + return fmt.Errorf("marshal steps: %w", err) + } + tagsJSON, err := json.Marshal(rb.Tags) + if err != nil { + return fmt.Errorf("marshal tags: %w", err) + } + + query := `INSERT OR REPLACE INTO runbooks + (id, project_id, name, description, steps, success_rate, last_used, usage_count, tags) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)` + + _, err = db.Exec(query, + rb.ID, + rb.ProjectID, + rb.Name, + rb.Description, + string(stepsJSON), + rb.SuccessRate, + rb.LastUsed.Unix(), + rb.UsageCount, + string(tagsJSON), + ) + return err +} + +func GetRunbookByID(db *sql.DB, id string) (*Runbook, error) { + query := `SELECT id, project_id, name, description, steps, success_rate, last_used, usage_count, tags + FROM runbooks WHERE id = ?` + + row := db.QueryRow(query, id) + return scanRunbook(row) +} + +func GetRunbooksForProject(db *sql.DB, projectID string) ([]Runbook, error) { + query := `SELECT id, project_id, name, description, steps, success_rate, last_used, usage_count, tags + FROM runbooks WHERE project_id = ? ORDER BY success_rate DESC` + + rows, err := db.Query(query, projectID) + if err != nil { + return nil, err + } + defer rows.Close() + + var results []Runbook + for rows.Next() { + rb, err := scanRunbookRow(rows) + if err != nil { + return nil, err + } + results = append(results, *rb) + } + return results, nil +} + +func UpdateRunbookStats(db *sql.DB, id string, success bool) error { + rb, err := GetRunbookByID(db, id) + if err != nil { + return err + } + if rb == nil { + return fmt.Errorf("runbook not found: %s", id) + } + + rb.UsageCount++ + if success { + rb.SuccessRate = ((rb.SuccessRate * float64(rb.UsageCount-1)) + 1.0) / float64(rb.UsageCount) + } else { + rb.SuccessRate = (rb.SuccessRate * float64(rb.UsageCount-1)) / float64(rb.UsageCount) + } + rb.LastUsed = time.Now() + + query := `UPDATE runbooks SET success_rate = ?, last_used = ?, usage_count = ? WHERE id = ?` + _, err = db.Exec(query, rb.SuccessRate, rb.LastUsed.Unix(), rb.UsageCount, id) + return err +} + +func scanRunbook(row *sql.Row) (*Runbook, error) { + var rb Runbook + var lastUsed int64 + var stepsJSON, tagsJSON string + + err := row.Scan(&rb.ID, &rb.ProjectID, &rb.Name, &rb.Description, &stepsJSON, &rb.SuccessRate, &lastUsed, &rb.UsageCount, &tagsJSON) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + + rb.LastUsed = time.Unix(lastUsed, 0) + if err := json.Unmarshal([]byte(stepsJSON), &rb.Steps); err != nil { + rb.Steps = []RunbookStep{} + } + if err := json.Unmarshal([]byte(tagsJSON), &rb.Tags); err != nil { + rb.Tags = []string{} + } + + return &rb, nil +} + +func scanRunbookRow(rows *sql.Rows) (*Runbook, error) { + var rb Runbook + var lastUsed int64 + var stepsJSON, tagsJSON string + + err := rows.Scan(&rb.ID, &rb.ProjectID, &rb.Name, &rb.Description, &stepsJSON, &rb.SuccessRate, &lastUsed, &rb.UsageCount, &tagsJSON) + if err != nil { + return nil, err + } + + rb.LastUsed = time.Unix(lastUsed, 0) + if err := json.Unmarshal([]byte(stepsJSON), &rb.Steps); err != nil { + rb.Steps = []RunbookStep{} + } + if err := json.Unmarshal([]byte(tagsJSON), &rb.Tags); err != nil { + rb.Tags = []string{} + } + + return &rb, nil +} + +func SaveProjectFingerprint(db *sql.DB, fp ProjectFingerprint) error { + issuesJSON, err := json.Marshal(fp.CommonIssues) + if err != nil { + return fmt.Errorf("marshal common_issues: %w", err) + } + runbooksJSON, err := json.Marshal(fp.AssociatedRunbooks) + if err != nil { + return fmt.Errorf("marshal associated_runbooks: %w", err) + } + + query := `INSERT OR REPLACE INTO project_fingerprints + (id, project_type, package_manager, common_issues, associated_runbooks, detected_at, detected_time) + VALUES (?, ?, ?, ?, ?, ?, ?)` + + _, err = db.Exec(query, + fp.ID, + fp.ProjectType, + fp.PackageManager, + string(issuesJSON), + string(runbooksJSON), + fp.DetectedAt, + fp.DetectedTime.Unix(), + ) + return err +} + +func GetProjectFingerprint(db *sql.DB, path string) (*ProjectFingerprint, error) { + query := `SELECT id, project_type, package_manager, common_issues, associated_runbooks, detected_at, detected_time + FROM project_fingerprints WHERE detected_at = ?` + + row := db.QueryRow(query, path) + return scanProjectFingerprint(row) +} + +func GetProjectFingerprintByType(db *sql.DB, projectType string) ([]ProjectFingerprint, error) { + query := `SELECT id, project_type, package_manager, common_issues, associated_runbooks, detected_at, detected_time + FROM project_fingerprints WHERE project_type = ?` + + rows, err := db.Query(query, projectType) + if err != nil { + return nil, err + } + defer rows.Close() + + var results []ProjectFingerprint + for rows.Next() { + fp, err := scanProjectFingerprintRow(rows) + if err != nil { + return nil, err + } + results = append(results, *fp) + } + return results, nil +} + +func scanProjectFingerprint(row *sql.Row) (*ProjectFingerprint, error) { + var fp ProjectFingerprint + var detectedTime int64 + var issuesJSON, runbooksJSON string + + err := row.Scan(&fp.ID, &fp.ProjectType, &fp.PackageManager, &issuesJSON, &runbooksJSON, &fp.DetectedAt, &detectedTime) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + + fp.DetectedTime = time.Unix(detectedTime, 0) + if err := json.Unmarshal([]byte(issuesJSON), &fp.CommonIssues); err != nil { + fp.CommonIssues = []string{} + } + if err := json.Unmarshal([]byte(runbooksJSON), &fp.AssociatedRunbooks); err != nil { + fp.AssociatedRunbooks = []string{} + } + + return &fp, nil +} + +func scanProjectFingerprintRow(rows *sql.Rows) (*ProjectFingerprint, error) { + var fp ProjectFingerprint + var detectedTime int64 + var issuesJSON, runbooksJSON string + + err := rows.Scan(&fp.ID, &fp.ProjectType, &fp.PackageManager, &issuesJSON, &runbooksJSON, &fp.DetectedAt, &detectedTime) + if err != nil { + return nil, err + } + + fp.DetectedTime = time.Unix(detectedTime, 0) + if err := json.Unmarshal([]byte(issuesJSON), &fp.CommonIssues); err != nil { + fp.CommonIssues = []string{} + } + if err := json.Unmarshal([]byte(runbooksJSON), &fp.AssociatedRunbooks); err != nil { + fp.AssociatedRunbooks = []string{} + } + + return &fp, nil +} diff --git a/internal/hook/zsh.go b/internal/hook/zsh.go index 9951711..55d1127 100644 --- a/internal/hook/zsh.go +++ b/internal/hook/zsh.go @@ -7,6 +7,8 @@ typeset -g __DEVOPS_CMD="" typeset -g __DEVOPS_START_TIME=0 typeset -g __DEVOPS_SKIP_LOG=0 typeset -g __DEVOPS_LAST_OUTPUT="" +typeset -g __DEVOPS_LAST_FAILURE_ID="" +typeset -g __DEVOPS_LAST_FAILURE_CMD="" typeset -ga __DEVOPS_SKIP_CMDS=(vim vi nvim nano less more top htop man ssh tmux screen) __devops_is_interactive() { @@ -83,6 +85,48 @@ __devops_suggest_fix() { return 1 } +# Check for unresolved failure and prompt user +__devops_check_resolution() { + # Check if there's an unresolved failure + local failure_info=$(dev-cli check-last-failure 2>/dev/null) + if [[ -z "$failure_info" ]]; then + __DEVOPS_LAST_FAILURE_ID="" + __DEVOPS_LAST_FAILURE_CMD="" + return + fi + + __DEVOPS_LAST_FAILURE_ID="${failure_info%%|*}" + __DEVOPS_LAST_FAILURE_CMD="${failure_info#*|}" +} + +__devops_prompt_resolution() { + local failure_id="$1" + local failure_cmd="$2" + + echo "" + echo "\033[32m✓\033[0m Success after failure: \033[90m$failure_cmd\033[0m" + echo -n "\033[33m❓ Did this fix the issue? [y/n/skip]: \033[0m" + read -r response + + case "$response" in + [Yy]*) + dev-cli mark-resolved --id "$failure_id" --resolution solution 2>/dev/null + echo "\033[32m✓\033[0m Marked as solution!" + ;; + [Nn]*) + dev-cli mark-resolved --id "$failure_id" --resolution unrelated 2>/dev/null + echo "\033[90m○ Marked as unrelated\033[0m" + ;; + *) + dev-cli mark-resolved --id "$failure_id" --resolution skipped 2>/dev/null + echo "\033[90m○ Skipped\033[0m" + ;; + esac + + __DEVOPS_LAST_FAILURE_ID="" + __DEVOPS_LAST_FAILURE_CMD="" +} + __devops_precmd() { local exit_code=$? [[ -z "$__DEVOPS_CMD" || $__DEVOPS_SKIP_LOG -eq 1 ]] && return 0 @@ -97,11 +141,17 @@ __devops_precmd() { --duration-ms "$duration_ms" 2>/dev/null &! if [[ $exit_code -ne 0 && $exit_code -ne 130 ]]; then + # Command failed - check for unresolved failure after a short delay + (sleep 0.1 && __devops_check_resolution) &! + # Try smart suggestion first if ! __devops_suggest_fix "$__DEVOPS_CMD" "$exit_code" "$__DEVOPS_LAST_OUTPUT"; then # Fallback to generic message echo "\033[90m× Failure logged. For AI analysis:\033[0m dcap \"$__DEVOPS_CMD\"" fi + elif [[ $exit_code -eq 0 && -n "$__DEVOPS_LAST_FAILURE_ID" ]]; then + # Command succeeded and there was a prior unresolved failure + __devops_prompt_resolution "$__DEVOPS_LAST_FAILURE_ID" "$__DEVOPS_LAST_FAILURE_CMD" fi __DEVOPS_CMD="" @@ -132,6 +182,10 @@ dcap() { rm -f "$tmpfile" if [[ $exit_code -ne 0 && $exit_code -ne 130 ]]; then + # Check for unresolved failure after logging + sleep 0.1 + __devops_check_resolution + # Show smart suggestion __devops_suggest_fix "$*" "$exit_code" "$output" echo "" @@ -141,11 +195,17 @@ dcap() { if [[ "$response" =~ ^[Yy]$ ]]; then dev-cli explain --last 1 --interactive 2>/dev/null fi + elif [[ $exit_code -eq 0 && -n "$__DEVOPS_LAST_FAILURE_ID" ]]; then + # Success after failure - prompt for resolution + __devops_prompt_resolution "$__DEVOPS_LAST_FAILURE_ID" "$__DEVOPS_LAST_FAILURE_CMD" fi return $exit_code } +# Initialize by checking for any pending unresolved failures +__devops_check_resolution + autoload -Uz add-zsh-hook add-zsh-hook preexec __devops_preexec add-zsh-hook precmd __devops_precmd diff --git a/internal/infra/docker.go b/internal/infra/docker.go index 9cd47dc..4a16fd1 100644 --- a/internal/infra/docker.go +++ b/internal/infra/docker.go @@ -94,9 +94,6 @@ type DockerClient struct { cli *client.Client } -// Note: Shared client management has moved to Registry. -// Use GetRegistry().Docker() instead of the deprecated GetSharedDockerClient(). - func NewDockerClient() (*DockerClient, error) { cli, err := client.NewClientWithOpts(client.FromEnv, client.WithAPIVersionNegotiation()) if err != nil { @@ -214,8 +211,6 @@ func (d *DockerClient) Close() error { return nil } -// Container Control Methods - func (d *DockerClient) StartContainer(ctx context.Context, containerID string) error { return d.cli.ContainerStart(ctx, containerID, container.StartOptions{}) } @@ -249,8 +244,6 @@ func (d *DockerClient) UnpauseContainer(ctx context.Context, containerID string) return d.cli.ContainerUnpause(ctx, containerID) } -// Stats Streaming - func (d *DockerClient) GetContainerStats(ctx context.Context, containerID string) (*ContainerStatsSnapshot, error) { stats, err := d.cli.ContainerStats(ctx, containerID, false) if err != nil { @@ -302,14 +295,12 @@ func (d *DockerClient) GetContainerStats(ctx context.Context, containerID string PIDs: v.PidsStats.Current, } - // Calculate CPU percentage cpuDelta := float64(v.CPUStats.CPUUsage.TotalUsage - v.PreCPUStats.CPUUsage.TotalUsage) systemDelta := float64(v.CPUStats.SystemUsage - v.PreCPUStats.SystemUsage) if systemDelta > 0 && cpuDelta > 0 { snapshot.CPUPercent = (cpuDelta / systemDelta) * float64(v.CPUStats.OnlineCPUs) * 100.0 } - // Memory cacheVal := uint64(0) if v.MemoryStats.Stats != nil { cacheVal = v.MemoryStats.Stats["cache"] @@ -320,13 +311,11 @@ func (d *DockerClient) GetContainerStats(ctx context.Context, containerID string snapshot.MemPercent = float64(snapshot.MemUsed) / float64(snapshot.MemLimit) * 100.0 } - // Network I/O for _, netStats := range v.Networks { snapshot.NetRx += netStats.RxBytes snapshot.NetTx += netStats.TxBytes } - // Block I/O for _, bioEntry := range v.BlkioStats.IoServiceBytesRecursive { switch bioEntry.Op { case "Read", "read": @@ -339,8 +328,6 @@ func (d *DockerClient) GetContainerStats(ctx context.Context, containerID string return snapshot, nil } -// Container Inspection - func (d *DockerClient) InspectContainer(ctx context.Context, containerID string) (*ContainerDetail, error) { info, err := d.cli.ContainerInspect(ctx, containerID) if err != nil { @@ -367,13 +354,11 @@ func (d *DockerClient) InspectContainer(ctx context.Context, containerID string) Cmd: info.Config.Cmd, } - // Calculate uptime if info.State.Running { startTime, _ := time.Parse(time.RFC3339Nano, info.State.StartedAt) detail.Uptime = time.Since(startTime).Round(time.Second).String() } - // Mounts for _, m := range info.Mounts { detail.Mounts = append(detail.Mounts, Mount{ Source: m.Source, @@ -383,13 +368,11 @@ func (d *DockerClient) InspectContainer(ctx context.Context, containerID string) }) } - // Network for netName := range info.NetworkSettings.Networks { detail.NetworkID = netName break } - // Ports for portProto, bindings := range info.NetworkSettings.Ports { for _, b := range bindings { var publicPort uint16 @@ -408,8 +391,6 @@ func (d *DockerClient) InspectContainer(ctx context.Context, containerID string) return detail, nil } -// Images - func (d *DockerClient) ListImages(ctx context.Context) ([]ImageInfo, error) { images, err := d.cli.ImageList(ctx, image.ListOptions{}) if err != nil { @@ -420,7 +401,7 @@ func (d *DockerClient) ListImages(ctx context.Context) ([]ImageInfo, error) { for _, img := range images { id := img.ID if len(id) > 19 { - id = id[7:19] // Remove "sha256:" prefix and truncate + id = id[7:19] } result = append(result, ImageInfo{ ID: id, @@ -440,8 +421,6 @@ func (d *DockerClient) RemoveImage(ctx context.Context, imageID string, force bo return err } -// Volumes - func (d *DockerClient) ListVolumes(ctx context.Context) ([]VolumeInfo, error) { volumes, err := d.cli.VolumeList(ctx, volume.ListOptions{}) if err != nil { @@ -465,15 +444,12 @@ func (d *DockerClient) RemoveVolume(ctx context.Context, volumeName string, forc return d.cli.VolumeRemove(ctx, volumeName, force) } -// Container Processes (docker top) - func (d *DockerClient) TopContainer(ctx context.Context, containerID string) ([]ProcessInfo, error) { top, err := d.cli.ContainerTop(ctx, containerID, []string{}) if err != nil { return nil, fmt.Errorf("top failed: %w", err) } - // Find column indices pidIdx, userIdx, cmdIdx := -1, -1, -1 for i, title := range top.Titles { switch title { @@ -503,8 +479,6 @@ func (d *DockerClient) TopContainer(ctx context.Context, containerID string) ([] return result, nil } -// Bulk operations - func (d *DockerClient) PruneContainers(ctx context.Context) (uint64, error) { report, err := d.cli.ContainersPrune(ctx, filters.Args{}) if err != nil { @@ -529,10 +503,6 @@ func (d *DockerClient) PruneVolumes(ctx context.Context) (uint64, error) { return report.SpaceReclaimed, nil } -// ============================================================================= -// Log Streaming -// ============================================================================= - // StreamLogs streams container logs to a LogSink. // Returns when context is cancelled or an error occurs. func (d *DockerClient) StreamLogs(ctx context.Context, containerID string, containerName string, sink LogSink) error { @@ -567,7 +537,6 @@ func (d *DockerClient) StreamLogsToWriter(ctx context.Context, containerID strin } defer reader.Close() - // Docker multiplexes stdout/stderr with 8-byte header buf := make([]byte, 8192) for { select { @@ -578,10 +547,10 @@ func (d *DockerClient) StreamLogsToWriter(ctx context.Context, containerID strin n, err := reader.Read(buf) if n > 0 { - // Skip 8-byte header and write content + data := buf[:n] for len(data) > 8 { - // Header: [stream_type][0][0][0][size1][size2][size3][size4] + size := int(data[4])<<24 | int(data[5])<<16 | int(data[6])<<8 | int(data[7]) if size <= 0 || size > len(data)-8 { break @@ -638,8 +607,8 @@ func (d *DockerClient) processLogStream(ctx context.Context, reader io.ReadClose if n > 0 { data := buf[:n] for len(data) > 8 { - // Parse Docker log header - streamType := data[0] // 1 = stdout, 2 = stderr + + streamType := data[0] size := int(data[4])<<24 | int(data[5])<<16 | int(data[6])<<8 | int(data[7]) if size <= 0 || size > len(data)-8 { break @@ -651,7 +620,6 @@ func (d *DockerClient) processLogStream(ctx context.Context, reader io.ReadClose stream = "stderr" } - // Parse timestamp from content (RFC3339 format at start) timestamp := time.Now() if len(content) > 30 && content[4] == '-' { if t, perr := time.Parse(time.RFC3339Nano, strings.TrimSpace(content[:30])); perr == nil { @@ -667,13 +635,12 @@ func (d *DockerClient) processLogStream(ctx context.Context, reader io.ReadClose Message: strings.TrimSuffix(content, "\n"), } - // Add snapshots at interval if snapshotInterval != nil && time.Since(lastSnapshot) >= *snapshotInterval { if gpu != nil { gpuStats := gpu.GetStats() entry.GPUSnapshot = &gpuStats } - // Could add container stats here too + lastSnapshot = time.Now() } diff --git a/internal/infra/docker_test.go b/internal/infra/docker_test.go index 5d0f4f1..4c9d59d 100644 --- a/internal/infra/docker_test.go +++ b/internal/infra/docker_test.go @@ -10,10 +10,6 @@ import ( "github.com/testcontainers/testcontainers-go/wait" ) -// Integration tests using Testcontainers -// Run with: go test -v -race -run Integration ./internal/infra/... -// Skip with: go test -short ./... - func TestIntegration_DockerClient_CheckHealth(t *testing.T) { if testing.Short() { t.Skip("skipping integration test in short mode") @@ -21,7 +17,6 @@ func TestIntegration_DockerClient_CheckHealth(t *testing.T) { ctx := context.Background() - // Start a simple alpine container req := testcontainers.ContainerRequest{ Image: "alpine:latest", Cmd: []string{"sleep", "30"}, @@ -38,7 +33,6 @@ func TestIntegration_DockerClient_CheckHealth(t *testing.T) { } defer container.Terminate(ctx) - // Test our DockerClient can see this container client, err := NewDockerClient() if err != nil { t.Fatalf("failed to create docker client: %v", err) @@ -50,12 +44,10 @@ func TestIntegration_DockerClient_CheckHealth(t *testing.T) { t.Fatalf("docker should be available: %v", health.Error) } - // Verify we can see at least one container if len(health.Containers) == 0 { t.Error("expected at least one container to be visible") } - // Find our test container containerID := container.GetContainerID() found := false @@ -81,7 +73,6 @@ func TestIntegration_DockerClient_StartStop(t *testing.T) { ctx := context.Background() - // Start a container that stays running req := testcontainers.ContainerRequest{ Image: "alpine:latest", Cmd: []string{"sleep", "60"}, @@ -105,12 +96,10 @@ func TestIntegration_DockerClient_StartStop(t *testing.T) { } defer client.Close() - // Test stop if err := client.StopContainer(ctx, containerID); err != nil { t.Fatalf("failed to stop container: %v", err) } - // Verify stopped time.Sleep(500 * time.Millisecond) health := client.CheckHealth(ctx) for _, c := range health.Containers { @@ -122,12 +111,10 @@ func TestIntegration_DockerClient_StartStop(t *testing.T) { } } - // Test start if err := client.StartContainer(ctx, containerID); err != nil { t.Fatalf("failed to start container: %v", err) } - // Verify running time.Sleep(500 * time.Millisecond) health = client.CheckHealth(ctx) for _, c := range health.Containers { @@ -147,7 +134,6 @@ func TestIntegration_DockerClient_ContainerLogs(t *testing.T) { ctx := context.Background() - // Container that outputs something req := testcontainers.ContainerRequest{ Image: "alpine:latest", Cmd: []string{"sh", "-c", "echo 'test-log-output' && sleep 5"}, @@ -171,7 +157,6 @@ func TestIntegration_DockerClient_ContainerLogs(t *testing.T) { } defer client.Close() - // Give container time to output logs time.Sleep(1 * time.Second) logs, err := client.GetContainerLogs(ctx, containerID, 10) @@ -183,7 +168,6 @@ func TestIntegration_DockerClient_ContainerLogs(t *testing.T) { t.Error("expected at least one log line") } - // Check for our expected output found := false for _, line := range logs { if line == "test-log-output" { @@ -204,7 +188,6 @@ func TestIntegration_MockOllamaAPI(t *testing.T) { ctx := context.Background() - // Use nginx to mock the Ollama API with a static response nginxConf := ` events {} http { @@ -241,7 +224,6 @@ http { } defer container.Terminate(ctx) - // Get the mapped port mappedPort, err := container.MappedPort(ctx, "11434") if err != nil { t.Fatalf("failed to get mapped port: %v", err) @@ -254,15 +236,12 @@ http { baseURL := fmt.Sprintf("http://%s:%s", host, mappedPort.Port()) - // Test OllamaClient against mock ollamaClient := NewOllamaClient(nil, baseURL) - // Test Ping if err := ollamaClient.Ping(ctx); err != nil { t.Fatalf("ping failed: %v", err) } - // Test ListModels models, err := ollamaClient.ListModels(ctx) if err != nil { t.Fatalf("list models failed: %v", err) @@ -276,7 +255,6 @@ http { t.Errorf("expected model name 'qwen2.5-coder:3b-instruct', got '%s'", models[0].Name) } - // Test HasModel hasModel, err := ollamaClient.HasModel(ctx, "qwen2.5-coder") if err != nil { t.Fatalf("has model failed: %v", err) @@ -286,7 +264,6 @@ http { t.Error("expected HasModel to return true for 'qwen2.5-coder'") } - // Test HasModel for non-existent model hasModel, err = ollamaClient.HasModel(ctx, "nonexistent-model") if err != nil { t.Fatalf("has model failed: %v", err) diff --git a/internal/infra/services_test.go b/internal/infra/services_test.go index 7a8769e..c7704c8 100644 --- a/internal/infra/services_test.go +++ b/internal/infra/services_test.go @@ -6,8 +6,6 @@ import ( ) func TestCheckServices(t *testing.T) { - // This test depends on actual services running, which might be flaky. - // So we can mock, or just check that it returns a list of 3 items (Postgres, Redis, Ollama). results := CheckServices() @@ -29,18 +27,16 @@ func TestCheckServices(t *testing.T) { if res.Port != port { t.Errorf("expected port %d for %s, got %d", port, res.Name, res.Port) } - // We don't verify Available field as it depends on environment + } } func TestCheckServices_LocalListener(t *testing.T) { - // Start a dummy listener to simulate a running service + l, err := net.Listen("tcp", "localhost:0") if err != nil { t.Skip("could not listen on local port") } defer l.Close() - // We can't easily inject this into CheckServices without refactoring it to accept a list. - // So for now, we just rely on previous test for structure. } diff --git a/internal/llm/cache.go b/internal/llm/cache.go new file mode 100644 index 0000000..9ca0b46 --- /dev/null +++ b/internal/llm/cache.go @@ -0,0 +1,202 @@ +package llm + +import ( + "sync" + "time" +) + +// CachedAnalysis represents a cached RCA result. +type CachedAnalysis struct { + Signature string + RootCauseNodes []string + RemediationSteps []string + Explanation string + Fix string + Confidence float64 + HitCount int + LastHit time.Time + CreatedAt time.Time +} + +// ErrorCache provides fast lookup for previously analyzed errors. +// Uses LRU eviction to limit memory usage. +type ErrorCache struct { + cache map[string]*CachedAnalysis + order []string // For LRU ordering + maxSize int + mu sync.RWMutex + hits int64 + misses int64 +} + +// NewErrorCache creates a cache with configurable size. +func NewErrorCache(maxSize int) *ErrorCache { + if maxSize <= 0 { + maxSize = 100 + } + return &ErrorCache{ + cache: make(map[string]*CachedAnalysis), + order: make([]string, 0, maxSize), + maxSize: maxSize, + } +} + +// Get retrieves cached analysis by error signature. +// Returns nil if not found. +func (c *ErrorCache) Get(signature string) *CachedAnalysis { + c.mu.Lock() + defer c.mu.Unlock() + + analysis, ok := c.cache[signature] + if !ok { + c.misses++ + return nil + } + + c.hits++ + analysis.HitCount++ + analysis.LastHit = time.Now() + + c.moveToFront(signature) + + return analysis +} + +// Put stores analysis result. +func (c *ErrorCache) Put(signature string, analysis *CachedAnalysis) { + c.mu.Lock() + defer c.mu.Unlock() + + if _, exists := c.cache[signature]; exists { + c.cache[signature] = analysis + c.moveToFront(signature) + return + } + + if len(c.cache) >= c.maxSize { + c.evictOldest() + } + + analysis.CreatedAt = time.Now() + analysis.LastHit = time.Now() + analysis.HitCount = 1 + analysis.Signature = signature + c.cache[signature] = analysis + c.order = append([]string{signature}, c.order...) +} + +// Contains checks if a signature is cached without updating LRU order. +func (c *ErrorCache) Contains(signature string) bool { + c.mu.RLock() + defer c.mu.RUnlock() + _, ok := c.cache[signature] + return ok +} + +// Delete removes an entry from the cache. +func (c *ErrorCache) Delete(signature string) { + c.mu.Lock() + defer c.mu.Unlock() + + delete(c.cache, signature) + c.removeFromOrder(signature) +} + +// Clear removes all entries from the cache. +func (c *ErrorCache) Clear() { + c.mu.Lock() + defer c.mu.Unlock() + + c.cache = make(map[string]*CachedAnalysis) + c.order = make([]string, 0, c.maxSize) + c.hits = 0 + c.misses = 0 +} + +// Size returns the current number of cached entries. +func (c *ErrorCache) Size() int { + c.mu.RLock() + defer c.mu.RUnlock() + return len(c.cache) +} + +// Stats returns cache hit/miss statistics. +func (c *ErrorCache) Stats() CacheStats { + c.mu.RLock() + defer c.mu.RUnlock() + + hitRate := float64(0) + total := c.hits + c.misses + if total > 0 { + hitRate = float64(c.hits) / float64(total) + } + + return CacheStats{ + Hits: c.hits, + Misses: c.misses, + Size: len(c.cache), + MaxSize: c.maxSize, + HitRate: hitRate, + } +} + +// CacheStats contains cache performance metrics. +type CacheStats struct { + Hits int64 + Misses int64 + Size int + MaxSize int + HitRate float64 +} + +// GetTopHits returns the most frequently accessed cache entries. +func (c *ErrorCache) GetTopHits(limit int) []*CachedAnalysis { + c.mu.RLock() + defer c.mu.RUnlock() + + entries := make([]*CachedAnalysis, 0, len(c.cache)) + for _, v := range c.cache { + entries = append(entries, v) + } + + for i := 0; i < len(entries)-1; i++ { + for j := i + 1; j < len(entries); j++ { + if entries[j].HitCount > entries[i].HitCount { + entries[i], entries[j] = entries[j], entries[i] + } + } + } + + if limit > len(entries) { + limit = len(entries) + } + + return entries[:limit] +} + +// moveToFront moves a signature to the front of the LRU order. +func (c *ErrorCache) moveToFront(signature string) { + c.removeFromOrder(signature) + c.order = append([]string{signature}, c.order...) +} + +// removeFromOrder removes a signature from the order slice. +func (c *ErrorCache) removeFromOrder(signature string) { + for i, s := range c.order { + if s == signature { + c.order = append(c.order[:i], c.order[i+1:]...) + break + } + } +} + +// evictOldest removes the least recently used entry. +func (c *ErrorCache) evictOldest() { + if len(c.order) == 0 { + return + } + + oldest := c.order[len(c.order)-1] + delete(c.cache, oldest) + c.order = c.order[:len(c.order)-1] +} diff --git a/internal/llm/cache_test.go b/internal/llm/cache_test.go new file mode 100644 index 0000000..1fcf370 --- /dev/null +++ b/internal/llm/cache_test.go @@ -0,0 +1,173 @@ +package llm + +import ( + "testing" + "time" +) + +func TestErrorCache_BasicOperations(t *testing.T) { + cache := NewErrorCache(10) + + if cache.Size() != 0 { + t.Errorf("expected size 0, got %d", cache.Size()) + } + + analysis := &CachedAnalysis{ + RootCauseNodes: []string{"missing dependency"}, + RemediationSteps: []string{"npm install"}, + Explanation: "Package not found", + Fix: "npm install express", + Confidence: 0.9, + } + cache.Put("sig-001", analysis) + + if cache.Size() != 1 { + t.Errorf("expected size 1, got %d", cache.Size()) + } + + retrieved := cache.Get("sig-001") + if retrieved == nil { + t.Fatal("expected to find cached entry") + } + if retrieved.Fix != "npm install express" { + t.Errorf("expected fix 'npm install express', got '%s'", retrieved.Fix) + } + if retrieved.HitCount != 2 { + t.Errorf("expected hit count 2, got %d", retrieved.HitCount) + } + + miss := cache.Get("sig-nonexistent") + if miss != nil { + t.Error("expected nil for cache miss") + } + + if !cache.Contains("sig-001") { + t.Error("contains should return true for existing key") + } + if cache.Contains("sig-nonexistent") { + t.Error("contains should return false for non-existing key") + } +} + +func TestErrorCache_LRUEviction(t *testing.T) { + cache := NewErrorCache(3) + + cache.Put("sig-1", &CachedAnalysis{Explanation: "first"}) + cache.Put("sig-2", &CachedAnalysis{Explanation: "second"}) + cache.Put("sig-3", &CachedAnalysis{Explanation: "third"}) + + if cache.Size() != 3 { + t.Errorf("expected size 3, got %d", cache.Size()) + } + + cache.Get("sig-1") + + cache.Put("sig-4", &CachedAnalysis{Explanation: "fourth"}) + + if cache.Size() != 3 { + t.Errorf("expected size still 3 after eviction, got %d", cache.Size()) + } + + if !cache.Contains("sig-1") { + t.Error("sig-1 should not be evicted (recently used)") + } + + if !cache.Contains("sig-4") { + t.Error("sig-4 should exist") + } +} + +func TestErrorCache_Stats(t *testing.T) { + cache := NewErrorCache(10) + + stats := cache.Stats() + if stats.Hits != 0 || stats.Misses != 0 { + t.Error("initial stats should be zero") + } + + cache.Put("sig-1", &CachedAnalysis{Explanation: "test"}) + cache.Get("sig-1") + cache.Get("sig-1") + cache.Get("sig-2") + + stats = cache.Stats() + if stats.Hits != 2 { + t.Errorf("expected 2 hits, got %d", stats.Hits) + } + if stats.Misses != 1 { + t.Errorf("expected 1 miss, got %d", stats.Misses) + } + if stats.HitRate < 0.66 || stats.HitRate > 0.67 { + t.Errorf("expected hit rate ~0.67, got %f", stats.HitRate) + } +} + +func TestErrorCache_Clear(t *testing.T) { + cache := NewErrorCache(10) + + cache.Put("sig-1", &CachedAnalysis{Explanation: "test1"}) + cache.Put("sig-2", &CachedAnalysis{Explanation: "test2"}) + cache.Get("sig-1") + + cache.Clear() + + if cache.Size() != 0 { + t.Errorf("expected size 0 after clear, got %d", cache.Size()) + } + stats := cache.Stats() + if stats.Hits != 0 || stats.Misses != 0 { + t.Error("stats should be reset after clear") + } +} + +func TestErrorCache_Delete(t *testing.T) { + cache := NewErrorCache(10) + + cache.Put("sig-1", &CachedAnalysis{Explanation: "test"}) + cache.Delete("sig-1") + + if cache.Contains("sig-1") { + t.Error("sig-1 should be deleted") + } + if cache.Size() != 0 { + t.Errorf("expected size 0, got %d", cache.Size()) + } +} + +func TestErrorCache_GetTopHits(t *testing.T) { + cache := NewErrorCache(10) + + cache.Put("sig-1", &CachedAnalysis{Explanation: "low hits"}) + cache.Put("sig-2", &CachedAnalysis{Explanation: "high hits"}) + cache.Put("sig-3", &CachedAnalysis{Explanation: "medium hits"}) + + cache.Get("sig-2") + cache.Get("sig-2") + cache.Get("sig-2") + + cache.Get("sig-3") + + topHits := cache.GetTopHits(2) + if len(topHits) != 2 { + t.Fatalf("expected 2 top hits, got %d", len(topHits)) + } + if topHits[0].Signature != "sig-2" { + t.Errorf("expected sig-2 to be top hit, got %s", topHits[0].Signature) + } +} + +func TestErrorCache_Timestamps(t *testing.T) { + cache := NewErrorCache(10) + + before := time.Now() + cache.Put("sig-1", &CachedAnalysis{Explanation: "test"}) + after := time.Now() + + entry := cache.Get("sig-1") + if entry.CreatedAt.Before(before) || entry.CreatedAt.After(after) { + t.Error("CreatedAt should be within test bounds") + } + if entry.LastHit.Before(entry.CreatedAt) { + t.Error("LastHit should be at or after CreatedAt") + } +} diff --git a/internal/llm/monitor_test.go b/internal/llm/monitor_test.go index 1e75053..ffb0296 100644 --- a/internal/llm/monitor_test.go +++ b/internal/llm/monitor_test.go @@ -8,7 +8,7 @@ import ( ) func TestAnalyzeLog_Parsing(t *testing.T) { - // 1. Mock Server to simulate LLM response + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // Verify request format (input format check) var req generateRequest @@ -18,7 +18,6 @@ func TestAnalyzeLog_Parsing(t *testing.T) { return } - // Send back valid JSON response (output format simulation) resp := generateResponse{ Response: `{"explanation": "Root cause is a missing env var", "fix": "export DB_URL=..."}`, Done: true, @@ -27,20 +26,17 @@ func TestAnalyzeLog_Parsing(t *testing.T) { })) defer server.Close() - // 2. Setup Client client := &Client{ baseURL: server.URL, model: "test-model", httpClient: server.Client(), } - // 3. Run Analysis result, err := client.AnalyzeLog("Error: Connection failed") if err != nil { t.Fatalf("AnalyzeLog failed: %v", err) } - // 4. Verify Result if result.Explanation != "Root cause is a missing env var" { t.Errorf("Expected explanation 'Root cause is a missing env var', got '%s'", result.Explanation) } @@ -50,10 +46,10 @@ func TestAnalyzeLog_Parsing(t *testing.T) { } func TestAnalyzeLog_MalformedJSON(t *testing.T) { - // Test how it handles "broken" output from LLM + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { resp := generateResponse{ - Response: `This is not JSON`, // raw text response case + Response: `This is not JSON`, Done: true, } json.NewEncoder(w).Encode(resp) diff --git a/internal/llm/ollama.go b/internal/llm/ollama.go index dd6047e..dcba0b7 100644 --- a/internal/llm/ollama.go +++ b/internal/llm/ollama.go @@ -350,3 +350,80 @@ COMMAND:`, goal, goal) return strings.TrimSpace(genResp.Response), nil } + +// ToolCallResult represents the result of a tool-aware LLM generation. +type ToolCallResult struct { + ToolName string `json:"tool_name"` + Parameters map[string]any `json:"parameters"` + Reasoning string `json:"reasoning,omitempty"` +} + +// GenerateWithTools calls the LLM with tool definitions and expects a tool call response. +func (c *Client) GenerateWithTools(prompt string, toolSchemas string) (*ToolCallResult, error) { + systemPrompt := fmt.Sprintf(`You are an AI assistant with access to tools. Based on the user's request, determine which tool to use and with what parameters. + +AVAILABLE TOOLS: +%s + +RULES: +1. Analyze the user's request carefully +2. Select the most appropriate tool +3. Determine the correct parameters +4. Respond with ONLY valid JSON in this exact format: +{ + "tool_name": "name_of_tool", + "parameters": { + "param1": "value1", + "param2": "value2" + }, + "reasoning": "brief explanation of why this tool was chosen" +} + +Do NOT include any text outside the JSON object.`, toolSchemas) + + fullPrompt := fmt.Sprintf(`%s + +USER REQUEST: %s + +JSON RESPONSE:`, systemPrompt, prompt) + + req := generateRequest{ + Model: c.model, + Prompt: fullPrompt, + Stream: false, + Format: "json", + } + + if os.Getenv("DEV_CLI_OLLAMA_UNLOAD") == "true" { + req.KeepAlive = "0m" + } + + reqBody, err := json.Marshal(req) + if err != nil { + return nil, fmt.Errorf("marshal request: %w", err) + } + + resp, err := c.httpClient.Post(c.baseURL+"/api/generate", "application/json", bytes.NewReader(reqBody)) + if err != nil { + return nil, fmt.Errorf("call Ollama: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("ollama status %d: %s", resp.StatusCode, string(body)) + } + + var genResp generateResponse + if err := json.NewDecoder(resp.Body).Decode(&genResp); err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + + var result ToolCallResult + responseText := strings.TrimSpace(genResp.Response) + if err := json.Unmarshal([]byte(responseText), &result); err != nil { + return nil, fmt.Errorf("parse tool call: %w (response: %s)", err, responseText) + } + + return &result, nil +} diff --git a/internal/llm/perplexity_test.go b/internal/llm/perplexity_test.go index 5f58a25..2990262 100644 --- a/internal/llm/perplexity_test.go +++ b/internal/llm/perplexity_test.go @@ -8,16 +8,14 @@ import ( ) func TestPerplexityConfig(t *testing.T) { - // Set up environment variables + os.Setenv("DEV_CLI_PERPLEXITY_KEY", "test-key") os.Setenv("DEV_CLI_PERPLEXITY_MODEL", "sonar-pro") defer os.Unsetenv("DEV_CLI_PERPLEXITY_KEY") defer os.Unsetenv("DEV_CLI_PERPLEXITY_MODEL") - // Load config cfg := config.Load() - // Verify config loading if cfg.PerplexityKey != "test-key" { t.Errorf("expected PerplexityKey to be 'test-key', got '%s'", cfg.PerplexityKey) } @@ -25,13 +23,11 @@ func TestPerplexityConfig(t *testing.T) { t.Errorf("expected PerplexityModel to be 'sonar-pro', got '%s'", cfg.PerplexityModel) } - // Create client client := NewPerplexityClient(cfg) if client == nil { t.Fatal("expected client to be non-nil") } - // We can't access private fields directly in test unless we are in the same package (which we are: package llm) if client.apiKey != "test-key" { t.Errorf("expected client.apiKey to be 'test-key', got '%s'", client.apiKey) } @@ -43,11 +39,7 @@ func TestPerplexityConfig(t *testing.T) { func TestPerplexityDefaultConfig(t *testing.T) { os.Unsetenv("DEV_CLI_PERPLEXITY_KEY") os.Unsetenv("DEV_CLI_PERPLEXITY_MODEL") - // Start with empty env for these vars - // but Load() might pick up other env vars if they are set in the system, so we should be careful. - // However, for defaults check, we just want to see if 'sonar' is default if not set. - // We need to set key to get a client os.Setenv("PERPLEXITY_API_KEY", "legacy-key") defer os.Unsetenv("PERPLEXITY_API_KEY") diff --git a/internal/llm/sanitizer_test.go b/internal/llm/sanitizer_test.go new file mode 100644 index 0000000..5e0b475 --- /dev/null +++ b/internal/llm/sanitizer_test.go @@ -0,0 +1,349 @@ +package llm + +import ( + "strings" + "testing" +) + +// NOTE: All "secrets" in this file are INTENTIONALLY FAKE test patterns. +// They are designed to test the sanitizer and are NOT real credentials. + +func TestSanitize_APIKey(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + {"api_key with equals", "api_key=sk-abc123def456ghi789jkl", "api_key=[REDACTED_API_KEY]"}, + {"apikey no separator", "apikey=abcdefghijklmnopqrstuvwxyz", "apikey=[REDACTED_API_KEY]"}, + {"API-KEY quoted", `API-KEY="my_super_secret_key_123456"`, `API-KEY=[REDACTED_API_KEY]`}, + {"mixed case", "Api_Key:verylongsecretkeyvalue123456", "Api_Key=[REDACTED_API_KEY]"}, + } + + sanitizer := DefaultSanitizer() + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := sanitizer.Sanitize(tt.input) + if got != tt.want { + t.Errorf("Sanitize(%q) = %q, want %q", tt.input, got, tt.want) + } + }) + } +} + +func TestSanitize_BearerToken(t *testing.T) { + // Using obviously fake bearer tokens for testing + tests := []struct { + input string + want string + }{ + {"Authorization: Bearer FAKE-test-token.for.testing", "Authorization: Bearer [REDACTED_TOKEN]"}, + {"bearer FAKE-access-token-12345", "Bearer [REDACTED_TOKEN]"}, + } + + sanitizer := DefaultSanitizer() + for _, tt := range tests { + got := sanitizer.Sanitize(tt.input) + if got != tt.want { + t.Errorf("Sanitize(%q) = %q, want %q", tt.input, got, tt.want) + } + } +} + +func TestSanitize_AWSCredentials(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + // Using AWS example keys that are clearly fake but match regex format + {"AWS Access Key", "AKIAFAKEEXAMPLEKEY12", "[REDACTED_AWS_KEY]"}, + {"AWS Secret", "aws_secret_access_key=FAKEexampleSECRETkeyVALUE123456789abcdef", "aws_secret_access_key=[REDACTED_AWS_SECRET]"}, + {"Asia prefix", "ASIAFAKEACCESSKEY123", "[REDACTED_AWS_KEY]"}, + } + + sanitizer := DefaultSanitizer() + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := sanitizer.Sanitize(tt.input) + if got != tt.want { + t.Errorf("Sanitize(%q) = %q, want %q", tt.input, got, tt.want) + } + }) + } +} + +func TestSanitize_PrivateKey(t *testing.T) { + // Using obviously fake private key for testing + input := `-----BEGIN PRIVATE KEY----- +FAKE-TEST-KEY-NOT-REAL-AAAABBBBCCCCDDDD1111222233334444 +-----END PRIVATE KEY-----` + + sanitizer := DefaultSanitizer() + got := sanitizer.Sanitize(input) + + if got != "[REDACTED_PRIVATE_KEY]" { + t.Errorf("Expected private key to be redacted, got: %s", got) + } +} + +func TestSanitize_RSAPrivateKey(t *testing.T) { + // Using obviously fake RSA key for testing + input := `-----BEGIN RSA PRIVATE KEY----- +FAKE-RSA-TEST-KEY-NOT-REAL-AAAABBBBCCCC +-----END RSA PRIVATE KEY-----` + + sanitizer := DefaultSanitizer() + got := sanitizer.Sanitize(input) + + if got != "[REDACTED_PRIVATE_KEY]" { + t.Errorf("Expected RSA private key to be redacted, got: %s", got) + } +} + +func TestSanitize_GitHubToken(t *testing.T) { + // Using fake GitHub PAT pattern for testing (ghp_ + 36 alphanumeric chars) + input := "ghp_FAKETOKENabcdefghij1234567890abcdefX" + sanitizer := DefaultSanitizer() + got := sanitizer.Sanitize(input) + + if !strings.Contains(got, "[REDACTED_GITHUB_TOKEN]") { + t.Errorf("Expected GitHub token to be redacted, got: %s", got) + } +} + +func TestSanitize_DatabaseURL(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + // Using obviously fake database passwords for testing + {"PostgreSQL", "postgresql://user:FAKEPASS@localhost:5432/db", "postgresql://[user]:[REDACTED]@localhost:5432/db"}, + {"MongoDB", "mongodb://admin:FAKEPASS@cluster.mongodb.net/mydb", "mongodb://[user]:[REDACTED]@cluster.mongodb.net/mydb"}, + {"MySQL", "mysql://root:FAKEPASS@127.0.0.1:3306/app", "mysql://[user]:[REDACTED]@127.0.0.1:3306/app"}, + {"Redis", "redis://default:FAKEPASS@redis.io:6379", "redis://[user]:[REDACTED]@redis.io:6379"}, + } + + sanitizer := DefaultSanitizer() + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := sanitizer.Sanitize(tt.input) + if got != tt.want { + t.Errorf("Sanitize(%q) = %q, want %q", tt.input, got, tt.want) + } + }) + } +} + +func TestSanitize_JWTToken(t *testing.T) { + // Using obviously fake JWT pattern for testing (not a valid JWT) + input := "eyJGQUtFIjoiVEVTVCJ9.eyJGQUtFIjoiREFUQSJ9.FAKESIGNATURE" + sanitizer := DefaultSanitizer() + got := sanitizer.Sanitize(input) + + if !strings.Contains(got, "[REDACTED_JWT]") { + t.Errorf("Expected JWT to be redacted, got: %s", got) + } +} + +func TestSanitize_SlackToken(t *testing.T) { + tests := []struct { + input string + }{ + // Using clearly fake placeholder values to avoid secret scanning + {"xoxb-FAKE-PLACEHOLDER-TOKEN"}, + {"xoxp-FAKE-TEST-VALUE"}, + {"xoxa-FAKE-TOKEN"}, + } + + sanitizer := DefaultSanitizer() + for _, tt := range tests { + got := sanitizer.Sanitize(tt.input) + if !strings.Contains(got, "[REDACTED_SLACK_TOKEN]") { + t.Errorf("Expected Slack token to be redacted, got: %s", got) + } + } +} + +func TestSanitize_Password(t *testing.T) { + tests := []struct { + input string + want string + }{ + {"password=mysecret123", "password=[REDACTED]"}, + {"secret:verysecretvalue", "secret=[REDACTED]"}, + {"TOKEN=abc123xyz789", "TOKEN=[REDACTED]"}, + } + + sanitizer := DefaultSanitizer() + for _, tt := range tests { + got := sanitizer.Sanitize(tt.input) + if got != tt.want { + t.Errorf("Sanitize(%q) = %q, want %q", tt.input, got, tt.want) + } + } +} + +func TestSanitizeWithReport(t *testing.T) { + input := "api_key=supersecretapikey12345 and password=secret123" + sanitizer := DefaultSanitizer() + + sanitized, found := sanitizer.SanitizeWithReport(input) + + if len(found) == 0 { + t.Error("Expected to find secrets in report") + } + if strings.Contains(sanitized, "supersecretapikey12345") { + t.Error("API key should be redacted") + } + if strings.Contains(sanitized, "secret123") { + t.Error("Password should be redacted") + } +} + +func TestContainsSecrets(t *testing.T) { + sanitizer := DefaultSanitizer() + + tests := []struct { + input string + want bool + }{ + {"api_key=secretkey12345678901234", true}, + {"Bearer mytoken123", true}, + {"Just some normal text", false}, + {"password=secretpassword", true}, + {"public data only", false}, + } + + for _, tt := range tests { + got := sanitizer.ContainsSecrets(tt.input) + if got != tt.want { + t.Errorf("ContainsSecrets(%q) = %v, want %v", tt.input, got, tt.want) + } + } +} + +func TestTruncateForLLM(t *testing.T) { + tests := []struct { + name string + input string + maxLen int + check func(string) bool + }{ + { + name: "short input unchanged", + input: "short", + maxLen: 100, + check: func(s string) bool { return s == "short" }, + }, + { + name: "long input truncated", + input: strings.Repeat("a", 1000), + maxLen: 100, + check: func(s string) bool { return len(s) <= 100 && strings.Contains(s, "[truncated]") }, + }, + { + name: "preserves start and end", + input: "START" + strings.Repeat("x", 500) + "END", + maxLen: 100, + check: func(s string) bool { return strings.HasPrefix(s, "START") && strings.HasSuffix(s, "END") }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := TruncateForLLM(tt.input, tt.maxLen) + if !tt.check(got) { + t.Errorf("TruncateForLLM failed check: got %q (len=%d)", got, len(got)) + } + }) + } +} + +func TestPrepareForLLM(t *testing.T) { + + input := " api_key=secretkey12345678901234 " + strings.Repeat("x", 200) + " " + + got := PrepareForLLM(input, 100) + + if strings.Contains(got, "secretkey12345678901234") { + t.Error("API key should be redacted") + } + + if len(got) > 100 { + t.Errorf("Output should be <= 100 chars, got %d", len(got)) + } + + if strings.HasPrefix(got, " ") || strings.HasSuffix(got, " ") { + t.Error("Output should be trimmed") + } +} + +func TestMaskEnvVars(t *testing.T) { + tests := []struct { + input string + want string + }{ + {"export API_KEY=mysecret", "export API_KEY=[REDACTED]"}, + {"SECRET=topsecret", "SECRET=[REDACTED]"}, + {"PASSWORD=pass123", "PASSWORD=[REDACTED]"}, + } + + for _, tt := range tests { + got := MaskEnvVars(tt.input) + if got != tt.want { + t.Errorf("MaskEnvVars(%q) = %q, want %q", tt.input, got, tt.want) + } + } +} + +func TestAddPattern(t *testing.T) { + sanitizer := DefaultSanitizer() + initialCount := sanitizer.PatternCount() + + err := sanitizer.AddPattern("Custom", `custom-secret-\d+`, "[REDACTED_CUSTOM]") + if err != nil { + t.Fatalf("AddPattern failed: %v", err) + } + + if sanitizer.PatternCount() != initialCount+1 { + t.Error("Pattern count should increase by 1") + } + + got := sanitizer.Sanitize("Found custom-secret-12345 in logs") + if !strings.Contains(got, "[REDACTED_CUSTOM]") { + t.Errorf("Custom pattern should match, got: %s", got) + } +} + +func TestAddPattern_InvalidRegex(t *testing.T) { + sanitizer := DefaultSanitizer() + + err := sanitizer.AddPattern("Invalid", "[invalid", "replacement") + if err == nil { + t.Error("Expected error for invalid regex") + } +} + +func TestSanitizeOutput(t *testing.T) { + input := "api_key=secretkey12345678901234 export SECRET=mysecret" + + got := SanitizeOutput(input) + + if strings.Contains(got, "secretkey12345678901234") { + t.Error("API key should be redacted") + } + if strings.Contains(got, "mysecret") { + t.Error("Exported secret should be redacted") + } +} + +func TestSanitizeForLLM(t *testing.T) { + input := "password=mysecret123" + got := SanitizeForLLM(input) + + if strings.Contains(got, "mysecret123") { + t.Error("Password should be redacted by global sanitizer") + } +} diff --git a/internal/pipeline/events.go b/internal/pipeline/events.go index a8ae4ba..40c4ea2 100644 --- a/internal/pipeline/events.go +++ b/internal/pipeline/events.go @@ -25,6 +25,26 @@ const ( EventSystemAlert EventType = "system.alert" EventSystemStats EventType = "system.stats" + + // Workflow events + EventWorkflowStart EventType = "workflow.start" + EventWorkflowStep EventType = "workflow.step" + EventWorkflowCheckpoint EventType = "workflow.checkpoint" + EventWorkflowComplete EventType = "workflow.complete" + EventWorkflowRollback EventType = "workflow.rollback" + + // RCA (Root Cause Analysis) events + EventRCAStart EventType = "rca.start" + EventRCANodeFound EventType = "rca.node_found" + EventRCAComplete EventType = "rca.complete" + EventRCACacheHit EventType = "rca.cache_hit" + + // Remediation events + EventRemediationPending EventType = "remediation.pending" + EventRemediationApproved EventType = "remediation.approved" + EventRemediationExecuted EventType = "remediation.executed" + EventRemediationRolledBack EventType = "remediation.rollback" + EventRemediationSkipped EventType = "remediation.skipped" ) type Event struct { diff --git a/internal/pipeline/events_test.go b/internal/pipeline/events_test.go new file mode 100644 index 0000000..00d4a81 --- /dev/null +++ b/internal/pipeline/events_test.go @@ -0,0 +1,208 @@ +package pipeline + +import ( + "sync" + "sync/atomic" + "testing" + "time" +) + +func TestEventBus_Subscribe(t *testing.T) { + bus := NewEventBus() + received := false + + bus.Subscribe(EventCommandStart, func(e Event) { + received = true + }) + + bus.Publish(Event{ + Type: EventCommandStart, + Timestamp: time.Now(), + Source: "test", + }) + + if !received { + t.Error("handler should have received the event") + } +} + +func TestEventBus_SubscribeAll(t *testing.T) { + bus := NewEventBus() + count := 0 + + bus.SubscribeAll(func(e Event) { + count++ + }) + + bus.Publish(Event{Type: EventCommandStart}) + bus.Publish(Event{Type: EventCommandComplete}) + bus.Publish(Event{Type: EventContainerLog}) + + if count != 3 { + t.Errorf("expected 3 events, got %d", count) + } +} + +func TestEventBus_Publish_ToCorrectHandlers(t *testing.T) { + bus := NewEventBus() + startCount := 0 + completeCount := 0 + + bus.Subscribe(EventCommandStart, func(e Event) { + startCount++ + }) + bus.Subscribe(EventCommandComplete, func(e Event) { + completeCount++ + }) + + bus.Publish(Event{Type: EventCommandStart}) + bus.Publish(Event{Type: EventCommandStart}) + bus.Publish(Event{Type: EventCommandComplete}) + + if startCount != 2 { + t.Errorf("expected 2 start events, got %d", startCount) + } + if completeCount != 1 { + t.Errorf("expected 1 complete event, got %d", completeCount) + } +} + +func TestEventBus_PublishMultipleHandlers(t *testing.T) { + bus := NewEventBus() + handler1Called := false + handler2Called := false + + bus.Subscribe(EventCommandError, func(e Event) { + handler1Called = true + }) + bus.Subscribe(EventCommandError, func(e Event) { + handler2Called = true + }) + + bus.Publish(Event{Type: EventCommandError}) + + if !handler1Called || !handler2Called { + t.Error("both handlers should be called") + } +} + +func TestEventBus_RecentEvents(t *testing.T) { + bus := NewEventBus() + + for i := 0; i < 5; i++ { + bus.Publish(Event{ + Type: EventCommandOutput, + BlockID: string(rune('a' + i)), + }) + } + + recent := bus.RecentEvents(3) + if len(recent) != 3 { + t.Errorf("expected 3 recent events, got %d", len(recent)) + } + + if recent[0].BlockID != "c" || recent[2].BlockID != "e" { + t.Error("should return most recent events in order") + } +} + +func TestEventBus_RecentByType(t *testing.T) { + bus := NewEventBus() + + bus.Publish(Event{Type: EventCommandStart, BlockID: "1"}) + bus.Publish(Event{Type: EventCommandError, BlockID: "2"}) + bus.Publish(Event{Type: EventCommandStart, BlockID: "3"}) + bus.Publish(Event{Type: EventCommandComplete, BlockID: "4"}) + + results := bus.RecentByType(EventCommandStart, 10) + if len(results) != 2 { + t.Errorf("expected 2 start events, got %d", len(results)) + } +} + +func TestEventBus_HistoryLimit(t *testing.T) { + bus := NewEventBus() + bus.maxHistory = 5 + + for i := 0; i < 10; i++ { + bus.Publish(Event{Type: EventCommandOutput, BlockID: string(rune('0' + i))}) + } + + all := bus.RecentEvents(100) + if len(all) != 5 { + t.Errorf("expected 5 events (maxHistory), got %d", len(all)) + } + + if all[0].BlockID != "5" { + t.Errorf("oldest event should be '5', got '%s'", all[0].BlockID) + } +} + +func TestEventBus_ConcurrentPublish(t *testing.T) { + bus := NewEventBus() + var count int64 + var wg sync.WaitGroup + + bus.SubscribeAll(func(e Event) { + atomic.AddInt64(&count, 1) + }) + + for i := 0; i < 100; i++ { + wg.Add(1) + go func() { + defer wg.Done() + bus.Publish(Event{Type: EventSystemStats}) + }() + } + + wg.Wait() + + if count != 100 { + t.Errorf("expected 100 events handled, got %d", count) + } +} + +func TestEventBus_EventData(t *testing.T) { + bus := NewEventBus() + var receivedData interface{} + + bus.Subscribe(EventAISuggestion, func(e Event) { + receivedData = e.Data + }) + + testData := map[string]string{"suggestion": "try npm install"} + bus.Publish(Event{ + Type: EventAISuggestion, + Data: testData, + }) + + if receivedData == nil { + t.Error("event data should be passed to handler") + } + + data, ok := receivedData.(map[string]string) + if !ok || data["suggestion"] != "try npm install" { + t.Error("event data should match what was published") + } +} + +func TestEventBus_Timestamp(t *testing.T) { + bus := NewEventBus() + + before := time.Now() + bus.Publish(Event{ + Type: EventSystemAlert, + Timestamp: time.Now(), + }) + after := time.Now() + + recent := bus.RecentEvents(1) + if len(recent) != 1 { + t.Fatal("expected 1 event") + } + + ts := recent[0].Timestamp + if ts.Before(before) || ts.After(after) { + t.Error("event timestamp should be preserved") + } +} diff --git a/internal/pipeline/graph_pager.go b/internal/pipeline/graph_pager.go new file mode 100644 index 0000000..87c2ba1 --- /dev/null +++ b/internal/pipeline/graph_pager.go @@ -0,0 +1,212 @@ +package pipeline + +import ( + "context" + "time" +) + +// CausalNode represents a node in the causal dependency graph. +type CausalNode struct { + ID string `json:"id"` + Type string `json:"type"` // "error", "service", "dependency", "config" + Name string `json:"name"` + Description string `json:"description"` + Level int `json:"level"` // Depth in causal chain (0 = root cause) + Metadata map[string]string `json:"metadata,omitempty"` + Children []string `json:"children,omitempty"` // Child node IDs +} + +// GraphPager handles paginated traversal of causal graphs. +// Useful for large dependency graphs during RCA. +type GraphPager struct { + PageSize int + MaxDepth int + nodes map[string]*CausalNode +} + +// NewGraphPager creates a pager with configurable page size and max depth. +func NewGraphPager(pageSize, maxDepth int) *GraphPager { + if pageSize <= 0 { + pageSize = 20 + } + if maxDepth <= 0 { + maxDepth = 10 + } + return &GraphPager{ + PageSize: pageSize, + MaxDepth: maxDepth, + nodes: make(map[string]*CausalNode), + } +} + +// Page represents a subset of graph nodes for analysis. +type Page struct { + Nodes []CausalNode `json:"nodes"` + TotalNodes int `json:"total_nodes"` + HasMore bool `json:"has_more"` + NextCursor string `json:"next_cursor,omitempty"` + Level int `json:"level"` +} + +// AddNode adds a node to the graph. +func (p *GraphPager) AddNode(node CausalNode) { + p.nodes[node.ID] = &node +} + +// AddNodes adds multiple nodes to the graph. +func (p *GraphPager) AddNodes(nodes []CausalNode) { + for _, node := range nodes { + p.AddNode(node) + } +} + +// GetNode retrieves a node by ID. +func (p *GraphPager) GetNode(id string) *CausalNode { + return p.nodes[id] +} + +// NodeCount returns the total number of nodes. +func (p *GraphPager) NodeCount() int { + return len(p.nodes) +} + +// Clear removes all nodes from the graph. +func (p *GraphPager) Clear() { + p.nodes = make(map[string]*CausalNode) +} + +// GetPage returns a page of nodes at or below a specific level. +func (p *GraphPager) GetPage(ctx context.Context, level int, cursor string) (*Page, error) { + + nodesAtLevel := make([]CausalNode, 0) + for _, node := range p.nodes { + if node.Level == level { + nodesAtLevel = append(nodesAtLevel, *node) + } + } + + startIdx := 0 + if cursor != "" { + for i, n := range nodesAtLevel { + if n.ID == cursor { + startIdx = i + break + } + } + } + + endIdx := startIdx + p.PageSize + if endIdx > len(nodesAtLevel) { + endIdx = len(nodesAtLevel) + } + + pageNodes := nodesAtLevel[startIdx:endIdx] + + page := &Page{ + Nodes: pageNodes, + TotalNodes: len(nodesAtLevel), + HasMore: endIdx < len(nodesAtLevel), + Level: level, + } + + if page.HasMore && len(pageNodes) > 0 { + page.NextCursor = pageNodes[len(pageNodes)-1].ID + } + + return page, nil +} + +// GetNodesAtLevel returns all nodes at a specific depth level. +func (p *GraphPager) GetNodesAtLevel(level int) []CausalNode { + result := make([]CausalNode, 0) + for _, node := range p.nodes { + if node.Level == level { + result = append(result, *node) + } + } + return result +} + +// GetRootCauses returns all level-0 (root cause) nodes. +func (p *GraphPager) GetRootCauses() []CausalNode { + return p.GetNodesAtLevel(0) +} + +// GetChildren returns all direct children of a node. +func (p *GraphPager) GetChildren(nodeID string) []CausalNode { + parent := p.nodes[nodeID] + if parent == nil { + return nil + } + + children := make([]CausalNode, 0, len(parent.Children)) + for _, childID := range parent.Children { + if child := p.nodes[childID]; child != nil { + children = append(children, *child) + } + } + return children +} + +// TraverseBFS performs breadth-first traversal from a starting node. +// Calls the visitor function for each node, stopping if it returns false. +func (p *GraphPager) TraverseBFS(ctx context.Context, startID string, visitor func(CausalNode) bool) error { + visited := make(map[string]bool) + queue := []string{startID} + + for len(queue) > 0 { + select { + case <-ctx.Done(): + return ctx.Err() + default: + } + + nodeID := queue[0] + queue = queue[1:] + + if visited[nodeID] { + continue + } + visited[nodeID] = true + + node := p.nodes[nodeID] + if node == nil { + continue + } + + if !visitor(*node) { + return nil + } + + queue = append(queue, node.Children...) + } + + return nil +} + +// BuildFromFailure creates a causal graph from a failure block. +// This is a starting point - actual causal analysis would involve LLM. +func (p *GraphPager) BuildFromFailure(failure *Block) { + + root := CausalNode{ + ID: failure.ID, + Type: "error", + Name: failure.Command, + Description: truncateString(failure.Output, 200), + Level: 0, + Metadata: map[string]string{ + "exit_code": string(rune('0' + failure.ExitCode)), + "working_dir": failure.WorkingDir, + "timestamp": failure.Timestamp.Format(time.RFC3339), + }, + } + p.AddNode(root) +} + +// truncateString truncates a string to maxLen characters. +func truncateString(s string, maxLen int) string { + if len(s) <= maxLen { + return s + } + return s[:maxLen] + "..." +} diff --git a/internal/pipeline/state_test.go b/internal/pipeline/state_test.go new file mode 100644 index 0000000..2a6f2c7 --- /dev/null +++ b/internal/pipeline/state_test.go @@ -0,0 +1,325 @@ +package pipeline + +import ( + "sync" + "testing" + "time" + + "dev-cli/internal/infra" +) + +func TestStateStore_AddBlock(t *testing.T) { + store := NewStateStore() + + block := Block{ + ID: "block-1", + Type: BlockTypeCommand, + Timestamp: time.Now(), + Command: "ls -la", + } + store.AddBlock(block) + + if len(store.Blocks) != 1 { + t.Errorf("expected 1 block, got %d", len(store.Blocks)) + } + if store.SelectedIdx != 0 { + t.Errorf("expected SelectedIdx 0, got %d", store.SelectedIdx) + } +} + +func TestStateStore_GetBlock(t *testing.T) { + store := NewStateStore() + + store.AddBlock(Block{ID: "block-1", Command: "cmd1"}) + store.AddBlock(Block{ID: "block-2", Command: "cmd2"}) + + block := store.GetBlock("block-1") + if block == nil { + t.Fatal("expected to find block-1") + } + if block.Command != "cmd1" { + t.Errorf("expected cmd1, got %s", block.Command) + } + + notFound := store.GetBlock("nonexistent") + if notFound != nil { + t.Error("should return nil for nonexistent block") + } +} + +func TestStateStore_GetRecentBlocks(t *testing.T) { + store := NewStateStore() + + for i := 0; i < 5; i++ { + store.AddBlock(Block{ID: string(rune('a' + i))}) + } + + recent := store.GetRecentBlocks(3) + if len(recent) != 3 { + t.Errorf("expected 3 blocks, got %d", len(recent)) + } + if recent[0].ID != "c" || recent[2].ID != "e" { + t.Error("should return most recent blocks") + } +} + +func TestStateStore_GetBlocks(t *testing.T) { + store := NewStateStore() + + store.AddBlock(Block{ID: "1"}) + store.AddBlock(Block{ID: "2"}) + + blocks := store.GetBlocks() + if len(blocks) != 2 { + t.Errorf("expected 2 blocks, got %d", len(blocks)) + } + + blocks[0].ID = "modified" + if store.Blocks[0].ID == "modified" { + t.Error("GetBlocks should return a copy") + } +} + +func TestStateStore_UpdateBlock(t *testing.T) { + store := NewStateStore() + + store.AddBlock(Block{ID: "block-1", Output: ""}) + + store.UpdateBlock("block-1", func(b *Block) { + b.Output = "new output" + b.ExitCode = 1 + }) + + block := store.GetBlock("block-1") + if block.Output != "new output" { + t.Errorf("expected 'new output', got '%s'", block.Output) + } + if block.ExitCode != 1 { + t.Errorf("expected exit code 1, got %d", block.ExitCode) + } +} + +func TestStateStore_MaxBlocks(t *testing.T) { + store := NewStateStore() + store.MaxBlocks = 5 + + for i := 0; i < 10; i++ { + store.AddBlock(Block{ID: string(rune('0' + i))}) + } + + if len(store.Blocks) != 5 { + t.Errorf("expected 5 blocks (MaxBlocks), got %d", len(store.Blocks)) + } + + if store.GetBlock("0") != nil { + t.Error("block '0' should have been evicted") + } + if store.GetBlock("9") == nil { + t.Error("block '9' should still exist") + } +} + +func TestStateStore_AddSuggestion(t *testing.T) { + store := NewStateStore() + + store.AddSuggestion(Suggestion{ + ForBlockID: "block-1", + Title: "Try this", + Command: "npm install", + Explanation: "Missing dependencies", + }) + + if len(store.Suggestions) != 1 { + t.Errorf("expected 1 suggestion, got %d", len(store.Suggestions)) + } +} + +func TestStateStore_GetSuggestionsForBlock(t *testing.T) { + store := NewStateStore() + + store.AddSuggestion(Suggestion{ForBlockID: "block-1", Title: "Sug1"}) + store.AddSuggestion(Suggestion{ForBlockID: "block-2", Title: "Sug2"}) + store.AddSuggestion(Suggestion{ForBlockID: "block-1", Title: "Sug3"}) + + sugs := store.GetSuggestionsForBlock("block-1") + if len(sugs) != 2 { + t.Errorf("expected 2 suggestions for block-1, got %d", len(sugs)) + } +} + +func TestStateStore_SuggestionLimit(t *testing.T) { + store := NewStateStore() + + for i := 0; i < 15; i++ { + store.AddSuggestion(Suggestion{ + ForBlockID: string(rune('a' + i)), + Title: "Suggestion", + }) + } + + if len(store.Suggestions) != 10 { + t.Errorf("expected 10 suggestions (limit), got %d", len(store.Suggestions)) + } +} + +func TestStateStore_ClearBlocks(t *testing.T) { + store := NewStateStore() + + store.AddBlock(Block{ID: "1"}) + store.AddBlock(Block{ID: "2"}) + + store.ClearBlocks() + + if len(store.Blocks) != 0 { + t.Errorf("expected 0 blocks after clear, got %d", len(store.Blocks)) + } + if store.SelectedIdx != -1 { + t.Errorf("expected SelectedIdx -1 after clear, got %d", store.SelectedIdx) + } +} + +func TestStateStore_LastError(t *testing.T) { + store := NewStateStore() + + store.AddBlock(Block{ID: "1", ExitCode: 0}) + if store.LastError != nil { + t.Error("LastError should be nil for successful command") + } + + store.AddBlock(Block{ID: "2", ExitCode: 1}) + if store.LastError == nil { + t.Fatal("LastError should be set for failed command") + } + if store.LastError.ID != "2" { + t.Error("LastError should point to the failed block") + } +} + +func TestStateStore_SetDockerHealth(t *testing.T) { + store := NewStateStore() + + health := infra.DockerHealth{ + Available: true, + Containers: []infra.ContainerInfo{ + {ID: "abc123", Name: "test", State: "running"}, + }, + } + + store.SetDockerHealth(health) + + if !store.DockerHealth.Available { + t.Error("DockerHealth.Available should be true") + } + if len(store.DockerHealth.Containers) != 1 { + t.Error("DockerHealth.Containers should have 1 container") + } +} + +func TestStateStore_SetGPUStats(t *testing.T) { + store := NewStateStore() + + stats := infra.GPUStats{ + Available: true, + } + + store.SetGPUStats(stats) + + if !store.GPUStats.Available { + t.Error("GPUStats.Available should be true") + } +} + +func TestStateStore_SetStarshipLine(t *testing.T) { + store := NewStateStore() + + store.SetStarshipLine(" on main [!?]") + + if store.StarshipLine != " on main [!?]" { + t.Errorf("StarshipLine not set correctly: %s", store.StarshipLine) + } +} + +func TestStateStore_SetCwd(t *testing.T) { + store := NewStateStore() + + store.SetCwd("/home/user/project") + + if store.Cwd != "/home/user/project" { + t.Errorf("Cwd not set correctly: %s", store.Cwd) + } +} + +func TestStateStore_GetContext(t *testing.T) { + store := NewStateStore() + + store.SetCwd("/test") + store.AddBlock(Block{ID: "1"}) + store.AddBlock(Block{ID: "2"}) + store.SetDockerHealth(infra.DockerHealth{ + Containers: []infra.ContainerInfo{{ID: "c1"}}, + }) + + ctx := store.GetContext() + + if ctx["cwd"] != "/test" { + t.Error("context should include cwd") + } + if ctx["recent_commands"] != 2 { + t.Error("context should include recent_commands count") + } + if ctx["container_count"] != 1 { + t.Error("context should include container_count") + } +} + +func TestStateStore_ConcurrentAccess(t *testing.T) { + store := NewStateStore() + var wg sync.WaitGroup + + for i := 0; i < 50; i++ { + wg.Add(1) + go func(id int) { + defer wg.Done() + store.AddBlock(Block{ID: string(rune('a' + (id % 26)))}) + }(i) + } + + for i := 0; i < 50; i++ { + wg.Add(1) + go func() { + defer wg.Done() + _ = store.GetBlocks() + _ = store.GetRecentBlocks(5) + }() + } + + wg.Wait() + + if len(store.Blocks) == 0 { + t.Error("expected some blocks after concurrent access") + } +} + +func TestRebuildIndex(t *testing.T) { + store := NewStateStore() + store.MaxBlocks = 3 + + store.AddBlock(Block{ID: "a"}) + store.AddBlock(Block{ID: "b"}) + store.AddBlock(Block{ID: "c"}) + store.AddBlock(Block{ID: "d"}) + + if store.GetBlock("a") != nil { + t.Error("block 'a' should have been evicted") + } + if store.GetBlock("d") == nil { + t.Error("block 'd' should exist") + } + + for id, idx := range store.blockIndex { + if store.Blocks[idx].ID != id { + t.Errorf("index mismatch: index[%s]=%d but Blocks[%d].ID=%s", + id, idx, idx, store.Blocks[idx].ID) + } + } +} diff --git a/internal/storage/db.go b/internal/storage/db.go index b1d66c9..7bd809e 100644 --- a/internal/storage/db.go +++ b/internal/storage/db.go @@ -58,14 +58,92 @@ func migrate(db *sql.DB) error { duration_ms INTEGER, directory TEXT, session_id TEXT, - details TEXT + details TEXT, + resolution TEXT ); CREATE INDEX IF NOT EXISTS idx_history_timestamp ON history(timestamp); CREATE INDEX IF NOT EXISTS idx_history_exit_code ON history(exit_code); CREATE INDEX IF NOT EXISTS idx_history_session ON history(session_id); + + -- Workflow automation tables + CREATE TABLE IF NOT EXISTS workflow_runs ( + id TEXT PRIMARY KEY, + workflow_id TEXT NOT NULL, + workflow_name TEXT, + status TEXT NOT NULL, + current_step INTEGER DEFAULT 0, + started_at DATETIME, + updated_at DATETIME, + completed_at DATETIME, + error TEXT + ); + + CREATE TABLE IF NOT EXISTS workflow_step_results ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + run_id TEXT NOT NULL, + step_id TEXT NOT NULL, + status TEXT NOT NULL, + exit_code INTEGER, + output TEXT, + error TEXT, + retries INTEGER DEFAULT 0, + started_at DATETIME, + completed_at DATETIME, + duration_ms INTEGER, + FOREIGN KEY (run_id) REFERENCES workflow_runs(id) + ); + + CREATE INDEX IF NOT EXISTS idx_workflow_runs_status ON workflow_runs(status); + CREATE INDEX IF NOT EXISTS idx_step_results_run_id ON workflow_step_results(run_id); + + -- RCA (Root Cause Analysis) tables + CREATE TABLE IF NOT EXISTS root_causes ( + id TEXT PRIMARY KEY, + error_signature TEXT NOT NULL, + timestamp INTEGER NOT NULL, + root_cause_nodes TEXT, + remediation_steps TEXT, + confidence REAL DEFAULT 0.0, + history_item_id INTEGER, + FOREIGN KEY (history_item_id) REFERENCES history(id) + ); + + CREATE TABLE IF NOT EXISTS runbooks ( + id TEXT PRIMARY KEY, + project_id TEXT, + name TEXT NOT NULL, + description TEXT, + steps TEXT NOT NULL, + success_rate REAL DEFAULT 0.0, + last_used INTEGER, + usage_count INTEGER DEFAULT 0, + tags TEXT + ); + + CREATE TABLE IF NOT EXISTS project_fingerprints ( + id TEXT PRIMARY KEY, + project_type TEXT NOT NULL, + package_manager TEXT, + common_issues TEXT, + associated_runbooks TEXT, + detected_at TEXT NOT NULL, + detected_time INTEGER + ); + + CREATE INDEX IF NOT EXISTS idx_root_cause_signature ON root_causes(error_signature); + CREATE INDEX IF NOT EXISTS idx_root_cause_history ON root_causes(history_item_id); + CREATE INDEX IF NOT EXISTS idx_runbook_project ON runbooks(project_id); + CREATE INDEX IF NOT EXISTS idx_fingerprint_type ON project_fingerprints(project_type); + CREATE INDEX IF NOT EXISTS idx_fingerprint_path ON project_fingerprints(detected_at); ` _, err := db.Exec(schema) - return err + if err != nil { + return err + } + + _, _ = db.Exec("ALTER TABLE history ADD COLUMN resolution TEXT") + + return nil } diff --git a/internal/storage/db_test.go b/internal/storage/db_test.go index 2120f89..d97cd01 100644 --- a/internal/storage/db_test.go +++ b/internal/storage/db_test.go @@ -8,7 +8,7 @@ import ( ) func TestStorage(t *testing.T) { - // Setup temp DB + tmpDir, err := os.MkdirTemp("", "dev-cli-test") if err != nil { t.Fatal(err) @@ -22,7 +22,6 @@ func TestStorage(t *testing.T) { } defer db.Close() - // Test Insert entry := LogEntry{ Command: "git status", ExitCode: 0, @@ -35,7 +34,6 @@ func TestStorage(t *testing.T) { t.Errorf("SaveCommand failed: %v", err) } - // Test GetRecentHistory items, err := GetRecentHistory(db, 10) if err != nil { t.Errorf("GetRecentHistory failed: %v", err) @@ -48,8 +46,6 @@ func TestStorage(t *testing.T) { } } - // Test FTS Search - // Add another item specifically for search entry2 := LogEntry{ Command: "docker run hello-world", ExitCode: 0, @@ -60,9 +56,6 @@ func TestStorage(t *testing.T) { t.Errorf("SaveCommand failed: %v", err) } - // FTS triggers might be asynchronous if using FTS5? No, usually sync within transaction. - // But sqlite pure driver might handle it fine. - results, err := SearchHistory(db, "docker") if err != nil { t.Errorf("SearchHistory failed: %v", err) @@ -73,7 +66,6 @@ func TestStorage(t *testing.T) { t.Errorf("Got wrong result: %s", results[0].Command) } - // Search in output (details) results2, err := SearchHistory(db, "Hello") if err != nil { t.Errorf("SearchHistory failed: %v", err) @@ -82,3 +74,70 @@ func TestStorage(t *testing.T) { t.Errorf("Expected 1 search result for 'Hello' (in output), got %d", len(results2)) } } + +func TestResolutionTracking(t *testing.T) { + + tmpDir, err := os.MkdirTemp("", "dev-cli-test-resolution") + if err != nil { + t.Fatal(err) + } + defer os.RemoveAll(tmpDir) + + dbPath := filepath.Join(tmpDir, "history.db") + db, err := OpenDB(dbPath) + if err != nil { + t.Fatalf("OpenDB failed: %v", err) + } + defer db.Close() + + failEntry := LogEntry{ + Command: "npm run build", + ExitCode: 1, + Output: "Error: Module not found", + Cwd: "/tmp/project", + DurationMs: 500, + Timestamp: time.Now().Format(time.RFC3339), + } + if err := SaveCommand(db, failEntry); err != nil { + t.Fatalf("SaveCommand failed: %v", err) + } + + failure, err := GetLastUnresolvedFailure(db) + if err != nil { + t.Fatalf("GetLastUnresolvedFailure failed: %v", err) + } + if failure == nil { + t.Fatal("Expected to find an unresolved failure, got nil") + } + if failure.Command != "npm run build" { + t.Errorf("Expected 'npm run build', got '%s'", failure.Command) + } + if failure.Resolution != "" { + t.Errorf("Expected empty resolution, got '%s'", failure.Resolution) + } + + if err := MarkResolution(db, failure.ID, "solution"); err != nil { + t.Fatalf("MarkResolution failed: %v", err) + } + + item, err := GetHistoryByID(db, failure.ID) + if err != nil { + t.Fatalf("GetHistoryByID failed: %v", err) + } + if item.Resolution != "solution" { + t.Errorf("Expected resolution 'solution', got '%s'", item.Resolution) + } + + failure2, err := GetLastUnresolvedFailure(db) + if err != nil { + t.Fatalf("GetLastUnresolvedFailure failed: %v", err) + } + if failure2 != nil { + t.Errorf("Expected no unresolved failures, but got: %+v", failure2) + } + + err = MarkResolution(db, 99999, "solution") + if err == nil { + t.Error("Expected error for invalid ID, got nil") + } +} diff --git a/internal/storage/rca_models.go b/internal/storage/rca_models.go new file mode 100644 index 0000000..e8cc683 --- /dev/null +++ b/internal/storage/rca_models.go @@ -0,0 +1,88 @@ +package storage + +import ( + "fmt" + "hash/fnv" + "time" +) + +// RootCause represents a diagnosed failure with remediation steps. +// It captures the causal chain and suggested fixes for an error pattern. +type RootCause struct { + ID string `json:"id"` + ErrorSignature string `json:"error_signature"` // Hash/pattern of error + Timestamp time.Time `json:"timestamp"` + RootCauseNodes []string `json:"root_cause_nodes"` // Causal chain nodes + RemediationSteps []string `json:"remediation_steps"` // Ordered fix steps + Confidence float64 `json:"confidence"` // 0.0-1.0 confidence score + HistoryItemID int64 `json:"history_item_id"` // Link to original failure +} + +// Runbook represents a reusable remediation workflow. +// Runbooks are learned from successful fixes and can be reapplied. +type Runbook struct { + ID string `json:"id"` + ProjectID string `json:"project_id"` + Name string `json:"name"` + Description string `json:"description"` + Steps []RunbookStep `json:"steps"` // Ordered workflow steps + SuccessRate float64 `json:"success_rate"` // Historical success percentage + LastUsed time.Time `json:"last_used"` + UsageCount int `json:"usage_count"` + Tags []string `json:"tags"` // For categorization +} + +// RunbookStep represents a single step in a runbook. +type RunbookStep struct { + ID string `json:"id"` + Name string `json:"name"` + Command string `json:"command"` + Description string `json:"description"` + Rollback string `json:"rollback,omitempty"` // Optional rollback command + Condition string `json:"condition,omitempty"` // Optional execution condition +} + +// ProjectFingerprint identifies project characteristics for targeted RCA. +// Used to match errors to relevant runbooks based on project type. +type ProjectFingerprint struct { + ID string `json:"id"` + ProjectType string `json:"project_type"` // "nodejs", "go", "python", etc. + PackageManager string `json:"package_manager"` // "npm", "go mod", "pip", etc. + CommonIssues []string `json:"common_issues"` // Frequent error patterns + AssociatedRunbooks []string `json:"associated_runbooks"` // Runbook IDs + DetectedAt string `json:"detected_at"` // Directory path + DetectedTime time.Time `json:"detected_time"` +} + +// GenerateErrorSignature generates a normalized signature for an error. +// Used for cache lookups and pattern matching. +func GenerateErrorSignature(command string, exitCode int, output string) string { + + firstLine := output + if idx := indexOf(output, '\n'); idx > 0 { + firstLine = output[:idx] + } + if len(firstLine) > 100 { + firstLine = firstLine[:100] + } + + combined := fmt.Sprintf("%s|%d|%s", command, exitCode, firstLine) + return hashString(combined) +} + +// indexOf finds the first occurrence of a rune in a string +func indexOf(s string, r rune) int { + for i, c := range s { + if c == r { + return i + } + } + return -1 +} + +// hashString creates a hex hash string using FNV-1a algorithm +func hashString(s string) string { + h := fnv.New64a() + h.Write([]byte(s)) + return fmt.Sprintf("%016x", h.Sum64()) +} diff --git a/internal/storage/rca_repository.go b/internal/storage/rca_repository.go new file mode 100644 index 0000000..67ec7b8 --- /dev/null +++ b/internal/storage/rca_repository.go @@ -0,0 +1,354 @@ +package storage + +import ( + "database/sql" + "encoding/json" + "fmt" + "time" +) + +// SaveRootCause persists a root cause analysis result. +func SaveRootCause(db *sql.DB, rc RootCause) error { + nodesJSON, err := json.Marshal(rc.RootCauseNodes) + if err != nil { + return fmt.Errorf("marshal root_cause_nodes: %w", err) + } + stepsJSON, err := json.Marshal(rc.RemediationSteps) + if err != nil { + return fmt.Errorf("marshal remediation_steps: %w", err) + } + + query := `INSERT OR REPLACE INTO root_causes + (id, error_signature, timestamp, root_cause_nodes, remediation_steps, confidence, history_item_id) + VALUES (?, ?, ?, ?, ?, ?, ?)` + + _, err = db.Exec(query, + rc.ID, + rc.ErrorSignature, + rc.Timestamp.Unix(), + string(nodesJSON), + string(stepsJSON), + rc.Confidence, + rc.HistoryItemID, + ) + return err +} + +// GetRootCauseBySignature retrieves a root cause by its error signature. +func GetRootCauseBySignature(db *sql.DB, signature string) (*RootCause, error) { + query := `SELECT id, error_signature, timestamp, root_cause_nodes, remediation_steps, confidence, history_item_id + FROM root_causes WHERE error_signature = ? ORDER BY timestamp DESC LIMIT 1` + + row := db.QueryRow(query, signature) + return scanRootCause(row) +} + +// GetRootCauseByID retrieves a root cause by its ID. +func GetRootCauseByID(db *sql.DB, id string) (*RootCause, error) { + query := `SELECT id, error_signature, timestamp, root_cause_nodes, remediation_steps, confidence, history_item_id + FROM root_causes WHERE id = ?` + + row := db.QueryRow(query, id) + return scanRootCause(row) +} + +// GetRecentRootCauses retrieves the most recent root cause analyses. +func GetRecentRootCauses(db *sql.DB, limit int) ([]RootCause, error) { + query := `SELECT id, error_signature, timestamp, root_cause_nodes, remediation_steps, confidence, history_item_id + FROM root_causes ORDER BY timestamp DESC LIMIT ?` + + rows, err := db.Query(query, limit) + if err != nil { + return nil, err + } + defer rows.Close() + + var results []RootCause + for rows.Next() { + rc, err := scanRootCauseRow(rows) + if err != nil { + return nil, err + } + results = append(results, *rc) + } + return results, nil +} + +func scanRootCause(row *sql.Row) (*RootCause, error) { + var rc RootCause + var ts int64 + var nodesJSON, stepsJSON string + + err := row.Scan(&rc.ID, &rc.ErrorSignature, &ts, &nodesJSON, &stepsJSON, &rc.Confidence, &rc.HistoryItemID) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + + rc.Timestamp = time.Unix(ts, 0) + if err := json.Unmarshal([]byte(nodesJSON), &rc.RootCauseNodes); err != nil { + rc.RootCauseNodes = []string{} + } + if err := json.Unmarshal([]byte(stepsJSON), &rc.RemediationSteps); err != nil { + rc.RemediationSteps = []string{} + } + + return &rc, nil +} + +func scanRootCauseRow(rows *sql.Rows) (*RootCause, error) { + var rc RootCause + var ts int64 + var nodesJSON, stepsJSON string + + err := rows.Scan(&rc.ID, &rc.ErrorSignature, &ts, &nodesJSON, &stepsJSON, &rc.Confidence, &rc.HistoryItemID) + if err != nil { + return nil, err + } + + rc.Timestamp = time.Unix(ts, 0) + if err := json.Unmarshal([]byte(nodesJSON), &rc.RootCauseNodes); err != nil { + rc.RootCauseNodes = []string{} + } + if err := json.Unmarshal([]byte(stepsJSON), &rc.RemediationSteps); err != nil { + rc.RemediationSteps = []string{} + } + + return &rc, nil +} + +// SaveRunbook persists a runbook. +func SaveRunbook(db *sql.DB, rb Runbook) error { + stepsJSON, err := json.Marshal(rb.Steps) + if err != nil { + return fmt.Errorf("marshal steps: %w", err) + } + tagsJSON, err := json.Marshal(rb.Tags) + if err != nil { + return fmt.Errorf("marshal tags: %w", err) + } + + query := `INSERT OR REPLACE INTO runbooks + (id, project_id, name, description, steps, success_rate, last_used, usage_count, tags) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)` + + _, err = db.Exec(query, + rb.ID, + rb.ProjectID, + rb.Name, + rb.Description, + string(stepsJSON), + rb.SuccessRate, + rb.LastUsed.Unix(), + rb.UsageCount, + string(tagsJSON), + ) + return err +} + +// GetRunbookByID retrieves a runbook by its ID. +func GetRunbookByID(db *sql.DB, id string) (*Runbook, error) { + query := `SELECT id, project_id, name, description, steps, success_rate, last_used, usage_count, tags + FROM runbooks WHERE id = ?` + + row := db.QueryRow(query, id) + return scanRunbook(row) +} + +// GetRunbooksForProject retrieves all runbooks for a project. +func GetRunbooksForProject(db *sql.DB, projectID string) ([]Runbook, error) { + query := `SELECT id, project_id, name, description, steps, success_rate, last_used, usage_count, tags + FROM runbooks WHERE project_id = ? ORDER BY success_rate DESC` + + rows, err := db.Query(query, projectID) + if err != nil { + return nil, err + } + defer rows.Close() + + var results []Runbook + for rows.Next() { + rb, err := scanRunbookRow(rows) + if err != nil { + return nil, err + } + results = append(results, *rb) + } + return results, nil +} + +// UpdateRunbookStats updates a runbook's success rate after execution. +func UpdateRunbookStats(db *sql.DB, id string, success bool) error { + + rb, err := GetRunbookByID(db, id) + if err != nil { + return err + } + if rb == nil { + return fmt.Errorf("runbook not found: %s", id) + } + + rb.UsageCount++ + if success { + + rb.SuccessRate = ((rb.SuccessRate * float64(rb.UsageCount-1)) + 1.0) / float64(rb.UsageCount) + } else { + + rb.SuccessRate = (rb.SuccessRate * float64(rb.UsageCount-1)) / float64(rb.UsageCount) + } + rb.LastUsed = time.Now() + + query := `UPDATE runbooks SET success_rate = ?, last_used = ?, usage_count = ? WHERE id = ?` + _, err = db.Exec(query, rb.SuccessRate, rb.LastUsed.Unix(), rb.UsageCount, id) + return err +} + +func scanRunbook(row *sql.Row) (*Runbook, error) { + var rb Runbook + var lastUsed int64 + var stepsJSON, tagsJSON string + + err := row.Scan(&rb.ID, &rb.ProjectID, &rb.Name, &rb.Description, &stepsJSON, &rb.SuccessRate, &lastUsed, &rb.UsageCount, &tagsJSON) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + + rb.LastUsed = time.Unix(lastUsed, 0) + if err := json.Unmarshal([]byte(stepsJSON), &rb.Steps); err != nil { + rb.Steps = []RunbookStep{} + } + if err := json.Unmarshal([]byte(tagsJSON), &rb.Tags); err != nil { + rb.Tags = []string{} + } + + return &rb, nil +} + +func scanRunbookRow(rows *sql.Rows) (*Runbook, error) { + var rb Runbook + var lastUsed int64 + var stepsJSON, tagsJSON string + + err := rows.Scan(&rb.ID, &rb.ProjectID, &rb.Name, &rb.Description, &stepsJSON, &rb.SuccessRate, &lastUsed, &rb.UsageCount, &tagsJSON) + if err != nil { + return nil, err + } + + rb.LastUsed = time.Unix(lastUsed, 0) + if err := json.Unmarshal([]byte(stepsJSON), &rb.Steps); err != nil { + rb.Steps = []RunbookStep{} + } + if err := json.Unmarshal([]byte(tagsJSON), &rb.Tags); err != nil { + rb.Tags = []string{} + } + + return &rb, nil +} + +// SaveProjectFingerprint persists a project fingerprint. +func SaveProjectFingerprint(db *sql.DB, fp ProjectFingerprint) error { + issuesJSON, err := json.Marshal(fp.CommonIssues) + if err != nil { + return fmt.Errorf("marshal common_issues: %w", err) + } + runbooksJSON, err := json.Marshal(fp.AssociatedRunbooks) + if err != nil { + return fmt.Errorf("marshal associated_runbooks: %w", err) + } + + query := `INSERT OR REPLACE INTO project_fingerprints + (id, project_type, package_manager, common_issues, associated_runbooks, detected_at, detected_time) + VALUES (?, ?, ?, ?, ?, ?, ?)` + + _, err = db.Exec(query, + fp.ID, + fp.ProjectType, + fp.PackageManager, + string(issuesJSON), + string(runbooksJSON), + fp.DetectedAt, + fp.DetectedTime.Unix(), + ) + return err +} + +// GetProjectFingerprint retrieves a project fingerprint by directory path. +func GetProjectFingerprint(db *sql.DB, path string) (*ProjectFingerprint, error) { + query := `SELECT id, project_type, package_manager, common_issues, associated_runbooks, detected_at, detected_time + FROM project_fingerprints WHERE detected_at = ?` + + row := db.QueryRow(query, path) + return scanProjectFingerprint(row) +} + +// GetProjectFingerprintByType retrieves project fingerprints by type. +func GetProjectFingerprintByType(db *sql.DB, projectType string) ([]ProjectFingerprint, error) { + query := `SELECT id, project_type, package_manager, common_issues, associated_runbooks, detected_at, detected_time + FROM project_fingerprints WHERE project_type = ?` + + rows, err := db.Query(query, projectType) + if err != nil { + return nil, err + } + defer rows.Close() + + var results []ProjectFingerprint + for rows.Next() { + fp, err := scanProjectFingerprintRow(rows) + if err != nil { + return nil, err + } + results = append(results, *fp) + } + return results, nil +} + +func scanProjectFingerprint(row *sql.Row) (*ProjectFingerprint, error) { + var fp ProjectFingerprint + var detectedTime int64 + var issuesJSON, runbooksJSON string + + err := row.Scan(&fp.ID, &fp.ProjectType, &fp.PackageManager, &issuesJSON, &runbooksJSON, &fp.DetectedAt, &detectedTime) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + + fp.DetectedTime = time.Unix(detectedTime, 0) + if err := json.Unmarshal([]byte(issuesJSON), &fp.CommonIssues); err != nil { + fp.CommonIssues = []string{} + } + if err := json.Unmarshal([]byte(runbooksJSON), &fp.AssociatedRunbooks); err != nil { + fp.AssociatedRunbooks = []string{} + } + + return &fp, nil +} + +func scanProjectFingerprintRow(rows *sql.Rows) (*ProjectFingerprint, error) { + var fp ProjectFingerprint + var detectedTime int64 + var issuesJSON, runbooksJSON string + + err := rows.Scan(&fp.ID, &fp.ProjectType, &fp.PackageManager, &issuesJSON, &runbooksJSON, &fp.DetectedAt, &detectedTime) + if err != nil { + return nil, err + } + + fp.DetectedTime = time.Unix(detectedTime, 0) + if err := json.Unmarshal([]byte(issuesJSON), &fp.CommonIssues); err != nil { + fp.CommonIssues = []string{} + } + if err := json.Unmarshal([]byte(runbooksJSON), &fp.AssociatedRunbooks); err != nil { + fp.AssociatedRunbooks = []string{} + } + + return &fp, nil +} diff --git a/internal/storage/rca_test.go b/internal/storage/rca_test.go new file mode 100644 index 0000000..dd047fa --- /dev/null +++ b/internal/storage/rca_test.go @@ -0,0 +1,194 @@ +package storage + +import ( + "database/sql" + "os" + "testing" + "time" +) + +func setupTestDB(t *testing.T) *sql.DB { + + tmpFile, err := os.CreateTemp("", "rca_test_*.db") + if err != nil { + t.Fatalf("failed to create temp file: %v", err) + } + tmpFile.Close() + + db, err := OpenDB(tmpFile.Name()) + if err != nil { + t.Fatalf("failed to open DB: %v", err) + } + + t.Cleanup(func() { + db.Close() + os.Remove(tmpFile.Name()) + }) + + return db +} + +func TestRootCause_CRUD(t *testing.T) { + db := setupTestDB(t) + + rc := RootCause{ + ID: "rc-001", + ErrorSignature: "npm-enoent-001", + Timestamp: time.Now(), + RootCauseNodes: []string{"missing package.json", "npm install required"}, + RemediationSteps: []string{"npm install", "retry command"}, + Confidence: 0.85, + HistoryItemID: 1, + } + + err := SaveRootCause(db, rc) + if err != nil { + t.Fatalf("SaveRootCause failed: %v", err) + } + + retrieved, err := GetRootCauseBySignature(db, "npm-enoent-001") + if err != nil { + t.Fatalf("GetRootCauseBySignature failed: %v", err) + } + if retrieved == nil { + t.Fatal("expected to find root cause, got nil") + } + if retrieved.ID != "rc-001" { + t.Errorf("expected ID 'rc-001', got '%s'", retrieved.ID) + } + if len(retrieved.RootCauseNodes) != 2 { + t.Errorf("expected 2 root cause nodes, got %d", len(retrieved.RootCauseNodes)) + } + if retrieved.Confidence != 0.85 { + t.Errorf("expected confidence 0.85, got %f", retrieved.Confidence) + } + + byID, err := GetRootCauseByID(db, "rc-001") + if err != nil { + t.Fatalf("GetRootCauseByID failed: %v", err) + } + if byID == nil || byID.ID != "rc-001" { + t.Error("GetRootCauseByID should return the correct root cause") + } + + rcList, err := GetRecentRootCauses(db, 10) + if err != nil { + t.Fatalf("GetRecentRootCauses failed: %v", err) + } + if len(rcList) != 1 { + t.Errorf("expected 1 root cause, got %d", len(rcList)) + } +} + +func TestRunbook_CRUD(t *testing.T) { + db := setupTestDB(t) + + rb := Runbook{ + ID: "rb-001", + ProjectID: "proj-nodejs", + Name: "NPM Install Fix", + Description: "Fixes missing dependency issues", + Steps: []RunbookStep{ + {ID: "s1", Name: "Check package.json", Command: "cat package.json", Description: "Verify package.json exists"}, + {ID: "s2", Name: "Install dependencies", Command: "npm install", Rollback: "rm -rf node_modules"}, + }, + SuccessRate: 0.9, + LastUsed: time.Now(), + UsageCount: 10, + Tags: []string{"npm", "dependencies"}, + } + + err := SaveRunbook(db, rb) + if err != nil { + t.Fatalf("SaveRunbook failed: %v", err) + } + + retrieved, err := GetRunbookByID(db, "rb-001") + if err != nil { + t.Fatalf("GetRunbookByID failed: %v", err) + } + if retrieved == nil { + t.Fatal("expected to find runbook, got nil") + } + if retrieved.Name != "NPM Install Fix" { + t.Errorf("expected name 'NPM Install Fix', got '%s'", retrieved.Name) + } + if len(retrieved.Steps) != 2 { + t.Errorf("expected 2 steps, got %d", len(retrieved.Steps)) + } + if retrieved.Steps[0].Command != "cat package.json" { + t.Errorf("expected first step command 'cat package.json', got '%s'", retrieved.Steps[0].Command) + } + + projectRunbooks, err := GetRunbooksForProject(db, "proj-nodejs") + if err != nil { + t.Fatalf("GetRunbooksForProject failed: %v", err) + } + if len(projectRunbooks) != 1 { + t.Errorf("expected 1 runbook for project, got %d", len(projectRunbooks)) + } + + err = UpdateRunbookStats(db, "rb-001", true) + if err != nil { + t.Fatalf("UpdateRunbookStats failed: %v", err) + } + + updated, _ := GetRunbookByID(db, "rb-001") + if updated.UsageCount != 11 { + t.Errorf("expected usage count 11, got %d", updated.UsageCount) + } +} + +func TestProjectFingerprint_CRUD(t *testing.T) { + db := setupTestDB(t) + + fp := ProjectFingerprint{ + ID: "fp-001", + ProjectType: "nodejs", + PackageManager: "npm", + CommonIssues: []string{"ENOENT", "MODULE_NOT_FOUND"}, + AssociatedRunbooks: []string{"rb-001", "rb-002"}, + DetectedAt: "/home/user/project", + DetectedTime: time.Now(), + } + + err := SaveProjectFingerprint(db, fp) + if err != nil { + t.Fatalf("SaveProjectFingerprint failed: %v", err) + } + + retrieved, err := GetProjectFingerprint(db, "/home/user/project") + if err != nil { + t.Fatalf("GetProjectFingerprint failed: %v", err) + } + if retrieved == nil { + t.Fatal("expected to find fingerprint, got nil") + } + if retrieved.ProjectType != "nodejs" { + t.Errorf("expected project type 'nodejs', got '%s'", retrieved.ProjectType) + } + if len(retrieved.CommonIssues) != 2 { + t.Errorf("expected 2 common issues, got %d", len(retrieved.CommonIssues)) + } + + byType, err := GetProjectFingerprintByType(db, "nodejs") + if err != nil { + t.Fatalf("GetProjectFingerprintByType failed: %v", err) + } + if len(byType) != 1 { + t.Errorf("expected 1 fingerprint, got %d", len(byType)) + } +} + +func TestErrorSignature(t *testing.T) { + sig1 := GenerateErrorSignature("npm install", 1, "ENOENT: no such file") + sig2 := GenerateErrorSignature("npm install", 1, "ENOENT: no such file") + sig3 := GenerateErrorSignature("npm install", 1, "EACCES: permission denied") + + if sig1 != sig2 { + t.Error("same input should produce same signature") + } + if sig1 == sig3 { + t.Error("different errors should produce different signatures") + } +} diff --git a/internal/storage/repository.go b/internal/storage/repository.go index 5c99846..62d3ca9 100644 --- a/internal/storage/repository.go +++ b/internal/storage/repository.go @@ -28,6 +28,7 @@ type HistoryItem struct { Directory string SessionID string Details string // Raw JSON + Resolution string // "solution", "unrelated", "skipped", or "" (empty) } func SaveCommand(db *sql.DB, entry LogEntry) error { @@ -53,7 +54,7 @@ func SaveCommand(db *sql.DB, entry LogEntry) error { } func GetRecentHistory(db *sql.DB, limit int) ([]HistoryItem, error) { - query := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details + query := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details, COALESCE(resolution, '') FROM history ORDER BY id DESC LIMIT ?` rows, err := db.Query(query, limit) @@ -66,7 +67,7 @@ func GetRecentHistory(db *sql.DB, limit int) ([]HistoryItem, error) { for rows.Next() { var item HistoryItem var ts int64 - if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details); err != nil { + if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution); err != nil { return nil, err } item.Timestamp = time.Unix(ts, 0) @@ -76,7 +77,7 @@ func GetRecentHistory(db *sql.DB, limit int) ([]HistoryItem, error) { } func SearchHistory(db *sql.DB, query string) ([]HistoryItem, error) { - sqlQuery := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details + sqlQuery := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details, COALESCE(resolution, '') FROM history WHERE command LIKE ? OR details LIKE ? ORDER BY id DESC @@ -93,7 +94,7 @@ func SearchHistory(db *sql.DB, query string) ([]HistoryItem, error) { for rows.Next() { var item HistoryItem var ts int64 - if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details); err != nil { + if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution); err != nil { return nil, err } item.Timestamp = time.Unix(ts, 0) @@ -109,7 +110,7 @@ type QueryOpts struct { } func GetFailures(db *sql.DB, opts QueryOpts) ([]HistoryItem, error) { - queryBuilder := `SELECT h.id, h.timestamp, h.command, h.exit_code, h.duration_ms, h.directory, h.session_id, h.details + queryBuilder := `SELECT h.id, h.timestamp, h.command, h.exit_code, h.duration_ms, h.directory, h.session_id, h.details, COALESCE(h.resolution, '') FROM history h` var args []interface{} var whereClauses []string @@ -148,7 +149,7 @@ func GetFailures(db *sql.DB, opts QueryOpts) ([]HistoryItem, error) { for rows.Next() { var item HistoryItem var ts int64 - if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details); err != nil { + if err := rows.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution); err != nil { return nil, err } item.Timestamp = time.Unix(ts, 0) @@ -156,3 +157,61 @@ func GetFailures(db *sql.DB, opts QueryOpts) ([]HistoryItem, error) { } return items, nil } + +// GetLastUnresolvedFailure returns the most recent failed command that hasn't been resolved. +func GetLastUnresolvedFailure(db *sql.DB) (*HistoryItem, error) { + query := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details, COALESCE(resolution, '') + FROM history + WHERE exit_code != 0 AND exit_code != 130 AND (resolution IS NULL OR resolution = '') + ORDER BY id DESC LIMIT 1` + + row := db.QueryRow(query) + var item HistoryItem + var ts int64 + err := row.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + item.Timestamp = time.Unix(ts, 0) + return &item, nil +} + +// GetHistoryByID retrieves a specific history item by ID. +func GetHistoryByID(db *sql.DB, id int64) (*HistoryItem, error) { + query := `SELECT id, timestamp, command, exit_code, duration_ms, directory, session_id, details, COALESCE(resolution, '') + FROM history WHERE id = ?` + + row := db.QueryRow(query, id) + var item HistoryItem + var ts int64 + err := row.Scan(&item.ID, &ts, &item.Command, &item.ExitCode, &item.DurationMs, &item.Directory, &item.SessionID, &item.Details, &item.Resolution) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + item.Timestamp = time.Unix(ts, 0) + return &item, nil +} + +// MarkResolution updates the resolution status of a history entry. +// Valid values: "solution", "unrelated", "skipped" +func MarkResolution(db *sql.DB, id int64, resolution string) error { + query := `UPDATE history SET resolution = ? WHERE id = ?` + result, err := db.Exec(query, resolution, id) + if err != nil { + return err + } + rowsAffected, err := result.RowsAffected() + if err != nil { + return err + } + if rowsAffected == 0 { + return fmt.Errorf("history entry not found: %d", id) + } + return nil +} diff --git a/internal/tools/command_tools.go b/internal/tools/command_tools.go new file mode 100644 index 0000000..4fed1af --- /dev/null +++ b/internal/tools/command_tools.go @@ -0,0 +1,54 @@ +package tools + +import ( + "context" + "time" + + "dev-cli/internal/executor" +) + +// RunCommandTool executes shell commands with timeout. +type RunCommandTool struct{} + +func (t *RunCommandTool) Name() string { return "run_command" } +func (t *RunCommandTool) Description() string { + return "Execute shell command with timeout and capture output" +} + +func (t *RunCommandTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "command", Type: "string", Description: "Command to execute", Required: true}, + {Name: "timeout", Type: "duration", Description: "Timeout (e.g., '30s', '5m')", Required: false, Default: "60s"}, + {Name: "cwd", Type: "string", Description: "Working directory", Required: false}, + } +} + +// CommandResult contains the command execution output. +type CommandResult struct { + Command string `json:"command"` + Output string `json:"output"` + ExitCode int `json:"exit_code"` + Duration string `json:"duration"` + Cwd string `json:"cwd,omitempty"` +} + +func (t *RunCommandTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + command := GetString(params, "command", "") + if command == "" { + return NewErrorResult("command is required", time.Since(start)) + } + + timeout := GetDuration(params, "timeout", 60*time.Second) + + result := executor.ExecuteWithTimeout(command, timeout) + + return NewResult(CommandResult{ + Command: result.Command, + Output: result.Output, + ExitCode: result.ExitCode, + Duration: result.Duration.String(), + Cwd: result.Cwd, + }, time.Since(start)) +} diff --git a/internal/tools/docker_tools.go b/internal/tools/docker_tools.go new file mode 100644 index 0000000..6f03d02 --- /dev/null +++ b/internal/tools/docker_tools.go @@ -0,0 +1,204 @@ +package tools + +import ( + "context" + "fmt" + "strings" + "time" + + "dev-cli/internal/infra" +) + +// QueryDockerTool queries Docker containers for logs, stats, and inspection. +type QueryDockerTool struct{} + +func (t *QueryDockerTool) Name() string { return "query_docker" } +func (t *QueryDockerTool) Description() string { + return "Query Docker containers for logs, stats, and info" +} + +func (t *QueryDockerTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "action", Type: "string", Description: "Action: logs, stats, inspect, list", Required: true}, + {Name: "container", Type: "string", Description: "Container ID or name", Required: false}, + {Name: "tail", Type: "int", Description: "Number of log lines (for logs action)", Required: false, Default: 100}, + } +} + +// DockerLogsResult contains container logs output. +type DockerLogsResult struct { + Container string `json:"container"` + Lines []string `json:"lines"` + Count int `json:"count"` +} + +// DockerStatsResult contains container stats. +type DockerStatsResult struct { + Container string `json:"container"` + CPUPercent float64 `json:"cpu_percent"` + MemUsedMB uint64 `json:"mem_used_mb"` + MemLimitMB uint64 `json:"mem_limit_mb"` + MemPercent float64 `json:"mem_percent"` + NetRxMB float64 `json:"net_rx_mb"` + NetTxMB float64 `json:"net_tx_mb"` + PIDs uint64 `json:"pids"` +} + +// DockerInspectResult contains container inspection details. +type DockerInspectResult struct { + ID string `json:"id"` + Name string `json:"name"` + Image string `json:"image"` + State string `json:"state"` + Status string `json:"status"` + Ports []string `json:"ports"` + Mounts []string `json:"mounts"` + EnvVars []string `json:"env_vars"` + Cmd []string `json:"cmd"` + NetworkID string `json:"network_id"` + Uptime string `json:"uptime"` +} + +// DockerListResult contains list of containers. +type DockerListResult struct { + Containers []DockerContainerInfo `json:"containers"` + Count int `json:"count"` +} + +// DockerContainerInfo contains basic container info. +type DockerContainerInfo struct { + ID string `json:"id"` + Name string `json:"name"` + Image string `json:"image"` + State string `json:"state"` + Status string `json:"status"` +} + +func (t *QueryDockerTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + action := GetString(params, "action", "") + if action == "" { + return NewErrorResult("action is required (logs, stats, inspect, list)", time.Since(start)) + } + + docker, err := infra.GetRegistry().Docker() + if err != nil { + return NewErrorResult(fmt.Sprintf("Docker not available: %v", err), time.Since(start)) + } + + switch action { + case "logs": + return t.getLogs(ctx, docker, params, start) + case "stats": + return t.getStats(ctx, docker, params, start) + case "inspect": + return t.inspect(ctx, docker, params, start) + case "list": + return t.list(ctx, docker, start) + default: + return NewErrorResult(fmt.Sprintf("unknown action: %s", action), time.Since(start)) + } +} + +func (t *QueryDockerTool) getLogs(ctx context.Context, docker *infra.DockerClient, params map[string]any, start time.Time) ToolResult { + container := GetString(params, "container", "") + if container == "" { + return NewErrorResult("container is required for logs action", time.Since(start)) + } + + tail := GetInt(params, "tail", 100) + + lines, err := docker.GetContainerLogs(ctx, container, tail) + if err != nil { + return NewErrorResult(fmt.Sprintf("failed to get logs: %v", err), time.Since(start)) + } + + return NewResult(DockerLogsResult{ + Container: container, + Lines: lines, + Count: len(lines), + }, time.Since(start)) +} + +func (t *QueryDockerTool) getStats(ctx context.Context, docker *infra.DockerClient, params map[string]any, start time.Time) ToolResult { + container := GetString(params, "container", "") + if container == "" { + return NewErrorResult("container is required for stats action", time.Since(start)) + } + + stats, err := docker.GetContainerStats(ctx, container) + if err != nil { + return NewErrorResult(fmt.Sprintf("failed to get stats: %v", err), time.Since(start)) + } + + return NewResult(DockerStatsResult{ + Container: container, + CPUPercent: stats.CPUPercent, + MemUsedMB: stats.MemUsed / (1024 * 1024), + MemLimitMB: stats.MemLimit / (1024 * 1024), + MemPercent: stats.MemPercent, + NetRxMB: float64(stats.NetRx) / (1024 * 1024), + NetTxMB: float64(stats.NetTx) / (1024 * 1024), + PIDs: stats.PIDs, + }, time.Since(start)) +} + +func (t *QueryDockerTool) inspect(ctx context.Context, docker *infra.DockerClient, params map[string]any, start time.Time) ToolResult { + container := GetString(params, "container", "") + if container == "" { + return NewErrorResult("container is required for inspect action", time.Since(start)) + } + + detail, err := docker.InspectContainer(ctx, container) + if err != nil { + return NewErrorResult(fmt.Sprintf("failed to inspect: %v", err), time.Since(start)) + } + + ports := make([]string, 0, len(detail.Ports)) + for _, p := range detail.Ports { + ports = append(ports, fmt.Sprintf("%d:%d/%s", p.Public, p.Private, p.Protocol)) + } + + mounts := make([]string, 0, len(detail.Mounts)) + for _, m := range detail.Mounts { + mounts = append(mounts, fmt.Sprintf("%s:%s", m.Source, m.Destination)) + } + + return NewResult(DockerInspectResult{ + ID: detail.ID, + Name: detail.Name, + Image: detail.Image, + State: detail.State, + Status: detail.Status, + Ports: ports, + Mounts: mounts, + EnvVars: detail.EnvVars, + Cmd: detail.Cmd, + NetworkID: detail.NetworkID, + Uptime: detail.Uptime, + }, time.Since(start)) +} + +func (t *QueryDockerTool) list(ctx context.Context, docker *infra.DockerClient, start time.Time) ToolResult { + health := docker.CheckHealth(ctx) + if !health.Available { + return NewErrorResult("Docker not available", time.Since(start)) + } + + containers := make([]DockerContainerInfo, 0, len(health.Containers)) + for _, c := range health.Containers { + containers = append(containers, DockerContainerInfo{ + ID: c.ID[:12], + Name: strings.TrimPrefix(c.Name, "/"), + Image: c.Image, + State: c.State, + Status: c.Status, + }) + } + + return NewResult(DockerListResult{ + Containers: containers, + Count: len(containers), + }, time.Since(start)) +} diff --git a/internal/tools/file_tools.go b/internal/tools/file_tools.go new file mode 100644 index 0000000..c349e1a --- /dev/null +++ b/internal/tools/file_tools.go @@ -0,0 +1,406 @@ +package tools + +import ( + "context" + "fmt" + "io" + "os" + "path/filepath" + "strings" + "time" +) + +// ReadFileTool reads file contents. +type ReadFileTool struct{} + +func (t *ReadFileTool) Name() string { return "read_file" } +func (t *ReadFileTool) Description() string { return "Read file contents with optional line range" } + +func (t *ReadFileTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "path", Type: "string", Description: "Path to file", Required: true}, + {Name: "start_line", Type: "int", Description: "Start line (1-indexed)", Required: false, Default: 0}, + {Name: "end_line", Type: "int", Description: "End line (1-indexed, 0 = all)", Required: false, Default: 0}, + {Name: "max_size", Type: "int", Description: "Max bytes to read (0 = 1MB default)", Required: false, Default: 0}, + } +} + +// ReadFileResult contains the file reading output. +type ReadFileResult struct { + Path string `json:"path"` + Content string `json:"content"` + Lines int `json:"lines"` + Size int64 `json:"size"` + Truncated bool `json:"truncated"` + StartLine int `json:"start_line,omitempty"` + EndLine int `json:"end_line,omitempty"` +} + +func (t *ReadFileTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + path := GetString(params, "path", "") + if path == "" { + return NewErrorResult("path is required", time.Since(start)) + } + + if strings.HasPrefix(path, "~/") { + if home, err := os.UserHomeDir(); err == nil { + path = filepath.Join(home, path[2:]) + } + } + + absPath, err := filepath.Abs(path) + if err != nil { + return NewErrorResult(fmt.Sprintf("invalid path: %v", err), time.Since(start)) + } + + info, err := os.Stat(absPath) + if err != nil { + if os.IsNotExist(err) { + return NewErrorResult(fmt.Sprintf("file not found: %s", absPath), time.Since(start)) + } + return NewErrorResult(fmt.Sprintf("cannot access file: %v", err), time.Since(start)) + } + + if info.IsDir() { + return NewErrorResult("path is a directory, not a file", time.Since(start)) + } + + maxSize := GetInt(params, "max_size", 0) + if maxSize <= 0 { + maxSize = 1024 * 1024 + } + + truncated := false + if info.Size() > int64(maxSize) { + truncated = true + } + + file, err := os.Open(absPath) + if err != nil { + return NewErrorResult(fmt.Sprintf("cannot open file: %v", err), time.Since(start)) + } + defer file.Close() + + reader := io.LimitReader(file, int64(maxSize)) + data, err := io.ReadAll(reader) + if err != nil { + return NewErrorResult(fmt.Sprintf("read error: %v", err), time.Since(start)) + } + + content := string(data) + lines := strings.Split(content, "\n") + totalLines := len(lines) + + startLine := GetInt(params, "start_line", 0) + endLine := GetInt(params, "end_line", 0) + + if startLine > 0 || endLine > 0 { + if startLine < 1 { + startLine = 1 + } + if endLine < 1 || endLine > totalLines { + endLine = totalLines + } + if startLine > totalLines { + startLine = totalLines + } + if startLine > endLine { + startLine = endLine + } + + lines = lines[startLine-1 : endLine] + content = strings.Join(lines, "\n") + } + + result := ReadFileResult{ + Path: absPath, + Content: content, + Lines: len(lines), + Size: info.Size(), + Truncated: truncated, + } + if startLine > 0 { + result.StartLine = startLine + result.EndLine = endLine + } + + return NewResult(result, time.Since(start)) +} + +// WriteFileTool writes content to a file. +type WriteFileTool struct{} + +func (t *WriteFileTool) Name() string { return "write_file" } +func (t *WriteFileTool) Description() string { return "Write content to a file with optional backup" } + +func (t *WriteFileTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "path", Type: "string", Description: "Path to file", Required: true}, + {Name: "content", Type: "string", Description: "Content to write", Required: true}, + {Name: "mode", Type: "string", Description: "File mode (e.g., '0644')", Required: false, Default: "0644"}, + {Name: "backup", Type: "bool", Description: "Create backup of existing file", Required: false, Default: false}, + {Name: "create_dirs", Type: "bool", Description: "Create parent directories", Required: false, Default: true}, + } +} + +// WriteFileResult contains the file writing output. +type WriteFileResult struct { + Path string `json:"path"` + Size int `json:"size"` + BackupPath string `json:"backup_path,omitempty"` + Created bool `json:"created"` +} + +func (t *WriteFileTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + path := GetString(params, "path", "") + if path == "" { + return NewErrorResult("path is required", time.Since(start)) + } + + content := GetString(params, "content", "") + + if strings.HasPrefix(path, "~/") { + if home, err := os.UserHomeDir(); err == nil { + path = filepath.Join(home, path[2:]) + } + } + + absPath, err := filepath.Abs(path) + if err != nil { + return NewErrorResult(fmt.Sprintf("invalid path: %v", err), time.Since(start)) + } + + _, err = os.Stat(absPath) + exists := err == nil + created := !exists + + // Create backup if requested + var backupPath string + if exists && GetBool(params, "backup", false) { + backupPath = absPath + ".bak" + if err := copyFile(absPath, backupPath); err != nil { + return NewErrorResult(fmt.Sprintf("backup failed: %v", err), time.Since(start)) + } + } + + if GetBool(params, "create_dirs", true) { + dir := filepath.Dir(absPath) + if err := os.MkdirAll(dir, 0755); err != nil { + return NewErrorResult(fmt.Sprintf("cannot create directory: %v", err), time.Since(start)) + } + } + + mode := os.FileMode(0644) + modeStr := GetString(params, "mode", "0644") + if _, err := fmt.Sscanf(modeStr, "%o", &mode); err != nil { + mode = 0644 + } + + if err := os.WriteFile(absPath, []byte(content), mode); err != nil { + return NewErrorResult(fmt.Sprintf("write failed: %v", err), time.Since(start)) + } + + return NewResult(WriteFileResult{ + Path: absPath, + Size: len(content), + BackupPath: backupPath, + Created: created, + }, time.Since(start)) +} + +func copyFile(src, dst string) error { + source, err := os.Open(src) + if err != nil { + return err + } + defer source.Close() + + dest, err := os.Create(dst) + if err != nil { + return err + } + defer dest.Close() + + _, err = io.Copy(dest, source) + return err +} + +// ReadDirTool lists directory contents. +type ReadDirTool struct{} + +func (t *ReadDirTool) Name() string { return "read_dir" } +func (t *ReadDirTool) Description() string { + return "List directory contents with optional recursive traversal" +} + +func (t *ReadDirTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "path", Type: "string", Description: "Path to directory", Required: true}, + {Name: "recursive", Type: "bool", Description: "Recursively list subdirectories", Required: false, Default: false}, + {Name: "max_depth", Type: "int", Description: "Max depth for recursive listing (0 = unlimited)", Required: false, Default: 0}, + {Name: "include_hidden", Type: "bool", Description: "Include hidden files (starting with .)", Required: false, Default: false}, + {Name: "max_entries", Type: "int", Description: "Max entries to return (0 = 1000 default)", Required: false, Default: 0}, + } +} + +// DirEntry represents a single directory entry. +type DirEntry struct { + Name string `json:"name"` + Path string `json:"path"` + Type string `json:"type"` // "file" or "dir" + Size int64 `json:"size,omitempty"` + ModTime string `json:"mod_time,omitempty"` +} + +// ReadDirResult contains the directory listing output. +type ReadDirResult struct { + Path string `json:"path"` + Entries []DirEntry `json:"entries"` + TotalCount int `json:"total_count"` + Truncated bool `json:"truncated"` +} + +func (t *ReadDirTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + path := GetString(params, "path", "") + if path == "" { + return NewErrorResult("path is required", time.Since(start)) + } + + if strings.HasPrefix(path, "~/") { + if home, err := os.UserHomeDir(); err == nil { + path = filepath.Join(home, path[2:]) + } + } + + absPath, err := filepath.Abs(path) + if err != nil { + return NewErrorResult(fmt.Sprintf("invalid path: %v", err), time.Since(start)) + } + + info, err := os.Stat(absPath) + if err != nil { + if os.IsNotExist(err) { + return NewErrorResult(fmt.Sprintf("directory not found: %s", absPath), time.Since(start)) + } + return NewErrorResult(fmt.Sprintf("cannot access directory: %v", err), time.Since(start)) + } + + if !info.IsDir() { + return NewErrorResult("path is a file, not a directory", time.Since(start)) + } + + recursive := GetBool(params, "recursive", false) + maxDepth := GetInt(params, "max_depth", 0) + includeHidden := GetBool(params, "include_hidden", false) + maxEntries := GetInt(params, "max_entries", 0) + if maxEntries <= 0 { + maxEntries = 1000 + } + + entries := make([]DirEntry, 0) + truncated := false + + if recursive { + entries, truncated = t.walkDir(absPath, absPath, 0, maxDepth, includeHidden, maxEntries) + } else { + entries, truncated = t.listDir(absPath, includeHidden, maxEntries) + } + + return NewResult(ReadDirResult{ + Path: absPath, + Entries: entries, + TotalCount: len(entries), + Truncated: truncated, + }, time.Since(start)) +} + +func (t *ReadDirTool) listDir(dir string, includeHidden bool, maxEntries int) ([]DirEntry, bool) { + files, err := os.ReadDir(dir) + if err != nil { + return nil, false + } + + entries := make([]DirEntry, 0, len(files)) + for _, f := range files { + if !includeHidden && strings.HasPrefix(f.Name(), ".") { + continue + } + + if len(entries) >= maxEntries { + return entries, true + } + + entry := DirEntry{ + Name: f.Name(), + Path: filepath.Join(dir, f.Name()), + } + + if f.IsDir() { + entry.Type = "dir" + } else { + entry.Type = "file" + if info, err := f.Info(); err == nil { + entry.Size = info.Size() + entry.ModTime = info.ModTime().Format(time.RFC3339) + } + } + + entries = append(entries, entry) + } + + return entries, false +} + +func (t *ReadDirTool) walkDir(basePath, currentPath string, currentDepth, maxDepth int, includeHidden bool, maxEntries int) ([]DirEntry, bool) { + if maxDepth > 0 && currentDepth >= maxDepth { + return nil, false + } + + files, err := os.ReadDir(currentPath) + if err != nil { + return nil, false + } + + entries := make([]DirEntry, 0) + for _, f := range files { + if !includeHidden && strings.HasPrefix(f.Name(), ".") { + continue + } + + if len(entries) >= maxEntries { + return entries, true + } + + fullPath := filepath.Join(currentPath, f.Name()) + entry := DirEntry{ + Name: f.Name(), + Path: fullPath, + } + + if f.IsDir() { + entry.Type = "dir" + entries = append(entries, entry) + + subEntries, truncated := t.walkDir(basePath, fullPath, currentDepth+1, maxDepth, includeHidden, maxEntries-len(entries)) + entries = append(entries, subEntries...) + if truncated { + return entries, true + } + } else { + entry.Type = "file" + if info, err := f.Info(); err == nil { + entry.Size = info.Size() + entry.ModTime = info.ModTime().Format(time.RFC3339) + } + entries = append(entries, entry) + } + } + + return entries, false +} diff --git a/internal/tools/git_inspector.go b/internal/tools/git_inspector.go new file mode 100644 index 0000000..56e628b --- /dev/null +++ b/internal/tools/git_inspector.go @@ -0,0 +1,120 @@ +package tools + +import ( + "context" + "strconv" + "strings" + "time" + + "dev-cli/internal/executor" +) + +// GitInspectorTool gathers git context for error diagnosis. +// Runs git status and git log -n 5 to provide repository context. +type GitInspectorTool struct{} + +func (t *GitInspectorTool) Name() string { return "git_inspector" } +func (t *GitInspectorTool) Description() string { + return "Gather git repository context: status and recent commits for error diagnosis" +} + +func (t *GitInspectorTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "commit_count", Type: "int", Description: "Number of recent commits to include", Required: false, Default: 5}, + } +} + +// GitInspectorResult contains combined git status and log output. +type GitInspectorResult struct { + InGitRepo bool `json:"in_git_repo"` + Branch string `json:"branch,omitempty"` + Clean bool `json:"clean"` + Staged []FileChange `json:"staged,omitempty"` + Unstaged []FileChange `json:"unstaged,omitempty"` + RecentCommits []GitCommit `json:"recent_commits,omitempty"` + Error string `json:"error,omitempty"` +} + +func (t *GitInspectorTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + commitCount := GetInt(params, "commit_count", 5) + + checkResult := executor.ExecuteSimple("git rev-parse --is-inside-work-tree") + if checkResult.ExitCode != 0 { + return NewResult(GitInspectorResult{ + InGitRepo: false, + Error: "not a git repository", + }, time.Since(start)) + } + + result := GitInspectorResult{InGitRepo: true} + + branchResult := executor.ExecuteSimple("git branch --show-current") + result.Branch = strings.TrimSpace(branchResult.Output) + + statusResult := executor.ExecuteSimple("git status --porcelain") + result.Clean = strings.TrimSpace(statusResult.Output) == "" + + for _, line := range strings.Split(statusResult.Output, "\n") { + if len(line) < 3 { + continue + } + indexStatus := line[0] + workStatus := line[1] + path := strings.TrimSpace(line[3:]) + + if indexStatus != ' ' && indexStatus != '?' { + result.Staged = append(result.Staged, FileChange{ + Status: string(indexStatus), + Path: path, + }) + } + if workStatus != ' ' { + result.Unstaged = append(result.Unstaged, FileChange{ + Status: string(workStatus), + Path: path, + }) + } + } + + logCmd := "git log --pretty=format:'%h|%an|%ad|%s' --date=short -n " + strconv.Itoa(commitCount) + logResult := executor.ExecuteSimple(logCmd) + + if logResult.ExitCode == 0 { + for _, line := range strings.Split(logResult.Output, "\n") { + line = strings.Trim(line, "'") + if line == "" { + continue + } + parts := strings.SplitN(line, "|", 4) + if len(parts) == 4 { + result.RecentCommits = append(result.RecentCommits, GitCommit{ + Hash: parts[0], + Author: parts[1], + Date: parts[2], + Subject: parts[3], + }) + } + } + } + + return NewResult(result, time.Since(start)) +} + +// InspectOnError is a helper that returns git context when a command fails. +// This can be called after command execution to enrich error context. +func InspectOnError(exitCode int) *GitInspectorResult { + if exitCode == 0 { + return nil + } + + tool := &GitInspectorTool{} + result := tool.Execute(context.Background(), map[string]any{"commit_count": 5}) + if result.Success { + if data, ok := result.Data.(GitInspectorResult); ok { + return &data + } + } + return nil +} diff --git a/internal/tools/git_tools.go b/internal/tools/git_tools.go new file mode 100644 index 0000000..b840264 --- /dev/null +++ b/internal/tools/git_tools.go @@ -0,0 +1,261 @@ +package tools + +import ( + "context" + "strconv" + "strings" + "time" + + "dev-cli/internal/executor" +) + +// GitInfoTool retrieves Git repository information. +type GitInfoTool struct{} + +func (t *GitInfoTool) Name() string { return "git_info" } +func (t *GitInfoTool) Description() string { + return "Get Git repository info: commits, blame, diff, status" +} + +func (t *GitInfoTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "action", Type: "string", Description: "Action: log, blame, diff, status, branch", Required: true}, + {Name: "path", Type: "string", Description: "File path (for blame)", Required: false}, + {Name: "count", Type: "int", Description: "Number of commits (for log)", Required: false, Default: 10}, + {Name: "ref", Type: "string", Description: "Git ref (branch, commit, tag)", Required: false, Default: "HEAD"}, + } +} + +// GitLogResult contains git log output. +type GitLogResult struct { + Commits []GitCommit `json:"commits"` + Count int `json:"count"` +} + +// GitCommit represents a single commit. +type GitCommit struct { + Hash string `json:"hash"` + Author string `json:"author"` + Date string `json:"date"` + Subject string `json:"subject"` +} + +// GitBlameResult contains git blame output. +type GitBlameResult struct { + Path string `json:"path"` + Lines []BlameLine `json:"lines"` +} + +// BlameLine represents a line with blame info. +type BlameLine struct { + LineNum int `json:"line"` + Hash string `json:"hash"` + Author string `json:"author"` + Content string `json:"content"` +} + +// GitDiffResult contains git diff output. +type GitDiffResult struct { + Ref string `json:"ref"` + Diff string `json:"diff"` + Stats string `json:"stats"` + Changed int `json:"changed"` +} + +// GitStatusResult contains git status output. +type GitStatusResult struct { + Branch string `json:"branch"` + Clean bool `json:"clean"` + Staged []FileChange `json:"staged"` + Unstaged []FileChange `json:"unstaged"` +} + +// FileChange represents a changed file. +type FileChange struct { + Status string `json:"status"` + Path string `json:"path"` +} + +// GitBranchResult contains git branch info. +type GitBranchResult struct { + Current string `json:"current"` + Branches []string `json:"branches"` +} + +func (t *GitInfoTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + action := GetString(params, "action", "") + if action == "" { + return NewErrorResult("action is required (log, blame, diff, status, branch)", time.Since(start)) + } + + switch action { + case "log": + return t.getLog(params, start) + case "blame": + return t.getBlame(params, start) + case "diff": + return t.getDiff(params, start) + case "status": + return t.getStatus(start) + case "branch": + return t.getBranch(start) + default: + return NewErrorResult("unknown action: "+action, time.Since(start)) + } +} + +func (t *GitInfoTool) getLog(params map[string]any, start time.Time) ToolResult { + count := GetInt(params, "count", 10) + ref := GetString(params, "ref", "HEAD") + + cmd := "git log --pretty=format:'%h|%an|%ad|%s' --date=short -n " + strconv.Itoa(count) + " " + ref + result := executor.ExecuteSimple(cmd) + + if result.ExitCode != 0 { + return NewErrorResult("git log failed: "+result.Output, time.Since(start)) + } + + commits := make([]GitCommit, 0) + for _, line := range strings.Split(result.Output, "\n") { + line = strings.Trim(line, "'") + if line == "" { + continue + } + parts := strings.SplitN(line, "|", 4) + if len(parts) == 4 { + commits = append(commits, GitCommit{ + Hash: parts[0], + Author: parts[1], + Date: parts[2], + Subject: parts[3], + }) + } + } + + return NewResult(GitLogResult{Commits: commits, Count: len(commits)}, time.Since(start)) +} + +func (t *GitInfoTool) getBlame(params map[string]any, start time.Time) ToolResult { + path := GetString(params, "path", "") + if path == "" { + return NewErrorResult("path is required for blame action", time.Since(start)) + } + + cmd := "git blame --line-porcelain " + path + result := executor.ExecuteSimple(cmd) + + if result.ExitCode != 0 { + return NewErrorResult("git blame failed: "+result.Output, time.Since(start)) + } + + lines := parseBlameOutput(result.Output) + + return NewResult(GitBlameResult{Path: path, Lines: lines}, time.Since(start)) +} + +func parseBlameOutput(output string) []BlameLine { + var lines []BlameLine + var current BlameLine + lineNum := 0 + + for _, line := range strings.Split(output, "\n") { + if len(line) >= 40 && !strings.HasPrefix(line, "\t") { + + parts := strings.Fields(line) + if len(parts) >= 1 { + current.Hash = parts[0][:8] + } + } else if strings.HasPrefix(line, "author ") { + current.Author = strings.TrimPrefix(line, "author ") + } else if strings.HasPrefix(line, "\t") { + lineNum++ + current.LineNum = lineNum + current.Content = strings.TrimPrefix(line, "\t") + lines = append(lines, current) + current = BlameLine{} + } + } + + return lines +} + +func (t *GitInfoTool) getDiff(params map[string]any, start time.Time) ToolResult { + ref := GetString(params, "ref", "HEAD") + + diffResult := executor.ExecuteSimple("git diff " + ref) + + statsResult := executor.ExecuteSimple("git diff --stat " + ref) + + changedResult := executor.ExecuteSimple("git diff --name-only " + ref + " | wc -l") + changed, _ := strconv.Atoi(strings.TrimSpace(changedResult.Output)) + + return NewResult(GitDiffResult{ + Ref: ref, + Diff: diffResult.Output, + Stats: statsResult.Output, + Changed: changed, + }, time.Since(start)) +} + +func (t *GitInfoTool) getStatus(start time.Time) ToolResult { + + branchResult := executor.ExecuteSimple("git branch --show-current") + branch := strings.TrimSpace(branchResult.Output) + + statusResult := executor.ExecuteSimple("git status --porcelain") + + staged := make([]FileChange, 0) + unstaged := make([]FileChange, 0) + + for _, line := range strings.Split(statusResult.Output, "\n") { + if len(line) < 3 { + continue + } + indexStatus := line[0] + workStatus := line[1] + path := strings.TrimSpace(line[3:]) + + if indexStatus != ' ' && indexStatus != '?' { + staged = append(staged, FileChange{ + Status: string(indexStatus), + Path: path, + }) + } + if workStatus != ' ' { + unstaged = append(unstaged, FileChange{ + Status: string(workStatus), + Path: path, + }) + } + } + + return NewResult(GitStatusResult{ + Branch: branch, + Clean: len(staged) == 0 && len(unstaged) == 0, + Staged: staged, + Unstaged: unstaged, + }, time.Since(start)) +} + +func (t *GitInfoTool) getBranch(start time.Time) ToolResult { + + currentResult := executor.ExecuteSimple("git branch --show-current") + current := strings.TrimSpace(currentResult.Output) + + branchesResult := executor.ExecuteSimple("git branch --format='%(refname:short)'") + + branches := make([]string, 0) + for _, line := range strings.Split(branchesResult.Output, "\n") { + line = strings.Trim(strings.TrimSpace(line), "'") + if line != "" { + branches = append(branches, line) + } + } + + return NewResult(GitBranchResult{ + Current: current, + Branches: branches, + }, time.Since(start)) +} diff --git a/internal/tools/network_tools.go b/internal/tools/network_tools.go new file mode 100644 index 0000000..1b0ef7b --- /dev/null +++ b/internal/tools/network_tools.go @@ -0,0 +1,103 @@ +package tools + +import ( + "context" + "time" + + "dev-cli/internal/infra" +) + +// CheckPortsTool checks port availability and conflicts. +type CheckPortsTool struct{} + +func (t *CheckPortsTool) Name() string { return "check_ports" } +func (t *CheckPortsTool) Description() string { + return "Check port availability and find processes using ports" +} + +func (t *CheckPortsTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "ports", Type: "[]int", Description: "Ports to check", Required: true}, + {Name: "action", Type: "string", Description: "Action: check, suggest", Required: false, Default: "check"}, + } +} + +// PortStatus represents the status of a single port. +type PortStatus struct { + Port int `json:"port"` + Available bool `json:"available"` + Process string `json:"process,omitempty"` + PID int `json:"pid,omitempty"` + Suggested int `json:"suggested,omitempty"` +} + +// PortCheckResult contains port check results. +type PortCheckResult struct { + Ports []PortStatus `json:"ports"` + Conflicts int `json:"conflicts"` + AllFree bool `json:"all_free"` +} + +func (t *CheckPortsTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + ports := getIntSlice(params, "ports") + if len(ports) == 0 { + return NewErrorResult("ports is required", time.Since(start)) + } + + action := GetString(params, "action", "check") + + results := make([]PortStatus, 0, len(ports)) + conflicts := 0 + + for _, port := range ports { + status := PortStatus{Port: port} + + conflict := infra.CheckPortAvailable(port) + if conflict == nil { + status.Available = true + } else { + status.Available = false + status.Process = conflict.Process + status.PID = conflict.PID + conflicts++ + + if action == "suggest" { + status.Suggested = infra.FindAvailablePort(port) + } + } + + results = append(results, status) + } + + return NewResult(PortCheckResult{ + Ports: results, + Conflicts: conflicts, + AllFree: conflicts == 0, + }, time.Since(start)) +} + +// getIntSlice extracts an int slice from params. +func getIntSlice(params map[string]any, key string) []int { + if v, ok := params[key]; ok { + switch s := v.(type) { + case []int: + return s + case []any: + result := make([]int, 0, len(s)) + for _, item := range s { + switch n := item.(type) { + case int: + result = append(result, n) + case int64: + result = append(result, int(n)) + case float64: + result = append(result, int(n)) + } + } + return result + } + } + return nil +} diff --git a/internal/tools/package_tools.go b/internal/tools/package_tools.go new file mode 100644 index 0000000..cbed948 --- /dev/null +++ b/internal/tools/package_tools.go @@ -0,0 +1,328 @@ +package tools + +import ( + "context" + "encoding/json" + "os" + "path/filepath" + "strings" + "time" + + "dev-cli/internal/executor" +) + +// PackageInfoTool analyzes project dependencies. +type PackageInfoTool struct{} + +func (t *PackageInfoTool) Name() string { return "package_info" } +func (t *PackageInfoTool) Description() string { return "Analyze project dependencies (Go, npm, pip)" } + +func (t *PackageInfoTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "type", Type: "string", Description: "Package type: auto, go, npm, pip", Required: false, Default: "auto"}, + {Name: "action", Type: "string", Description: "Action: list, outdated, check", Required: false, Default: "list"}, + {Name: "path", Type: "string", Description: "Project path", Required: false, Default: "."}, + } +} + +// PackageResult contains dependency analysis results. +type PackageResult struct { + Type string `json:"type"` + Path string `json:"path"` + Packages []PackageInfo `json:"packages,omitempty"` + Outdated []PackageInfo `json:"outdated,omitempty"` + TotalCount int `json:"total_count"` + DirectCount int `json:"direct_count"` +} + +// PackageInfo represents a single package/dependency. +type PackageInfo struct { + Name string `json:"name"` + Version string `json:"version"` + Latest string `json:"latest,omitempty"` + Direct bool `json:"direct,omitempty"` + Indirect bool `json:"indirect,omitempty"` +} + +func (t *PackageInfoTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + pkgType := GetString(params, "type", "auto") + action := GetString(params, "action", "list") + path := GetString(params, "path", ".") + + if pkgType == "auto" { + pkgType = detectPackageType(path) + if pkgType == "" { + return NewErrorResult("could not detect package type, specify 'type' parameter", time.Since(start)) + } + } + + switch pkgType { + case "go": + return t.analyzeGo(path, action, start) + case "npm": + return t.analyzeNpm(path, action, start) + case "pip": + return t.analyzePip(path, action, start) + default: + return NewErrorResult("unknown package type: "+pkgType, time.Since(start)) + } +} + +func detectPackageType(path string) string { + if _, err := os.Stat(filepath.Join(path, "go.mod")); err == nil { + return "go" + } + if _, err := os.Stat(filepath.Join(path, "package.json")); err == nil { + return "npm" + } + if _, err := os.Stat(filepath.Join(path, "requirements.txt")); err == nil { + return "pip" + } + if _, err := os.Stat(filepath.Join(path, "pyproject.toml")); err == nil { + return "pip" + } + return "" +} + +func (t *PackageInfoTool) analyzeGo(path, action string, start time.Time) ToolResult { + absPath, err := filepath.Abs(path) + if err != nil { + absPath = path + } + + switch action { + case "list": + + result := executor.ExecuteSimple("cd " + absPath + " && go list -m -f '{{.Path}}@{{.Version}}' all 2>/dev/null | head -100") + + packages := make([]PackageInfo, 0) + directCount := 0 + + lines := strings.Split(result.Output, "\n") + for i, line := range lines { + line = strings.Trim(line, "'") + if line == "" { + continue + } + parts := strings.Split(line, "@") + if len(parts) == 2 { + pkg := PackageInfo{ + Name: parts[0], + Version: parts[1], + Direct: i == 0, + } + packages = append(packages, pkg) + if pkg.Direct { + directCount++ + } + } + } + + return NewResult(PackageResult{ + Type: "go", + Path: absPath, + Packages: packages, + TotalCount: len(packages), + DirectCount: directCount, + }, time.Since(start)) + + case "outdated": + + result := executor.ExecuteSimple("cd " + absPath + " && go list -u -m -f '{{if .Update}}{{.Path}}@{{.Version}}->{{.Update.Version}}{{end}}' all 2>/dev/null") + + outdated := make([]PackageInfo, 0) + for _, line := range strings.Split(result.Output, "\n") { + line = strings.Trim(line, "'") + if line == "" { + continue + } + + parts := strings.Split(line, "->") + if len(parts) == 2 { + nameParts := strings.Split(parts[0], "@") + if len(nameParts) == 2 { + outdated = append(outdated, PackageInfo{ + Name: nameParts[0], + Version: nameParts[1], + Latest: parts[1], + }) + } + } + } + + return NewResult(PackageResult{ + Type: "go", + Path: absPath, + Outdated: outdated, + TotalCount: len(outdated), + }, time.Since(start)) + + default: + return NewErrorResult("unknown action for go: "+action, time.Since(start)) + } +} + +func (t *PackageInfoTool) analyzeNpm(path, action string, start time.Time) ToolResult { + absPath, err := filepath.Abs(path) + if err != nil { + absPath = path + } + + switch action { + case "list": + + pkgPath := filepath.Join(absPath, "package.json") + data, err := os.ReadFile(pkgPath) + if err != nil { + return NewErrorResult("cannot read package.json: "+err.Error(), time.Since(start)) + } + + var pkg struct { + Dependencies map[string]string `json:"dependencies"` + DevDependencies map[string]string `json:"devDependencies"` + } + if err := json.Unmarshal(data, &pkg); err != nil { + return NewErrorResult("invalid package.json: "+err.Error(), time.Since(start)) + } + + packages := make([]PackageInfo, 0) + for name, version := range pkg.Dependencies { + packages = append(packages, PackageInfo{Name: name, Version: version, Direct: true}) + } + for name, version := range pkg.DevDependencies { + packages = append(packages, PackageInfo{Name: name, Version: version, Direct: true}) + } + + return NewResult(PackageResult{ + Type: "npm", + Path: absPath, + Packages: packages, + TotalCount: len(packages), + DirectCount: len(packages), + }, time.Since(start)) + + case "outdated": + result := executor.ExecuteSimple("cd " + absPath + " && npm outdated --json 2>/dev/null") + + var outdatedMap map[string]struct { + Current string `json:"current"` + Latest string `json:"latest"` + } + + outdated := make([]PackageInfo, 0) + if err := json.Unmarshal([]byte(result.Output), &outdatedMap); err == nil { + for name, info := range outdatedMap { + outdated = append(outdated, PackageInfo{ + Name: name, + Version: info.Current, + Latest: info.Latest, + }) + } + } + + return NewResult(PackageResult{ + Type: "npm", + Path: absPath, + Outdated: outdated, + TotalCount: len(outdated), + }, time.Since(start)) + + default: + return NewErrorResult("unknown action for npm: "+action, time.Since(start)) + } +} + +func (t *PackageInfoTool) analyzePip(path, action string, start time.Time) ToolResult { + absPath, err := filepath.Abs(path) + if err != nil { + absPath = path + } + + switch action { + case "list": + + reqPath := filepath.Join(absPath, "requirements.txt") + data, err := os.ReadFile(reqPath) + if err != nil { + + result := executor.ExecuteSimple("pip list --format=json 2>/dev/null") + return t.parsePipList(result.Output, absPath, start) + } + + packages := make([]PackageInfo, 0) + for _, line := range strings.Split(string(data), "\n") { + line = strings.TrimSpace(line) + if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, "-") { + continue + } + + for _, sep := range []string{"==", ">=", "<=", "~=", "!="} { + if parts := strings.SplitN(line, sep, 2); len(parts) == 2 { + packages = append(packages, PackageInfo{Name: parts[0], Version: parts[1], Direct: true}) + break + } + } + } + + return NewResult(PackageResult{ + Type: "pip", + Path: absPath, + Packages: packages, + TotalCount: len(packages), + DirectCount: len(packages), + }, time.Since(start)) + + case "outdated": + result := executor.ExecuteSimple("pip list --outdated --format=json 2>/dev/null") + + var outdatedList []struct { + Name string `json:"name"` + Version string `json:"version"` + Latest string `json:"latest_version"` + } + + outdated := make([]PackageInfo, 0) + if err := json.Unmarshal([]byte(result.Output), &outdatedList); err == nil { + for _, pkg := range outdatedList { + outdated = append(outdated, PackageInfo{ + Name: pkg.Name, + Version: pkg.Version, + Latest: pkg.Latest, + }) + } + } + + return NewResult(PackageResult{ + Type: "pip", + Path: absPath, + Outdated: outdated, + TotalCount: len(outdated), + }, time.Since(start)) + + default: + return NewErrorResult("unknown action for pip: "+action, time.Since(start)) + } +} + +func (t *PackageInfoTool) parsePipList(output, path string, start time.Time) ToolResult { + var pkgList []struct { + Name string `json:"name"` + Version string `json:"version"` + } + + packages := make([]PackageInfo, 0) + if err := json.Unmarshal([]byte(output), &pkgList); err == nil { + for _, pkg := range pkgList { + packages = append(packages, PackageInfo{Name: pkg.Name, Version: pkg.Version}) + } + } + + return NewResult(PackageResult{ + Type: "pip", + Path: path, + Packages: packages, + TotalCount: len(packages), + }, time.Since(start)) +} diff --git a/internal/tools/registry.go b/internal/tools/registry.go new file mode 100644 index 0000000..6a39079 --- /dev/null +++ b/internal/tools/registry.go @@ -0,0 +1,151 @@ +package tools + +import ( + "fmt" + "sort" + "sync" +) + +// Registry manages tool registration and lookup. +type Registry struct { + mu sync.RWMutex + tools map[string]Tool +} + +var ( + globalRegistry *Registry + globalRegistryOnce sync.Once +) + +// GetRegistry returns the global tool registry singleton. +func GetRegistry() *Registry { + globalRegistryOnce.Do(func() { + globalRegistry = NewRegistry() + }) + return globalRegistry +} + +// NewRegistry creates a new tool registry. +func NewRegistry() *Registry { + return &Registry{ + tools: make(map[string]Tool), + } +} + +// Register adds a tool to the registry. +func (r *Registry) Register(tool Tool) error { + r.mu.Lock() + defer r.mu.Unlock() + + name := tool.Name() + if _, exists := r.tools[name]; exists { + return fmt.Errorf("tool %q already registered", name) + } + r.tools[name] = tool + return nil +} + +// MustRegister adds a tool to the registry, panicking on error. +func (r *Registry) MustRegister(tool Tool) { + if err := r.Register(tool); err != nil { + panic(err) + } +} + +// Get retrieves a tool by name. +func (r *Registry) Get(name string) (Tool, bool) { + r.mu.RLock() + defer r.mu.RUnlock() + tool, ok := r.tools[name] + return tool, ok +} + +// List returns information about all registered tools. +func (r *Registry) List() []ToolInfo { + r.mu.RLock() + defer r.mu.RUnlock() + + infos := make([]ToolInfo, 0, len(r.tools)) + for _, tool := range r.tools { + infos = append(infos, ToolInfo{ + Name: tool.Name(), + Description: tool.Description(), + Parameters: tool.Parameters(), + }) + } + + sort.Slice(infos, func(i, j int) bool { + return infos[i].Name < infos[j].Name + }) + + return infos +} + +// Names returns the names of all registered tools. +func (r *Registry) Names() []string { + r.mu.RLock() + defer r.mu.RUnlock() + + names := make([]string, 0, len(r.tools)) + for name := range r.tools { + names = append(names, name) + } + sort.Strings(names) + return names +} + +// Count returns the number of registered tools. +func (r *Registry) Count() int { + r.mu.RLock() + defer r.mu.RUnlock() + return len(r.tools) +} + +// RegisterAll registers multiple tools. +func (r *Registry) RegisterAll(tools ...Tool) error { + for _, tool := range tools { + if err := r.Register(tool); err != nil { + return err + } + } + return nil +} + +// RegisterDefaults registers all default tools. +// Call this to populate the registry with standard RCA tools. +func (r *Registry) RegisterDefaults() { + r.MustRegister(&ReadFileTool{}) + r.MustRegister(&ReadDirTool{}) + r.MustRegister(&WriteFileTool{}) + r.MustRegister(&RunCommandTool{}) + r.MustRegister(&SearchCodebaseTool{}) + r.MustRegister(&QueryDockerTool{}) + r.MustRegister(&CheckPortsTool{}) + r.MustRegister(&GitInfoTool{}) + r.MustRegister(&PackageInfoTool{}) + r.MustRegister(&GitInspectorTool{}) +} + +// GetSchemas returns JSON schemas for all registered tools. +func (r *Registry) GetSchemas() []ToolSchema { + r.mu.RLock() + defer r.mu.RUnlock() + + tools := make([]Tool, 0, len(r.tools)) + for _, tool := range r.tools { + tools = append(tools, tool) + } + return GenerateToolsSchema(tools) +} + +// GetSchemasJSON returns JSON string of all tool schemas for LLM prompts. +func (r *Registry) GetSchemasJSON() (string, error) { + r.mu.RLock() + defer r.mu.RUnlock() + + tools := make([]Tool, 0, len(r.tools)) + for _, tool := range r.tools { + tools = append(tools, tool) + } + return ToolsPromptJSON(tools) +} diff --git a/internal/tools/schema.go b/internal/tools/schema.go new file mode 100644 index 0000000..c1ac533 --- /dev/null +++ b/internal/tools/schema.go @@ -0,0 +1,124 @@ +// Package tools provides a unified tool abstraction for the RCA agent. +// This file contains JSON Schema generation for LLM tool integration. +package tools + +import ( + "encoding/json" +) + +// ToolSchema represents a tool in JSON Schema format for LLM integration. +type ToolSchema struct { + Name string `json:"name"` + Description string `json:"description"` + Parameters ToolSchemaParams `json:"parameters"` +} + +// ToolSchemaParams defines the parameters schema for a tool. +type ToolSchemaParams struct { + Type string `json:"type"` + Properties map[string]ToolSchemaProperty `json:"properties"` + Required []string `json:"required"` +} + +// ToolSchemaProperty defines a single parameter property. +type ToolSchemaProperty struct { + Type string `json:"type"` + Description string `json:"description"` + Default any `json:"default,omitempty"` + Items *ToolSchemaItems `json:"items,omitempty"` +} + +// ToolSchemaItems defines array item schema. +type ToolSchemaItems struct { + Type string `json:"type"` +} + +// GenerateToolSchema converts a Tool to JSON Schema format. +func GenerateToolSchema(tool Tool) ToolSchema { + params := tool.Parameters() + properties := make(map[string]ToolSchemaProperty) + required := make([]string, 0) + + for _, p := range params { + prop := ToolSchemaProperty{ + Type: mapTypeToJSONSchema(p.Type), + Description: p.Description, + } + + if p.Type == "[]string" { + prop.Items = &ToolSchemaItems{Type: "string"} + } else if p.Type == "[]int" { + prop.Items = &ToolSchemaItems{Type: "integer"} + } + + if p.Default != nil { + prop.Default = p.Default + } + + properties[p.Name] = prop + + if p.Required { + required = append(required, p.Name) + } + } + + return ToolSchema{ + Name: tool.Name(), + Description: tool.Description(), + Parameters: ToolSchemaParams{ + Type: "object", + Properties: properties, + Required: required, + }, + } +} + +// GenerateToolsSchema converts multiple tools to JSON Schema format. +func GenerateToolsSchema(tools []Tool) []ToolSchema { + schemas := make([]ToolSchema, len(tools)) + for i, tool := range tools { + schemas[i] = GenerateToolSchema(tool) + } + return schemas +} + +// mapTypeToJSONSchema converts internal type names to JSON Schema types. +func mapTypeToJSONSchema(internalType string) string { + switch internalType { + case "string", "duration": + return "string" + case "int": + return "integer" + case "bool": + return "boolean" + case "[]string", "[]int": + return "array" + default: + return "string" + } +} + +// ToolCallRequest represents a tool call from the LLM. +type ToolCallRequest struct { + ToolName string `json:"tool_name"` + Parameters map[string]any `json:"parameters"` +} + +// ToolsPromptJSON returns a JSON string of all tool schemas for LLM prompts. +func ToolsPromptJSON(tools []Tool) (string, error) { + schemas := GenerateToolsSchema(tools) + data, err := json.MarshalIndent(schemas, "", " ") + if err != nil { + return "", err + } + return string(data), nil +} + +// ParseToolCall parses an LLM response into a ToolCallRequest. +func ParseToolCall(response string) (*ToolCallRequest, error) { + var call ToolCallRequest + if err := json.Unmarshal([]byte(response), &call); err != nil { + return nil, err + } + return &call, nil +} diff --git a/internal/tools/search_tools.go b/internal/tools/search_tools.go new file mode 100644 index 0000000..afccaae --- /dev/null +++ b/internal/tools/search_tools.go @@ -0,0 +1,189 @@ +package tools + +import ( + "context" + "encoding/json" + "os/exec" + "strconv" + "strings" + "time" +) + +// SearchCodebaseTool searches for patterns in code using ripgrep. +type SearchCodebaseTool struct{} + +func (t *SearchCodebaseTool) Name() string { return "search_codebase" } +func (t *SearchCodebaseTool) Description() string { return "Search for patterns in code using ripgrep" } + +func (t *SearchCodebaseTool) Parameters() []ToolParam { + return []ToolParam{ + {Name: "pattern", Type: "string", Description: "Search pattern (regex)", Required: true}, + {Name: "path", Type: "string", Description: "Path to search in", Required: false, Default: "."}, + {Name: "file_types", Type: "[]string", Description: "File types to include (e.g., 'go', 'py')", Required: false}, + {Name: "ignore_case", Type: "bool", Description: "Case-insensitive search", Required: false, Default: false}, + {Name: "max_results", Type: "int", Description: "Maximum results", Required: false, Default: 50}, + {Name: "context_lines", Type: "int", Description: "Context lines around match", Required: false, Default: 0}, + } +} + +// SearchMatch represents a single search match. +type SearchMatch struct { + File string `json:"file"` + Line int `json:"line"` + Column int `json:"column,omitempty"` + Content string `json:"content"` +} + +// SearchResult contains the search output. +type SearchResult struct { + Pattern string `json:"pattern"` + Path string `json:"path"` + Matches []SearchMatch `json:"matches"` + TotalCount int `json:"total_count"` + Truncated bool `json:"truncated"` +} + +func (t *SearchCodebaseTool) Execute(ctx context.Context, params map[string]any) ToolResult { + start := time.Now() + + pattern := GetString(params, "pattern", "") + if pattern == "" { + return NewErrorResult("pattern is required", time.Since(start)) + } + + searchPath := GetString(params, "path", ".") + ignoreCase := GetBool(params, "ignore_case", false) + maxResults := GetInt(params, "max_results", 50) + contextLines := GetInt(params, "context_lines", 0) + fileTypes := GetStringSlice(params, "file_types") + + if _, err := exec.LookPath("rg"); err != nil { + + return t.executeWithGrep(ctx, pattern, searchPath, ignoreCase, maxResults) + } + + args := []string{ + "--json", + "--max-count", strconv.Itoa(maxResults * 2), + } + + if ignoreCase { + args = append(args, "-i") + } + + if contextLines > 0 { + args = append(args, "-C", strconv.Itoa(contextLines)) + } + + for _, ft := range fileTypes { + args = append(args, "-t", ft) + } + + args = append(args, pattern, searchPath) + + cmd := exec.CommandContext(ctx, "rg", args...) + output, _ := cmd.Output() + + matches := parseRipgrepJSON(string(output)) + + truncated := false + if len(matches) > maxResults { + matches = matches[:maxResults] + truncated = true + } + + return NewResult(SearchResult{ + Pattern: pattern, + Path: searchPath, + Matches: matches, + TotalCount: len(matches), + Truncated: truncated, + }, time.Since(start)) +} + +func (t *SearchCodebaseTool) executeWithGrep(ctx context.Context, pattern, path string, ignoreCase bool, maxResults int) ToolResult { + start := time.Now() + + args := []string{"-rn"} + if ignoreCase { + args = append(args, "-i") + } + args = append(args, pattern, path) + + cmd := exec.CommandContext(ctx, "grep", args...) + output, _ := cmd.Output() + + lines := strings.Split(strings.TrimSpace(string(output)), "\n") + matches := make([]SearchMatch, 0, len(lines)) + + for _, line := range lines { + if line == "" { + continue + } + + parts := strings.SplitN(line, ":", 3) + if len(parts) >= 3 { + lineNum, _ := strconv.Atoi(parts[1]) + matches = append(matches, SearchMatch{ + File: parts[0], + Line: lineNum, + Content: parts[2], + }) + } + if len(matches) >= maxResults { + break + } + } + + return NewResult(SearchResult{ + Pattern: pattern, + Path: path, + Matches: matches, + TotalCount: len(matches), + Truncated: len(lines) > maxResults, + }, time.Since(start)) +} + +func parseRipgrepJSON(output string) []SearchMatch { + var matches []SearchMatch + + for _, line := range strings.Split(output, "\n") { + if line == "" { + continue + } + + var msg struct { + Type string `json:"type"` + Data struct { + Path struct { + Text string `json:"text"` + } `json:"path"` + LineNumber int `json:"line_number"` + Lines struct { + Text string `json:"text"` + } `json:"lines"` + Submatches []struct { + Start int `json:"start"` + } `json:"submatches"` + } `json:"data"` + } + + if err := json.Unmarshal([]byte(line), &msg); err != nil { + continue + } + + if msg.Type == "match" { + match := SearchMatch{ + File: msg.Data.Path.Text, + Line: msg.Data.LineNumber, + Content: strings.TrimRight(msg.Data.Lines.Text, "\n"), + } + if len(msg.Data.Submatches) > 0 { + match.Column = msg.Data.Submatches[0].Start + 1 + } + matches = append(matches, match) + } + } + + return matches +} diff --git a/internal/tools/tool.go b/internal/tools/tool.go new file mode 100644 index 0000000..dbb0a59 --- /dev/null +++ b/internal/tools/tool.go @@ -0,0 +1,141 @@ +// Package tools provides a unified tool abstraction for the RCA agent. +// Tools enable structured execution of diagnostic operations like file inspection, +// command execution, Docker queries, and Git analysis. +package tools + +import ( + "context" + "time" +) + +// Tool defines the interface for all agent tools. +type Tool interface { + // Name returns the unique identifier for this tool. + Name() string + + // Description returns a human-readable description of what this tool does. + Description() string + + // Parameters returns the parameter definitions for this tool. + Parameters() []ToolParam + + // Execute runs the tool with the given parameters. + Execute(ctx context.Context, params map[string]any) ToolResult +} + +// ToolParam defines a parameter for a tool. +type ToolParam struct { + Name string `json:"name"` + Type string `json:"type"` // string, int, bool, []string, []int + Description string `json:"description"` + Required bool `json:"required"` + Default any `json:"default,omitempty"` +} + +// ToolResult represents the outcome of a tool execution. +type ToolResult struct { + Success bool `json:"success"` + Data any `json:"data,omitempty"` + Error string `json:"error,omitempty"` + Duration time.Duration `json:"duration"` +} + +// ToolInfo provides metadata about a registered tool. +type ToolInfo struct { + Name string `json:"name"` + Description string `json:"description"` + Parameters []ToolParam `json:"parameters"` +} + +// NewResult creates a successful result with data. +func NewResult(data any, duration time.Duration) ToolResult { + return ToolResult{ + Success: true, + Data: data, + Duration: duration, + } +} + +// NewErrorResult creates a failed result with an error message. +func NewErrorResult(err string, duration time.Duration) ToolResult { + return ToolResult{ + Success: false, + Error: err, + Duration: duration, + } +} + +// GetString extracts a string parameter with a default value. +func GetString(params map[string]any, key string, defaultVal string) string { + if v, ok := params[key]; ok { + if s, ok := v.(string); ok { + return s + } + } + return defaultVal +} + +// GetInt extracts an int parameter with a default value. +func GetInt(params map[string]any, key string, defaultVal int) int { + if v, ok := params[key]; ok { + switch n := v.(type) { + case int: + return n + case int64: + return int(n) + case float64: + return int(n) + } + } + return defaultVal +} + +// GetBool extracts a bool parameter with a default value. +func GetBool(params map[string]any, key string, defaultVal bool) bool { + if v, ok := params[key]; ok { + if b, ok := v.(bool); ok { + return b + } + } + return defaultVal +} + +// GetStringSlice extracts a string slice parameter. +func GetStringSlice(params map[string]any, key string) []string { + if v, ok := params[key]; ok { + switch s := v.(type) { + case []string: + return s + case []any: + result := make([]string, 0, len(s)) + for _, item := range s { + if str, ok := item.(string); ok { + result = append(result, str) + } + } + return result + } + } + return nil +} + +// GetDuration extracts a duration parameter (accepts string like "30s" or seconds as int). +func GetDuration(params map[string]any, key string, defaultVal time.Duration) time.Duration { + if v, ok := params[key]; ok { + switch d := v.(type) { + case time.Duration: + return d + case string: + if parsed, err := time.ParseDuration(d); err == nil { + return parsed + } + case int: + return time.Duration(d) * time.Second + case int64: + return time.Duration(d) * time.Second + case float64: + return time.Duration(d) * time.Second + } + } + return defaultVal +} diff --git a/internal/tools/tools_test.go b/internal/tools/tools_test.go new file mode 100644 index 0000000..885dc99 --- /dev/null +++ b/internal/tools/tools_test.go @@ -0,0 +1,272 @@ +package tools + +import ( + "context" + "os" + "path/filepath" + "testing" + "time" +) + +func TestReadFileTool(t *testing.T) { + tool := &ReadFileTool{} + + t.Run("Name and Description", func(t *testing.T) { + if tool.Name() != "read_file" { + t.Errorf("expected name 'read_file', got %s", tool.Name()) + } + if tool.Description() == "" { + t.Error("expected non-empty description") + } + }) + + t.Run("Read existing file", func(t *testing.T) { + + tmpDir := t.TempDir() + testFile := filepath.Join(tmpDir, "test.txt") + content := "line1\nline2\nline3" + if err := os.WriteFile(testFile, []byte(content), 0644); err != nil { + t.Fatal(err) + } + + result := tool.Execute(context.Background(), map[string]any{ + "path": testFile, + }) + + if !result.Success { + t.Errorf("expected success, got error: %s", result.Error) + } + + data, ok := result.Data.(ReadFileResult) + if !ok { + t.Fatal("expected ReadFileResult data") + } + if data.Content != content { + t.Errorf("expected content %q, got %q", content, data.Content) + } + if data.Lines != 3 { + t.Errorf("expected 3 lines, got %d", data.Lines) + } + }) + + t.Run("Read with line range", func(t *testing.T) { + tmpDir := t.TempDir() + testFile := filepath.Join(tmpDir, "test.txt") + content := "line1\nline2\nline3\nline4\nline5" + if err := os.WriteFile(testFile, []byte(content), 0644); err != nil { + t.Fatal(err) + } + + result := tool.Execute(context.Background(), map[string]any{ + "path": testFile, + "start_line": 2, + "end_line": 4, + }) + + if !result.Success { + t.Errorf("expected success, got error: %s", result.Error) + } + + data := result.Data.(ReadFileResult) + if data.Lines != 3 { + t.Errorf("expected 3 lines, got %d", data.Lines) + } + }) + + t.Run("File not found", func(t *testing.T) { + result := tool.Execute(context.Background(), map[string]any{ + "path": "/nonexistent/file.txt", + }) + + if result.Success { + t.Error("expected error for non-existent file") + } + }) + + t.Run("Missing path parameter", func(t *testing.T) { + result := tool.Execute(context.Background(), map[string]any{}) + + if result.Success { + t.Error("expected error for missing path") + } + }) +} + +func TestWriteFileTool(t *testing.T) { + tool := &WriteFileTool{} + + t.Run("Write new file", func(t *testing.T) { + tmpDir := t.TempDir() + testFile := filepath.Join(tmpDir, "new.txt") + content := "hello world" + + result := tool.Execute(context.Background(), map[string]any{ + "path": testFile, + "content": content, + }) + + if !result.Success { + t.Errorf("expected success, got error: %s", result.Error) + } + + data := result.Data.(WriteFileResult) + if !data.Created { + t.Error("expected Created to be true") + } + + actual, _ := os.ReadFile(testFile) + if string(actual) != content { + t.Errorf("expected %q, got %q", content, string(actual)) + } + }) + + t.Run("Write with backup", func(t *testing.T) { + tmpDir := t.TempDir() + testFile := filepath.Join(tmpDir, "existing.txt") + + os.WriteFile(testFile, []byte("original"), 0644) + + result := tool.Execute(context.Background(), map[string]any{ + "path": testFile, + "content": "updated", + "backup": true, + }) + + if !result.Success { + t.Errorf("expected success, got error: %s", result.Error) + } + + data := result.Data.(WriteFileResult) + if data.BackupPath == "" { + t.Error("expected backup path") + } + + backup, _ := os.ReadFile(data.BackupPath) + if string(backup) != "original" { + t.Errorf("expected backup content 'original', got %q", string(backup)) + } + }) +} + +func TestRunCommandTool(t *testing.T) { + tool := &RunCommandTool{} + + t.Run("Execute simple command", func(t *testing.T) { + result := tool.Execute(context.Background(), map[string]any{ + "command": "echo hello", + }) + + if !result.Success { + t.Errorf("expected success, got error: %s", result.Error) + } + + data := result.Data.(CommandResult) + if data.ExitCode != 0 { + t.Errorf("expected exit code 0, got %d", data.ExitCode) + } + }) + + t.Run("Missing command", func(t *testing.T) { + result := tool.Execute(context.Background(), map[string]any{}) + + if result.Success { + t.Error("expected error for missing command") + } + }) +} + +func TestRegistry(t *testing.T) { + t.Run("Register and get tool", func(t *testing.T) { + reg := NewRegistry() + + tool := &ReadFileTool{} + if err := reg.Register(tool); err != nil { + t.Errorf("unexpected error: %v", err) + } + + got, ok := reg.Get("read_file") + if !ok { + t.Error("expected to find registered tool") + } + if got.Name() != "read_file" { + t.Errorf("expected 'read_file', got %s", got.Name()) + } + }) + + t.Run("Duplicate registration", func(t *testing.T) { + reg := NewRegistry() + + tool := &ReadFileTool{} + reg.Register(tool) + + err := reg.Register(tool) + if err == nil { + t.Error("expected error for duplicate registration") + } + }) + + t.Run("List tools", func(t *testing.T) { + reg := NewRegistry() + reg.Register(&ReadFileTool{}) + reg.Register(&WriteFileTool{}) + + list := reg.List() + if len(list) != 2 { + t.Errorf("expected 2 tools, got %d", len(list)) + } + }) + + t.Run("RegisterDefaults", func(t *testing.T) { + reg := NewRegistry() + reg.RegisterDefaults() + + if reg.Count() != 10 { + t.Errorf("expected 10 default tools, got %d", reg.Count()) + } + }) +} + +func TestParameterHelpers(t *testing.T) { + t.Run("GetString", func(t *testing.T) { + params := map[string]any{"key": "value"} + if GetString(params, "key", "") != "value" { + t.Error("expected 'value'") + } + if GetString(params, "missing", "default") != "default" { + t.Error("expected 'default'") + } + }) + + t.Run("GetInt", func(t *testing.T) { + params := map[string]any{"key": 42, "float": 3.14} + if GetInt(params, "key", 0) != 42 { + t.Error("expected 42") + } + if GetInt(params, "float", 0) != 3 { + t.Error("expected 3 from float") + } + if GetInt(params, "missing", 99) != 99 { + t.Error("expected 99") + } + }) + + t.Run("GetBool", func(t *testing.T) { + params := map[string]any{"key": true} + if !GetBool(params, "key", false) { + t.Error("expected true") + } + if GetBool(params, "missing", true) != true { + t.Error("expected true default") + } + }) + + t.Run("GetDuration", func(t *testing.T) { + params := map[string]any{"key": "30s", "seconds": 60} + if GetDuration(params, "key", 0) != 30*time.Second { + t.Error("expected 30s") + } + if GetDuration(params, "seconds", 0) != 60*time.Second { + t.Error("expected 60s from int") + } + }) +} diff --git a/internal/tui/app.go b/internal/tui/app.go index 5d8b0fe..f9daa5a 100644 --- a/internal/tui/app.go +++ b/internal/tui/app.go @@ -195,6 +195,10 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { if m.mode == ModeNormal { switch msg.String() { + case "tab": + m.activeTab = Tab((int(m.activeTab) + 1) % 3) + case "shift+tab": + m.activeTab = Tab((int(m.activeTab) + 2) % 3) case "1": m.activeTab = TabAgent case "2": @@ -314,26 +318,26 @@ func (m Model) viewMain() string { func (m Model) getFocusLabel() string { switch m.activeTab { case TabAgent: - return "agent" + return "Agent" case TabContainers: switch m.containers.Focus() { case monitor.FocusServices: - return "services" + return "Services" case monitor.FocusImages: - return "images" + return "Images" case monitor.FocusLogs: - return "logs" + return "Logs" case monitor.FocusStats: - return "stats" + return "Stats" } - return "containers" + return "Containers" case TabHistory: if m.history.Focus() == history.FocusSidebar { - return "list" + return "History" } - return "details" + return "Details" } - return "main" + return "Main" } type containerLogsMsg struct { diff --git a/internal/tui/app_test.go b/internal/tui/app_test.go new file mode 100644 index 0000000..9f53e13 --- /dev/null +++ b/internal/tui/app_test.go @@ -0,0 +1,212 @@ +package tui + +import ( + "strings" + "testing" + + tea "github.com/charmbracelet/bubbletea" +) + +func TestInitialModel(t *testing.T) { + model := InitialModel() + + if model.state != StateLoading { + t.Errorf("expected StateLoading, got %v", model.state) + } + if model.mode != ModeNormal { + t.Errorf("expected ModeNormal, got %v", model.mode) + } + if model.activeTab != TabAgent { + t.Errorf("expected TabAgent as initial tab, got %v", model.activeTab) + } + if model.quitting { + t.Error("quitting should be false initially") + } +} + +func TestModel_TabSwitching(t *testing.T) { + model := InitialModel() + model.state = StateMain + + tabMsg := tea.KeyMsg{Type: tea.KeyTab} + + newModel, _ := model.Update(tabMsg) + m := newModel.(Model) + + if m.activeTab != TabContainers { + t.Errorf("expected TabContainers after first tab, got %v", m.activeTab) + } + + newModel, _ = m.Update(tabMsg) + m = newModel.(Model) + + if m.activeTab != TabHistory { + t.Errorf("expected TabHistory after second tab, got %v", m.activeTab) + } + + newModel, _ = m.Update(tabMsg) + m = newModel.(Model) + + if m.activeTab != TabAgent { + t.Errorf("expected TabAgent after wrap, got %v", m.activeTab) + } +} + +func TestModel_QuitOnCtrlC(t *testing.T) { + model := InitialModel() + model.state = StateMain + + ctrlC := tea.KeyMsg{Type: tea.KeyCtrlC} + newModel, cmd := model.Update(ctrlC) + m := newModel.(Model) + + if !m.quitting { + t.Error("quitting should be true after Ctrl+C") + } + + if cmd == nil { + t.Error("expected quit command") + } +} + +func TestModel_WindowResize(t *testing.T) { + model := InitialModel() + + resizeMsg := tea.WindowSizeMsg{Width: 120, Height: 40} + newModel, _ := model.Update(resizeMsg) + m := newModel.(Model) + + if m.width != 120 { + t.Errorf("expected width 120, got %d", m.width) + } + if m.height != 40 { + t.Errorf("expected height 40, got %d", m.height) + } +} + +func TestModel_ModeFromTab(t *testing.T) { + tests := []struct { + tab Tab + expected AppMode + }{ + {TabAgent, ModeNormal}, // Agent starts in normal mode + {TabContainers, ModeNormal}, + {TabHistory, ModeNormal}, + } + + for _, tt := range tests { + model := InitialModel() + model.activeTab = tt.tab + + got := model.getModeFromTab() + if got != tt.expected { + t.Errorf("getModeFromTab() for tab %v = %v, want %v", tt.tab, got, tt.expected) + } + } +} + +func TestModel_ViewRendering(t *testing.T) { + model := InitialModel() + model.width = 80 + model.height = 24 + + model.state = StateLoading + loadingView := model.View() + if loadingView == "" { + t.Error("loading view should not be empty") + } + + model.state = StateMain + mainView := model.View() + if mainView == "" { + t.Error("main view should not be empty") + } +} + +func TestModel_ViewLoading(t *testing.T) { + model := InitialModel() + model.width = 80 + model.height = 24 + model.state = StateLoading + + view := model.viewLoading() + + if view == "" { + t.Error("viewLoading should return content") + } +} + +func TestModel_ViewMain(t *testing.T) { + model := InitialModel() + model.width = 80 + model.height = 24 + model.state = StateMain + + view := model.viewMain() + + if view == "" { + t.Error("viewMain should return content") + } +} + +func TestModel_GetFocusLabel(t *testing.T) { + model := InitialModel() + model.state = StateMain + + tests := []struct { + tab Tab + contains string + }{ + {TabAgent, "Agent"}, + {TabContainers, ""}, + {TabHistory, "History"}, + } + + for _, tt := range tests { + model.activeTab = tt.tab + label := model.getFocusLabel() + + if tt.contains != "" && !strings.Contains(label, tt.contains) { + t.Errorf("getFocusLabel() for tab %v should contain %q, got %q", + tt.tab, tt.contains, label) + } + } +} + +func TestModel_Init(t *testing.T) { + model := InitialModel() + + cmd := model.Init() + + if cmd == nil { + t.Error("Init should return a command") + } +} + +func TestModel_ShiftTabReverse(t *testing.T) { + model := InitialModel() + model.state = StateMain + model.activeTab = TabAgent + + shiftTabMsg := tea.KeyMsg{Type: tea.KeyShiftTab} + newModel, _ := model.Update(shiftTabMsg) + m := newModel.(Model) + + if m.activeTab != TabHistory { + t.Errorf("expected TabHistory after Shift+Tab from first tab, got %v", m.activeTab) + } +} + +func TestModel_EscapeKey(t *testing.T) { + model := InitialModel() + model.state = StateMain + model.mode = ModeInsert + + escMsg := tea.KeyMsg{Type: tea.KeyEsc} + newModel, _ := model.Update(escMsg) + m := newModel.(Model) + + if m.mode != ModeNormal { + t.Errorf("expected ModeNormal after Escape, got %v", m.mode) + } +} diff --git a/internal/tui/tabs/monitor/model.go b/internal/tui/tabs/monitor/model.go index 175c277..0fbebf5 100644 --- a/internal/tui/tabs/monitor/model.go +++ b/internal/tui/tabs/monitor/model.go @@ -61,7 +61,6 @@ func (d serviceDelegate) Render(w io.Writer, m list.Model, index int, listItem l return } - // Status indicator status := "●" statusColor := theme.Green if i.info.State != "running" { @@ -168,7 +167,7 @@ type Model struct { } func New() Model { - // Services list + sDelegate := serviceDelegate{} sList := list.New([]list.Item{}, sDelegate, 0, 0) sList.SetShowHelp(false) @@ -178,7 +177,6 @@ func New() Model { sList.DisableQuitKeybindings() sList.Styles.NoItems = lipgloss.NewStyle().Foreground(theme.Overlay0).Padding(1) - // Images list iDelegate := imageDelegate{} iList := list.New([]list.Item{}, iDelegate, 0, 0) iList.SetShowHelp(false) @@ -203,17 +201,15 @@ func (m Model) SetSize(w, h int) Model { m.width = w m.height = h - // Left sidebar width sidebarWidth := 28 if w < 100 { sidebarWidth = 24 } - // Calculate panel heights panelHeight := h - 4 - servicesHeight := (panelHeight - 8) / 2 // Half for services - imagesHeight := (panelHeight - 8) / 2 // Half for images - _ = 6 // Stats height (used in view.go) + servicesHeight := (panelHeight - 8) / 2 + imagesHeight := (panelHeight - 8) / 2 + _ = 6 if servicesHeight < 5 { servicesHeight = 5 @@ -222,13 +218,11 @@ func (m Model) SetSize(w, h int) Model { imagesHeight = 5 } - // Set list dimensions m.servicesList.SetWidth(sidebarWidth - 4) m.servicesList.SetHeight(servicesHeight - 2) m.imagesList.SetWidth(sidebarWidth - 4) m.imagesList.SetHeight(imagesHeight - 2) - // Viewport for logs logWidth := w - sidebarWidth - 4 if logWidth < 40 { logWidth = 40 @@ -269,7 +263,6 @@ func (m Model) SetImages(images []infra.ImageInfo) Model { func (m Model) SetLogLines(lines []string) Model { m.logLines = lines - // If recording, write to file if m.isRecording && m.recordingFile != nil { for _, line := range lines { m.recordingFile.WriteString(line + "\n") @@ -285,7 +278,6 @@ func (m Model) StartRecording() Model { return m } - // Create ~/.devlogs directory homeDir, err := os.UserHomeDir() if err != nil { return m @@ -296,7 +288,6 @@ func (m Model) StartRecording() Model { return m } - // Get selected service name serviceName := "unknown" if sel := m.servicesList.SelectedItem(); sel != nil { if s, ok := sel.(serviceItem); ok { @@ -304,7 +295,6 @@ func (m Model) StartRecording() Model { } } - // Create log file timestamp := time.Now().Format("2006-01-02_15-04-05") filename := fmt.Sprintf("docker-%s-%s.log", serviceName, timestamp) m.recordingPath = filepath.Join(logDir, filename) @@ -317,7 +307,6 @@ func (m Model) StartRecording() Model { m.recordingFile = file m.isRecording = true - // Write header m.recordingFile.WriteString(fmt.Sprintf("# Docker Log Recording: %s\n", serviceName)) m.recordingFile.WriteString(fmt.Sprintf("# Started: %s\n\n", time.Now().Format(time.RFC3339))) diff --git a/internal/tui/tabs/monitor/update.go b/internal/tui/tabs/monitor/update.go index 43fb80a..1359a13 100644 --- a/internal/tui/tabs/monitor/update.go +++ b/internal/tui/tabs/monitor/update.go @@ -86,7 +86,7 @@ func (m Model) Update(msg tea.Msg, keys KeyMap) (Model, tea.Cmd) { case tea.KeyMsg: switch { case key.Matches(msg, keys.Tab): - // Cycle focus: Services → Logs → Images → Stats → Services + switch m.focus { case FocusServices: m.focus = FocusLogs @@ -161,7 +161,7 @@ func (m Model) Update(msg tea.Msg, keys KeyMap) (Model, tea.Cmd) { } case key.Matches(msg, keys.Start): - // Will be handled by parent to call Docker client + if m.focus == FocusServices { if svc := m.SelectedService(); svc != nil { return m, func() tea.Msg { diff --git a/internal/tui/tabs/monitor/view.go b/internal/tui/tabs/monitor/view.go index c3b9e4f..6dc368a 100644 --- a/internal/tui/tabs/monitor/view.go +++ b/internal/tui/tabs/monitor/view.go @@ -11,7 +11,7 @@ import ( ) func (m Model) View() string { - // Left sidebar width + sidebarWidth := 28 if m.width < 100 { sidebarWidth = 24 @@ -24,7 +24,6 @@ func (m Model) View() string { panelHeight := m.height - 4 - // Calculate heights for left panels servicesHeight := (panelHeight - 8) / 2 imagesHeight := (panelHeight - 8) / 2 statsHeight := 6 @@ -36,14 +35,12 @@ func (m Model) View() string { imagesHeight = 5 } - // Render left column panels servicesPanel := m.renderServicesPanel(sidebarWidth, servicesHeight) imagesPanel := m.renderImagesPanel(sidebarWidth, imagesHeight) statsPanel := m.renderStatsPanel(sidebarWidth, statsHeight) leftColumn := lipgloss.JoinVertical(lipgloss.Left, servicesPanel, imagesPanel, statsPanel) - // Render logs panel logsPanel := m.renderLogsPanel(logWidth, panelHeight) return lipgloss.JoinHorizontal(lipgloss.Top, leftColumn, logsPanel) @@ -152,7 +149,6 @@ func (m Model) renderStatsPanel(width, height int) string { stats := m.GetSelectedServiceStats() labelStyle := lipgloss.NewStyle().Foreground(theme.Overlay0).Width(4) - // CPU sparkline content.WriteString(labelStyle.Render("CPU ")) sparkWidth := width - 12 if sparkWidth < 5 { @@ -170,7 +166,6 @@ func (m Model) renderStatsPanel(width, height int) string { } content.WriteString("\n") - // Memory bar content.WriteString(labelStyle.Render("MEM ")) if stats.MemTotal > 0 { memBar := components.NewProgressBar(stats.MemUsed, stats.MemTotal). @@ -182,7 +177,6 @@ func (m Model) renderStatsPanel(width, height int) string { } content.WriteString("\n") - // Network content.WriteString(labelStyle.Render("NET ")) netStyle := lipgloss.NewStyle().Foreground(theme.Overlay0) if stats.NetIn > 0 || stats.NetOut > 0 { @@ -218,7 +212,6 @@ func (m Model) renderLogsPanel(width, height int) string { header := headerStyle.Render("≡ Logs") - // Show selected service name if svc := m.SelectedService(); svc != nil { serviceName := svc.Name if len(serviceName) > 15 { @@ -227,7 +220,6 @@ func (m Model) renderLogsPanel(width, height int) string { header += dimStyle.Render(" (" + serviceName + ")") } - // Recording indicator if m.isRecording { recBadge := lipgloss.NewStyle(). Background(theme.Red). @@ -238,7 +230,6 @@ func (m Model) renderLogsPanel(width, height int) string { header += " " + recBadge } - // Follow mode indicator if m.followMode { followBadge := lipgloss.NewStyle(). Background(theme.Green). @@ -248,7 +239,6 @@ func (m Model) renderLogsPanel(width, height int) string { header += " " + followBadge } - // Log level filter indicator if m.logLevelFilter != "" { filterBadge := lipgloss.NewStyle(). Background(theme.Surface0). diff --git a/internal/workflow/checkpoint.go b/internal/workflow/checkpoint.go new file mode 100644 index 0000000..091d9b7 --- /dev/null +++ b/internal/workflow/checkpoint.go @@ -0,0 +1,313 @@ +package workflow + +import ( + "database/sql" + "encoding/json" + "fmt" + "time" +) + +// CheckpointStore handles persistence of workflow run states. +type CheckpointStore struct { + db *sql.DB +} + +// NewCheckpointStore creates a new checkpoint store. +func NewCheckpointStore(db *sql.DB) *CheckpointStore { + return &CheckpointStore{db: db} +} + +// InitSchema creates the workflow tables if they don't exist. +func (s *CheckpointStore) InitSchema() error { + schema := ` + CREATE TABLE IF NOT EXISTS workflow_runs ( + id TEXT PRIMARY KEY, + workflow_id TEXT NOT NULL, + workflow_name TEXT, + status TEXT NOT NULL, + current_step INTEGER DEFAULT 0, + started_at DATETIME, + updated_at DATETIME, + completed_at DATETIME, + error TEXT + ); + + CREATE TABLE IF NOT EXISTS workflow_step_results ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + run_id TEXT NOT NULL, + step_id TEXT NOT NULL, + status TEXT NOT NULL, + exit_code INTEGER, + output TEXT, + error TEXT, + retries INTEGER DEFAULT 0, + started_at DATETIME, + completed_at DATETIME, + duration_ms INTEGER, + FOREIGN KEY (run_id) REFERENCES workflow_runs(id) + ); + + CREATE INDEX IF NOT EXISTS idx_workflow_runs_status ON workflow_runs(status); + CREATE INDEX IF NOT EXISTS idx_step_results_run_id ON workflow_step_results(run_id); + ` + + _, err := s.db.Exec(schema) + return err +} + +// SaveRun persists or updates a workflow run state. +func (s *CheckpointStore) SaveRun(state *RunState) error { + query := ` + INSERT OR REPLACE INTO workflow_runs + (id, workflow_id, workflow_name, status, current_step, started_at, updated_at, completed_at, error) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + + var completedAt *time.Time + if !state.CompletedAt.IsZero() { + completedAt = &state.CompletedAt + } + + _, err := s.db.Exec(query, + state.RunID, + state.WorkflowID, + state.WorkflowName, + string(state.Status), + state.CurrentStepIdx, + state.StartedAt, + state.UpdatedAt, + completedAt, + state.Error, + ) + + return err +} + +// SaveStepResult persists a step execution result. +func (s *CheckpointStore) SaveStepResult(runID string, result *StepResult) error { + query := ` + INSERT OR REPLACE INTO workflow_step_results + (run_id, step_id, status, exit_code, output, error, retries, started_at, completed_at, duration_ms) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ` + + var completedAt *time.Time + if !result.CompletedAt.IsZero() { + completedAt = &result.CompletedAt + } + + _, err := s.db.Exec(query, + runID, + result.StepID, + string(result.Status), + result.ExitCode, + truncateString(result.Output, 10240), + result.Error, + result.Retries, + result.StartedAt, + completedAt, + result.Duration.Milliseconds(), + ) + + return err +} + +// LoadRun retrieves a workflow run state by ID. +func (s *CheckpointStore) LoadRun(runID string) (*RunState, error) { + query := ` + SELECT id, workflow_id, workflow_name, status, current_step, started_at, updated_at, completed_at, error + FROM workflow_runs WHERE id = ? + ` + + row := s.db.QueryRow(query, runID) + + state := &RunState{ + StepResults: make(map[string]*StepResult), + } + + var completedAt sql.NullTime + var errStr sql.NullString + var status string + + err := row.Scan( + &state.RunID, + &state.WorkflowID, + &state.WorkflowName, + &status, + &state.CurrentStepIdx, + &state.StartedAt, + &state.UpdatedAt, + &completedAt, + &errStr, + ) + + if err == sql.ErrNoRows { + return nil, fmt.Errorf("run not found: %s", runID) + } + if err != nil { + return nil, err + } + + state.Status = RunStatus(status) + if completedAt.Valid { + state.CompletedAt = completedAt.Time + } + if errStr.Valid { + state.Error = errStr.String + } + + stepResults, err := s.LoadStepResults(runID) + if err != nil { + return nil, err + } + state.StepResults = stepResults + + return state, nil +} + +// LoadStepResults retrieves all step results for a run. +func (s *CheckpointStore) LoadStepResults(runID string) (map[string]*StepResult, error) { + query := ` + SELECT step_id, status, exit_code, output, error, retries, started_at, completed_at, duration_ms + FROM workflow_step_results WHERE run_id = ? ORDER BY started_at + ` + + rows, err := s.db.Query(query, runID) + if err != nil { + return nil, err + } + defer rows.Close() + + results := make(map[string]*StepResult) + + for rows.Next() { + result := &StepResult{} + var status string + var completedAt sql.NullTime + var errStr sql.NullString + var durationMs int64 + + err := rows.Scan( + &result.StepID, + &status, + &result.ExitCode, + &result.Output, + &errStr, + &result.Retries, + &result.StartedAt, + &completedAt, + &durationMs, + ) + if err != nil { + return nil, err + } + + result.Status = StepStatus(status) + result.Duration = time.Duration(durationMs) * time.Millisecond + if completedAt.Valid { + result.CompletedAt = completedAt.Time + } + if errStr.Valid { + result.Error = errStr.String + } + + results[result.StepID] = result + } + + return results, rows.Err() +} + +// ListRuns returns recent workflow runs. +func (s *CheckpointStore) ListRuns(limit int) ([]*RunState, error) { + query := ` + SELECT id, workflow_id, workflow_name, status, current_step, started_at, updated_at, completed_at, error + FROM workflow_runs ORDER BY started_at DESC LIMIT ? + ` + + rows, err := s.db.Query(query, limit) + if err != nil { + return nil, err + } + defer rows.Close() + + var runs []*RunState + + for rows.Next() { + state := &RunState{ + StepResults: make(map[string]*StepResult), + } + + var completedAt sql.NullTime + var errStr sql.NullString + var status string + + err := rows.Scan( + &state.RunID, + &state.WorkflowID, + &state.WorkflowName, + &status, + &state.CurrentStepIdx, + &state.StartedAt, + &state.UpdatedAt, + &completedAt, + &errStr, + ) + if err != nil { + return nil, err + } + + state.Status = RunStatus(status) + if completedAt.Valid { + state.CompletedAt = completedAt.Time + } + if errStr.Valid { + state.Error = errStr.String + } + + runs = append(runs, state) + } + + return runs, rows.Err() +} + +// DeleteRun removes a workflow run and its step results. +func (s *CheckpointStore) DeleteRun(runID string) error { + tx, err := s.db.Begin() + if err != nil { + return err + } + defer tx.Rollback() + + _, err = tx.Exec("DELETE FROM workflow_step_results WHERE run_id = ?", runID) + if err != nil { + return err + } + + _, err = tx.Exec("DELETE FROM workflow_runs WHERE id = ?", runID) + if err != nil { + return err + } + + return tx.Commit() +} + +func truncateString(s string, max int) string { + if len(s) <= max { + return s + } + return s[:max-20] + "\n...[truncated]..." +} + +// MarshalRunState serializes run state to JSON. +func MarshalRunState(state *RunState) ([]byte, error) { + return json.Marshal(state) +} + +// UnmarshalRunState deserializes run state from JSON. +func UnmarshalRunState(data []byte) (*RunState, error) { + var state RunState + if err := json.Unmarshal(data, &state); err != nil { + return nil, err + } + return &state, nil +} diff --git a/internal/workflow/condition.go b/internal/workflow/condition.go new file mode 100644 index 0000000..ba893db --- /dev/null +++ b/internal/workflow/condition.go @@ -0,0 +1,104 @@ +package workflow + +import ( + "os" + "regexp" + "strings" +) + +// Evaluate checks if a condition is met based on the last step result. +// If no condition is specified, returns true (step should run). +func (c *Condition) Evaluate(lastResult *StepResult) bool { + if c == nil { + return true + } + + switch c.Type { + case CondExitCode: + if lastResult == nil { + return c.Value == "0" + } + return matchExitCode(lastResult.ExitCode, c.Value) + + case CondOutputContains: + if lastResult == nil { + return false + } + return strings.Contains(lastResult.Output, c.Value) + + case CondOutputMatches: + if lastResult == nil { + return false + } + matched, _ := regexp.MatchString(c.Value, lastResult.Output) + return matched + + case CondFileExists: + _, err := os.Stat(c.Value) + return err == nil + + case CondEnvSet: + _, exists := os.LookupEnv(c.Value) + return exists + + default: + return true + } +} + +// matchExitCode checks if an exit code matches the condition value. +// Supports: "0", "!0" (non-zero), or specific code like "1". +func matchExitCode(exitCode int, value string) bool { + if value == "!0" { + return exitCode != 0 + } + + // Parse as integer + var expected int + if _, err := parseIntFromString(value, &expected); err != nil { + return false + } + return exitCode == expected +} + +// parseIntFromString is a helper to parse int from string. +func parseIntFromString(s string, result *int) (bool, error) { + n := 0 + for _, ch := range s { + if ch < '0' || ch > '9' { + return false, nil + } + n = n*10 + int(ch-'0') + } + *result = n + return true, nil +} + +// EvaluateWithStepRef evaluates a condition against a specific step result. +func (c *Condition) EvaluateWithStepRef(results map[string]*StepResult) bool { + if c == nil { + return true + } + + var targetResult *StepResult + if c.StepRef != "" { + targetResult = results[c.StepRef] + } else { + + for _, r := range results { + if targetResult == nil || r.CompletedAt.After(targetResult.CompletedAt) { + targetResult = r + } + } + } + + return c.Evaluate(targetResult) +} + +// ShouldSkip returns true if the step should be skipped due to condition. +func ShouldSkip(step *Step, results map[string]*StepResult) bool { + if step.Condition == nil { + return false + } + return !step.Condition.EvaluateWithStepRef(results) +} diff --git a/internal/workflow/condition_test.go b/internal/workflow/condition_test.go new file mode 100644 index 0000000..5deb9fd --- /dev/null +++ b/internal/workflow/condition_test.go @@ -0,0 +1,159 @@ +package workflow + +import ( + "testing" + "time" +) + +func TestConditionEvaluate(t *testing.T) { + tests := []struct { + name string + cond *Condition + result *StepResult + expected bool + }{ + { + name: "nil condition returns true", + cond: nil, + result: nil, + expected: true, + }, + { + name: "exit_code 0 matches success", + cond: &Condition{ + Type: CondExitCode, + Value: "0", + }, + result: &StepResult{ + ExitCode: 0, + }, + expected: true, + }, + { + name: "exit_code 0 does not match failure", + cond: &Condition{ + Type: CondExitCode, + Value: "0", + }, + result: &StepResult{ + ExitCode: 1, + }, + expected: false, + }, + { + name: "exit_code !0 matches failure", + cond: &Condition{ + Type: CondExitCode, + Value: "!0", + }, + result: &StepResult{ + ExitCode: 1, + }, + expected: true, + }, + { + name: "output_contains matches", + cond: &Condition{ + Type: CondOutputContains, + Value: "success", + }, + result: &StepResult{ + Output: "build success completed", + }, + expected: true, + }, + { + name: "output_contains does not match", + cond: &Condition{ + Type: CondOutputContains, + Value: "error", + }, + result: &StepResult{ + Output: "build success completed", + }, + expected: false, + }, + { + name: "output_matches regex", + cond: &Condition{ + Type: CondOutputMatches, + Value: `version \d+\.\d+`, + }, + result: &StepResult{ + Output: "version 1.5 installed", + }, + expected: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := tt.cond.Evaluate(tt.result) + if got != tt.expected { + t.Errorf("Condition.Evaluate() = %v, want %v", got, tt.expected) + } + }) + } +} + +func TestShouldSkip(t *testing.T) { + results := map[string]*StepResult{ + "step1": { + StepID: "step1", + ExitCode: 0, + CompletedAt: time.Now(), + }, + } + + tests := []struct { + name string + step *Step + results map[string]*StepResult + expected bool + }{ + { + name: "no condition - should not skip", + step: &Step{ + ID: "step2", + Command: "echo test", + }, + results: results, + expected: false, + }, + { + name: "condition met - should not skip", + step: &Step{ + ID: "step2", + Command: "echo test", + Condition: &Condition{ + Type: CondExitCode, + Value: "0", + }, + }, + results: results, + expected: false, + }, + { + name: "condition not met - should skip", + step: &Step{ + ID: "step2", + Command: "echo test", + Condition: &Condition{ + Type: CondExitCode, + Value: "!0", + }, + }, + results: results, + expected: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := ShouldSkip(tt.step, tt.results) + if got != tt.expected { + t.Errorf("ShouldSkip() = %v, want %v", got, tt.expected) + } + }) + } +} diff --git a/internal/workflow/engine.go b/internal/workflow/engine.go new file mode 100644 index 0000000..c618b54 --- /dev/null +++ b/internal/workflow/engine.go @@ -0,0 +1,414 @@ +package workflow + +import ( + "context" + "fmt" + "time" + + "dev-cli/internal/executor" + "dev-cli/internal/pipeline" +) + +// Engine executes workflows with support for conditionals, rollback, and checkpointing. +type Engine struct { + store *CheckpointStore + bus *pipeline.EventBus + verbose bool + safeCtx *SafeModeContext + rollback *RollbackRegistry +} + +// NewEngine creates a new workflow execution engine. +func NewEngine(store *CheckpointStore, bus *pipeline.EventBus) *Engine { + return &Engine{ + store: store, + bus: bus, + safeCtx: NewSafeModeContext(), + rollback: NewRollbackRegistry(), + } +} + +// SetVerbose enables verbose logging. +func (e *Engine) SetVerbose(v bool) { + e.verbose = v +} + +// SetSafeMode configures the safe mode context. +func (e *Engine) SetSafeMode(ctx *SafeModeContext) { + e.safeCtx = ctx +} + +// GetSafeMode returns the current safe mode context. +func (e *Engine) GetSafeMode() *SafeModeContext { + return e.safeCtx +} + +// GetRollbackRegistry returns the rollback registry. +func (e *Engine) GetRollbackRegistry() *RollbackRegistry { + return e.rollback +} + +// RunResult contains the outcome of a workflow execution. +type RunResult struct { + RunID string + Status RunStatus + StepResults map[string]*StepResult + Error string + Duration time.Duration +} + +// Run executes a workflow from the beginning. +func (e *Engine) Run(ctx context.Context, wf *Workflow) (*RunResult, error) { + runID := GenerateRunID() + state := NewRunState(runID, wf) + state.Status = StatusRunning + + if e.store != nil { + if err := e.store.SaveRun(state); err != nil { + return nil, fmt.Errorf("failed to save initial state: %w", err) + } + } + + e.publishEvent(pipeline.Event{ + Type: pipeline.EventType("workflow.start"), + Timestamp: time.Now(), + Source: "workflow", + Data: map[string]interface{}{ + "run_id": runID, + "workflow_name": wf.Name, + "total_steps": len(wf.Steps), + }, + }) + + return e.executeSteps(ctx, wf, state) +} + +// Resume continues execution of a paused or failed workflow. +func (e *Engine) Resume(ctx context.Context, wf *Workflow, runID string) (*RunResult, error) { + if e.store == nil { + return nil, fmt.Errorf("checkpoint store required for resume") + } + + state, err := e.store.LoadRun(runID) + if err != nil { + return nil, fmt.Errorf("failed to load run state: %w", err) + } + + if state.Status != StatusPaused && state.Status != StatusFailed { + return nil, fmt.Errorf("cannot resume run with status: %s", state.Status) + } + + state.Status = StatusRunning + state.UpdatedAt = time.Now() + + if err := e.store.SaveRun(state); err != nil { + return nil, fmt.Errorf("failed to update state: %w", err) + } + + return e.executeSteps(ctx, wf, state) +} + +// Rollback executes rollback actions for a failed workflow. +func (e *Engine) Rollback(ctx context.Context, wf *Workflow, runID string) error { + if e.store == nil { + return fmt.Errorf("checkpoint store required for rollback") + } + + state, err := e.store.LoadRun(runID) + if err != nil { + return fmt.Errorf("failed to load run state: %w", err) + } + + return e.executeRollback(ctx, wf, state) +} + +// executeSteps runs workflow steps starting from the current position. +func (e *Engine) executeSteps(ctx context.Context, wf *Workflow, state *RunState) (*RunResult, error) { + startTime := time.Now() + + for i := state.CurrentStepIdx; i < len(wf.Steps); i++ { + select { + case <-ctx.Done(): + state.Status = StatusPaused + state.UpdatedAt = time.Now() + if e.store != nil { + e.store.SaveRun(state) + } + return &RunResult{ + RunID: state.RunID, + Status: StatusPaused, + StepResults: state.StepResults, + Error: "cancelled", + Duration: time.Since(startTime), + }, ctx.Err() + + default: + } + + step := wf.Steps[i] + state.CurrentStepIdx = i + + if ShouldSkip(&step, state.StepResults) { + result := &StepResult{ + StepID: step.ID, + Status: StepSkipped, + StartedAt: time.Now(), + CompletedAt: time.Now(), + } + state.SetStepResult(result) + + if e.store != nil { + e.store.SaveStepResult(state.RunID, result) + e.store.SaveRun(state) + } + + e.log("⏭ Skipping step: %s (condition not met)", step.Name) + continue + } + + result := e.executeStep(ctx, &step, wf.Env, state) + state.SetStepResult(result) + + if e.store != nil { + e.store.SaveStepResult(state.RunID, result) + e.store.SaveRun(state) + } + + e.publishEvent(pipeline.Event{ + Type: pipeline.EventType("workflow.step"), + Timestamp: time.Now(), + Source: "workflow", + BlockID: step.ID, + Data: map[string]interface{}{ + "run_id": state.RunID, + "step_id": step.ID, + "step_name": step.Name, + "status": string(result.Status), + "exit_code": result.ExitCode, + }, + }) + + if result.Status == StepFailed { + action := e.determineFailureAction(wf, &step) + + switch action { + case FailureRollback: + e.log("⚠ Step failed, initiating rollback...") + if err := e.executeRollback(ctx, wf, state); err != nil { + e.log("✗ Rollback failed: %v", err) + } + state.Status = StatusRolledBack + state.Error = result.Error + state.CompletedAt = time.Now() + if e.store != nil { + e.store.SaveRun(state) + } + return &RunResult{ + RunID: state.RunID, + Status: StatusRolledBack, + StepResults: state.StepResults, + Error: result.Error, + Duration: time.Since(startTime), + }, nil + + case FailureAbort: + state.Status = StatusFailed + state.Error = result.Error + state.CompletedAt = time.Now() + if e.store != nil { + e.store.SaveRun(state) + } + return &RunResult{ + RunID: state.RunID, + Status: StatusFailed, + StepResults: state.StepResults, + Error: result.Error, + Duration: time.Since(startTime), + }, nil + + case FailureContinue: + e.log("⚠ Step failed but continuing...") + continue + } + } + + if step.OnSuccess != "" { + nextIdx := e.findStepIndex(wf, step.OnSuccess) + if nextIdx >= 0 { + state.CurrentStepIdx = nextIdx - 1 + } + } + } + + state.Status = StatusCompleted + state.CompletedAt = time.Now() + if e.store != nil { + e.store.SaveRun(state) + } + + e.publishEvent(pipeline.Event{ + Type: pipeline.EventType("workflow.complete"), + Timestamp: time.Now(), + Source: "workflow", + Data: map[string]interface{}{ + "run_id": state.RunID, + "status": string(StatusCompleted), + "duration": time.Since(startTime).String(), + }, + }) + + return &RunResult{ + RunID: state.RunID, + Status: StatusCompleted, + StepResults: state.StepResults, + Duration: time.Since(startTime), + }, nil +} + +// executeStep runs a single step with retries. +func (e *Engine) executeStep(ctx context.Context, step *Step, env map[string]string, state *RunState) *StepResult { + result := &StepResult{ + StepID: step.ID, + Status: StepRunning, + StartedAt: time.Now(), + } + + maxRetries := step.Retries + if maxRetries == 0 { + maxRetries = 1 + } + + for attempt := 0; attempt < maxRetries; attempt++ { + result.Retries = attempt + + e.log("▶ Running step: %s (attempt %d/%d)", step.Name, attempt+1, maxRetries) + + stepCtx := ctx + if step.Timeout > 0 { + var cancel context.CancelFunc + stepCtx, cancel = context.WithTimeout(ctx, step.Timeout) + defer cancel() + } + + execResult := executor.ExecuteWithContext(stepCtx, step.Command) + + result.ExitCode = execResult.ExitCode + result.Output = execResult.Output + result.Duration = execResult.Duration + result.CompletedAt = time.Now() + + if execResult.ExitCode == 0 { + result.Status = StepSuccess + e.log("✓ Step completed: %s", step.Name) + return result + } + + e.log("✗ Step failed (exit %d): %s", execResult.ExitCode, step.Name) + + if attempt < maxRetries-1 { + e.log(" Retrying in 2 seconds...") + time.Sleep(2 * time.Second) + } + } + + result.Status = StepFailed + result.Error = fmt.Sprintf("step failed with exit code %d after %d attempts", result.ExitCode, maxRetries) + return result +} + +// executeRollback runs rollback commands in reverse order. +func (e *Engine) executeRollback(ctx context.Context, wf *Workflow, state *RunState) error { + e.publishEvent(pipeline.Event{ + Type: pipeline.EventType("workflow.rollback"), + Timestamp: time.Now(), + Source: "workflow", + Data: map[string]interface{}{ + "run_id": state.RunID, + }, + }) + + for i := state.CurrentStepIdx; i >= 0; i-- { + step := wf.Steps[i] + + stepResult := state.GetStepResult(step.ID) + if stepResult == nil || stepResult.Status == StepSkipped { + continue + } + + if step.Rollback == nil { + e.log("⏭ No rollback defined for: %s", step.Name) + continue + } + + e.log("↺ Rolling back: %s", step.Name) + + rollbackCtx := ctx + if step.Rollback.Timeout > 0 { + var cancel context.CancelFunc + rollbackCtx, cancel = context.WithTimeout(ctx, step.Rollback.Timeout) + defer cancel() + } + + result := executor.ExecuteWithContext(rollbackCtx, step.Rollback.Command) + + if result.ExitCode != 0 { + e.log("⚠ Rollback failed for %s: %s", step.Name, result.Output) + } else { + e.log("✓ Rolled back: %s", step.Name) + + if stepResult != nil { + stepResult.Status = StepRolledBack + if e.store != nil { + e.store.SaveStepResult(state.RunID, stepResult) + } + } + } + } + + return nil +} + +// determineFailureAction returns the action to take on step failure. +func (e *Engine) determineFailureAction(wf *Workflow, step *Step) FailureAction { + + if step.OnFailure != "" { + switch step.OnFailure { + case "abort": + return FailureAbort + case "rollback": + return FailureRollback + case "continue": + return FailureContinue + } + } + + if wf.OnFailure != nil { + return wf.OnFailure.Action + } + + return FailureAbort +} + +// findStepIndex returns the index of a step by ID. +func (e *Engine) findStepIndex(wf *Workflow, stepID string) int { + for i, step := range wf.Steps { + if step.ID == stepID { + return i + } + } + return -1 +} + +// publishEvent sends an event to the event bus if available. +func (e *Engine) publishEvent(event pipeline.Event) { + if e.bus != nil { + e.bus.Publish(event) + } +} + +// log outputs a message if verbose mode is enabled. +func (e *Engine) log(format string, args ...interface{}) { + if e.verbose { + fmt.Printf(format+"\n", args...) + } +} diff --git a/internal/workflow/parser.go b/internal/workflow/parser.go new file mode 100644 index 0000000..b75c0f1 --- /dev/null +++ b/internal/workflow/parser.go @@ -0,0 +1,190 @@ +package workflow + +import ( + "fmt" + "os" + "time" + + "gopkg.in/yaml.v3" +) + +// ParseFile reads a workflow definition from a YAML file. +func ParseFile(path string) (*Workflow, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("failed to read workflow file: %w", err) + } + + return Parse(data) +} + +// Parse parses a workflow definition from YAML bytes. +func Parse(data []byte) (*Workflow, error) { + var raw rawWorkflow + if err := yaml.Unmarshal(data, &raw); err != nil { + return nil, fmt.Errorf("failed to parse workflow YAML: %w", err) + } + + return raw.toWorkflow() +} + +// rawWorkflow is the YAML structure with string durations. +type rawWorkflow struct { + ID string `yaml:"id"` + Name string `yaml:"name"` + Description string `yaml:"description"` + Steps []rawStep `yaml:"steps"` + OnFailure *FailurePolicy `yaml:"on_failure"` + Env map[string]string `yaml:"env"` +} + +type rawStep struct { + ID string `yaml:"id"` + Name string `yaml:"name"` + Command string `yaml:"command"` + Condition *Condition `yaml:"condition"` + OnSuccess string `yaml:"on_success"` + OnFailure string `yaml:"on_failure"` + Rollback *rawRollback `yaml:"rollback"` + Timeout string `yaml:"timeout"` + Retries int `yaml:"retries"` + Env map[string]string `yaml:"env"` + WorkDir string `yaml:"workdir"` +} + +type rawRollback struct { + Command string `yaml:"command"` + Timeout string `yaml:"timeout"` +} + +// UnmarshalYAML allows rollback to be specified as just a string. +func (r *rawRollback) UnmarshalYAML(node *yaml.Node) error { + + if node.Kind == yaml.ScalarNode { + r.Command = node.Value + return nil + } + + // Otherwise, parse as struct + type plain rawRollback + return node.Decode((*plain)(r)) +} + +func (rw *rawWorkflow) toWorkflow() (*Workflow, error) { + wf := &Workflow{ + ID: rw.ID, + Name: rw.Name, + Description: rw.Description, + OnFailure: rw.OnFailure, + Env: rw.Env, + Steps: make([]Step, 0, len(rw.Steps)), + } + + if wf.ID == "" { + wf.ID = generateID() + } + + for i, rs := range rw.Steps { + step, err := rs.toStep(i) + if err != nil { + return nil, fmt.Errorf("step %d (%s): %w", i, rs.ID, err) + } + wf.Steps = append(wf.Steps, step) + } + + if err := validateWorkflow(wf); err != nil { + return nil, err + } + + return wf, nil +} + +func (rs *rawStep) toStep(index int) (Step, error) { + step := Step{ + ID: rs.ID, + Name: rs.Name, + Command: rs.Command, + Condition: rs.Condition, + OnSuccess: rs.OnSuccess, + OnFailure: rs.OnFailure, + Retries: rs.Retries, + Env: rs.Env, + WorkDir: rs.WorkDir, + } + + if step.ID == "" { + step.ID = fmt.Sprintf("step_%d", index) + } + + if rs.Timeout != "" { + d, err := time.ParseDuration(rs.Timeout) + if err != nil { + return step, fmt.Errorf("invalid timeout %q: %w", rs.Timeout, err) + } + step.Timeout = d + } else { + step.Timeout = 5 * time.Minute + } + + if rs.Rollback != nil { + step.Rollback = &RollbackAction{ + Command: rs.Rollback.Command, + } + if rs.Rollback.Timeout != "" { + d, err := time.ParseDuration(rs.Rollback.Timeout) + if err != nil { + return step, fmt.Errorf("invalid rollback timeout: %w", err) + } + step.Rollback.Timeout = d + } else { + step.Rollback.Timeout = 2 * time.Minute + } + } + + return step, nil +} + +func validateWorkflow(wf *Workflow) error { + if wf.Name == "" { + return fmt.Errorf("workflow name is required") + } + + if len(wf.Steps) == 0 { + return fmt.Errorf("workflow must have at least one step") + } + + stepIDs := make(map[string]bool) + for _, step := range wf.Steps { + if step.Command == "" { + return fmt.Errorf("step %q: command is required", step.ID) + } + + if stepIDs[step.ID] { + return fmt.Errorf("duplicate step ID: %s", step.ID) + } + stepIDs[step.ID] = true + } + + for _, step := range wf.Steps { + if step.OnSuccess != "" && !stepIDs[step.OnSuccess] { + return fmt.Errorf("step %q: on_success references unknown step %q", step.ID, step.OnSuccess) + } + if step.OnFailure != "" && step.OnFailure != "abort" && step.OnFailure != "rollback" && step.OnFailure != "continue" { + if !stepIDs[step.OnFailure] { + return fmt.Errorf("step %q: on_failure references unknown step %q", step.ID, step.OnFailure) + } + } + } + + return nil +} + +// generateID creates a simple unique ID based on timestamp. +func generateID() string { + return fmt.Sprintf("wf_%d", time.Now().UnixNano()) +} + +// GenerateRunID creates a unique run ID. +func GenerateRunID() string { + return fmt.Sprintf("run_%d", time.Now().UnixNano()) +} diff --git a/internal/workflow/parser_test.go b/internal/workflow/parser_test.go new file mode 100644 index 0000000..4b66f07 --- /dev/null +++ b/internal/workflow/parser_test.go @@ -0,0 +1,182 @@ +package workflow + +import ( + "strings" + "testing" + "time" +) + +func TestParse(t *testing.T) { + yaml := ` +name: test-workflow +description: A test workflow +steps: + - id: step1 + name: First Step + command: echo "hello" + timeout: 30s + rollback: echo "rollback step1" + - id: step2 + name: Second Step + command: echo "world" + condition: + type: exit_code + value: "0" + on_failure: abort +on_failure: + action: rollback +` + + wf, err := Parse([]byte(yaml)) + if err != nil { + t.Fatalf("Parse() error = %v", err) + } + + if wf.Name != "test-workflow" { + t.Errorf("Name = %q, want %q", wf.Name, "test-workflow") + } + + if wf.Description != "A test workflow" { + t.Errorf("Description = %q, want %q", wf.Description, "A test workflow") + } + + if len(wf.Steps) != 2 { + t.Fatalf("len(Steps) = %d, want 2", len(wf.Steps)) + } + + step1 := wf.Steps[0] + if step1.ID != "step1" { + t.Errorf("Step1.ID = %q, want %q", step1.ID, "step1") + } + if step1.Command != `echo "hello"` { + t.Errorf("Step1.Command = %q, want %q", step1.Command, `echo "hello"`) + } + if step1.Timeout != 30*time.Second { + t.Errorf("Step1.Timeout = %v, want %v", step1.Timeout, 30*time.Second) + } + if step1.Rollback == nil { + t.Error("Step1.Rollback is nil, expected non-nil") + } else if step1.Rollback.Command != `echo "rollback step1"` { + t.Errorf("Step1.Rollback.Command = %q", step1.Rollback.Command) + } + + step2 := wf.Steps[1] + if step2.Condition == nil { + t.Error("Step2.Condition is nil") + } else { + if step2.Condition.Type != CondExitCode { + t.Errorf("Step2.Condition.Type = %q, want %q", step2.Condition.Type, CondExitCode) + } + if step2.Condition.Value != "0" { + t.Errorf("Step2.Condition.Value = %q, want %q", step2.Condition.Value, "0") + } + } + if step2.OnFailure != "abort" { + t.Errorf("Step2.OnFailure = %q, want %q", step2.OnFailure, "abort") + } + + if wf.OnFailure == nil { + t.Error("OnFailure is nil") + } else if wf.OnFailure.Action != FailureRollback { + t.Errorf("OnFailure.Action = %q, want %q", wf.OnFailure.Action, FailureRollback) + } +} + +func TestParseValidation(t *testing.T) { + tests := []struct { + name string + yaml string + wantErr string + }{ + { + name: "missing name", + yaml: "steps:\n - command: echo test", + wantErr: "name is required", + }, + { + name: "missing steps", + yaml: "name: test", + wantErr: "at least one step", + }, + { + name: "missing command", + yaml: ` +name: test +steps: + - id: step1 + name: No Command`, + wantErr: "command is required", + }, + { + name: "duplicate step ID", + yaml: ` +name: test +steps: + - id: step1 + command: echo 1 + - id: step1 + command: echo 2`, + wantErr: "duplicate step ID", + }, + { + name: "invalid on_success reference", + yaml: ` +name: test +steps: + - id: step1 + command: echo 1 + on_success: nonexistent`, + wantErr: "unknown step", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := Parse([]byte(tt.yaml)) + if err == nil { + t.Error("Parse() expected error, got nil") + return + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Errorf("Parse() error = %q, want to contain %q", err.Error(), tt.wantErr) + } + }) + } +} + +func TestRollbackYAMLShorthand(t *testing.T) { + + yaml := ` +name: test +steps: + - id: step1 + command: touch /tmp/test + rollback: rm /tmp/test +` + + wf, err := Parse([]byte(yaml)) + if err != nil { + t.Fatalf("Parse() error = %v", err) + } + + if wf.Steps[0].Rollback == nil { + t.Fatal("Rollback is nil") + } + + if wf.Steps[0].Rollback.Command != "rm /tmp/test" { + t.Errorf("Rollback.Command = %q, want %q", wf.Steps[0].Rollback.Command, "rm /tmp/test") + } +} + +func TestGenerateRunID(t *testing.T) { + id1 := GenerateRunID() + id2 := GenerateRunID() + + if id1 == id2 { + t.Error("GenerateRunID() should return unique IDs") + } + + if !strings.HasPrefix(id1, "run_") { + t.Errorf("GenerateRunID() = %q, want prefix 'run_'", id1) + } +} diff --git a/internal/workflow/rollback.go b/internal/workflow/rollback.go new file mode 100644 index 0000000..fdd2acf --- /dev/null +++ b/internal/workflow/rollback.go @@ -0,0 +1,188 @@ +package workflow + +import ( + "context" + "fmt" + "sort" + "sync" + + "dev-cli/internal/executor" +) + +// RollbackHook defines undo logic for a remediation step. +type RollbackHook struct { + StepID string // ID of the step this rolls back + Name string // Human-readable name + Command string // Rollback command to execute + Validator func() bool // Optional: check if rollback succeeded + Priority int // Higher priority = execute first +} + +// RollbackRegistry tracks rollback hooks for active remediations. +type RollbackRegistry struct { + hooks []RollbackHook + mu sync.Mutex +} + +// NewRollbackRegistry creates a new rollback registry. +func NewRollbackRegistry() *RollbackRegistry { + return &RollbackRegistry{ + hooks: make([]RollbackHook, 0), + } +} + +// Register adds a rollback hook. +func (r *RollbackRegistry) Register(hook RollbackHook) { + r.mu.Lock() + defer r.mu.Unlock() + r.hooks = append(r.hooks, hook) +} + +// Unregister removes a rollback hook by step ID. +func (r *RollbackRegistry) Unregister(stepID string) { + r.mu.Lock() + defer r.mu.Unlock() + + filtered := make([]RollbackHook, 0, len(r.hooks)) + for _, hook := range r.hooks { + if hook.StepID != stepID { + filtered = append(filtered, hook) + } + } + r.hooks = filtered +} + +// Count returns the number of registered hooks. +func (r *RollbackRegistry) Count() int { + r.mu.Lock() + defer r.mu.Unlock() + return len(r.hooks) +} + +// Clear removes all registered hooks. +func (r *RollbackRegistry) Clear() { + r.mu.Lock() + defer r.mu.Unlock() + r.hooks = make([]RollbackHook, 0) +} + +// RollbackResult represents the outcome of a single rollback. +type RollbackResult struct { + StepID string + Name string + Success bool + Output string + Error error +} + +// ExecuteAll runs all rollbacks in priority order (highest first). +// Returns results for each rollback attempt. +func (r *RollbackRegistry) ExecuteAll(ctx context.Context) []RollbackResult { + r.mu.Lock() + + hooks := make([]RollbackHook, len(r.hooks)) + copy(hooks, r.hooks) + r.mu.Unlock() + + sort.Slice(hooks, func(i, j int) bool { + return hooks[i].Priority > hooks[j].Priority + }) + + results := make([]RollbackResult, 0, len(hooks)) + + for _, hook := range hooks { + select { + case <-ctx.Done(): + results = append(results, RollbackResult{ + StepID: hook.StepID, + Name: hook.Name, + Success: false, + Error: ctx.Err(), + }) + continue + default: + } + + result := r.executeRollback(ctx, hook) + results = append(results, result) + } + + return results +} + +// executeRollback runs a single rollback hook. +func (r *RollbackRegistry) executeRollback(ctx context.Context, hook RollbackHook) RollbackResult { + result := RollbackResult{ + StepID: hook.StepID, + Name: hook.Name, + } + + execResult := executor.ExecuteWithContext(ctx, hook.Command) + result.Output = execResult.Output + result.Success = execResult.ExitCode == 0 + + if !result.Success { + result.Error = fmt.Errorf("rollback failed with exit code %d", execResult.ExitCode) + } + + if result.Success && hook.Validator != nil { + if !hook.Validator() { + result.Success = false + result.Error = fmt.Errorf("rollback validation failed") + } + } + + return result +} + +// ExecuteForStep runs rollback only for a specific step. +func (r *RollbackRegistry) ExecuteForStep(ctx context.Context, stepID string) *RollbackResult { + r.mu.Lock() + var hook *RollbackHook + for _, h := range r.hooks { + if h.StepID == stepID { + hCopy := h + hook = &hCopy + break + } + } + r.mu.Unlock() + + if hook == nil { + return nil + } + + result := r.executeRollback(ctx, *hook) + return &result +} + +// GetHooks returns a copy of all registered hooks (for inspection). +func (r *RollbackRegistry) GetHooks() []RollbackHook { + r.mu.Lock() + defer r.mu.Unlock() + + hooks := make([]RollbackHook, len(r.hooks)) + copy(hooks, r.hooks) + return hooks +} + +// CreateRollbackHook is a helper to create a rollback hook from step info. +func CreateRollbackHook(stepID, name, command string, priority int) RollbackHook { + return RollbackHook{ + StepID: stepID, + Name: name, + Command: command, + Priority: priority, + } +} + +// CreateRollbackHookWithValidator creates a hook with a validation function. +func CreateRollbackHookWithValidator(stepID, name, command string, priority int, validator func() bool) RollbackHook { + return RollbackHook{ + StepID: stepID, + Name: name, + Command: command, + Priority: priority, + Validator: validator, + } +} diff --git a/internal/workflow/safemode.go b/internal/workflow/safemode.go new file mode 100644 index 0000000..31fc330 --- /dev/null +++ b/internal/workflow/safemode.go @@ -0,0 +1,183 @@ +package workflow + +import ( + "fmt" + "strings" +) + +// SafeMode controls whether remediation actions are executed or just previewed. +type SafeMode int + +const ( + // SafeModePreview is the default: shows what would happen without executing + SafeModePreview SafeMode = iota + // SafeModeExecute actually runs remediation commands + SafeModeExecute +) + +// String returns a human-readable representation of the SafeMode. +func (m SafeMode) String() string { + switch m { + case SafeModePreview: + return "preview" + case SafeModeExecute: + return "execute" + default: + return "unknown" + } +} + +// SafeModeContext wraps execution with preview/approval logic. +// It provides governance controls for automated remediation. +type SafeModeContext struct { + // Mode controls preview vs execute behavior + Mode SafeMode + + // ApprovalFunc is called to prompt for user confirmation + // Returns true if approved, false if denied + ApprovalFunc func(action string) bool + + // RollbackEnabled indicates whether rollback hooks should be registered + RollbackEnabled bool + + // DryRunOutput collects preview actions when in SafeModePreview + DryRunOutput []PreviewAction + + // DestructivePatterns are command patterns that require extra confirmation + DestructivePatterns []string +} + +// PreviewAction represents an action that would be taken in execute mode. +type PreviewAction struct { + Description string + Command string + Destructive bool + StepID string +} + +// DefaultDestructivePatterns returns common dangerous command patterns. +func DefaultDestructivePatterns() []string { + return []string{ + "rm -rf", + "rm -r /", + "dd if=", + "mkfs", + "> /dev/", + "chmod 777", + ":(){ :|:& };:", + "drop database", + "drop table", + "truncate table", + "delete from", + "git reset --hard", + "git clean -fdx", + "docker system prune", + } +} + +// NewSafeModeContext creates a preview-only context by default. +func NewSafeModeContext() *SafeModeContext { + return &SafeModeContext{ + Mode: SafeModePreview, + RollbackEnabled: true, + DryRunOutput: make([]PreviewAction, 0), + DestructivePatterns: DefaultDestructivePatterns(), + } +} + +// NewExecuteContext creates a context that will actually execute commands. +func NewExecuteContext(approvalFunc func(string) bool) *SafeModeContext { + return &SafeModeContext{ + Mode: SafeModeExecute, + ApprovalFunc: approvalFunc, + RollbackEnabled: true, + DryRunOutput: make([]PreviewAction, 0), + DestructivePatterns: DefaultDestructivePatterns(), + } +} + +// IsPreview returns true if in preview mode. +func (c *SafeModeContext) IsPreview() bool { + return c.Mode == SafeModePreview +} + +// PreviewAction records an action without executing (in preview mode). +func (c *SafeModeContext) PreviewAction(stepID, description, command string) { + destructive := c.isDestructive(command) + c.DryRunOutput = append(c.DryRunOutput, PreviewAction{ + Description: description, + Command: command, + Destructive: destructive, + StepID: stepID, + }) +} + +// RequireApproval prompts for confirmation before destructive operations. +// Returns true if approved or no approval function is set. +func (c *SafeModeContext) RequireApproval(action string) bool { + if c.ApprovalFunc == nil { + return true + } + return c.ApprovalFunc(action) +} + +// RequireApprovalForDestructive checks if the command is destructive and requires approval. +// Returns true if approved (or not destructive), false if denied. +func (c *SafeModeContext) RequireApprovalForDestructive(command string) bool { + if !c.isDestructive(command) { + return true + } + return c.RequireApproval(fmt.Sprintf("⚠️ Potentially destructive command:\n %s\n\nProceed?", command)) +} + +// isDestructive checks if a command matches any destructive pattern. +func (c *SafeModeContext) isDestructive(command string) bool { + lower := strings.ToLower(command) + for _, pattern := range c.DestructivePatterns { + if strings.Contains(lower, strings.ToLower(pattern)) { + return true + } + } + return false +} + +// GetPreviewSummary returns a formatted summary of all preview actions. +func (c *SafeModeContext) GetPreviewSummary() string { + if len(c.DryRunOutput) == 0 { + return "No actions would be taken." + } + + var sb strings.Builder + sb.WriteString("Preview of actions that would be taken:\n\n") + + for i, action := range c.DryRunOutput { + marker := " " + if action.Destructive { + marker = "⚠️" + } + sb.WriteString(fmt.Sprintf("%d. %s %s\n", i+1, marker, action.Description)) + if action.Command != "" { + sb.WriteString(fmt.Sprintf(" $ %s\n", action.Command)) + } + sb.WriteString("\n") + } + + destructiveCount := 0 + for _, a := range c.DryRunOutput { + if a.Destructive { + destructiveCount++ + } + } + + if destructiveCount > 0 { + sb.WriteString(fmt.Sprintf("⚠️ %d potentially destructive action(s) detected.\n", destructiveCount)) + sb.WriteString("Use --force to execute, or review commands carefully.\n") + } + + return sb.String() +} + +// ClearPreview resets the preview output. +func (c *SafeModeContext) ClearPreview() { + c.DryRunOutput = make([]PreviewAction, 0) +} diff --git a/internal/workflow/safemode_test.go b/internal/workflow/safemode_test.go new file mode 100644 index 0000000..2311335 --- /dev/null +++ b/internal/workflow/safemode_test.go @@ -0,0 +1,160 @@ +package workflow + +import ( + "testing" +) + +func TestSafeMode_DefaultIsPreview(t *testing.T) { + ctx := NewSafeModeContext() + + if ctx.Mode != SafeModePreview { + t.Errorf("expected default mode to be preview, got %v", ctx.Mode) + } + if !ctx.IsPreview() { + t.Error("IsPreview should return true by default") + } +} + +func TestSafeMode_ExecuteContext(t *testing.T) { + approvalCalled := false + ctx := NewExecuteContext(func(action string) bool { + approvalCalled = true + return true + }) + + if ctx.Mode != SafeModeExecute { + t.Errorf("expected execute mode, got %v", ctx.Mode) + } + if ctx.IsPreview() { + t.Error("IsPreview should return false for execute context") + } + + ctx.RequireApproval("test action") + if !approvalCalled { + t.Error("approval function should be called") + } +} + +func TestSafeMode_PreviewAction(t *testing.T) { + ctx := NewSafeModeContext() + + ctx.PreviewAction("step-1", "Install dependencies", "npm install") + ctx.PreviewAction("step-2", "Delete temp files", "rm -rf /tmp/test") + + if len(ctx.DryRunOutput) != 2 { + t.Errorf("expected 2 preview actions, got %d", len(ctx.DryRunOutput)) + } + if ctx.DryRunOutput[0].Command != "npm install" { + t.Errorf("expected first command 'npm install', got '%s'", ctx.DryRunOutput[0].Command) + } + if !ctx.DryRunOutput[1].Destructive { + t.Error("rm -rf should be detected as destructive") + } +} + +func TestSafeMode_DestructivePatternDetection(t *testing.T) { + ctx := NewSafeModeContext() + + tests := []struct { + command string + destructive bool + }{ + {"npm install", false}, + {"cat package.json", false}, + {"rm -rf /tmp/test", true}, + {"dd if=/dev/zero of=/dev/sda", true}, + {"DROP TABLE users", true}, + {"git reset --hard HEAD", true}, + {"docker system prune", true}, + {"echo hello", false}, + } + + for _, tt := range tests { + ctx.ClearPreview() + ctx.PreviewAction("test", "Test action", tt.command) + + if ctx.DryRunOutput[0].Destructive != tt.destructive { + t.Errorf("command '%s': expected destructive=%v, got %v", + tt.command, tt.destructive, ctx.DryRunOutput[0].Destructive) + } + } +} + +func TestSafeMode_RequireApprovalForDestructive(t *testing.T) { + approvalCount := 0 + ctx := NewExecuteContext(func(action string) bool { + approvalCount++ + return true + }) + + result := ctx.RequireApprovalForDestructive("npm install") + if !result { + t.Error("non-destructive command should be approved") + } + if approvalCount != 0 { + t.Errorf("expected 0 approval calls for non-destructive, got %d", approvalCount) + } + + result = ctx.RequireApprovalForDestructive("rm -rf /important") + if !result { + t.Error("destructive command should be approved when approval func returns true") + } + if approvalCount != 1 { + t.Errorf("expected 1 approval call for destructive, got %d", approvalCount) + } +} + +func TestSafeMode_GetPreviewSummary(t *testing.T) { + ctx := NewSafeModeContext() + + summary := ctx.GetPreviewSummary() + if summary != "No actions would be taken." { + t.Errorf("expected 'No actions would be taken.', got '%s'", summary) + } + + ctx.PreviewAction("s1", "Safe action", "npm install") + ctx.PreviewAction("s2", "Dangerous action", "rm -rf /") + + summary = ctx.GetPreviewSummary() + if len(summary) == 0 { + t.Error("summary should not be empty") + } + + if !contains(summary, "destructive") { + t.Error("summary should mention destructive actions") + } +} + +func TestSafeMode_ClearPreview(t *testing.T) { + ctx := NewSafeModeContext() + + ctx.PreviewAction("s1", "Action 1", "cmd1") + ctx.PreviewAction("s2", "Action 2", "cmd2") + ctx.ClearPreview() + + if len(ctx.DryRunOutput) != 0 { + t.Errorf("expected 0 actions after clear, got %d", len(ctx.DryRunOutput)) + } +} + +func TestSafeMode_String(t *testing.T) { + if SafeModePreview.String() != "preview" { + t.Errorf("expected 'preview', got '%s'", SafeModePreview.String()) + } + if SafeModeExecute.String() != "execute" { + t.Errorf("expected 'execute', got '%s'", SafeModeExecute.String()) + } +} + +func contains(s, substr string) bool { + return len(s) >= len(substr) && (s == substr || containsSubstr(s, substr)) +} + +func containsSubstr(s, substr string) bool { + for i := 0; i <= len(s)-len(substr); i++ { + if s[i:i+len(substr)] == substr { + return true + } + } + return false +} diff --git a/internal/workflow/workflow.go b/internal/workflow/workflow.go new file mode 100644 index 0000000..2f873b5 --- /dev/null +++ b/internal/workflow/workflow.go @@ -0,0 +1,157 @@ +// Package workflow provides multi-step workflow automation with conditional +// branching, rollback capabilities, and checkpoint/resume functionality. +package workflow + +import ( + "time" +) + +// RunStatus represents the current state of a workflow run. +type RunStatus string + +const ( + StatusPending RunStatus = "pending" + StatusRunning RunStatus = "running" + StatusPaused RunStatus = "paused" + StatusCompleted RunStatus = "completed" + StatusFailed RunStatus = "failed" + StatusRolledBack RunStatus = "rolledback" +) + +// StepStatus represents the current state of a step execution. +type StepStatus string + +const ( + StepPending StepStatus = "pending" + StepRunning StepStatus = "running" + StepSuccess StepStatus = "success" + StepFailed StepStatus = "failed" + StepSkipped StepStatus = "skipped" + StepRolledBack StepStatus = "rolledback" +) + +// ConditionType defines how a condition should be evaluated. +type ConditionType string + +const ( + CondExitCode ConditionType = "exit_code" + CondOutputContains ConditionType = "output_contains" + CondOutputMatches ConditionType = "output_matches" + CondFileExists ConditionType = "file_exists" + CondEnvSet ConditionType = "env_set" +) + +// FailureAction defines what to do when a workflow fails. +type FailureAction string + +const ( + FailureAbort FailureAction = "abort" + FailureRollback FailureAction = "rollback" + FailureContinue FailureAction = "continue" +) + +// Condition specifies when a step should execute. +type Condition struct { + Type ConditionType `yaml:"type"` + Value string `yaml:"value"` + // StepRef references a previous step's result (optional, defaults to previous step) + StepRef string `yaml:"step_ref,omitempty"` +} + +// RollbackAction defines how to undo a step. +type RollbackAction struct { + Command string `yaml:"command"` + Timeout time.Duration `yaml:"timeout,omitempty"` +} + +// Step represents a single executable action in a workflow. +type Step struct { + ID string `yaml:"id"` + Name string `yaml:"name"` + Command string `yaml:"command"` + Condition *Condition `yaml:"condition,omitempty"` + OnSuccess string `yaml:"on_success,omitempty"` // Next step ID (optional) + OnFailure string `yaml:"on_failure,omitempty"` // Step ID, "rollback", or "abort" + Rollback *RollbackAction `yaml:"rollback,omitempty"` + Timeout time.Duration `yaml:"timeout,omitempty"` + Retries int `yaml:"retries,omitempty"` + Env map[string]string `yaml:"env,omitempty"` + WorkDir string `yaml:"workdir,omitempty"` +} + +// FailurePolicy defines workflow-level failure handling. +type FailurePolicy struct { + Action FailureAction `yaml:"action"` +} + +// Workflow represents a complete multi-step automation definition. +type Workflow struct { + ID string `yaml:"id,omitempty"` + Name string `yaml:"name"` + Description string `yaml:"description,omitempty"` + Steps []Step `yaml:"steps"` + OnFailure *FailurePolicy `yaml:"on_failure,omitempty"` + Env map[string]string `yaml:"env,omitempty"` +} + +// StepResult holds the outcome of executing a single step. +type StepResult struct { + StepID string + Status StepStatus + ExitCode int + Output string + Error string + StartedAt time.Time + CompletedAt time.Time + Duration time.Duration + Retries int +} + +// RunState holds the complete state of a workflow execution. +type RunState struct { + RunID string + WorkflowID string + WorkflowName string + Status RunStatus + CurrentStepIdx int + StepResults map[string]*StepResult + StartedAt time.Time + UpdatedAt time.Time + CompletedAt time.Time + Error string +} + +// NewRunState creates a new RunState for a workflow execution. +func NewRunState(runID string, wf *Workflow) *RunState { + return &RunState{ + RunID: runID, + WorkflowID: wf.ID, + WorkflowName: wf.Name, + Status: StatusPending, + StepResults: make(map[string]*StepResult), + StartedAt: time.Now(), + UpdatedAt: time.Now(), + } +} + +// GetStepResult returns the result for a given step ID. +func (r *RunState) GetStepResult(stepID string) *StepResult { + return r.StepResults[stepID] +} + +// SetStepResult stores the result for a step. +func (r *RunState) SetStepResult(result *StepResult) { + r.StepResults[result.StepID] = result + r.UpdatedAt = time.Now() +} + +// LastStepResult returns the most recently completed step result. +func (r *RunState) LastStepResult() *StepResult { + var last *StepResult + for _, result := range r.StepResults { + if last == nil || result.CompletedAt.After(last.CompletedAt) { + last = result + } + } + return last +}