diff options
| author | Paul Buetow <paul@buetow.org> | 2026-04-21 22:58:44 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-04-21 22:58:44 +0300 |
| commit | c5856133c8f12e3fdc76de0fc0482bf072580252 (patch) | |
| tree | d85832e3a5c73a66f7e37d7e0bd04d5ecf7173e7 /cmd | |
| parent | 15c08b9e665ad7c11bffcb671ab1a8338243bf72 (diff) | |
t7 centralize provider registries
Diffstat (limited to 'cmd')
| -rw-r--r-- | cmd/comicforge/cli.go | 58 |
1 files changed, 23 insertions, 35 deletions
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 { |
