summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/apicircuit/apicircuit.go85
-rw-r--r--internal/apicircuit/apicircuit_test.go19
-rw-r--r--internal/config/config_test.go10
-rw-r--r--internal/provider/provider.go86
-rw-r--r--internal/provider/provider_test.go66
5 files changed, 221 insertions, 45 deletions
diff --git a/internal/apicircuit/apicircuit.go b/internal/apicircuit/apicircuit.go
index 05081e5..7f623cb 100644
--- a/internal/apicircuit/apicircuit.go
+++ b/internal/apicircuit/apicircuit.go
@@ -1,10 +1,12 @@
-// Package apicircuit wraps outbound Gemini API calls with sony/gobreaker
-// circuit breakers so repeated failures do not pile up unbounded work.
+// Package apicircuit wraps outbound API calls with circuit breakers so repeated
+// failures against any provider do not pile up unbounded work.
package apicircuit
import (
"context"
"errors"
+ "fmt"
+ "strings"
"sync"
"time"
@@ -12,6 +14,15 @@ import (
)
const (
+ // CapabilityText identifies text-generation traffic.
+ CapabilityText Capability = "text"
+ // CapabilityImage identifies image-generation traffic.
+ CapabilityImage Capability = "image"
+ // CapabilityTTS identifies text-to-speech traffic.
+ CapabilityTTS Capability = "tts"
+)
+
+const (
// breakerInterval clears rolling failure counts in the closed state so stale
// errors do not keep the breaker sensitive forever.
breakerInterval = 2 * time.Minute
@@ -24,12 +35,22 @@ const (
breakerTripAfterConsecutiveFailures uint32 = 5
)
-var (
- geminiTTSOnce sync.Once
- geminiTTSBreaker *gobreaker.CircuitBreaker
- geminiImageOnce sync.Once
- geminiImageBreaker *gobreaker.CircuitBreaker
-)
+// Capability identifies the class of API call protected by a circuit breaker.
+type Capability string
+
+// Registry stores lazily-created circuit breakers keyed by provider and capability.
+type Registry struct {
+ mu sync.Mutex
+ breakers map[string]*gobreaker.CircuitBreaker
+}
+
+// NewRegistry creates an empty breaker registry.
+func NewRegistry() *Registry {
+ return &Registry{breakers: make(map[string]*gobreaker.CircuitBreaker)}
+}
+
+// DefaultRegistry is used by package-level helpers.
+var DefaultRegistry = NewRegistry()
// isSuccessful counts only real API outcomes: nil is success; context.Canceled is
// treated as success so user abort does not trip the breaker. Timeouts and
@@ -57,18 +78,36 @@ func newBreaker(name string) *gobreaker.CircuitBreaker {
})
}
-func geminiBreaker(name string, slot **gobreaker.CircuitBreaker, once *sync.Once) *gobreaker.CircuitBreaker {
- once.Do(func() {
- *slot = newBreaker(name)
- })
+func breakerKey(providerName string, capability Capability) string {
+ return NormalizeName(providerName) + ":" + NormalizeName(string(capability))
+}
+
+func (r *Registry) breaker(providerName string, capability Capability) *gobreaker.CircuitBreaker {
+ key := breakerKey(providerName, capability)
+
+ r.mu.Lock()
+ defer r.mu.Unlock()
- return *slot
+ if r.breakers == nil {
+ r.breakers = make(map[string]*gobreaker.CircuitBreaker)
+ }
+ if breaker := r.breakers[key]; breaker != nil {
+ return breaker
+ }
+
+ breaker := newBreaker(key)
+ r.breakers[key] = breaker
+ return breaker
}
-func runValue[T any](cb *gobreaker.CircuitBreaker, fn func() (T, error)) (T, error) {
+// Execute runs fn through the breaker for providerName and capability.
+func Execute[T any](r *Registry, providerName string, capability Capability, fn func() (T, error)) (T, error) {
var zero T
+ if r == nil {
+ r = DefaultRegistry
+ }
- v, err := cb.Execute(func() (interface{}, error) {
+ v, err := r.breaker(providerName, capability).Execute(func() (interface{}, error) {
return fn()
})
if err != nil {
@@ -78,15 +117,15 @@ func runValue[T any](cb *gobreaker.CircuitBreaker, fn func() (T, error)) (T, err
return zero, nil
}
- return v.(T), nil
-}
+ result, ok := v.(T)
+ if !ok {
+ return zero, fmt.Errorf("unexpected breaker result type for %s", breakerKey(providerName, capability))
+ }
-// GeminiTTS runs one Gemini TTS GenerateContent call through its circuit breaker.
-func GeminiTTS[T any](fn func() (T, error)) (T, error) {
- return runValue(geminiBreaker("gemini-tts", &geminiTTSBreaker, &geminiTTSOnce), fn)
+ return result, nil
}
-// GeminiImage runs one Gemini image-generation call through its circuit breaker.
-func GeminiImage[T any](fn func() (T, error)) (T, error) {
- return runValue(geminiBreaker("gemini-image", &geminiImageBreaker, &geminiImageOnce), fn)
+// NormalizeName returns a canonical lower-case name for breaker keys.
+func NormalizeName(name string) string {
+ return strings.ToLower(strings.TrimSpace(name))
}
diff --git a/internal/apicircuit/apicircuit_test.go b/internal/apicircuit/apicircuit_test.go
index 842a0f2..91ee783 100644
--- a/internal/apicircuit/apicircuit_test.go
+++ b/internal/apicircuit/apicircuit_test.go
@@ -6,25 +6,28 @@ import (
"testing"
)
-func TestGeminiTTS_Success(t *testing.T) {
+func TestExecute_Success(t *testing.T) {
t.Parallel()
- v, err := GeminiTTS(func() (string, error) {
+ v, err := Execute(NewRegistry(), "gemini", CapabilityText, func() (string, error) {
return "ok", nil
})
if err != nil || v != "ok" {
- t.Fatalf("GeminiTTS() = %q, %v; want ok, nil", v, err)
+ t.Fatalf("Execute() = %q, %v; want ok, nil", v, err)
}
}
-func TestGeminiImage_Success(t *testing.T) {
+func TestExecute_ContextCanceled(t *testing.T) {
t.Parallel()
- v, err := GeminiImage(func() (int, error) {
- return 42, nil
+ _, err := Execute(NewRegistry(), "gemini", CapabilityText, func() (string, error) {
+ return "", context.Canceled
})
- if err != nil || v != 42 {
- t.Fatalf("GeminiImage() = %d, %v; want 42, nil", v, err)
+ if err == nil {
+ t.Fatal("expected cancellation to be returned")
+ }
+ if !errors.Is(err, context.Canceled) {
+ t.Fatalf("error = %v, want context.Canceled", err)
}
}
diff --git a/internal/config/config_test.go b/internal/config/config_test.go
index d0a4efe..fe370f6 100644
--- a/internal/config/config_test.go
+++ b/internal/config/config_test.go
@@ -11,8 +11,6 @@ import (
)
func TestLoadReturnsDefaultsWhenConfigMissing(t *testing.T) {
- t.Parallel()
-
cfg, err := Load("")
if err != nil {
t.Fatalf("Load() error = %v", err)
@@ -63,8 +61,6 @@ comic:
}
func TestRenderPrompt(t *testing.T) {
- t.Parallel()
-
tmpDir := t.TempDir()
cfg := DefaultConfig()
cfg.PromptsDir = tmpDir
@@ -86,8 +82,6 @@ func TestRenderPrompt(t *testing.T) {
}
func TestRenderPromptMissingKeyReturnsError(t *testing.T) {
- t.Parallel()
-
tmpDir := t.TempDir()
cfg := DefaultConfig()
cfg.PromptsDir = tmpDir
@@ -106,8 +100,6 @@ func TestRenderPromptMissingKeyReturnsError(t *testing.T) {
}
func TestLoadRejectsUnknownProvider(t *testing.T) {
- t.Parallel()
-
tmpDir := t.TempDir()
configPath := filepath.Join(tmpDir, "config.yaml")
if err := os.WriteFile(configPath, []byte(strings.TrimSpace(`
@@ -127,8 +119,6 @@ provider:
}
func TestHomeDirReturnsFallbackWhenResolutionFails(t *testing.T) {
- t.Parallel()
-
oldUserHomeDir := userHomeDir
t.Cleanup(func() {
userHomeDir = oldUserHomeDir
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
+}