summaryrefslogtreecommitdiff
path: root/internal/image
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 /internal/image
parent15c08b9e665ad7c11bffcb671ab1a8338243bf72 (diff)
t7 centralize provider registries
Diffstat (limited to 'internal/image')
-rw-r--r--internal/image/registry.go29
-rw-r--r--internal/image/types_test.go57
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