1"""Circuit breaker pattern for provider health management."""
2
3from __future__ import annotations
4
5from dataclasses import dataclass
6from datetime import UTC, datetime, timedelta
7from enum import StrEnum
8
9from lexigram.logging import (
10 get_logger,
11)
12
13logger = get_logger(__name__)
14
15
16class CircuitState(StrEnum):
17 """Circuit breaker state enumeration."""
18
19 CLOSED = "closed" # Normal operation
20 OPEN = "open" # Failing, block requests
21 HALF_OPEN = "half_open" # Testing recovery
22
23
24@dataclass
25class CircuitBreakerConfig:
26 """Configuration for provider circuit breaker.
27
28 Attributes:
29 failure_threshold: Errors before opening circuit
30 success_threshold: Successes before closing from half-open
31 timeout_seconds: Time in open state before half-open
32 window_size: Number of recent requests to track
33 """
34
35 failure_threshold: int = 5
36 success_threshold: int = 2
37 timeout_seconds: int = 60
38 window_size: int = 20
39
40
41class ProviderCircuitBreaker:
42 """Implements circuit breaker pattern for provider failures.
43
44 Automatically stops requests to a provider after repeated failures,
45 and gradually recovers when it becomes healthy again.
46 """
47
48 def __init__(
49 self,
50 provider_name: str,
51 config: CircuitBreakerConfig | None = None,
52 ) -> None:
53 """Initialize circuit breaker.
54
55 Args:
56 provider_name: Provider identifier
57 config: Circuit breaker configuration
58 """
59 self.provider_name = provider_name
60 self.config = config or CircuitBreakerConfig()
61
62 self.state = CircuitState.CLOSED
63 self.failure_count = 0
64 self.success_count = 0
65 self.last_failure_time: datetime | None = None
66 self.opened_at: datetime | None = None
67 self.request_times: list[datetime] = []
68
69 async def record_success(self) -> None:
70 """Record a successful request.
71
72 May transition from HALF_OPEN → CLOSED.
73 """
74 now = datetime.now(UTC)
75 self.request_times.append(now)
76
77 # Prune old requests
78 cutoff = now - timedelta(seconds=self.config.timeout_seconds)
79 self.request_times = [t for t in self.request_times if t > cutoff]
80
81 if self.state == CircuitState.HALF_OPEN:
82 self.success_count += 1
83 if self.success_count >= self.config.success_threshold:
84 self._close()
85 elif self.state == CircuitState.CLOSED:
86 self.failure_count = 0
87
88 async def record_failure(self, error: Exception) -> None:
89 """Record a failed request.
90
91 May transition CLOSED → OPEN.
92
93 Args:
94 error: Exception that occurred
95 """
96 now = datetime.now(UTC)
97 self.last_failure_time = now
98 self.failure_count += 1
99
100 if self.failure_count >= self.config.failure_threshold:
101 if self.state != CircuitState.OPEN:
102 self._open(now)
103
104 logger.warning(
105 "provider_failure_recorded",
106 provider=self.provider_name,
107 failure_count=self.failure_count,
108 circuit_state=self.state.value,
109 error=str(error)[:100],
110 )
111
112 async def is_available(self) -> bool:
113 """Check if provider is available for requests.
114
115 Returns:
116 False if circuit is OPEN, True otherwise
117 """
118 if self.state == CircuitState.OPEN:
119 # Check if we should transition to HALF_OPEN
120 if self.opened_at:
121 elapsed = datetime.now(UTC) - self.opened_at
122 if elapsed.total_seconds() >= self.config.timeout_seconds:
123 self._half_open()
124 return True # Allow test request
125 return False
126
127 return True
128
129 def _open(self, now: datetime) -> None:
130 """Transition to OPEN state."""
131 self.state = CircuitState.OPEN
132 self.opened_at = now
133 self.success_count = 0
134
135 logger.error(
136 "circuit_breaker_opened",
137 provider=self.provider_name,
138 failure_count=self.failure_count,
139 )
140
141 def _half_open(self) -> None:
142 """Transition to HALF_OPEN state."""
143 self.state = CircuitState.HALF_OPEN
144 self.failure_count = 0
145 self.success_count = 0
146
147 logger.info(
148 "circuit_breaker_half_open",
149 provider=self.provider_name,
150 )
151
152 def _close(self) -> None:
153 """Transition to CLOSED state."""
154 self.state = CircuitState.CLOSED
155 self.failure_count = 0
156 self.success_count = 0
157 self.opened_at = None
158
159 logger.info(
160 "circuit_breaker_closed",
161 provider=self.provider_name,
162 )
163
164 def reset(self) -> None:
165 """Reset circuit breaker to initial state."""
166 self.state = CircuitState.CLOSED
167 self.failure_count = 0
168 self.success_count = 0
169 self.last_failure_time = None
170 self.opened_at = None
171 self.request_times = []