From d97741966a6c97b5db7df62873bd2ee9d99f0565 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Sun, 19 Apr 2026 23:13:06 +0300 Subject: y4: wire ComicForge CLI providers and config --- .gitignore | 3 +- cmd/comicforge/cli.go | 278 ++++++++++++++++++++++++++++++++++++++++++++ cmd/comicforge/cli_test.go | 279 +++++++++++++++++++++++++++++++++++++++++++++ cmd/comicforge/main.go | 13 +++ config.yaml.example | 61 ++++++++++ internal/comic/runner.go | 2 + internal/image/gemini.go | 42 +++++++ internal/text/gemini.go | 110 ++++++++++++++++++ 8 files changed, 786 insertions(+), 2 deletions(-) create mode 100644 cmd/comicforge/cli.go create mode 100644 cmd/comicforge/cli_test.go create mode 100644 cmd/comicforge/main.go create mode 100644 config.yaml.example create mode 100644 internal/text/gemini.go diff --git a/.gitignore b/.gitignore index 4e40665..ac0606d 100644 --- a/.gitignore +++ b/.gitignore @@ -1,7 +1,6 @@ -comicforge +/comicforge comics/ *.pdf *.mp3 .idea/ .vscode/ - diff --git a/cmd/comicforge/cli.go b/cmd/comicforge/cli.go new file mode 100644 index 0000000..b0c539b --- /dev/null +++ b/cmd/comicforge/cli.go @@ -0,0 +1,278 @@ +package main + +import ( + "context" + "fmt" + "strings" + + "github.com/spf13/cobra" + + version "codeberg.org/snonux/comicforge/internal" + "codeberg.org/snonux/comicforge/internal/comic" + "codeberg.org/snonux/comicforge/internal/config" + "codeberg.org/snonux/comicforge/internal/image" + "codeberg.org/snonux/comicforge/internal/provider" + textprovider "codeberg.org/snonux/comicforge/internal/text" + "codeberg.org/snonux/comicforge/internal/tts" +) + +type commandDeps struct { + loadConfig func(string) (*config.Config, error) + newTextProvider func(*config.Config) (provider.TextProvider, error) + newImageProvider func(*config.Config) (provider.ImageProvider, error) + newTTSProvider func(*config.Config, string) (provider.TTSProvider, error) + newRunner func(*comic.RunnerConfig) comic.StoryRunner +} + +type cliFlags struct { + vocab string + configPath string + promptsDir string + outputDir string + style string + theme string + slug string + narratorVoice string + textModel string + imageModel string + imageTextModel string + ttsModel string + textProvider string + imageProvider string + ttsProvider string + narrateEnabled bool + ultraRealistic bool + noUltraRealistic bool + version bool +} + +func defaultCommandDeps() commandDeps { + return commandDeps{ + loadConfig: config.Load, + newTextProvider: func(cfg *config.Config) (provider.TextProvider, error) { + return buildTextProvider(cfg) + }, + newImageProvider: func(cfg *config.Config) (provider.ImageProvider, error) { + return buildImageProvider(cfg) + }, + newTTSProvider: func(cfg *config.Config, voice string) (provider.TTSProvider, error) { + return buildTTSProvider(cfg, voice) + }, + newRunner: func(cfg *comic.RunnerConfig) comic.StoryRunner { + return comic.NewRunner(cfg) + }, + } +} + +func newRootCommand() *cobra.Command { + return newRootCommandWithDeps(defaultCommandDeps()) +} + +func newRootCommandWithDeps(deps commandDeps) *cobra.Command { + flags := cliFlags{} + cmd := &cobra.Command{ + Use: "comicforge", + Short: "ComicForge generates and manages comic output", + SilenceUsage: true, + SilenceErrors: true, + RunE: func(cmd *cobra.Command, args []string) error { + if flags.version { + fmt.Fprintln(cmd.OutOrStdout(), version.Version) + return nil + } + if strings.TrimSpace(flags.vocab) == "" { + return fmt.Errorf("--vocab is required") + } + return runCommand(cmd.Context(), cmd, deps, flags) + }, + } + + 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.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", ".", "output directory for generated comic assets") + cmd.Flags().StringVar(&flags.style, "style", "", "comic art style override") + cmd.Flags().StringVar(&flags.theme, "theme", "", "story theme override") + cmd.Flags().BoolVar(&flags.ultraRealistic, "ultra-realistic", false, "force photorealistic rendering") + cmd.Flags().BoolVar(&flags.noUltraRealistic, "no-ultra-realistic", false, "disable photorealistic rendering") + cmd.Flags().BoolVar(&flags.narrateEnabled, "narrate", false, "generate narration after the comic") + cmd.Flags().StringVar(&flags.narratorVoice, "narrator-voice", "", "Gemini voice for narration") + cmd.Flags().StringVar(&flags.slug, "slug", "", "force the output slug for generated assets") + cmd.Flags().StringVar(&flags.textProvider, "text-provider", provider.Gemini, "text provider to use") + cmd.Flags().StringVar(&flags.imageProvider, "image-provider", provider.Gemini, "image provider to use") + cmd.Flags().StringVar(&flags.ttsProvider, "tts-provider", provider.Gemini, "text-to-speech provider to use") + cmd.Flags().StringVar(&flags.textModel, "text-model", "", "text model override") + cmd.Flags().StringVar(&flags.imageModel, "image-model", "", "image model override") + cmd.Flags().StringVar(&flags.imageTextModel, "image-text-model", "", "image text model override") + cmd.Flags().StringVar(&flags.ttsModel, "tts-model", "", "text-to-speech model override") + + return cmd +} + +func runCommand(ctx context.Context, cmd *cobra.Command, deps commandDeps, flags cliFlags) error { + if ctx == nil { + ctx = context.Background() + } + if flags.noUltraRealistic && flags.ultraRealistic { + return fmt.Errorf("only one of --ultra-realistic and --no-ultra-realistic may be set") + } + + cfg, err := deps.loadConfig(flags.configPath) + if err != nil { + return err + } + + applyConfigOverrides(cmd, cfg, flags) + voice := resolveNarratorVoice(flags, cfg) + + textProvider, err := deps.newTextProvider(cfg) + if err != nil { + return fmt.Errorf("build text provider: %w", err) + } + imageProvider, err := deps.newImageProvider(cfg) + if err != nil { + return fmt.Errorf("build image provider: %w", err) + } + mainTTSProvider, err := deps.newTTSProvider(cfg, voice) + if err != nil { + return fmt.Errorf("build TTS provider: %w", err) + } + + runner := deps.newRunner(&comic.RunnerConfig{ + TextProvider: textProvider, + ImageProvider: imageProvider, + MainTTSProvider: mainTTSProvider, + ConclusionTTSProvider: mainTTSProvider, + Prompts: cfg, + OutputDir: flags.outputDir, + Style: flags.style, + Theme: flags.theme, + Language: cfg.Language.Story, + Script: cfg.Language.Script, + NarratorVoice: voice, + Slug: flags.slug, + NarrateEnabled: flags.narrateEnabled, + UltraRealistic: resolveUltraRealistic(flags), + StoryPages: cfg.Comic.StoryPages, + GalleryPages: cfg.Comic.GalleryPages, + PanelsPerPage: cfg.Comic.PanelsPerPage, + }) + + return runner.Run(ctx, flags.vocab) +} + +func applyConfigOverrides(cmd *cobra.Command, cfg *config.Config, flags cliFlags) { + if cfg == nil { + return + } + if cmd.Flags().Changed("text-provider") { + cfg.Provider.Text = flags.textProvider + } + if cmd.Flags().Changed("image-provider") { + cfg.Provider.Image = flags.imageProvider + } + if cmd.Flags().Changed("tts-provider") { + cfg.Provider.TTS = flags.ttsProvider + } + if cmd.Flags().Changed("text-model") { + cfg.Models.Text = flags.textModel + } + if cmd.Flags().Changed("image-model") { + cfg.Models.Image = flags.imageModel + } + if cmd.Flags().Changed("image-text-model") { + cfg.Models.ImageText = flags.imageTextModel + } + if cmd.Flags().Changed("tts-model") { + cfg.Models.TTS = flags.ttsModel + } + if cmd.Flags().Changed("prompts-dir") { + cfg.PromptsDir = flags.promptsDir + } +} + +func buildTextProvider(cfg *config.Config) (provider.TextProvider, error) { + if cfg == nil { + return nil, fmt.Errorf("config is required") + } + switch provider.NormalizeName(cfg.Provider.Text) { + case provider.Gemini: + p := textprovider.NewGeminiProvider(&textprovider.GeminiConfig{ + APIKey: cfg.API.GoogleAPIKey, + Model: cfg.Models.Text, + }) + if err := p.IsAvailable(); err != nil { + return nil, err + } + return p, nil + default: + return nil, fmt.Errorf("text provider %q is not implemented", cfg.Provider.Text) + } +} + +func buildImageProvider(cfg *config.Config) (provider.ImageProvider, error) { + if cfg == nil { + return nil, fmt.Errorf("config is required") + } + switch provider.NormalizeName(cfg.Provider.Image) { + case provider.Gemini: + p := image.NewGeminiProvider(&image.GeminiConfig{ + APIKey: cfg.API.GoogleAPIKey, + Model: cfg.Models.Image, + TextModel: cfg.Models.ImageText, + }) + if err := p.IsAvailable(); err != nil { + return nil, err + } + return p, nil + default: + return nil, fmt.Errorf("image provider %q is not implemented", cfg.Provider.Image) + } +} + +func buildTTSProvider(cfg *config.Config, voice string) (provider.TTSProvider, error) { + if cfg == nil { + return nil, fmt.Errorf("config is required") + } + switch provider.NormalizeName(cfg.Provider.TTS) { + case provider.Gemini: + p := tts.NewGeminiProvider(&tts.GeminiConfig{ + APIKey: cfg.API.GoogleAPIKey, + Model: cfg.Models.TTS, + Voice: voice, + }) + if err := p.IsAvailable(); err != nil { + return nil, err + } + return p, nil + default: + return nil, fmt.Errorf("tts provider %q is not implemented", cfg.Provider.TTS) + } +} + +func resolveUltraRealistic(flags cliFlags) *bool { + switch { + case flags.ultraRealistic: + v := true + return &v + case flags.noUltraRealistic: + v := false + return &v + default: + return nil + } +} + +func resolveNarratorVoice(flags cliFlags, cfg *config.Config) string { + if strings.TrimSpace(flags.narratorVoice) != "" { + return strings.TrimSpace(flags.narratorVoice) + } + if cfg == nil { + return "" + } + if len(cfg.Narration.Voices) == 0 { + return "" + } + return strings.TrimSpace(cfg.Narration.Voices[0]) +} diff --git a/cmd/comicforge/cli_test.go b/cmd/comicforge/cli_test.go new file mode 100644 index 0000000..19a2f19 --- /dev/null +++ b/cmd/comicforge/cli_test.go @@ -0,0 +1,279 @@ +package main + +import ( + "bytes" + "context" + "os" + "path/filepath" + "strings" + "testing" + + version "codeberg.org/snonux/comicforge/internal" + "codeberg.org/snonux/comicforge/internal/comic" + "codeberg.org/snonux/comicforge/internal/config" + "codeberg.org/snonux/comicforge/internal/provider" +) + +func TestRootCommandAppliesConfigAndFlagOverrides(t *testing.T) { + tmpDir := t.TempDir() + configPath := filepath.Join(tmpDir, "config.yaml") + if err := os.WriteFile(configPath, []byte(strings.TrimSpace(` +provider: + text: openai + image: openai + tts: openai +models: + text: config-text + image: config-image + image_text: config-image-text + tts: config-tts +comic: + story_pages: 7 + gallery_pages: 3 + panels_per_page: 2 +language: + story_language: Bulgarian + script: Cyrillic +narration: + voices: + - Fenrir + - Charon +prompts_dir: ./config-prompts +`)), 0o644); err != nil { + t.Fatalf("write config: %v", err) + } + + vocabPath := filepath.Join(tmpDir, "vocab.txt") + if err := os.WriteFile(vocabPath, []byte("ябълка = apple\nкнига = book\n"), 0o644); err != nil { + t.Fatalf("write vocab: %v", err) + } + + var gotTextCfg *config.Config + var gotImageCfg *config.Config + var gotTTSCfg *config.Config + var gotTTSVoice string + var gotRunnerCfg *comic.RunnerConfig + runner := &recordingRunner{} + + cmd := newRootCommandWithDeps(commandDeps{ + loadConfig: config.Load, + newTextProvider: func(cfg *config.Config) (provider.TextProvider, error) { + gotTextCfg = cfg + return noopProvider{}, nil + }, + newImageProvider: func(cfg *config.Config) (provider.ImageProvider, error) { + gotImageCfg = cfg + return noopProvider{}, nil + }, + newTTSProvider: func(cfg *config.Config, voice string) (provider.TTSProvider, error) { + gotTTSCfg = cfg + gotTTSVoice = voice + 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{ + "--config", configPath, + "--vocab", vocabPath, + "--text-provider", provider.Gemini, + "--image-provider", provider.Gemini, + "--tts-provider", provider.Gemini, + "--text-model", "flag-text", + "--image-model", "flag-image", + "--image-text-model", "flag-image-text", + "--tts-model", "flag-tts", + "--prompts-dir", filepath.Join(tmpDir, "prompts"), + "--output", filepath.Join(tmpDir, "out"), + "--style", "noir", + "--theme", "mystery", + "--slug", "forced-slug", + "--narrate", + "--narrator-voice", "Aoede", + "--ultra-realistic", + }) + + if err := cmd.ExecuteContext(context.Background()); err != nil { + t.Fatalf("ExecuteContext() error = %v\noutput:\n%s", err, buf.String()) + } + + if runner.runs != 1 { + t.Fatalf("runner.runs = %d, want 1", runner.runs) + } + if got, want := runner.batchFile, vocabPath; got != want { + t.Fatalf("runner.batchFile = %q, want %q", got, want) + } + if gotRunnerCfg == nil { + t.Fatal("runner config was not captured") + } + if gotTextCfg == nil || gotImageCfg == nil || gotTTSCfg == nil { + t.Fatal("provider configs were not captured") + } + if got, want := gotTextCfg.Provider.Text, provider.Gemini; got != want { + t.Fatalf("text provider = %q, want %q", got, want) + } + if got, want := gotImageCfg.Provider.Image, provider.Gemini; got != want { + t.Fatalf("image provider = %q, want %q", got, want) + } + if got, want := gotTTSCfg.Provider.TTS, provider.Gemini; got != want { + t.Fatalf("tts provider = %q, want %q", got, want) + } + if got, want := gotTextCfg.Models.Text, "flag-text"; got != want { + t.Fatalf("text model = %q, want %q", got, want) + } + if got, want := gotImageCfg.Models.Image, "flag-image"; got != want { + t.Fatalf("image model = %q, want %q", got, want) + } + if got, want := gotImageCfg.Models.ImageText, "flag-image-text"; got != want { + t.Fatalf("image text model = %q, want %q", got, want) + } + if got, want := gotTTSCfg.Models.TTS, "flag-tts"; got != want { + t.Fatalf("tts model = %q, want %q", got, want) + } + if got, want := gotTextCfg.PromptsDir, filepath.Join(tmpDir, "prompts"); got != want { + t.Fatalf("prompts dir = %q, want %q", got, want) + } + if got, want := gotTTSVoice, "Aoede"; got != want { + t.Fatalf("tts voice = %q, want %q", 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.Style, "noir"; got != want { + t.Fatalf("style = %q, want %q", got, want) + } + if got, want := gotRunnerCfg.Theme, "mystery"; got != want { + t.Fatalf("theme = %q, want %q", got, want) + } + if got, want := gotRunnerCfg.Slug, "forced-slug"; got != want { + t.Fatalf("slug = %q, want %q", got, want) + } + if got, want := gotRunnerCfg.NarrateEnabled, true; got != want { + t.Fatalf("narrate = %t, want %t", got, want) + } + if got, want := gotRunnerCfg.NarratorVoice, "Aoede"; got != want { + t.Fatalf("narrator voice = %q, want %q", got, want) + } + if got, want := gotRunnerCfg.Language, "Bulgarian"; got != want { + t.Fatalf("language = %q, want %q", got, want) + } + if got, want := gotRunnerCfg.Script, "Cyrillic"; got != want { + t.Fatalf("script = %q, want %q", got, want) + } + if got, want := gotRunnerCfg.StoryPages, 7; got != want { + t.Fatalf("story pages = %d, want %d", got, want) + } + if got, want := gotRunnerCfg.GalleryPages, 3; got != want { + t.Fatalf("gallery pages = %d, want %d", got, want) + } + if got, want := gotRunnerCfg.PanelsPerPage, 2; got != want { + t.Fatalf("panels per page = %d, want %d", got, want) + } + if gotRunnerCfg.UltraRealistic == nil || !*gotRunnerCfg.UltraRealistic { + t.Fatalf("ultra realistic = %#v, want true", gotRunnerCfg.UltraRealistic) + } +} + +func TestRootCommandVersionSkipsRequiredFlags(t *testing.T) { + cmd := newRootCommandWithDeps(commandDeps{ + loadConfig: func(string) (*config.Config, error) { + t.Fatal("loadConfig should not be called for --version") + return nil, nil + }, + newTextProvider: func(*config.Config) (provider.TextProvider, error) { + t.Fatal("newTextProvider should not be called for --version") + return noopProvider{}, nil + }, + newImageProvider: func(*config.Config) (provider.ImageProvider, error) { + t.Fatal("newImageProvider should not be called for --version") + return noopProvider{}, nil + }, + newTTSProvider: func(*config.Config, string) (provider.TTSProvider, error) { + t.Fatal("newTTSProvider should not be called for --version") + return noopProvider{}, nil + }, + newRunner: func(*comic.RunnerConfig) comic.StoryRunner { + t.Fatal("newRunner should not be called for --version") + return &recordingRunner{} + }, + }) + buf := &bytes.Buffer{} + cmd.SetOut(buf) + cmd.SetErr(buf) + cmd.SetArgs([]string{"--version"}) + + if err := cmd.ExecuteContext(context.Background()); err != nil { + t.Fatalf("ExecuteContext() error = %v", err) + } + if got, want := strings.TrimSpace(buf.String()), version.Version; got != want { + t.Fatalf("version output = %q, want %q", got, want) + } +} + +func TestRootCommandRejectsConflictingUltraFlags(t *testing.T) { + cmd := newRootCommandWithDeps(commandDeps{ + loadConfig: func(string) (*config.Config, error) { + return config.DefaultConfig(), nil + }, + newTextProvider: func(*config.Config) (provider.TextProvider, error) { + return noopProvider{}, nil + }, + newImageProvider: func(*config.Config) (provider.ImageProvider, error) { + return noopProvider{}, nil + }, + newTTSProvider: func(*config.Config, string) (provider.TTSProvider, error) { + return noopProvider{}, nil + }, + newRunner: func(*comic.RunnerConfig) comic.StoryRunner { + t.Fatal("newRunner should not be called when ultra-realistic flags conflict") + return &recordingRunner{} + }, + }) + cmd.SetArgs([]string{ + "--vocab", "words.txt", + "--ultra-realistic", + "--no-ultra-realistic", + }) + + err := cmd.ExecuteContext(context.Background()) + if err == nil { + t.Fatal("ExecuteContext() error = nil, want conflict error") + } + if !strings.Contains(err.Error(), "only one of --ultra-realistic and --no-ultra-realistic may be set") { + t.Fatalf("ExecuteContext() error = %v, want conflict error", err) + } +} + +type recordingRunner struct { + cfg *comic.RunnerConfig + batchFile string + runs int +} + +func (r *recordingRunner) Run(_ context.Context, batchFile string) error { + r.batchFile = batchFile + r.runs++ + return nil +} + +type noopProvider struct{} + +func (noopProvider) Name() string { return "noop" } +func (noopProvider) IsAvailable() error { + return nil +} +func (noopProvider) GenerateText(context.Context, string) (string, error) { + return "noop", nil +} +func (noopProvider) GenerateImage(context.Context, string, string) error { + return nil +} +func (noopProvider) GenerateAudio(context.Context, string, string) error { + return nil +} diff --git a/cmd/comicforge/main.go b/cmd/comicforge/main.go new file mode 100644 index 0000000..f85be6c --- /dev/null +++ b/cmd/comicforge/main.go @@ -0,0 +1,13 @@ +package main + +import ( + "fmt" + "os" +) + +func main() { + if err := newRootCommand().Execute(); err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } +} diff --git a/config.yaml.example b/config.yaml.example new file mode 100644 index 0000000..95d12dc --- /dev/null +++ b/config.yaml.example @@ -0,0 +1,61 @@ +# ComicForge configuration example. +# +# Copy this file to config.yaml and adjust the values you want to override. +# Command-line flags always take precedence over config values. + +provider: + # Supported today: gemini. + # Future providers can be added here when ComicForge grows new backends. + text: gemini + image: gemini + tts: gemini + +api: + # Required for Gemini-backed providers. + google_api_key: "" + +models: + text: gemini-2.5-flash + image: gemini-3.1-flash-image-preview + image_text: gemini-2.5-flash + tts: gemini-2.5-flash-preview-tts + +comic: + story_pages: 5 + gallery_pages: 5 + panels_per_page: 4 + aspect_ratio: "16:9" + prompt_max_chars: 900 + page_max_retries: 5 + page_retry_base_seconds: 15 + +language: + input: Vocabulary + output: Story + # Language used for story and narration prompts. + story_language: Bulgarian + script: Latin + +story: + genres: + - a warm slice-of-life story + - a heartfelt family drama + - an exciting science-fiction adventure + realistic_weight: 0.4 + +styles: + comic: + - classic comic book with bold ink outlines + - graphic novel with dramatic shadows + realistic: + - ultra-realistic DSLR photography, cinematic 35mm lens + - cinematic realism with natural light + +narration: + voices: + - Charon + - Fenrir + chunk_words: 100 + +# Directory that ComicForge uses for prompt templates. +prompts_dir: ./prompts diff --git a/internal/comic/runner.go b/internal/comic/runner.go index 6ade637..7e4e9bb 100644 --- a/internal/comic/runner.go +++ b/internal/comic/runner.go @@ -32,6 +32,7 @@ type RunnerConfig struct { Theme string Language string Script string + NarratorVoice string Slug string NarrateEnabled bool UltraRealistic *bool @@ -88,6 +89,7 @@ func NewRunner(cfg *RunnerConfig) *Runner { MainProvider: cfg.MainTTSProvider, ConclusionProvider: cfg.ConclusionTTSProvider, Prompts: cfg.Prompts, + VoiceName: cfg.NarratorVoice, Language: cfg.Language, Script: cfg.Script, }) diff --git a/internal/image/gemini.go b/internal/image/gemini.go index 57fa61a..85da3a8 100644 --- a/internal/image/gemini.go +++ b/internal/image/gemini.go @@ -12,6 +12,7 @@ import ( "image/png" "io" "net/http" + "os" "strings" "time" @@ -165,6 +166,47 @@ func (c *GeminiProvider) Search(ctx context.Context, opts *SearchOptions) ([]Sea return []SearchResult{result}, nil } +// IsAvailable reports whether the provider was initialized successfully. +func (c *GeminiProvider) IsAvailable() error { + return c.ensureReady() +} + +// GenerateImage renders the first generated image to outputFile. +func (c *GeminiProvider) GenerateImage(ctx context.Context, prompt, outputFile string) error { + if c == nil { + return fmt.Errorf("image provider is nil") + } + if ctx == nil { + ctx = context.Background() + } + if strings.TrimSpace(prompt) == "" { + return fmt.Errorf("prompt is required") + } + if strings.TrimSpace(outputFile) == "" { + return fmt.Errorf("output file is required") + } + results, err := c.Search(ctx, &SearchOptions{Query: prompt}) + if err != nil { + return err + } + if len(results) == 0 { + return fmt.Errorf("no image results returned") + } + rc, err := c.Download(ctx, results[0].URL) + if err != nil { + return err + } + defer rc.Close() + data, err := io.ReadAll(rc) + if err != nil { + return fmt.Errorf("read image data: %w", err) + } + if err := os.WriteFile(outputFile, data, 0o644); err != nil { + return fmt.Errorf("write image: %w", err) + } + return nil +} + // Download returns the image bytes for a data URI or a remote URL. func (c *GeminiProvider) Download(ctx context.Context, url string) (io.ReadCloser, error) { if strings.HasPrefix(url, geminiDataPrefix) { diff --git a/internal/text/gemini.go b/internal/text/gemini.go new file mode 100644 index 0000000..94a8c16 --- /dev/null +++ b/internal/text/gemini.go @@ -0,0 +1,110 @@ +package text + +import ( + "context" + "fmt" + "strings" + + "google.golang.org/genai" + + "codeberg.org/snonux/comicforge/internal/httpctx" + "codeberg.org/snonux/comicforge/internal/provider" +) + +const ( + // DefaultModel is the Gemini text model used for story generation. + DefaultModel = "gemini-2.5-flash" +) + +// GeminiConfig holds the settings needed to build a Gemini-backed text provider. +type GeminiConfig struct { + APIKey string + Model string +} + +// GeminiProvider implements TextProvider for Google Gemini text generation. +type GeminiProvider struct { + client *genai.Client + model string + err error +} + +var _ provider.TextProvider = (*GeminiProvider)(nil) + +var newGeminiClient = httpctx.NewGenAIClient +var geminiGenerateText = func(ctx context.Context, p *GeminiProvider, prompt string) (string, error) { + return p.generateText(ctx, prompt) +} + +// NewGeminiProvider creates a Gemini text provider. +func NewGeminiProvider(cfg *GeminiConfig) *GeminiProvider { + g := &GeminiProvider{model: DefaultModel} + if cfg == nil { + g.err = fmt.Errorf("text config is required") + return g + } + g.model = defaultOr(cfg.Model, DefaultModel) + if strings.TrimSpace(cfg.APIKey) == "" { + g.err = fmt.Errorf("Google API key is required for text generation") + return g + } + client, err := newGeminiClient(context.Background(), &genai.ClientConfig{ + APIKey: cfg.APIKey, + Backend: genai.BackendGeminiAPI, + }) + if err != nil { + g.err = fmt.Errorf("create Gemini client: %w", err) + return g + } + g.client = client + return g +} + +// Name returns the provider name. +func (g *GeminiProvider) Name() string { return provider.Gemini } + +// IsAvailable reports whether the provider was initialized successfully. +func (g *GeminiProvider) IsAvailable() error { + if g == nil { + return fmt.Errorf("text provider is nil") + } + return g.err +} + +// GenerateText generates a text response for the provided prompt. +func (g *GeminiProvider) GenerateText(ctx context.Context, prompt string) (string, error) { + if g == nil { + return "", fmt.Errorf("text provider is nil") + } + if ctx == nil { + ctx = context.Background() + } + if g.err != nil { + return "", g.err + } + if strings.TrimSpace(prompt) == "" { + return "", fmt.Errorf("prompt is required") + } + return geminiGenerateText(ctx, g, prompt) +} + +func (g *GeminiProvider) generateText(ctx context.Context, prompt string) (string, error) { + resp, err := g.client.Models.GenerateContent(ctx, g.model, []*genai.Content{ + genai.NewContentFromText(prompt, genai.RoleUser), + }, nil) + if err != nil { + return "", fmt.Errorf("generate text: %w", err) + } + text := strings.TrimSpace(resp.Text()) + if text == "" { + return "", fmt.Errorf("no text content returned") + } + return text, nil +} + +func defaultOr(value, fallback string) string { + if strings.TrimSpace(value) != "" { + return value + } + return fallback +} -- cgit v1.2.3