summaryrefslogtreecommitdiff
path: root/internal/image/registry.go
blob: abf79177e34a26d2306e0624ffbd94f5d51ffe0f (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
package image

import (
	"fmt"
	"reflect"
	"sync"
)

// Factory builds a provider instance.
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
	ComicAspectRatio() string
}

// Registry resolves provider names to factories.
type Registry[C Config] struct {
	mu        sync.RWMutex
	factories map[string]Factory[C]
}

// NewRegistry creates an empty provider registry.
func NewRegistry[C Config]() *Registry[C] {
	return &Registry[C]{factories: make(map[string]Factory[C])}
}

// Register associates name with a factory. Later registrations replace earlier ones.
func (r *Registry[C]) Register(name string, factory Factory[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[C])
	}
	r.factories[normalized] = factory
}

// Resolve returns the factory registered for name.
func (r *Registry[C]) Resolve(name string) (Factory[C], bool) {
	if r == nil {
		return nil, false
	}

	r.mu.RLock()
	defer r.mu.RUnlock()
	factory, ok := r.factories[NormalizeName(name)]
	return factory, ok
}

// Has reports whether a factory is registered for name.
func (r *Registry[C]) Has(name string) bool {
	_, ok := r.Resolve(name)
	return ok
}

// New constructs a provider for name.
func (r *Registry[C]) New(name string, cfg C) (ImageProvider, error) {
	var zero ImageProvider
	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[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")
	}
	return r.New(cfg.ImageProviderName(), cfg)
}

// DefaultRegistry returns the built-in image provider registry.
func DefaultRegistry() *Registry[GeminiRegistryConfig] {
	registry := NewRegistry[GeminiRegistryConfig]()
	registry.Register(Gemini, func(cfg GeminiRegistryConfig) (ImageProvider, error) {
		geminiProvider := NewGeminiProvider(&GeminiConfig{
			APIKey:      cfg.GoogleAPIKey(),
			Model:       cfg.ImageModel(),
			TextModel:   cfg.ImageTextModel(),
			AspectRatio: cfg.ComicAspectRatio(),
		})
		if err := geminiProvider.IsAvailable(); err != nil {
			return nil, err
		}
		return geminiProvider, nil
	})
	registry.Register(OpenAI, func(GeminiRegistryConfig) (ImageProvider, error) {
		return nil, fmt.Errorf("image provider %q is not implemented", OpenAI)
	})
	return registry
}

func isNilValue[T any](value T) bool {
	v := reflect.ValueOf(value)
	if !v.IsValid() {
		return true
	}

	switch v.Kind() {
	case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice:
		return v.IsNil()
	default:
		return false
	}
}