diff options
| author | Paul Buetow <paul@buetow.org> | 2026-04-01 15:36:50 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-04-01 15:36:50 +0300 |
| commit | 3967d8dca4ebe87e43d243c5bbffe34ca0dcab51 (patch) | |
| tree | c2c974cb06ad5f65514f492ea8e21f98291a52e2 | |
| parent | ba04917670aac3fb4dff524fcdf614f13bbbfebe (diff) | |
z7: wire Nano Banana CLI config
| -rw-r--r-- | internal/cli/command.go | 40 | ||||
| -rw-r--r-- | internal/cli/command_test.go | 40 | ||||
| -rw-r--r-- | internal/cli/flags.go | 30 | ||||
| -rw-r--r-- | internal/cli/flags_test.go | 5 |
4 files changed, 90 insertions, 25 deletions
diff --git a/internal/cli/command.go b/internal/cli/command.go index a5a8c13..7bf57d6 100644 --- a/internal/cli/command.go +++ b/internal/cli/command.go @@ -20,7 +20,7 @@ func CreateRootCommand(flags *Flags) *cobra.Command { Long: `totalrecall generates Anki flashcard materials from Bulgarian words. It creates audio pronunciation files using OpenAI TTS and downloads -representative images from web search APIs. +representative images using OpenAI or Gemini Nano Banana. Examples: totalrecall # Launch interactive GUI (default) @@ -55,7 +55,7 @@ func setupFlags(cmd *cobra.Command, flags *Flags) { // Local flags cmd.Flags().StringVarP(&flags.OutputDir, "output", "o", defaultOutputDir, "Output directory") cmd.Flags().StringVarP(&flags.AudioFormat, "format", "f", flags.AudioFormat, "Audio format (wav or mp3)") - cmd.Flags().StringVar(&flags.ImageAPI, "image-api", flags.ImageAPI, "Image source (only openai supported)") + cmd.Flags().StringVar(&flags.ImageAPI, "image-api", flags.ImageAPI, "Image source (openai or nanobanana; default: nanobanana)") cmd.Flags().StringVar(&flags.BatchFile, "batch", "", "Process words from file (one per line)") cmd.Flags().BoolVar(&flags.SkipAudio, "skip-audio", false, "Skip audio generation") cmd.Flags().BoolVar(&flags.SkipImages, "skip-images", false, "Skip image download") @@ -79,6 +79,10 @@ func setupFlags(cmd *cobra.Command, flags *Flags) { cmd.Flags().StringVar(&flags.OpenAIImageQuality, "openai-image-quality", flags.OpenAIImageQuality, "Image quality: standard or hd (dall-e-3 only)") cmd.Flags().StringVar(&flags.OpenAIImageStyle, "openai-image-style", flags.OpenAIImageStyle, "Image style: natural or vivid (dall-e-3 only)") + // Nano Banana Image Generation flags + cmd.Flags().StringVar(&flags.NanoBananaModel, "nanobanana-model", flags.NanoBananaModel, "Nano Banana image model (Gemini image preview model)") + cmd.Flags().StringVar(&flags.NanoBananaTextModel, "nanobanana-text-model", flags.NanoBananaTextModel, "Nano Banana text model for prompt generation") + // Bind flags to viper if err := bindFlagsToViper(cmd); err != nil { fmt.Fprintf(os.Stderr, "Warning: failed to bind flags to config: %v\n", err) @@ -87,17 +91,19 @@ func setupFlags(cmd *cobra.Command, flags *Flags) { func bindFlagsToViper(cmd *cobra.Command) error { bindings := map[string]string{ - "audio.format": "format", - "audio.openai_model": "openai-model", - "audio.openai_voice": "openai-voice", - "audio.openai_speed": "openai-speed", - "audio.openai_instruction": "openai-instruction", - "output.directory": "output", - "image.provider": "image-api", - "image.openai_model": "openai-image-model", - "image.openai_size": "openai-image-size", - "image.openai_quality": "openai-image-quality", - "image.openai_style": "openai-image-style", + "audio.format": "format", + "audio.openai_model": "openai-model", + "audio.openai_voice": "openai-voice", + "audio.openai_speed": "openai-speed", + "audio.openai_instruction": "openai-instruction", + "output.directory": "output", + "image.provider": "image-api", + "image.openai_model": "openai-image-model", + "image.openai_size": "openai-image-size", + "image.openai_quality": "openai-image-quality", + "image.openai_style": "openai-image-style", + "image.nanobanana_model": "nanobanana-model", + "image.nanobanana_text_model": "nanobanana-text-model", } var errs []error @@ -156,7 +162,8 @@ func GetOpenAIKey() string { return viper.GetString("audio.openai_key") } -// GetGoogleAPIKey retrieves the Google API key from environment or config. +// GetGoogleAPIKey retrieves the Google API key from GOOGLE_API_KEY or config. +// It prefers image.google_api_key and falls back to google.api_key for older configs. func GetGoogleAPIKey() string { // First check environment variable if key := os.Getenv("GOOGLE_API_KEY"); key != "" { @@ -164,5 +171,10 @@ func GetGoogleAPIKey() string { } // Then check config file + if key := viper.GetString("image.google_api_key"); key != "" { + return key + } + + // Fall back to the legacy key for compatibility with older configs. return viper.GetString("google.api_key") } diff --git a/internal/cli/command_test.go b/internal/cli/command_test.go index 2361cd2..ab21c49 100644 --- a/internal/cli/command_test.go +++ b/internal/cli/command_test.go @@ -50,6 +50,8 @@ func TestCreateRootCommand(t *testing.T) { {"openai-image-size", true}, {"openai-image-quality", true}, {"openai-image-style", true}, + {"nanobanana-model", true}, + {"nanobanana-text-model", true}, } for _, tt := range flagTests { @@ -93,6 +95,30 @@ func TestSetupFlags(t *testing.T) { if formatFlag.DefValue != "mp3" { t.Errorf("Expected default format to be mp3, got %s", formatFlag.DefValue) } + + imageAPIFlag := cmd.Flags().Lookup("image-api") + if imageAPIFlag == nil { + t.Fatal("image-api flag not found") + } + if imageAPIFlag.DefValue != "nanobanana" { + t.Errorf("Expected default image-api to be nanobanana, got %s", imageAPIFlag.DefValue) + } + + nanoBananaModelFlag := cmd.Flags().Lookup("nanobanana-model") + if nanoBananaModelFlag == nil { + t.Fatal("nanobanana-model flag not found") + } + if nanoBananaModelFlag.DefValue != "gemini-3.1-flash-image-preview" { + t.Errorf("Expected default nanobanana-model to be gemini-3.1-flash-image-preview, got %s", nanoBananaModelFlag.DefValue) + } + + nanoBananaTextModelFlag := cmd.Flags().Lookup("nanobanana-text-model") + if nanoBananaTextModelFlag == nil { + t.Fatal("nanobanana-text-model flag not found") + } + if nanoBananaTextModelFlag.DefValue != "gemini-2.5-flash" { + t.Errorf("Expected default nanobanana-text-model to be gemini-2.5-flash, got %s", nanoBananaTextModelFlag.DefValue) + } } func TestInitConfig(t *testing.T) { @@ -291,7 +317,7 @@ func TestGetGoogleAPIKey(t *testing.T) { } if tt.configKey != "" { - viper.Set("google.api_key", tt.configKey) + viper.Set("image.google_api_key", tt.configKey) } got := GetGoogleAPIKey() @@ -327,6 +353,12 @@ func TestBindFlagsToViper(t *testing.T) { if err := cmd.Flags().Set("openai-model", "tts-1-hd"); err != nil { t.Fatalf("Failed to set openai-model flag: %v", err) } + if err := cmd.Flags().Set("nanobanana-model", "gemini-3.1-flash-image-preview"); err != nil { + t.Fatalf("Failed to set nanobanana-model flag: %v", err) + } + if err := cmd.Flags().Set("nanobanana-text-model", "gemini-2.5-flash"); err != nil { + t.Fatalf("Failed to set nanobanana-text-model flag: %v", err) + } if err := bindFlagsToViper(cmd); err != nil { t.Fatalf("bindFlagsToViper() failed: %v", err) @@ -344,4 +376,10 @@ func TestBindFlagsToViper(t *testing.T) { if viper.GetString("audio.openai_model") != "tts-1-hd" { t.Errorf("Expected audio.openai_model to be tts-1-hd, got %s", viper.GetString("audio.openai_model")) } + if viper.GetString("image.nanobanana_model") != "gemini-3.1-flash-image-preview" { + t.Errorf("Expected image.nanobanana_model to be gemini-3.1-flash-image-preview, got %s", viper.GetString("image.nanobanana_model")) + } + if viper.GetString("image.nanobanana_text_model") != "gemini-2.5-flash" { + t.Errorf("Expected image.nanobanana_text_model to be gemini-2.5-flash, got %s", viper.GetString("image.nanobanana_text_model")) + } } diff --git a/internal/cli/flags.go b/internal/cli/flags.go index 904059f..172b854 100644 --- a/internal/cli/flags.go +++ b/internal/cli/flags.go @@ -1,5 +1,10 @@ package cli +const ( + defaultNanoBananaModel = "gemini-3.1-flash-image-preview" + defaultNanoBananaTextModel = "gemini-2.5-flash" +) + // Flags holds all command-line flag values type Flags struct { // General flags @@ -29,19 +34,26 @@ type Flags struct { OpenAIImageSize string OpenAIImageQuality string OpenAIImageStyle string + + // NanoBananaModel is the Gemini image model used for Nano Banana generation. + NanoBananaModel string + // NanoBananaTextModel is the Gemini text model used for Nano Banana prompt generation. + NanoBananaTextModel string } // NewFlags creates a new Flags instance with default values func NewFlags() *Flags { return &Flags{ - AudioFormat: "mp3", - ImageAPI: "openai", - DeckName: "Bulgarian Vocabulary", - OpenAIModel: "gpt-4o-mini-tts", - OpenAISpeed: 0.9, - OpenAIImageModel: "dall-e-2", - OpenAIImageSize: "512x512", - OpenAIImageQuality: "standard", - OpenAIImageStyle: "natural", + AudioFormat: "mp3", + ImageAPI: "nanobanana", + DeckName: "Bulgarian Vocabulary", + OpenAIModel: "gpt-4o-mini-tts", + OpenAISpeed: 0.9, + OpenAIImageModel: "dall-e-2", + OpenAIImageSize: "512x512", + OpenAIImageQuality: "standard", + OpenAIImageStyle: "natural", + NanoBananaModel: defaultNanoBananaModel, + NanoBananaTextModel: defaultNanoBananaTextModel, } } diff --git a/internal/cli/flags_test.go b/internal/cli/flags_test.go index ccc50a0..d5c2ccb 100644 --- a/internal/cli/flags_test.go +++ b/internal/cli/flags_test.go @@ -15,7 +15,7 @@ func TestNewFlags(t *testing.T) { expected interface{} }{ {"AudioFormat", flags.AudioFormat, "mp3"}, - {"ImageAPI", flags.ImageAPI, "openai"}, + {"ImageAPI", flags.ImageAPI, "nanobanana"}, {"DeckName", flags.DeckName, "Bulgarian Vocabulary"}, {"OpenAIModel", flags.OpenAIModel, "gpt-4o-mini-tts"}, {"OpenAISpeed", flags.OpenAISpeed, 0.9}, @@ -23,6 +23,8 @@ func TestNewFlags(t *testing.T) { {"OpenAIImageSize", flags.OpenAIImageSize, "512x512"}, {"OpenAIImageQuality", flags.OpenAIImageQuality, "standard"}, {"OpenAIImageStyle", flags.OpenAIImageStyle, "natural"}, + {"NanoBananaModel", flags.NanoBananaModel, "gemini-3.1-flash-image-preview"}, + {"NanoBananaTextModel", flags.NanoBananaTextModel, "gemini-2.5-flash"}, } for _, tt := range tests { @@ -87,6 +89,7 @@ func TestFlagsStructure(t *testing.T) { "ListModels", "AllVoices", "NoAutoPlay", "OpenAIModel", "OpenAIVoice", "OpenAISpeed", "OpenAIInstruction", "OpenAIImageModel", "OpenAIImageSize", "OpenAIImageQuality", "OpenAIImageStyle", + "NanoBananaModel", "NanoBananaTextModel", } for _, fieldName := range expectedFields { |
