diff --git a/cmd/survey/init.go b/cmd/survey/init.go index 6c67f65..ef9e108 100644 --- a/cmd/survey/init.go +++ b/cmd/survey/init.go @@ -26,15 +26,15 @@ func initInterfaces() map[string]task.Interface { return interfaces } -func initTasks(app *service.App) []task.Tasker { - var taskList []task.Tasker +func initTasks(app *service.App) map[string]task.Tasker { + taskMap := make(map[string]task.Tasker) for _, name := range task.ListTasks() { - t, err := task.GetTasks(name, app) + t, err := task.GetTask(name, app) if err != nil { continue } - taskList = append(taskList, t) + taskMap[name] = t } - return taskList + return taskMap } diff --git a/cmd/survey/runner.go b/cmd/survey/runner.go index 3b8fd1b..0fb8150 100644 --- a/cmd/survey/runner.go +++ b/cmd/survey/runner.go @@ -32,15 +32,14 @@ func runStartCmd(cmd *cobra.Command, args []string) error { } } - if cmd.Flags().Changed("stage") { - if surveyCmd.Task < 1 || surveyCmd.Task > len(surveyTasks) { - return fmt.Errorf("invalid stage") + if cmd.Flags().Changed("task") { + if surveyTask := surveyTasks[surveyCmd.Task]; surveyTask == nil { + return fmt.Errorf("invalid task, valid are %v", task.ListTasks()) } - if err := task.RunTask(cmd.Context(), surveyTasks[surveyCmd.Task-1], - interfaces[surveyCmd.InterfaceType], tasks.InterfaceToType(interfaces[surveyCmd.InterfaceType])); err != nil { - return err + + surveyTasks = map[string]task.Tasker{ + surveyCmd.Task: surveyTasks[surveyCmd.Task], } - return nil } if err := newIntroModel().Run(); err != nil { @@ -53,7 +52,7 @@ func runStartCmd(cmd *cobra.Command, args []string) error { return nil } -func taskLoop(ctx context.Context, surveyTasks []task.Tasker, interfaces map[string]task.Interface) error { +func taskLoop(ctx context.Context, surveyTasks map[string]task.Tasker, interfaces map[string]task.Interface) error { var iNames []string for name := range interfaces { iNames = append(iNames, name) @@ -63,14 +62,16 @@ func taskLoop(ctx context.Context, surveyTasks []task.Tasker, interfaces map[str iNames[i], iNames[j] = iNames[j], iNames[i] }) - for i, t := range surveyTasks { - idx := i % len(iNames) - selected := interfaces[iNames[idx]] + idx := 0 + for _, t := range surveyTasks { + iIdx := idx % len(iNames) + selected := interfaces[iNames[iIdx]] if err := task.RunTask(ctx, t, selected, tasks.InterfaceToType(selected)); err != nil { return err } } + idx++ return nil } diff --git a/internal/commands/survey/root.go b/internal/commands/survey/root.go index b5e00a9..df5d9a2 100644 --- a/internal/commands/survey/root.go +++ b/internal/commands/survey/root.go @@ -4,11 +4,6 @@ import ( "github.com/spf13/cobra" ) -var ( - InterfaceType string - Task int -) - // RootCmd is the base command for the survey CLI. var RootCmd = &cobra.Command{ Use: "survey", @@ -24,6 +19,4 @@ Your responses will be kept confidential and used solely for research purposes.` func init() { RootCmd.CompletionOptions.DisableDefaultCmd = true - StartCmd.Flags().StringVarP(&InterfaceType, "interface", "i", "tui", "Specify interface.") - StartCmd.Flags().IntVarP(&Task, "task", "t", 1, "Run task directly") } diff --git a/internal/commands/survey/start.go b/internal/commands/survey/start.go index c987941..1a1de58 100644 --- a/internal/commands/survey/start.go +++ b/internal/commands/survey/start.go @@ -2,8 +2,18 @@ package surveyCmd import "github.com/spf13/cobra" +var ( + InterfaceType string + Task string +) + // StartCmd is the start command - RunE is set in cmd/survey/ var StartCmd = &cobra.Command{ Use: "start", Short: "Start the user survey", } + +func init() { + StartCmd.Flags().StringVarP(&Task, "task", "t", "", "Specify task.") + StartCmd.Flags().StringVarP(&InterfaceType, "interface", "i", "", "Specify interface.") +} diff --git a/pkg/task/register.go b/pkg/task/register.go index c98d75e..37067fa 100644 --- a/pkg/task/register.go +++ b/pkg/task/register.go @@ -40,7 +40,7 @@ func RegisterTask(name string, constructor func(*service.App) Tasker) { taskRegistry[name] = constructor } -func GetTasks(name string, app *service.App) (Tasker, error) { +func GetTask(name string, app *service.App) (Tasker, error) { constructor, ok := taskRegistry[name] if !ok { return nil, fmt.Errorf("task %q not found", name)