summaryrefslogtreecommitdiff
path: root/internal/cli
diff options
context:
space:
mode:
Diffstat (limited to 'internal/cli')
-rw-r--r--internal/cli/command.go40
-rw-r--r--internal/cli/command_test.go40
-rw-r--r--internal/cli/flags.go30
-rw-r--r--internal/cli/flags_test.go5
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 {