summaryrefslogtreecommitdiff
path: root/internal/provider/provider_test.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/provider/provider_test.go')
-rw-r--r--internal/provider/provider_test.go66
1 files changed, 65 insertions, 1 deletions
diff --git a/internal/provider/provider_test.go b/internal/provider/provider_test.go
index 46931d1..147bf44 100644
--- a/internal/provider/provider_test.go
+++ b/internal/provider/provider_test.go
@@ -1,6 +1,10 @@
package provider
-import "testing"
+import (
+ "context"
+ "errors"
+ "testing"
+)
func TestNormalizeName(t *testing.T) {
t.Parallel()
@@ -20,3 +24,63 @@ func TestIsKnownName(t *testing.T) {
t.Fatal("expected bogus provider to be rejected")
}
}
+
+func TestRegistryNewFromConfig(t *testing.T) {
+ t.Parallel()
+
+ registry := NewTextRegistry()
+ registry.Register(Gemini, func(cfg TextConfig) (TextProvider, error) {
+ return fakeTextProvider{name: cfg.TextProviderName()}, nil
+ })
+
+ provider, err := registry.NewFromConfig(fakeTextConfig{name: "gemini"}, func(cfg TextConfig) string {
+ return cfg.TextProviderName()
+ })
+ if err != nil {
+ t.Fatalf("NewFromConfig() error = %v", err)
+ }
+ if got, want := provider.Name(), "gemini"; got != want {
+ t.Fatalf("provider.Name() = %q, want %q", got, want)
+ }
+}
+
+func TestRegistryUnknownProvider(t *testing.T) {
+ t.Parallel()
+
+ registry := NewTTSRegistry()
+ _, err := registry.New("missing", fakeTTSConfig{})
+ if err == nil {
+ t.Fatal("expected unknown provider error")
+ }
+ if !errors.Is(err, ErrUnknownProvider) {
+ t.Fatalf("error = %v, want ErrUnknownProvider", err)
+ }
+}
+
+type fakeTextConfig struct {
+ name string
+}
+
+func (f fakeTextConfig) TextProviderName() string { return f.name }
+
+type fakeTTSConfig struct{}
+
+func (fakeTTSConfig) TTSProviderName() string { return "missing" }
+
+type fakeTextProvider struct {
+ name string
+}
+
+func (f fakeTextProvider) Name() string { return f.name }
+func (f fakeTextProvider) IsAvailable() error { return nil }
+func (f fakeTextProvider) GenerateText(_ context.Context, _ string) (string, error) {
+ return "", nil
+}
+
+type fakeTTSProvider struct{}
+
+func (fakeTTSProvider) Name() string { return "missing" }
+func (fakeTTSProvider) IsAvailable() error { return nil }
+func (fakeTTSProvider) GenerateAudio(_ context.Context, _ string, _ string) error {
+ return nil
+}