From c5856133c8f12e3fdc76de0fc0482bf072580252 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Tue, 21 Apr 2026 22:58:44 +0300 Subject: t7 centralize provider registries --- cmd/comicforge/cli.go | 58 ++++++++++++++++++++------------------------------- 1 file changed, 23 insertions(+), 35 deletions(-) (limited to 'cmd') diff --git a/cmd/comicforge/cli.go b/cmd/comicforge/cli.go index 9da2eb3..e472108 100644 --- a/cmd/comicforge/cli.go +++ b/cmd/comicforge/cli.go @@ -2,6 +2,7 @@ package main import ( "context" + "errors" "fmt" "strings" @@ -199,59 +200,46 @@ 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 + p, err := textprovider.DefaultRegistry().NewFromConfig(cfg) + if err != nil { + if errors.Is(err, provider.ErrUnknownProvider) { + return nil, fmt.Errorf("text provider %q is not implemented", cfg.Provider.Text) } - return p, nil - default: - return nil, fmt.Errorf("text provider %q is not implemented", cfg.Provider.Text) + return nil, err } + return p, nil } 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 + p, err := image.DefaultRegistry().NewFromConfig(cfg) + if err != nil { + if errors.Is(err, image.ErrUnknownProvider) { + return nil, fmt.Errorf("image provider %q is not implemented", cfg.Provider.Image) } - return p, nil - default: - return nil, fmt.Errorf("image provider %q is not implemented", cfg.Provider.Image) + return nil, err } + sharedProvider, ok := p.(provider.ImageProvider) + if !ok { + return nil, fmt.Errorf("image provider %q does not satisfy the shared provider interface", cfg.Provider.Image) + } + return sharedProvider, nil } 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 + p, err := tts.DefaultRegistry().NewFromConfig(cfg, voice) + if err != nil { + if errors.Is(err, provider.ErrUnknownProvider) { + return nil, fmt.Errorf("tts provider %q is not implemented", cfg.Provider.TTS) } - return p, nil - default: - return nil, fmt.Errorf("tts provider %q is not implemented", cfg.Provider.TTS) + return nil, err } + return p, nil } func resolveUltraRealistic(flags cliFlags) *bool { -- cgit v1.2.3