summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--internal/image/registry.go16
-rw-r--r--internal/image/types_test.go36
2 files changed, 48 insertions, 4 deletions
diff --git a/internal/image/registry.go b/internal/image/registry.go
index 7b2a465..2d446e8 100644
--- a/internal/image/registry.go
+++ b/internal/image/registry.go
@@ -12,6 +12,11 @@ type Factory[C any] func(C) (ImageProvider, error)
// Config exposes the configured image provider name.
type Config interface {
ImageProviderName() string
+}
+
+// GeminiRegistryConfig exposes the settings required by the built-in Gemini image provider.
+type GeminiRegistryConfig interface {
+ Config
GoogleAPIKey() string
ImageModel() string
ImageTextModel() string
@@ -77,6 +82,9 @@ func (r *Registry[C]) New(name string, cfg C) (ImageProvider, error) {
// NewFromConfig resolves the provider name from cfg and constructs it.
func (r *Registry[C]) NewFromConfig(cfg C) (ImageProvider, error) {
+ if r == nil {
+ return nil, fmt.Errorf("image registry is required")
+ }
if isNilValue(cfg) {
return nil, fmt.Errorf("image config is required")
}
@@ -84,9 +92,9 @@ func (r *Registry[C]) NewFromConfig(cfg C) (ImageProvider, error) {
}
// DefaultRegistry returns the built-in image provider registry.
-func DefaultRegistry() *Registry[Config] {
- registry := NewRegistry[Config]()
- registry.Register(Gemini, func(cfg Config) (ImageProvider, error) {
+func DefaultRegistry() *Registry[GeminiRegistryConfig] {
+ registry := NewRegistry[GeminiRegistryConfig]()
+ registry.Register(Gemini, func(cfg GeminiRegistryConfig) (ImageProvider, error) {
geminiProvider := NewGeminiProvider(&GeminiConfig{
APIKey: cfg.GoogleAPIKey(),
Model: cfg.ImageModel(),
@@ -97,7 +105,7 @@ func DefaultRegistry() *Registry[Config] {
}
return geminiProvider, nil
})
- registry.Register(OpenAI, func(Config) (ImageProvider, error) {
+ registry.Register(OpenAI, func(GeminiRegistryConfig) (ImageProvider, error) {
return nil, fmt.Errorf("image provider %q is not implemented", OpenAI)
})
return registry
diff --git a/internal/image/types_test.go b/internal/image/types_test.go
index 4416870..5916389 100644
--- a/internal/image/types_test.go
+++ b/internal/image/types_test.go
@@ -71,6 +71,36 @@ func TestRegistryNewFromConfig(t *testing.T) {
}
}
+func TestRegistryNewFromConfigNilRegistry(t *testing.T) {
+ t.Parallel()
+
+ var registry *Registry[fakeConfig]
+ _, err := registry.NewFromConfig(fakeConfig{name: Gemini})
+ if err == nil {
+ t.Fatal("expected nil registry error")
+ }
+ if !strings.Contains(err.Error(), "image registry is required") {
+ t.Fatalf("error = %v, want nil registry error", err)
+ }
+}
+
+func TestRegistrySupportsCustomConfigWithoutGeminiFields(t *testing.T) {
+ t.Parallel()
+
+ registry := NewRegistry[customConfig]()
+ registry.Register(Gemini, func(cfg customConfig) (ImageProvider, error) {
+ return fakeProvider{name: cfg.name}, nil
+ })
+
+ gotProvider, err := registry.NewFromConfig(customConfig{name: Gemini})
+ 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 TestDefaultRegistryNewFromConfig(t *testing.T) {
t.Parallel()
@@ -133,6 +163,12 @@ 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 customConfig struct {
+ name string
+}
+
+func (f customConfig) ImageProviderName() string { return f.name }
+
type fakeProvider struct {
name string
token string