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 /internal/image | |
| parent | 15c08b9e665ad7c11bffcb671ab1a8338243bf72 (diff) | |
t7 centralize provider registries
Diffstat (limited to 'internal/image')
| -rw-r--r-- | internal/image/registry.go | 29 | ||||
| -rw-r--r-- | internal/image/types_test.go | 57 |
2 files changed, 80 insertions, 6 deletions
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 |
