summaryrefslogtreecommitdiff
path: root/cmd
diff options
context:
space:
mode:
Diffstat (limited to 'cmd')
-rw-r--r--cmd/comicforge/cli.go51
-rw-r--r--cmd/comicforge/cli_test.go148
2 files changed, 195 insertions, 4 deletions
diff --git a/cmd/comicforge/cli.go b/cmd/comicforge/cli.go
index c241b2c..1a7f1ec 100644
--- a/cmd/comicforge/cli.go
+++ b/cmd/comicforge/cli.go
@@ -28,6 +28,7 @@ type commandDeps struct {
type cliFlags struct {
vocab string
+ prompt string
configPath string
promptsDir string
outputDir string
@@ -82,8 +83,17 @@ func newRootCommandWithDeps(deps commandDeps) *cobra.Command {
_, _ = fmt.Fprintln(cmd.OutOrStdout(), version.Version)
return nil
}
+ if cmd.Flags().Changed("prompt") && strings.TrimSpace(flags.prompt) == "" {
+ return fmt.Errorf("--prompt is required when set")
+ }
+ if strings.TrimSpace(flags.prompt) != "" && strings.TrimSpace(flags.vocab) != "" {
+ return fmt.Errorf("only one of --prompt and --vocab may be set")
+ }
+ if strings.TrimSpace(flags.prompt) != "" {
+ return runPromptCommand(cmd.Context(), cmd, deps, flags)
+ }
if strings.TrimSpace(flags.vocab) == "" {
- return fmt.Errorf("--vocab is required")
+ return fmt.Errorf("--vocab is required unless --prompt is set")
}
return runCommand(cmd.Context(), cmd, deps, flags)
},
@@ -91,6 +101,7 @@ func newRootCommandWithDeps(deps commandDeps) *cobra.Command {
cmd.Flags().BoolVarP(&flags.version, "version", "v", false, "print version information")
cmd.Flags().StringVar(&flags.vocab, "vocab", "", "path to the vocabulary input file")
+ cmd.Flags().StringVar(&flags.prompt, "prompt", "", "manual prompt for generating a single image")
cmd.Flags().StringVar(&flags.configPath, "config", "", "config file (default: search ~/.config/comicforge, $HOME, and .)")
cmd.Flags().StringVar(&flags.promptsDir, "prompts-dir", "", "directory containing prompt templates")
cmd.Flags().StringVar(&flags.outputDir, "output", ".", "root output directory for generated comic data")
@@ -175,6 +186,44 @@ func runCommand(ctx context.Context, cmd *cobra.Command, deps commandDeps, flags
return runner.Run(ctx, flags.vocab)
}
+func runPromptCommand(ctx context.Context, cmd *cobra.Command, deps commandDeps, flags cliFlags) error {
+ if ctx == nil {
+ ctx = context.Background()
+ }
+
+ cfg, err := deps.loadConfig(flags.configPath)
+ if err != nil {
+ return err
+ }
+
+ applyConfigOverrides(cmd, cfg, flags)
+ imageProvider, err := deps.newImageProvider(cfg)
+ if err != nil {
+ return fmt.Errorf("build image provider: %w", err)
+ }
+
+ runner := deps.newRunner(&comic.RunnerConfig{
+ ImageProvider: imageProvider,
+ Prompts: cfg,
+ OutputDir: flags.outputDir,
+ Style: flags.style,
+ ComicStyles: cfg.Styles.Comic,
+ RealisticStyles: cfg.Styles.Realistic,
+ Theme: flags.theme,
+ Language: cfg.Language.Story,
+ Script: cfg.Language.Script,
+ Slug: flags.slug,
+ UltraRealistic: resolveUltraRealistic(flags),
+ RealisticWeight: cfg.Story.RealisticWeight,
+ AspectRatio: cfg.Comic.AspectRatio,
+ PromptMaxChars: cfg.Comic.PromptMaxChars,
+ PageMaxRetries: cfg.Comic.PageMaxRetries,
+ PageRetryBase: time.Duration(cfg.Comic.PageRetryBaseSeconds) * time.Second,
+ })
+
+ return runner.RunPrompt(ctx, flags.prompt)
+}
+
func applyConfigOverrides(cmd *cobra.Command, cfg *config.Config, flags cliFlags) {
if cfg == nil {
return
diff --git a/cmd/comicforge/cli_test.go b/cmd/comicforge/cli_test.go
index fd7904c..048d22e 100644
--- a/cmd/comicforge/cli_test.go
+++ b/cmd/comicforge/cli_test.go
@@ -363,6 +363,140 @@ func TestRootCommandSkipsTTSProviderWhenNarrationDisabled(t *testing.T) {
}
}
+func TestRootCommandPromptModeSkipsVocabFlow(t *testing.T) {
+ tmpDir := t.TempDir()
+
+ var textCalled bool
+ var ttsCalled bool
+ var gotRunnerCfg *comic.RunnerConfig
+ runner := &recordingRunner{}
+
+ cmd := newRootCommandWithDeps(commandDeps{
+ loadConfig: func(string) (*config.Config, error) {
+ return config.DefaultConfig(), nil
+ },
+ newTextProvider: func(*config.Config) (provider.TextProvider, error) {
+ textCalled = true
+ return noopProvider{}, nil
+ },
+ newImageProvider: func(cfg *config.Config) (provider.ImageProvider, error) {
+ return noopProvider{}, nil
+ },
+ newTTSProvider: func(*config.Config, string) (provider.TTSProvider, error) {
+ ttsCalled = true
+ return noopProvider{}, nil
+ },
+ newRunner: func(cfg *comic.RunnerConfig) comic.StoryRunner {
+ gotRunnerCfg = cfg
+ runner.cfg = cfg
+ return runner
+ },
+ })
+ buf := &bytes.Buffer{}
+ cmd.SetOut(buf)
+ cmd.SetErr(buf)
+ cmd.SetArgs([]string{
+ "--prompt", "a robot reading a newspaper",
+ "--output", filepath.Join(tmpDir, "out"),
+ "--slug", "manual-robot",
+ })
+
+ if err := cmd.ExecuteContext(context.Background()); err != nil {
+ t.Fatalf("ExecuteContext() error = %v\noutput:\n%s", err, buf.String())
+ }
+ if textCalled {
+ t.Fatal("newTextProvider was called, want prompt mode to skip story flow")
+ }
+ if ttsCalled {
+ t.Fatal("newTTSProvider was called, want prompt mode to skip narration setup")
+ }
+ if gotRunnerCfg == nil {
+ t.Fatal("runner config was not captured")
+ }
+ if got, want := runner.prompt, "a robot reading a newspaper"; got != want {
+ t.Fatalf("prompt = %q, want %q", got, want)
+ }
+ if got, want := runner.promptRuns, 1; got != want {
+ t.Fatalf("prompt runs = %d, want %d", got, want)
+ }
+ if got, want := gotRunnerCfg.OutputDir, filepath.Join(tmpDir, "out"); got != want {
+ t.Fatalf("output dir = %q, want %q", got, want)
+ }
+ if got, want := gotRunnerCfg.Slug, "manual-robot"; got != want {
+ t.Fatalf("slug = %q, want %q", got, want)
+ }
+}
+
+func TestRootCommandRejectsPromptAndVocabTogether(t *testing.T) {
+ cmd := newRootCommandWithDeps(commandDeps{
+ loadConfig: func(string) (*config.Config, error) {
+ return config.DefaultConfig(), nil
+ },
+ newTextProvider: func(*config.Config) (provider.TextProvider, error) {
+ t.Fatal("newTextProvider should not be called when flags conflict")
+ return noopProvider{}, nil
+ },
+ newImageProvider: func(*config.Config) (provider.ImageProvider, error) {
+ t.Fatal("newImageProvider should not be called when flags conflict")
+ return noopProvider{}, nil
+ },
+ newTTSProvider: func(*config.Config, string) (provider.TTSProvider, error) {
+ t.Fatal("newTTSProvider should not be called when flags conflict")
+ return noopProvider{}, nil
+ },
+ newRunner: func(*comic.RunnerConfig) comic.StoryRunner {
+ t.Fatal("newRunner should not be called when flags conflict")
+ return &recordingRunner{}
+ },
+ })
+ cmd.SetArgs([]string{
+ "--vocab", "words.txt",
+ "--prompt", "draw a robot",
+ })
+
+ err := cmd.ExecuteContext(context.Background())
+ if err == nil {
+ t.Fatal("ExecuteContext() error = nil, want conflict error")
+ }
+ if !strings.Contains(err.Error(), "only one of --prompt and --vocab may be set") {
+ t.Fatalf("ExecuteContext() error = %v, want conflict error", err)
+ }
+}
+
+func TestRootCommandRejectsEmptyPromptFlag(t *testing.T) {
+ cmd := newRootCommandWithDeps(commandDeps{
+ loadConfig: func(string) (*config.Config, error) {
+ t.Fatal("loadConfig should not be called for empty prompt validation")
+ return nil, nil
+ },
+ newTextProvider: func(*config.Config) (provider.TextProvider, error) {
+ t.Fatal("newTextProvider should not be called for empty prompt validation")
+ return noopProvider{}, nil
+ },
+ newImageProvider: func(*config.Config) (provider.ImageProvider, error) {
+ t.Fatal("newImageProvider should not be called for empty prompt validation")
+ return noopProvider{}, nil
+ },
+ newTTSProvider: func(*config.Config, string) (provider.TTSProvider, error) {
+ t.Fatal("newTTSProvider should not be called for empty prompt validation")
+ return noopProvider{}, nil
+ },
+ newRunner: func(*comic.RunnerConfig) comic.StoryRunner {
+ t.Fatal("newRunner should not be called for empty prompt validation")
+ return &recordingRunner{}
+ },
+ })
+ cmd.SetArgs([]string{"--prompt", ""})
+
+ err := cmd.ExecuteContext(context.Background())
+ if err == nil {
+ t.Fatal("ExecuteContext() error = nil, want empty prompt error")
+ }
+ if !strings.Contains(err.Error(), "--prompt is required when set") {
+ t.Fatalf("ExecuteContext() error = %v, want empty prompt error", err)
+ }
+}
+
func mustLookupFlag(t *testing.T, cmd *cobra.Command, name string) *pflag.Flag {
t.Helper()
@@ -374,9 +508,11 @@ func mustLookupFlag(t *testing.T, cmd *cobra.Command, name string) *pflag.Flag {
}
type recordingRunner struct {
- cfg *comic.RunnerConfig
- batchFile string
- runs int
+ cfg *comic.RunnerConfig
+ batchFile string
+ prompt string
+ runs int
+ promptRuns int
}
func (r *recordingRunner) Run(_ context.Context, batchFile string) error {
@@ -385,6 +521,12 @@ func (r *recordingRunner) Run(_ context.Context, batchFile string) error {
return nil
}
+func (r *recordingRunner) RunPrompt(_ context.Context, prompt string) error {
+ r.prompt = prompt
+ r.promptRuns++
+ return nil
+}
+
type noopProvider struct{}
func (noopProvider) Name() string { return "noop" }