From 621e24039b527a368a00a7d5164d8e004c2cf001 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Sun, 19 Apr 2026 22:03:02 +0300 Subject: u4: make provider/runtime boundary generic --- internal/provider/provider.go | 86 ++++++++++++++++++++++++++++++++++++-- internal/provider/provider_test.go | 66 ++++++++++++++++++++++++++++- 2 files changed, 148 insertions(+), 4 deletions(-) (limited to 'internal/provider') diff --git a/internal/provider/provider.go b/internal/provider/provider.go index 6239ecf..651688c 100644 --- a/internal/provider/provider.go +++ b/internal/provider/provider.go @@ -1,12 +1,13 @@ -// Package provider defines capability-specific AI provider interfaces and shared -// provider naming helpers. The concrete Gemini implementations will satisfy -// these interfaces once the comic pipeline is wired up. +// Package provider defines capability-specific AI provider interfaces and +// provider-neutral registries for runtime selection. package provider import ( "context" "errors" + "fmt" "strings" + "sync" ) const ( @@ -53,6 +54,85 @@ type TTSConfig interface { TTSProviderName() string } +// Factory builds a provider implementation from configuration. +type Factory[T any, C any] func(C) (T, error) + +// Registry resolves provider names to factories. +type Registry[T any, C any] struct { + mu sync.RWMutex + factories map[string]Factory[T, C] +} + +// NewRegistry creates an empty provider registry. +func NewRegistry[T any, C any]() *Registry[T, C] { + return &Registry[T, C]{factories: make(map[string]Factory[T, C])} +} + +// NewTextRegistry creates an empty text-provider registry. +func NewTextRegistry() *Registry[TextProvider, TextConfig] { + return NewRegistry[TextProvider, TextConfig]() +} + +// NewImageRegistry creates an empty image-provider registry. +func NewImageRegistry() *Registry[ImageProvider, ImageConfig] { + return NewRegistry[ImageProvider, ImageConfig]() +} + +// NewTTSRegistry creates an empty text-to-speech provider registry. +func NewTTSRegistry() *Registry[TTSProvider, TTSConfig] { + return NewRegistry[TTSProvider, TTSConfig]() +} + +// Register associates name with a factory. Later registrations replace earlier ones. +func (r *Registry[T, C]) Register(name string, factory Factory[T, C]) { + if r == nil || factory == nil { + return + } + + normalized := NormalizeName(name) + if normalized == "" { + return + } + + r.mu.Lock() + defer r.mu.Unlock() + if r.factories == nil { + r.factories = make(map[string]Factory[T, C]) + } + r.factories[normalized] = factory +} + +// Resolve returns the factory registered for name. +func (r *Registry[T, C]) Resolve(name string) (Factory[T, C], bool) { + if r == nil { + return nil, false + } + + r.mu.RLock() + defer r.mu.RUnlock() + factory, ok := r.factories[NormalizeName(name)] + return factory, ok +} + +// New constructs a provider for name using cfg. +func (r *Registry[T, C]) New(name string, cfg C) (T, error) { + var zero T + factory, ok := r.Resolve(name) + if !ok { + return zero, fmt.Errorf("%w: %s", ErrUnknownProvider, name) + } + return factory(cfg) +} + +// NewFromConfig resolves the provider name from cfg and constructs it. +func (r *Registry[T, C]) NewFromConfig(cfg C, providerName func(C) string) (T, error) { + var zero T + if providerName == nil { + return zero, errors.New("provider name resolver is required") + } + return r.New(providerName(cfg), cfg) +} + // NormalizeName returns a canonical lower-case provider name. func NormalizeName(name string) string { return strings.ToLower(strings.TrimSpace(name)) 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 +} -- cgit v1.2.3