diff options
| -rw-r--r-- | internal/audio/provider.go | 41 | ||||
| -rw-r--r-- | internal/gui/app.go | 7 | ||||
| -rw-r--r-- | internal/gui/app_test.go | 4 | ||||
| -rw-r--r-- | internal/gui/card_service.go | 3 | ||||
| -rw-r--r-- | internal/gui/generator_test.go | 6 | ||||
| -rw-r--r-- | internal/gui/orchestrator.go | 83 | ||||
| -rw-r--r-- | internal/image/search.go | 8 | ||||
| -rw-r--r-- | internal/processor/image_downloader.go | 20 | ||||
| -rw-r--r-- | internal/registry/registry.go | 40 | ||||
| -rw-r--r-- | internal/registry/registry_test.go | 29 |
10 files changed, 177 insertions, 64 deletions
diff --git a/internal/audio/provider.go b/internal/audio/provider.go index 296521d..0a9a4b8 100644 --- a/internal/audio/provider.go +++ b/internal/audio/provider.go @@ -4,6 +4,8 @@ import ( "context" "fmt" "strings" + + "codeberg.org/snonux/totalrecall/internal/registry" ) // Provider defines the interface for text-to-speech providers. @@ -137,6 +139,29 @@ func DefaultProviderConfig() *Config { // (processor, gui). The production default is audio.NewProvider itself. type ProviderFactory func(*Config) (Provider, error) +// defaultAudioProviders maps provider name to constructor. New providers are +// registered here so NewProvider does not grow a new switch branch each time. +var defaultAudioProviders = func() *registry.Registry[string, func(*Config) (Provider, error)] { + r := registry.New[string, func(*Config) (Provider, error)]() + r.Register("openai", newOpenAIProviderFromConfig) + r.Register("gemini", newGeminiProviderFromConfig) + return r +}() + +func newOpenAIProviderFromConfig(config *Config) (Provider, error) { + if config.OpenAIKey == "" { + return nil, fmt.Errorf("OpenAI API key is required") + } + return NewOpenAIProvider(openAIAudioConfigFrom(config), config.OutputFormat) +} + +func newGeminiProviderFromConfig(config *Config) (Provider, error) { + if config.GoogleAPIKey == "" { + return nil, fmt.Errorf("google API key is required") + } + return NewGeminiProvider(geminiAudioConfigFrom(config), config.OutputFormat) +} + // NewProvider creates the appropriate audio provider based on configuration. // It extracts provider-specific sub-configs so each implementation only // receives the fields it needs (ISP). @@ -145,18 +170,10 @@ func NewProvider(config *Config) (Provider, error) { config = DefaultProviderConfig() } - switch config.Provider { - case "openai": - if config.OpenAIKey == "" { - return nil, fmt.Errorf("OpenAI API key is required") - } - return NewOpenAIProvider(openAIAudioConfigFrom(config), config.OutputFormat) - case "gemini": - if config.GoogleAPIKey == "" { - return nil, fmt.Errorf("google API key is required") - } - return NewGeminiProvider(geminiAudioConfigFrom(config), config.OutputFormat) - default: + name := strings.ToLower(strings.TrimSpace(config.Provider)) + fn, ok := defaultAudioProviders.Get(name) + if !ok { return nil, fmt.Errorf("unknown audio provider: %s", config.Provider) } + return fn(config) } diff --git a/internal/gui/app.go b/internal/gui/app.go index f7a7d51..6c1cf21 100644 --- a/internal/gui/app.go +++ b/internal/gui/app.go @@ -146,11 +146,6 @@ type Config struct { Translator *translation.Translator } -const ( - imageProviderOpenAI = "openai" - imageProviderNanoBanana = "nanobanana" -) - // DefaultConfig returns default GUI configuration func DefaultConfig() *Config { homeDir, err := appconfig.HomeDir() @@ -168,7 +163,7 @@ func DefaultConfig() *Config { NanoBananaModel: image.DefaultNanoBananaModel, NanoBananaTextModel: image.DefaultNanoBananaTextModel, GeminiTTSModel: audioDefaults.GeminiTTSModel, - ImageProvider: imageProviderNanoBanana, + ImageProvider: image.ImageProviderNanoBanana, TranslationProvider: translation.ProviderGemini, PhoneticProvider: phonetic.ProviderGemini, AutoPlay: true, // Auto-play enabled by default diff --git a/internal/gui/app_test.go b/internal/gui/app_test.go index 9ac0125..86e8ec2 100644 --- a/internal/gui/app_test.go +++ b/internal/gui/app_test.go @@ -19,8 +19,8 @@ func TestDefaultConfigUsesGeminiLanguageProviders(t *testing.T) { if config.PhoneticProvider != phonetic.ProviderGemini { t.Fatalf("DefaultConfig() phonetic provider = %q, want %q", config.PhoneticProvider, phonetic.ProviderGemini) } - if config.ImageProvider != imageProviderNanoBanana { - t.Fatalf("DefaultConfig() image provider = %q, want %q", config.ImageProvider, imageProviderNanoBanana) + if config.ImageProvider != image.ImageProviderNanoBanana { + t.Fatalf("DefaultConfig() image provider = %q, want %q", config.ImageProvider, image.ImageProviderNanoBanana) } if config.AudioProvider != audioDefaults.Provider { t.Fatalf("DefaultConfig() audio provider = %q, want %q", config.AudioProvider, audioDefaults.Provider) diff --git a/internal/gui/card_service.go b/internal/gui/card_service.go index 395c741..cd4d549 100644 --- a/internal/gui/card_service.go +++ b/internal/gui/card_service.go @@ -9,6 +9,7 @@ import ( "codeberg.org/snonux/totalrecall/internal" "codeberg.org/snonux/totalrecall/internal/anki" + "codeberg.org/snonux/totalrecall/internal/image" "codeberg.org/snonux/totalrecall/internal/store" ) @@ -249,7 +250,7 @@ func (cs *CardService) loadImageFile(wordDir string, cf *CardFiles) { // Try to load the image prompt from the attribution file as a fallback // when the image provider is AI-based (OpenAI DALL-E or Nano Banana). - if cs.config.ImageProvider == imageProviderOpenAI || cs.config.ImageProvider == imageProviderNanoBanana { + if cs.config.ImageProvider == image.ImageProviderOpenAI || cs.config.ImageProvider == image.ImageProviderNanoBanana { cs.loadPromptFromAttribution(cf) } } diff --git a/internal/gui/generator_test.go b/internal/gui/generator_test.go index 40c0e2c..5ee238b 100644 --- a/internal/gui/generator_test.go +++ b/internal/gui/generator_test.go @@ -34,7 +34,7 @@ func (f *fakePromptAwareImageClient) Search(_ context.Context, opts *image.Searc Height: 1, Description: "fake result", Attribution: "fake attribution", - Source: imageProviderNanoBanana, + Source: image.ImageProviderNanoBanana, }, }, nil } @@ -48,7 +48,7 @@ func (f *fakePromptAwareImageClient) GetAttribution(*image.SearchResult) string } func (f *fakePromptAwareImageClient) Name() string { - return imageProviderNanoBanana + return image.ImageProviderNanoBanana } func (f *fakePromptAwareImageClient) SetPromptCallback(callback func(prompt string)) { @@ -99,7 +99,7 @@ func TestGenerateImagesWithPromptUsesNanoBananaProvider(t *testing.T) { tempDir := t.TempDir() app := &Application{ config: &Config{ - ImageProvider: imageProviderNanoBanana, + ImageProvider: image.ImageProviderNanoBanana, GoogleAPIKey: "google-key", NanoBananaModel: "custom-image-model", NanoBananaTextModel: "custom-text-model", diff --git a/internal/gui/orchestrator.go b/internal/gui/orchestrator.go index 191e919..ef4a167 100644 --- a/internal/gui/orchestrator.go +++ b/internal/gui/orchestrator.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "strings" "time" "fyne.io/fyne/v2" @@ -12,6 +13,7 @@ import ( "codeberg.org/snonux/totalrecall/internal/audio" "codeberg.org/snonux/totalrecall/internal/image" "codeberg.org/snonux/totalrecall/internal/phonetic" + "codeberg.org/snonux/totalrecall/internal/registry" "codeberg.org/snonux/totalrecall/internal/translation" ) @@ -410,48 +412,61 @@ func (o *GenerationOrchestrator) imagePromptCallback(cardDir, word string) func( } } -// newImageSearcher constructs the appropriate image client based on the -// configured image provider. Returns image.PromptAwareClient so callers can -// call SetPromptCallback directly without a type-assertion. The factory -// functions are sourced from imageFactories (the shared image.ClientFactories -// value) to avoid duplicating the factory signatures in this package. -func (o *GenerationOrchestrator) newImageSearcher() (image.PromptAwareClient, error) { - switch o.config.ImageProvider { - case imageProviderOpenAI: - if o.config.OpenAIKey == "" { - return nil, fmt.Errorf("OpenAI API key is required for image generation") - } +// guiImageClientFactories maps provider name to image client builder. Add new +// providers by registering here instead of extending a switch in newImageSearcher. +var guiImageClientFactories = func() *registry.Registry[string, func(*GenerationOrchestrator) (image.PromptAwareClient, error)] { + r := registry.New[string, func(*GenerationOrchestrator) (image.PromptAwareClient, error)]() + r.Register(image.ImageProviderOpenAI, (*GenerationOrchestrator).buildOpenAIImageClient) + r.Register(image.ImageProviderNanoBanana, (*GenerationOrchestrator).buildNanoBananaImageClient) + return r +}() - openaiConfig := &image.OpenAIConfig{ - APIKey: o.config.OpenAIKey, - Model: "dall-e-2", // DALL-E 2 supports 512×512 - Size: "512x512", - Quality: "standard", - Style: "natural", - } +func (o *GenerationOrchestrator) buildOpenAIImageClient() (image.PromptAwareClient, error) { + if o.config.OpenAIKey == "" { + return nil, fmt.Errorf("OpenAI API key is required for image generation") + } - return o.imageFactories.NewOpenAIClient(openaiConfig), nil + openaiConfig := &image.OpenAIConfig{ + APIKey: o.config.OpenAIKey, + Model: "dall-e-2", // DALL-E 2 supports 512×512 + Size: "512x512", + Quality: "standard", + Style: "natural", + } - case imageProviderNanoBanana: - cfg := o.config - if cfg == nil { - cfg = DefaultConfig() - } - if cfg.GoogleAPIKey == "" { - return nil, fmt.Errorf("google API key is required for image generation") - } + return o.imageFactories.NewOpenAIClient(openaiConfig), nil +} - nanoBananaConfig := &image.NanoBananaConfig{ - APIKey: cfg.GoogleAPIKey, - Model: cfg.NanoBananaModel, - TextModel: cfg.NanoBananaTextModel, - } +func (o *GenerationOrchestrator) buildNanoBananaImageClient() (image.PromptAwareClient, error) { + cfg := o.config + if cfg == nil { + cfg = DefaultConfig() + } + if cfg.GoogleAPIKey == "" { + return nil, fmt.Errorf("google API key is required for image generation") + } - return o.imageFactories.NewNanoBananaClient(nanoBananaConfig), nil + nanoBananaConfig := &image.NanoBananaConfig{ + APIKey: cfg.GoogleAPIKey, + Model: cfg.NanoBananaModel, + TextModel: cfg.NanoBananaTextModel, + } + + return o.imageFactories.NewNanoBananaClient(nanoBananaConfig), nil +} - default: +// newImageSearcher constructs the appropriate image client based on the +// configured image provider. Returns image.PromptAwareClient so callers can +// call SetPromptCallback directly without a type-assertion. The factory +// functions are sourced from imageFactories (the shared image.ClientFactories +// value) to avoid duplicating the factory signatures in this package. +func (o *GenerationOrchestrator) newImageSearcher() (image.PromptAwareClient, error) { + key := strings.ToLower(strings.TrimSpace(o.config.ImageProvider)) + fn, ok := guiImageClientFactories.Get(key) + if !ok { return nil, fmt.Errorf("unknown image provider: %s", o.config.ImageProvider) } + return fn(o) } // --- Phonetics --- diff --git a/internal/image/search.go b/internal/image/search.go index 61176cd..73a620a 100644 --- a/internal/image/search.go +++ b/internal/image/search.go @@ -77,6 +77,14 @@ type ImageClient interface { AttributionProvider } +// Image-generation provider names (AI backends). Use these keys when +// registering GUI/processor image client factories so string literals are not +// scattered across packages. +const ( + ImageProviderOpenAI = "openai" + ImageProviderNanoBanana = "nanobanana" +) + // PromptAwareClient extends ImageClient with a callback for receiving the // generated image prompt before the actual image download begins. Both // OpenAIClient and NanoBananaClient implement this interface. It is the diff --git a/internal/processor/image_downloader.go b/internal/processor/image_downloader.go index 0fe9229..60fa921 100644 --- a/internal/processor/image_downloader.go +++ b/internal/processor/image_downloader.go @@ -16,6 +16,7 @@ import ( "codeberg.org/snonux/totalrecall/internal/cli" "codeberg.org/snonux/totalrecall/internal/image" + "codeberg.org/snonux/totalrecall/internal/registry" ) // downloadImagesWithTranslation downloads images for a word into its card @@ -113,19 +114,26 @@ func (p *Processor) saveImagePrompt(wordDir string, searcher image.PromptAwareCl return nil } +// processorImageClientFactories maps run-mode image provider name to builder. +// Register new backends here instead of extending a switch in newImageSearcher. +var processorImageClientFactories = func() *registry.Registry[string, func(*Processor) (image.PromptAwareClient, error)] { + r := registry.New[string, func(*Processor) (image.PromptAwareClient, error)]() + r.Register(image.ImageProviderOpenAI, (*Processor).newOpenAIImageSearcher) + r.Register(image.ImageProviderNanoBanana, (*Processor).newNanoBananaImageSearcher) + return r +}() + // newImageSearcher creates the appropriate PromptAwareClient based on the // configured image provider (openai or nanobanana). Returning PromptAwareClient // instead of ImageClient means callers can call SetPromptCallback directly // without a type-assertion. func (p *Processor) newImageSearcher() (image.PromptAwareClient, error) { - switch p.imageProviderForRunMode() { - case "openai": - return p.newOpenAIImageSearcher() - case "nanobanana": - return p.newNanoBananaImageSearcher() - default: + key := strings.ToLower(strings.TrimSpace(p.imageProviderForRunMode())) + fn, ok := processorImageClientFactories.Get(key) + if !ok { return nil, fmt.Errorf("unknown image provider: %s", p.imageProviderForRunMode()) } + return fn(p) } // imageProviderForRunMode resolves the image provider, giving precedence to diff --git a/internal/registry/registry.go b/internal/registry/registry.go new file mode 100644 index 0000000..d3e9c2a --- /dev/null +++ b/internal/registry/registry.go @@ -0,0 +1,40 @@ +// Package registry provides a small generic keyed map for wiring named factories +// without central type switches (Open/Closed). K is usually string; T is the +// constructor or builder function type for that key. +package registry + +// Registry maps comparable keys to values (typically factory functions). +// It is intentionally minimal: no mutex; callers register at init/package load. +type Registry[K comparable, T any] struct { + m map[K]T +} + +// New returns an empty Registry. Call Register for each supported key. +func New[K comparable, T any]() *Registry[K, T] { + return &Registry[K, T]{m: make(map[K]T)} +} + +// Register associates key with value. It panics if key is already registered +// so duplicate wiring is caught at startup. +func (r *Registry[K, T]) Register(key K, value T) { + if r == nil { + panic("registry: Register on nil Registry") + } + if r.m == nil { + r.m = make(map[K]T) + } + if _, exists := r.m[key]; exists { + panic("registry: duplicate registration for key") + } + r.m[key] = value +} + +// Get returns the value for key and whether it was found. +func (r *Registry[K, T]) Get(key K) (T, bool) { + var zero T + if r == nil || r.m == nil { + return zero, false + } + v, ok := r.m[key] + return v, ok +} diff --git a/internal/registry/registry_test.go b/internal/registry/registry_test.go new file mode 100644 index 0000000..08f6189 --- /dev/null +++ b/internal/registry/registry_test.go @@ -0,0 +1,29 @@ +package registry + +import "testing" + +func TestRegistryRegisterGet(t *testing.T) { + r := New[string, int]() + r.Register("a", 1) + r.Register("b", 2) + + v, ok := r.Get("a") + if !ok || v != 1 { + t.Fatalf("Get(a) = %v, %v, want 1, true", v, ok) + } + _, ok = r.Get("missing") + if ok { + t.Fatal("Get(missing) should be false") + } +} + +func TestRegistryDuplicatePanics(t *testing.T) { + r := New[string, int]() + r.Register("a", 1) + defer func() { + if recover() == nil { + t.Fatal("expected panic on duplicate Register") + } + }() + r.Register("a", 2) +} |
