diff options
Diffstat (limited to 'internal/provider/provider.go')
| -rw-r--r-- | internal/provider/provider.go | 86 |
1 files changed, 83 insertions, 3 deletions
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)) |
