diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index a7a68e4..45d65d9 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -29,6 +29,4 @@ jobs: run: go mod download - name: Run tests - run: go test - - + run: go test -race diff --git a/examples/custom/main.go b/examples/custom/main.go index 64c3aeb..b86eae1 100644 --- a/examples/custom/main.go +++ b/examples/custom/main.go @@ -16,7 +16,7 @@ func main() { // sleep for 3 seconds then return nil Task: func(t *taskin.Task) error { for i := 0; i < 2; i++ { - t.Title = fmt.Sprintf("Task 1 - [%d/%d]", i+1, 2) + t.SetTitle(fmt.Sprintf("Task 1 - [%d/%d]", i+1, 2)) time.Sleep(1 * time.Second) } return nil @@ -28,7 +28,7 @@ func main() { Task: func(t *taskin.Task) error { for i := 0; i < 5; i++ { t.Progress(i+1, 5) - t.Title = fmt.Sprintf("Task 2 - [%d/%d]", i+1, 5) + t.SetTitle(fmt.Sprintf("Task 2 - [%d/%d]", i+1, 5)) time.Sleep(1 * time.Second) } return nil @@ -39,7 +39,7 @@ func main() { // sleep for 3 seconds then return nil Task: func(t *taskin.Task) error { for i := 0; i < 2; i++ { - t.Title = fmt.Sprintf("Task 3 - [%d/%d]", i+1, 2) + t.SetTitle(fmt.Sprintf("Task 3 - [%d/%d]", i+1, 2)) time.Sleep(1 * time.Second) } return nil diff --git a/examples/disable-ui/main.go b/examples/disable-ui/main.go index 22583b8..5f42566 100644 --- a/examples/disable-ui/main.go +++ b/examples/disable-ui/main.go @@ -13,7 +13,7 @@ func main() { Title: "Task with UI disabled", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("Task with UI disabled: [%d/3] processing", i+1) + t.SetTitle(fmt.Sprintf("Task with UI disabled: [%d/3] processing", i+1)) time.Sleep(500 * time.Millisecond) } return nil @@ -26,7 +26,7 @@ func main() { Title: "Child task 1", Task: func(t *taskin.Task) error { for i := 0; i < 2; i++ { - t.Title = fmt.Sprintf("Child task 1: [%d/2] working", i+1) + t.SetTitle(fmt.Sprintf("Child task 1: [%d/2] working", i+1)) time.Sleep(300 * time.Millisecond) } return nil @@ -36,7 +36,7 @@ func main() { Title: "Child task 2", Task: func(t *taskin.Task) error { for i := 0; i < 2; i++ { - t.Title = fmt.Sprintf("Child task 2: [%d/2] working", i+1) + t.SetTitle(fmt.Sprintf("Child task 2: [%d/2] working", i+1)) time.Sleep(300 * time.Millisecond) } return nil diff --git a/examples/error/main.go b/examples/error/main.go index d7d4ef9..fe72737 100644 --- a/examples/error/main.go +++ b/examples/error/main.go @@ -13,7 +13,7 @@ func main() { Title: "Task 1", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("Task 1: [%d/3] seconds have passed", i+1) + t.SetTitle(fmt.Sprintf("Task 1: [%d/3] seconds have passed", i+1)) time.Sleep(500 * time.Millisecond) } return nil @@ -29,7 +29,7 @@ func main() { Title: "Task 3", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("Task 3: [%d/3] seconds have passed", i+1) + t.SetTitle(fmt.Sprintf("Task 3: [%d/3] seconds have passed", i+1)) time.Sleep(500 * time.Millisecond) } return nil diff --git a/examples/multi/main.go b/examples/multi/main.go index 164e737..505789c 100644 --- a/examples/multi/main.go +++ b/examples/multi/main.go @@ -13,7 +13,7 @@ func main() { Title: "Mow the lawn", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("Mow the lawn: [%d/3] passes", i+1) + t.SetTitle(fmt.Sprintf("Mow the lawn: [%d/3] passes", i+1)) time.Sleep(500 * time.Millisecond) } return nil @@ -26,7 +26,7 @@ func main() { Title: "Pluck the silkies [0/3]", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("Pluck the silkies [%d/3]", i+1) + t.SetTitle(fmt.Sprintf("Pluck the silkies [%d/3]", i+1)) time.Sleep(500 * time.Millisecond) } return nil @@ -40,7 +40,7 @@ func main() { Title: "[0/3] Pluck the Polish", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("[%d/3] Pluck the Polish", i+1) + t.SetTitle(fmt.Sprintf("[%d/3] Pluck the Polish", i+1)) time.Sleep(500 * time.Millisecond) } return nil @@ -50,7 +50,7 @@ func main() { Title: "[0/3] Pluck the Marans", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("[%d/3] Pluck the Marans", i+1) + t.SetTitle(fmt.Sprintf("[%d/3] Pluck the Marans", i+1)) time.Sleep(500 * time.Millisecond) } return nil @@ -63,7 +63,7 @@ func main() { Title: "[0/3] Pluck the leghorns", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("[%d/3] Pluck the Leghorns", i+1) + t.SetTitle(fmt.Sprintf("[%d/3] Pluck the Leghorns", i+1)) time.Sleep(500 * time.Millisecond) } return nil @@ -76,7 +76,7 @@ func main() { Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { t.Progress(i+1, 3) - t.Title = fmt.Sprintf("Paint the house: [%d/3] walls painted", i+1) + t.SetTitle(fmt.Sprintf("Paint the house: [%d/3] walls painted", i+1)) time.Sleep(500 * time.Millisecond) } return nil diff --git a/examples/progress/main.go b/examples/progress/main.go index f30e3a7..c49787e 100644 --- a/examples/progress/main.go +++ b/examples/progress/main.go @@ -14,7 +14,7 @@ func main() { Task: func(t *taskin.Task) error { for i := 0; i < 5; i++ { t.Progress(i+1, 5) - t.Title = fmt.Sprintf("Progress [%d/%d]", i+1, 5) + t.SetTitle(fmt.Sprintf("Progress [%d/%d]", i+1, 5)) time.Sleep(1 * time.Second) } return nil diff --git a/examples/simple/main.go b/examples/simple/main.go index 3ffaab3..f934385 100644 --- a/examples/simple/main.go +++ b/examples/simple/main.go @@ -13,7 +13,7 @@ func main() { Title: "Task 1", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("Task 1: [%d/3] seconds have passed", i+1) + t.SetTitle(fmt.Sprintf("Task 1: [%d/3] seconds have passed", i+1)) time.Sleep(500 * time.Millisecond) } return nil @@ -23,7 +23,7 @@ func main() { Title: "Task 2", Task: func(t *taskin.Task) error { for i := 0; i < 3; i++ { - t.Title = fmt.Sprintf("Task 2: [%d/3] seconds have passed", i+1) + t.SetTitle(fmt.Sprintf("Task 2: [%d/3] seconds have passed", i+1)) time.Sleep(500 * time.Millisecond) } return nil diff --git a/models.go b/models.go index 5466290..07361ca 100644 --- a/models.go +++ b/models.go @@ -3,12 +3,32 @@ package taskin import ( "github.com/charmbracelet/bubbles/progress" "github.com/charmbracelet/bubbles/spinner" + tea "github.com/charmbracelet/bubbletea" ) type TerminateWithError struct { Error error } +type taskStartedMsg struct { + Path []int +} + +type taskUpdatedMsg struct { + Path []int + Task Task +} + +type taskCompletedMsg struct { + Path []int + Task Task +} + +type taskFailedMsg struct { + Path []int + Task Task +} + type TaskState int const ( @@ -50,4 +70,5 @@ type Model struct { HideView bool Shutdown bool ShutdownError error + taskMessages chan tea.Msg } diff --git a/mvc.go b/mvc.go index d5f0039..1951e1b 100644 --- a/mvc.go +++ b/mvc.go @@ -19,17 +19,34 @@ func (m *Model) Init() tea.Cmd { } } } + if m.taskMessages != nil { + if len(m.Runners) == 0 { + return tea.Quit + } + cmds = append(cmds, runTasksCmd(cloneRunners(m.Runners), m.taskMessages), waitForTaskMessage(m.taskMessages)) + } return tea.Batch(cmds...) } +func waitForTaskMessage(messages <-chan tea.Msg) tea.Cmd { + return func() tea.Msg { + return <-messages + } +} + +func (m *Model) waitForTaskMessage() tea.Cmd { + if m.taskMessages == nil { + return nil + } + return waitForTaskMessage(m.taskMessages) +} + func (m *Model) SetShutdown(err error) { m.Shutdown = true m.ShutdownError = err } func (m *Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { - var cmds []tea.Cmd - if m.Shutdown && m.ShutdownError != nil { return m, tea.Quit } @@ -38,9 +55,36 @@ func (m *Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case TerminateWithError: m.SetShutdown(msg.Error) return m, tea.Quit + case taskStartedMsg: + if runner := m.runnerAtPath(msg.Path); runner != nil { + runner.State = Running + } + return m, m.waitForTaskMessage() + case taskUpdatedMsg: + if runner := m.runnerAtPath(msg.Path); runner != nil { + runner.Task = msg.Task + } + return m, m.waitForTaskMessage() + case taskCompletedMsg: + if runner := m.runnerAtPath(msg.Path); runner != nil { + runner.Task = msg.Task + runner.State = Completed + } + allDone, anyFailed := m.checkTasksState() + if allDone && !anyFailed { + return m, tea.Quit + } + return m, m.waitForTaskMessage() + case taskFailedMsg: + if runner := m.runnerAtPath(msg.Path); runner != nil { + runner.Task = msg.Task + runner.State = Failed + } + return m, m.waitForTaskMessage() case spinner.TickMsg: - // Helper function to update spinners recursively + var cmds []tea.Cmd + var updateSpinners func(runner *Runner) []tea.Cmd updateSpinners = func(runner *Runner) []tea.Cmd { var spinnerCmds []tea.Cmd @@ -53,7 +97,6 @@ func (m *Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { } } - // Recursively update all children's spinners for i := range runner.Children { spinnerCmds = append(spinnerCmds, updateSpinners(&runner.Children[i])...) } @@ -61,20 +104,12 @@ func (m *Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return spinnerCmds } - allDone := true for i := range m.Runners { cmds = append(cmds, updateSpinners(&m.Runners[i])...) - - if m.Runners[i].State == Failed { - return m, tea.Quit - } - - if m.Runners[i].State != Completed && m.Runners[i].State != Failed { - allDone = false - } } - if allDone { + allDone, anyFailed := m.checkTasksState() + if allDone && !anyFailed { return m, tea.Quit } @@ -84,6 +119,20 @@ func (m *Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return m, nil } +func (m *Model) runnerAtPath(path []int) *Runner { + if len(path) == 0 || path[0] < 0 || path[0] >= len(m.Runners) { + return nil + } + runner := &m.Runners[path[0]] + for _, index := range path[1:] { + if index < 0 || index >= len(runner.Children) { + return nil + } + runner = &runner.Children[index] + } + return runner +} + func (m *Model) checkTasksState() (allDone, anyFailed bool) { allDone = true for _, runner := range m.Runners { @@ -148,7 +197,6 @@ func renderTask(runner Runner, indent string) string { view = indent + lipgloss.NewStyle().Render(status) + "\n" } - // Recursively render children if len(runner.Children) > 0 && (runner.State == Running || IsCI() || runner.Config.DisableUI) { for _, child := range runner.Children { view += renderTask(child, indent+" ") diff --git a/readme.md b/readme.md index 3e0e77a..663e004 100644 --- a/readme.md +++ b/readme.md @@ -74,21 +74,21 @@ https://github.com/fumeapp/taskin/blob/3cd766c21e5eaba5edb33f38d3781d6cf814f9f9/ The `*taskin.Task` struct passeed into your task has some useful properties that you can use to customize the task view. ### Change the title -Already demonstrated in most of the examples, you can change `t.Title` at any time +Already demonstrated in most of the examples, you can change the title with `t.SetTitle("New title")`. ### Hide a view Sometimes you might need to temporarily hide you task view in order to prompt a user for input. -You can do this by toggling the task.HideView boolean. +You can do this by calling `t.SetHideView(true)` and showing it again with `t.SetHideView(false)`. ```go Task: func(T *taskin.Task ) error { - t.HideView = true + t.SetHideView(true) if err := PromptForInput(); err != nil { - t.HideView = false + t.SetHideView(false) return err } - t.HideView = false - t.Title = "Input received" + t.SetHideView(false) + t.SetTitle("Input received") return nil } diff --git a/taskin.go b/taskin.go index 3dae290..273aead 100644 --- a/taskin.go +++ b/taskin.go @@ -6,6 +6,7 @@ import ( "io" "os" "regexp" + "sync" "github.com/charmbracelet/bubbles/progress" "github.com/charmbracelet/bubbles/spinner" @@ -13,15 +14,19 @@ import ( "github.com/charmbracelet/lipgloss" ) -var program *tea.Program +type taskUpdateCallback func(Task) + +var taskUpdateCallbacks sync.Map + +var ansiEscapeCodePattern = regexp.MustCompile(` *\x1b\[[0-?]*[ -/]*[@-~]`) func NewRunner(task Task, cfg Config) Runner { var spinr *spinner.Model if !IsCI() && !cfg.DisableUI { - spinnerModel := spinner.New(spinner.WithSpinner(cfg.Spinner)) // Initialize with a spinner model - spinnerModel.Style = lipgloss.NewStyle().Foreground(cfg.Colors.Spinner) // Styling spinner + spinnerModel := spinner.New(spinner.WithSpinner(cfg.Spinner)) + spinnerModel.Style = lipgloss.NewStyle().Foreground(cfg.Colors.Spinner) spinr = &spinnerModel if task.ShowProgress.Total != 0 { @@ -43,17 +48,41 @@ func NewRunner(task Task, cfg Config) Runner { } func (task *Task) Progress(current, total int) { - task.ShowProgress = TaskProgress{Current: current, Total: total} - if IsCI() { + task.applyProgress(TaskProgress{Current: current, Total: total}) + task.notifyUpdate() +} + +func (task *Task) SetTitle(title string) { + task.Title = title + task.notifyUpdate() +} + +func (task *Task) SetHideView(hide bool) { + task.HideView = hide + task.notifyUpdate() +} + +func (task *Task) applyProgress(taskProgress TaskProgress) { + task.ShowProgress = taskProgress + if IsCI() || task.Config.DisableUI { return } if !task.Bar.IsAnimating() { task.Bar = progress.New(task.Config.ProgressOptions...) } - if total != 0 { // Check if TaskProgress is set - percent := float64(current) / float64(total) - task.Bar.SetPercent(percent) + if taskProgress.Total == 0 { + return } + percent := float64(taskProgress.Current) / float64(taskProgress.Total) + task.Bar.SetPercent(percent) +} + +func (task *Task) notifyUpdate() { + callback, ok := taskUpdateCallbacks.Load(task) + if !ok { + return + } + callback.(taskUpdateCallback)(snapshotTask(*task)) } type ansiEscapeCodeFilter struct { @@ -61,25 +90,21 @@ type ansiEscapeCodeFilter struct { } func (f *ansiEscapeCodeFilter) Write(p []byte) (n int, err error) { - // Corrected regular expression to match ANSI escape codes - re := regexp.MustCompile(` *\x1b\[[0-?]*[ -/]*[@-~]`) - // Remove the escape codes from the input - p = re.ReplaceAll(p, []byte{}) - // Write the filtered input to the original writer + p = ansiEscapeCodePattern.ReplaceAll(p, []byte{}) return f.writer.Write(p) } func (r *Runners) Run() error { - m := &Model{Runners: *r, Shutdown: false, ShutdownError: nil} + m := &Model{Runners: cloneRunners(*r), Shutdown: false, ShutdownError: nil, taskMessages: make(chan tea.Msg, 64)} var out io.Writer = os.Stdout - // Check if we need to disable UI features or are in CI mode if IsCI() || (len(*r) > 0 && (*r)[0].Config.DisableUI) { out = &ansiEscapeCodeFilter{writer: out} } - program = tea.NewProgram(m, tea.WithInput(nil), tea.WithOutput(out)) + program := tea.NewProgram(m, tea.WithInput(nil), tea.WithOutput(out)) _, err := program.Run() + *r = m.Runners if err != nil { return fmt.Errorf("program run error: %w", err) } @@ -99,54 +124,113 @@ func New(tasks Tasks, cfg Config) Runners { runners = append(runners, NewRunner(task, cfg)) } - // Helper function to run a task and its children recursively - var runTaskAndChildren func(runner *Runner) error - runTaskAndChildren = func(runner *Runner) error { - runner.State = Running + return runners +} - // Run the task itself first if it has a function - if runner.Task.Task != nil { - err := runner.Task.Task(&runner.Task) - if err != nil { - runner.Task.Title = fmt.Sprintf("%s - %s", runner.Task.Title, err.Error()) - runner.State = Failed - return err +func runTasksCmd(runners Runners, messages chan<- tea.Msg) tea.Cmd { + return func() tea.Msg { + go func() { + if err := runRunners(runners, messages); err != nil { + messages <- TerminateWithError{Error: err} } - } + }() + return nil + } +} - // Run all children recursively - for i := range runner.Children { - err := runTaskAndChildren(&runner.Children[i]) - if err != nil { - runner.State = Failed - return err +func runRunners(runners Runners, messages chan<- tea.Msg) error { + var firstErr error + for i := range runners { + err := runTaskAndChildren(&runners[i], pathWithIndex(nil, i), messages) + if err != nil { + if firstErr == nil { + firstErr = err + } + if runners[i].Config.Options.ExitOnFailure { + return firstErr } } + } + return firstErr +} + +func runTaskAndChildren(runner *Runner, path []int, messages chan<- tea.Msg) error { + runner.State = Running + messages <- taskStartedMsg{Path: path} - runner.State = Completed - if program != nil { - program.Send(spinner.TickMsg{}) + task := runner.Task + callback := taskUpdateCallback(func(updated Task) { + messages <- taskUpdatedMsg{Path: path, Task: updated} + }) + taskUpdateCallbacks.Store(&task, callback) + defer taskUpdateCallbacks.Delete(&task) + + var err error + if task.Task != nil { + err = task.Task(&task) + } + + runner.Task = snapshotTask(task) + if err != nil { + runner.Task.Title = fmt.Sprintf("%s - %s", runner.Task.Title, err.Error()) + runner.State = Failed + messages <- taskFailedMsg{Path: path, Task: runner.Task} + return err + } + messages <- taskUpdatedMsg{Path: path, Task: runner.Task} + + for i := range runner.Children { + if err := runTaskAndChildren(&runner.Children[i], pathWithIndex(path, i), messages); err != nil { + runner.State = Failed + messages <- taskFailedMsg{Path: path, Task: runner.Task} + return err } - return nil } - go func() { - for i := range runners { - // Check for previous failures - for _, prev := range runners[:i] { - if prev.State == Failed && prev.Config.Options.ExitOnFailure { - return - } - } + runner.State = Completed + messages <- taskCompletedMsg{Path: path, Task: runner.Task} + return nil +} - err := runTaskAndChildren(&runners[i]) - if err != nil && program != nil { - program.Send(TerminateWithError{Error: err}) - } +func cloneRunners(runners Runners) Runners { + if runners == nil { + return nil + } + cloned := make(Runners, len(runners)) + for i := range runners { + cloned[i] = runners[i] + cloned[i].Task = snapshotTask(runners[i].Task) + cloned[i].Children = cloneRunners(runners[i].Children) + if runners[i].Spinner != nil { + spinnerCopy := *runners[i].Spinner + cloned[i].Spinner = &spinnerCopy } - }() + } + return cloned +} - return runners +func snapshotTask(task Task) Task { + task.Tasks = cloneTasks(task.Tasks) + return task +} + +func cloneTasks(tasks Tasks) Tasks { + if tasks == nil { + return nil + } + cloned := make(Tasks, len(tasks)) + for i := range tasks { + cloned[i] = tasks[i] + cloned[i].Tasks = cloneTasks(tasks[i].Tasks) + } + return cloned +} + +func pathWithIndex(path []int, index int) []int { + next := make([]int, len(path)+1) + copy(next, path) + next[len(path)] = index + return next } func IsCI() bool { diff --git a/taskin_test.go b/taskin_test.go index 65bd20e..6ac7cde 100644 --- a/taskin_test.go +++ b/taskin_test.go @@ -46,10 +46,14 @@ func TestRunnersRun(t *testing.T) { } func TestNew(t *testing.T) { + ran := false tasks := Tasks{ Task{ Title: "Test Task", - Task: func(t *Task) error { return nil }, + Task: func(t *Task) error { + ran = true + return nil + }, }, } cfg := Config{ @@ -61,6 +65,10 @@ func TestNew(t *testing.T) { if len(runners) != 1 { t.Errorf("Expected New to return 1 runner, got '%d'", len(runners)) } + + if ran { + t.Error("Expected New not to start tasks") + } } func TestTaskProgress(t *testing.T) { @@ -79,3 +87,33 @@ func TestTaskProgress(t *testing.T) { t.Errorf("Expected TaskProgress to be 1/10, got %d/%d", runner.Task.ShowProgress.Current, runner.Task.ShowProgress.Total) } } + +func TestRunAppliesTaskUpdates(t *testing.T) { + tasks := Tasks{ + { + Title: "Test Task", + Task: func(t *Task) error { + t.SetTitle("Updated Task") + t.Progress(2, 4) + return nil + }, + }, + } + cfg := Defaults + cfg.DisableUI = true + + runners := New(tasks, cfg) + if err := runners.Run(); err != nil { + t.Fatalf("Expected Run to return nil, got %s", err.Error()) + } + + if runners[0].State != Completed { + t.Fatalf("Expected runner state to be completed, got %d", runners[0].State) + } + if runners[0].Task.Title != "Updated Task" { + t.Fatalf("Expected task title update, got %q", runners[0].Task.Title) + } + if runners[0].Task.ShowProgress.Current != 2 || runners[0].Task.ShowProgress.Total != 4 { + t.Fatalf("Expected progress 2/4, got %d/%d", runners[0].Task.ShowProgress.Current, runners[0].Task.ShowProgress.Total) + } +}