diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/image/registry.go | 16 | ||||
| -rw-r--r-- | internal/image/types_test.go | 36 |
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 |
