summaryrefslogtreecommitdiff
path: root/cmd
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2026-04-21 22:58:44 +0300
committerPaul Buetow <paul@buetow.org>2026-04-21 22:58:44 +0300
commitc5856133c8f12e3fdc76de0fc0482bf072580252 (patch)
treed85832e3a5c73a66f7e37d7e0bd04d5ecf7173e7 /cmd
parent15c08b9e665ad7c11bffcb671ab1a8338243bf72 (diff)
t7 centralize provider registries
Diffstat (limited to 'cmd')
-rw-r--r--cmd/comicforge/cli.go58
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 {