summaryrefslogtreecommitdiff
path: root/internal/apicircuit/apicircuit.go
blob: 7f623cb34eb096fc19430fd915341ad9a6914fda (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
129
130
131
// 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"

	"github.com/sony/gobreaker"
)

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
	// breakerOpenTimeout is how long the breaker stays open before trying half-open.
	breakerOpenTimeout = 45 * time.Second
	// breakerMaxHalfOpenRequests limits trial traffic while recovering.
	breakerMaxHalfOpenRequests = 3
	// breakerTripAfterConsecutiveFailures opens the circuit after this many
	// consecutive failed requests in the closed state.
	breakerTripAfterConsecutiveFailures uint32 = 5
)

// 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
// remote errors still count as failures.
func isSuccessful(err error) bool {
	if err == nil {
		return true
	}

	return errors.Is(err, context.Canceled)
}

func readyToTrip(counts gobreaker.Counts) bool {
	return counts.ConsecutiveFailures >= breakerTripAfterConsecutiveFailures
}

func newBreaker(name string) *gobreaker.CircuitBreaker {
	return gobreaker.NewCircuitBreaker(gobreaker.Settings{
		Name:         name,
		MaxRequests:  breakerMaxHalfOpenRequests,
		Interval:     breakerInterval,
		Timeout:      breakerOpenTimeout,
		ReadyToTrip:  readyToTrip,
		IsSuccessful: isSuccessful,
	})
}

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()

	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
}

// 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 := r.breaker(providerName, capability).Execute(func() (interface{}, error) {
		return fn()
	})
	if err != nil {
		return zero, err
	}
	if v == nil {
		return zero, nil
	}

	result, ok := v.(T)
	if !ok {
		return zero, fmt.Errorf("unexpected breaker result type for %s", breakerKey(providerName, capability))
	}

	return result, nil
}

// NormalizeName returns a canonical lower-case name for breaker keys.
func NormalizeName(name string) string {
	return strings.ToLower(strings.TrimSpace(name))
}