diff options
| -rw-r--r-- | internal/apicircuit/apicircuit.go | 85 | ||||
| -rw-r--r-- | internal/apicircuit/apicircuit_test.go | 19 | ||||
| -rw-r--r-- | internal/config/config_test.go | 10 | ||||
| -rw-r--r-- | internal/provider/provider.go | 86 | ||||
| -rw-r--r-- | internal/provider/provider_test.go | 66 |
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 +} |
