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 | |
| parent | 15c08b9e665ad7c11bffcb671ab1a8338243bf72 (diff) | |
t7 centralize provider registries
| -rw-r--r-- | cmd/comicforge/cli.go | 58 | ||||
| -rw-r--r-- | internal/config/config.go | 52 | ||||
| -rw-r--r-- | internal/image/registry.go | 29 | ||||
| -rw-r--r-- | internal/image/types_test.go | 57 | ||||
| -rw-r--r-- | internal/provider/provider.go | 6 | ||||
| -rw-r--r-- | internal/text/registry.go | 80 | ||||
| -rw-r--r-- | internal/text/registry_test.go | 66 | ||||
| -rw-r--r-- | internal/tts/registry.go | 89 | ||||
| -rw-r--r-- | internal/tts/registry_test.go | 66 |
9 files changed, 459 insertions, 44 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 { diff --git a/internal/config/config.go b/internal/config/config.go index ba8a56e..5b4ab18 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -13,7 +13,10 @@ import ( "github.com/spf13/viper" + "codeberg.org/snonux/comicforge/internal/image" "codeberg.org/snonux/comicforge/internal/provider" + "codeberg.org/snonux/comicforge/internal/text" + "codeberg.org/snonux/comicforge/internal/tts" "codeberg.org/snonux/comicforge/prompts" ) @@ -40,6 +43,9 @@ var ( _ provider.TextConfig = (*Config)(nil) _ provider.ImageConfig = (*Config)(nil) _ provider.TTSConfig = (*Config)(nil) + _ text.Config = (*Config)(nil) + _ image.Config = (*Config)(nil) + _ tts.Config = (*Config)(nil) ) // ProviderConfig stores the selected provider name for each capability. @@ -208,6 +214,46 @@ func (c *Config) ImageProviderName() string { return provider.NormalizeName(c.Provider.Image) } +// GoogleAPIKey returns the configured Google API key. +func (c *Config) GoogleAPIKey() string { + if c == nil { + return "" + } + return c.API.GoogleAPIKey +} + +// TextModel returns the configured text model name. +func (c *Config) TextModel() string { + if c == nil { + return "" + } + return c.Models.Text +} + +// ImageModel returns the configured image model name. +func (c *Config) ImageModel() string { + if c == nil { + return "" + } + return c.Models.Image +} + +// ImageTextModel returns the configured image-text model name. +func (c *Config) ImageTextModel() string { + if c == nil { + return "" + } + return c.Models.ImageText +} + +// TTSModel returns the configured text-to-speech model name. +func (c *Config) TTSModel() string { + if c == nil { + return "" + } + return c.Models.TTS +} + // TTSProviderName returns the configured TTS provider name. func (c *Config) TTSProviderName() string { return provider.NormalizeName(c.Provider.TTS) @@ -290,13 +336,13 @@ func (c *Config) normalize() { } func (c *Config) validate() error { - if !provider.IsKnownName(c.Provider.Text) { + if !text.DefaultRegistry().Has(c.Provider.Text) { return fmt.Errorf("unknown text provider: %s", c.Provider.Text) } - if !provider.IsKnownName(c.Provider.Image) { + if !image.DefaultRegistry().Has(c.Provider.Image) { return fmt.Errorf("unknown image provider: %s", c.Provider.Image) } - if !provider.IsKnownName(c.Provider.TTS) { + if !tts.DefaultRegistry().Has(c.Provider.TTS) { return fmt.Errorf("unknown TTS provider: %s", c.Provider.TTS) } diff --git a/internal/image/registry.go b/internal/image/registry.go index 0338115..7b2a465 100644 --- a/internal/image/registry.go +++ b/internal/image/registry.go @@ -12,6 +12,9 @@ type Factory[C any] func(C) (ImageProvider, error) // Config exposes the configured image provider name. type Config interface { ImageProviderName() string + GoogleAPIKey() string + ImageModel() string + ImageTextModel() string } // Registry resolves provider names to factories. @@ -56,6 +59,12 @@ func (r *Registry[C]) Resolve(name string) (Factory[C], bool) { return factory, ok } +// Has reports whether a factory is registered for name. +func (r *Registry[C]) Has(name string) bool { + _, ok := r.Resolve(name) + return ok +} + // New constructs a provider for name. func (r *Registry[C]) New(name string, cfg C) (ImageProvider, error) { var zero ImageProvider @@ -74,6 +83,26 @@ func (r *Registry[C]) NewFromConfig(cfg C) (ImageProvider, error) { return r.New(cfg.ImageProviderName(), cfg) } +// DefaultRegistry returns the built-in image provider registry. +func DefaultRegistry() *Registry[Config] { + registry := NewRegistry[Config]() + registry.Register(Gemini, func(cfg Config) (ImageProvider, error) { + geminiProvider := NewGeminiProvider(&GeminiConfig{ + APIKey: cfg.GoogleAPIKey(), + Model: cfg.ImageModel(), + TextModel: cfg.ImageTextModel(), + }) + if err := geminiProvider.IsAvailable(); err != nil { + return nil, err + } + return geminiProvider, nil + }) + registry.Register(OpenAI, func(Config) (ImageProvider, error) { + return nil, fmt.Errorf("image provider %q is not implemented", OpenAI) + }) + return registry +} + func isNilValue[T any](value T) bool { v := reflect.ValueOf(value) if !v.IsValid() { diff --git a/internal/image/types_test.go b/internal/image/types_test.go index 9257392..4416870 100644 --- a/internal/image/types_test.go +++ b/internal/image/types_test.go @@ -53,21 +53,60 @@ func TestRegistryNewFromConfig(t *testing.T) { registry := NewRegistry[fakeConfig]() registry.Register(Gemini, func(cfg fakeConfig) (ImageProvider, error) { - return fakeProvider(cfg), nil + return fakeProvider{ + name: cfg.name, + token: cfg.token, + }, nil }) - provider, err := registry.NewFromConfig(fakeConfig{name: Gemini, token: "secret"}) + gotProvider, err := registry.NewFromConfig(fakeConfig{name: Gemini, token: "secret"}) if err != nil { t.Fatalf("NewFromConfig() error = %v", err) } - if got, want := provider.Name(), Gemini; got != want { + if got, want := gotProvider.Name(), Gemini; got != want { t.Fatalf("provider.Name() = %q, want %q", got, want) } - if got, want := provider.(fakeProvider).token, "secret"; got != want { + if got, want := gotProvider.(fakeProvider).token, "secret"; got != want { t.Fatalf("provider token = %q, want %q", got, want) } } +func TestDefaultRegistryNewFromConfig(t *testing.T) { + t.Parallel() + + registry := DefaultRegistry() + gotProvider, err := registry.NewFromConfig(fakeConfig{ + name: Gemini, + apiKey: "test-key", + model: DefaultGeminiImageModel, + textModel: DefaultGeminiTextModel, + }) + if err != nil { + t.Fatalf("NewFromConfig() error = %v", err) + } + if got, want := gotProvider.Name(), Gemini; got != want { + t.Fatalf("provider.Name() = %q, want %q", got, want) + } +} + +func TestDefaultRegistryUnsupportedProvider(t *testing.T) { + t.Parallel() + + registry := DefaultRegistry() + _, err := registry.NewFromConfig(fakeConfig{ + name: OpenAI, + apiKey: "test-key", + model: "unused", + textModel: "unused", + }) + if err == nil { + t.Fatal("expected unsupported provider error") + } + if !strings.Contains(err.Error(), "image provider \"openai\" is not implemented") { + t.Fatalf("error = %v, want unsupported provider error", err) + } +} + func TestRegistryUnknownProvider(t *testing.T) { t.Parallel() @@ -82,11 +121,17 @@ func TestRegistryUnknownProvider(t *testing.T) { } type fakeConfig struct { - name string - token string + name string + token string + apiKey string + model string + textModel string } func (f fakeConfig) ImageProviderName() string { return f.name } +func (f fakeConfig) GoogleAPIKey() string { return f.apiKey } +func (f fakeConfig) ImageModel() string { return f.model } +func (f fakeConfig) ImageTextModel() string { return f.textModel } type fakeProvider struct { name string diff --git a/internal/provider/provider.go b/internal/provider/provider.go index 651688c..256dc4c 100644 --- a/internal/provider/provider.go +++ b/internal/provider/provider.go @@ -114,6 +114,12 @@ func (r *Registry[T, C]) Resolve(name string) (Factory[T, C], bool) { return factory, ok } +// Has reports whether a factory is registered for name. +func (r *Registry[T, C]) Has(name string) bool { + _, ok := r.Resolve(name) + return ok +} + // New constructs a provider for name using cfg. func (r *Registry[T, C]) New(name string, cfg C) (T, error) { var zero T diff --git a/internal/text/registry.go b/internal/text/registry.go new file mode 100644 index 0000000..e337b16 --- /dev/null +++ b/internal/text/registry.go @@ -0,0 +1,80 @@ +package text + +import ( + "fmt" + "reflect" + + "codeberg.org/snonux/comicforge/internal/provider" +) + +// Config exposes the configured text provider settings. +type Config interface { + provider.TextConfig + GoogleAPIKey() string + TextModel() string +} + +// Registry resolves provider names to factories. +type Registry struct { + *provider.Registry[provider.TextProvider, Config] +} + +// NewRegistry creates an empty text-provider registry. +func NewRegistry() *Registry { + return &Registry{ + Registry: provider.NewRegistry[provider.TextProvider, Config](), + } +} + +// DefaultRegistry returns the built-in text provider registry. +func DefaultRegistry() *Registry { + registry := NewRegistry() + registry.Register(provider.Gemini, func(cfg Config) (provider.TextProvider, error) { + geminiProvider := NewGeminiProvider(&GeminiConfig{ + APIKey: cfg.GoogleAPIKey(), + Model: cfg.TextModel(), + }) + if err := geminiProvider.IsAvailable(); err != nil { + return nil, err + } + return geminiProvider, nil + }) + registry.Register(provider.OpenAI, func(Config) (provider.TextProvider, error) { + return nil, fmt.Errorf("text provider %q is not implemented", provider.OpenAI) + }) + return registry +} + +// Has reports whether a factory is registered for name. +func (r *Registry) Has(name string) bool { + if r == nil || r.Registry == nil { + return false + } + _, ok := r.Resolve(name) + return ok +} + +// NewFromConfig resolves the provider name from cfg and constructs it. +func (r *Registry) NewFromConfig(cfg Config) (provider.TextProvider, error) { + if r == nil || r.Registry == nil { + return nil, fmt.Errorf("text registry is required") + } + if isNilValue(cfg) { + return nil, fmt.Errorf("text config is required") + } + return r.New(cfg.TextProviderName(), cfg) +} + +func isNilValue[T any](value T) bool { + v := reflect.ValueOf(value) + if !v.IsValid() { + return true + } + + switch v.Kind() { + case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice: + return v.IsNil() + default: + return false + } +} diff --git a/internal/text/registry_test.go b/internal/text/registry_test.go new file mode 100644 index 0000000..3f5f72a --- /dev/null +++ b/internal/text/registry_test.go @@ -0,0 +1,66 @@ +package text + +import ( + "errors" + "strings" + "testing" + + "codeberg.org/snonux/comicforge/internal/provider" +) + +func TestDefaultRegistryNewFromConfig(t *testing.T) { + t.Parallel() + + registry := DefaultRegistry() + gotProvider, err := registry.NewFromConfig(fakeConfig{ + name: provider.Gemini, + apiKey: "test-key", + model: "gemini-2.5-flash", + }) + if err != nil { + t.Fatalf("NewFromConfig() error = %v", err) + } + if got, want := gotProvider.Name(), provider.Gemini; got != want { + t.Fatalf("provider.Name() = %q, want %q", got, want) + } +} + +func TestDefaultRegistryUnsupportedProvider(t *testing.T) { + t.Parallel() + + registry := DefaultRegistry() + _, err := registry.NewFromConfig(fakeConfig{ + name: provider.OpenAI, + apiKey: "test-key", + model: "unused", + }) + if err == nil { + t.Fatal("expected unsupported provider error") + } + if !strings.Contains(err.Error(), "text provider \"openai\" is not implemented") { + t.Fatalf("error = %v, want unsupported provider error", err) + } +} + +func TestRegistryUnknownProvider(t *testing.T) { + t.Parallel() + + registry := DefaultRegistry() + _, err := registry.New("missing", fakeConfig{}) + if err == nil { + t.Fatal("expected unknown provider error") + } + if !errors.Is(err, provider.ErrUnknownProvider) { + t.Fatalf("error = %v, want ErrUnknownProvider", err) + } +} + +type fakeConfig struct { + name string + apiKey string + model string +} + +func (f fakeConfig) TextProviderName() string { return f.name } +func (f fakeConfig) GoogleAPIKey() string { return f.apiKey } +func (f fakeConfig) TextModel() string { return f.model } diff --git a/internal/tts/registry.go b/internal/tts/registry.go new file mode 100644 index 0000000..0b38a5e --- /dev/null +++ b/internal/tts/registry.go @@ -0,0 +1,89 @@ +package tts + +import ( + "fmt" + "reflect" + + "codeberg.org/snonux/comicforge/internal/provider" +) + +// Config exposes the configured TTS provider settings. +type Config interface { + provider.TTSConfig + GoogleAPIKey() string + TTSModel() string +} + +type registryConfig struct { + Config + voice string +} + +// Registry resolves provider names to factories. +type Registry struct { + *provider.Registry[provider.TTSProvider, registryConfig] +} + +// NewRegistry creates an empty TTS-provider registry. +func NewRegistry() *Registry { + return &Registry{ + Registry: provider.NewRegistry[provider.TTSProvider, registryConfig](), + } +} + +// DefaultRegistry returns the built-in TTS provider registry. +func DefaultRegistry() *Registry { + registry := NewRegistry() + registry.Register(provider.Gemini, func(cfg registryConfig) (provider.TTSProvider, error) { + geminiProvider := NewGeminiProvider(&GeminiConfig{ + APIKey: cfg.GoogleAPIKey(), + Model: cfg.TTSModel(), + Voice: cfg.voice, + }) + if err := geminiProvider.IsAvailable(); err != nil { + return nil, err + } + return geminiProvider, nil + }) + registry.Register(provider.OpenAI, func(registryConfig) (provider.TTSProvider, error) { + return nil, fmt.Errorf("tts provider %q is not implemented", provider.OpenAI) + }) + return registry +} + +// Has reports whether a factory is registered for name. +func (r *Registry) Has(name string) bool { + if r == nil || r.Registry == nil { + return false + } + _, ok := r.Resolve(name) + return ok +} + +// NewFromConfig resolves the provider name from cfg and constructs it. +func (r *Registry) NewFromConfig(cfg Config, voice string) (provider.TTSProvider, error) { + if r == nil || r.Registry == nil { + return nil, fmt.Errorf("tts registry is required") + } + if isNilValue(cfg) { + return nil, fmt.Errorf("tts config is required") + } + return r.New(cfg.TTSProviderName(), registryConfig{ + Config: cfg, + voice: voice, + }) +} + +func isNilValue[T any](value T) bool { + v := reflect.ValueOf(value) + if !v.IsValid() { + return true + } + + switch v.Kind() { + case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice: + return v.IsNil() + default: + return false + } +} diff --git a/internal/tts/registry_test.go b/internal/tts/registry_test.go new file mode 100644 index 0000000..f53fa17 --- /dev/null +++ b/internal/tts/registry_test.go @@ -0,0 +1,66 @@ +package tts + +import ( + "errors" + "strings" + "testing" + + "codeberg.org/snonux/comicforge/internal/provider" +) + +func TestDefaultRegistryNewFromConfig(t *testing.T) { + t.Parallel() + + registry := DefaultRegistry() + gotProvider, err := registry.NewFromConfig(fakeConfig{ + name: provider.Gemini, + apiKey: "test-key", + model: "gemini-2.5-flash-preview-tts", + }, "Aoede") + if err != nil { + t.Fatalf("NewFromConfig() error = %v", err) + } + if got, want := gotProvider.Name(), provider.Gemini; got != want { + t.Fatalf("provider.Name() = %q, want %q", got, want) + } +} + +func TestDefaultRegistryUnsupportedProvider(t *testing.T) { + t.Parallel() + + registry := DefaultRegistry() + _, err := registry.NewFromConfig(fakeConfig{ + name: provider.OpenAI, + apiKey: "test-key", + model: "unused", + }, "Aoede") + if err == nil { + t.Fatal("expected unsupported provider error") + } + if !strings.Contains(err.Error(), "tts provider \"openai\" is not implemented") { + t.Fatalf("error = %v, want unsupported provider error", err) + } +} + +func TestRegistryUnknownProvider(t *testing.T) { + t.Parallel() + + registry := DefaultRegistry() + _, err := registry.New("missing", registryConfig{Config: fakeConfig{}}) + if err == nil { + t.Fatal("expected unknown provider error") + } + if !errors.Is(err, provider.ErrUnknownProvider) { + t.Fatalf("error = %v, want ErrUnknownProvider", err) + } +} + +type fakeConfig struct { + name string + apiKey string + model string +} + +func (f fakeConfig) TTSProviderName() string { return f.name } +func (f fakeConfig) GoogleAPIKey() string { return f.apiKey } +func (f fakeConfig) TTSModel() string { return f.model } |
