Coverage for src / osiris_cli / resilience.py: 0%
115 statements
« prev ^ index » next coverage.py v7.13.0, created at 2025-12-31 05:01 +0200
« prev ^ index » next coverage.py v7.13.0, created at 2025-12-31 05:01 +0200
1"""
2Resilience and retry logic for Osiris CLI
4Implements exponential backoff, circuit breaker pattern,
5and request rate limiting for reliable API calls.
6"""
8import asyncio
9import time
10import logging
11from typing import Callable, Any, Optional, Dict, List
12from functools import wraps
13from datetime import datetime, timedelta
14from collections import deque
16from tenacity import (
17 retry,
18 stop_after_attempt,
19 wait_exponential,
20 retry_if_exception_type,
21 before_sleep_log,
22 after_log
23)
25from .errors import (
26 APITimeoutError,
27 APIRateLimitError,
28 APIServerError,
29 APIError
30)
31from .logger import get_logger
33logger = get_logger()
36# Retry decorator for API calls
37def resilient_api_call(
38 max_attempts: int = 3,
39 min_wait: float = 1.0,
40 max_wait: float = 10.0,
41 multiplier: float = 2.0
42):
43 """
44 Decorator for resilient API calls with exponential backoff.
46 Args:
47 max_attempts: Maximum number of retry attempts
48 min_wait: Minimum wait time between retries (seconds)
49 max_wait: Maximum wait time between retries (seconds)
50 multiplier: Multiplier for exponential backoff
52 Example:
53 @resilient_api_call(max_attempts=3)
54 async def call_api():
55 ...
56 """
57 def decorator(func: Callable):
58 @retry(
59 stop=stop_after_attempt(max_attempts),
60 wait=wait_exponential(
61 multiplier=multiplier,
62 min=min_wait,
63 max=max_wait
64 ),
65 retry=retry_if_exception_type((
66 APITimeoutError,
67 APIRateLimitError,
68 APIServerError,
69 ConnectionError,
70 TimeoutError
71 )),
72 before_sleep=before_sleep_log(logger, logging.WARNING),
73 after=after_log(logger, logging.INFO)
74 )
75 @wraps(func)
76 async def wrapper(*args, **kwargs):
77 try:
78 return await func(*args, **kwargs)
79 except Exception as e:
80 logger.error(f"API call failed after {max_attempts} attempts: {e}")
81 raise
83 return wrapper
84 return decorator
87class CircuitBreaker:
88 """
89 Circuit breaker pattern for failing services.
91 States:
92 - CLOSED: Normal operation, requests pass through
93 - OPEN: Too many failures, requests blocked
94 - HALF_OPEN: Testing if service recovered
95 """
97 def __init__(
98 self,
99 failure_threshold: int = 5,
100 recovery_timeout: int = 60,
101 expected_exception: type = Exception
102 ):
103 """
104 Args:
105 failure_threshold: Number of failures before opening circuit
106 recovery_timeout: Seconds to wait before trying again
107 expected_exception: Exception type to catch
108 """
109 self.failure_threshold = failure_threshold
110 self.recovery_timeout = recovery_timeout
111 self.expected_exception = expected_exception
113 self.failure_count = 0
114 self.last_failure_time: Optional[datetime] = None
115 self.state = "CLOSED" # CLOSED, OPEN, HALF_OPEN
117 def call(self, func: Callable, *args, **kwargs) -> Any:
118 """
119 Execute function through circuit breaker.
121 Raises:
122 Exception: If circuit is open or function fails
123 """
124 if self.state == "OPEN":
125 # Check if we should try recovery
126 if datetime.now() - self.last_failure_time > timedelta(seconds=self.recovery_timeout):
127 self.state = "HALF_OPEN"
128 logger.info("Circuit breaker entering HALF_OPEN state")
129 else:
130 raise Exception(f"Circuit breaker is OPEN, blocking requests for {self.recovery_timeout}s")
132 try:
133 result = func(*args, **kwargs)
134 self._on_success()
135 return result
136 except self.expected_exception as e:
137 self._on_failure()
138 raise
140 def _on_success(self):
141 """Handle successful call"""
142 self.failure_count = 0
143 if self.state == "HALF_OPEN":
144 self.state = "CLOSED"
145 logger.info("Circuit breaker CLOSED (service recovered)")
147 def _on_failure(self):
148 """Handle failed call"""
149 self.failure_count += 1
150 self.last_failure_time = datetime.now()
152 if self.failure_count >= self.failure_threshold:
153 self.state = "OPEN"
154 logger.warning(
155 f"Circuit breaker OPEN after {self.failure_count} failures, "
156 f"blocking for {self.recovery_timeout}s"
157 )
160class RateLimiter:
161 """
162 Token bucket rate limiter for API calls.
164 Enforces rate limits like:
165 - 60 requests per minute
166 - 1000 requests per day
167 """
169 def __init__(self, requests_per_minute: int = 60, requests_per_day: int = 1000):
170 """
171 Args:
172 requests_per_minute: Maximum requests per minute
173 requests_per_day: Maximum requests per day
174 """
175 self.requests_per_minute = requests_per_minute
176 self.requests_per_day = requests_per_day
178 # Sliding window for minute tracking
179 self.minute_window: deque = deque(maxlen=requests_per_minute)
181 # Counter for daily tracking
182 self.day_start = datetime.now()
183 self.daily_count = 0
185 async def acquire(self):
186 """
187 Acquire permission to make request.
188 Blocks if rate limit exceeded.
189 """
190 now = datetime.now()
192 # Reset daily counter if new day
193 if now - self.day_start > timedelta(days=1):
194 self.day_start = now
195 self.daily_count = 0
197 # Check daily limit
198 if self.daily_count >= self.requests_per_day:
199 wait_seconds = (timedelta(days=1) - (now - self.day_start)).total_seconds()
200 logger.warning(f"Daily rate limit reached, waiting {wait_seconds:.0f}s")
201 raise APIRateLimitError(
202 provider="global",
203 retry_after=int(wait_seconds)
204 )
206 # Check per-minute limit
207 # Remove old entries (>60s ago)
208 cutoff = now - timedelta(seconds=60)
209 while self.minute_window and self.minute_window[0] < cutoff:
210 self.minute_window.popleft()
212 if len(self.minute_window) >= self.requests_per_minute:
213 # Calculate wait time until oldest entry expires
214 wait_seconds = 60 - (now - self.minute_window[0]).total_seconds()
215 if wait_seconds > 0:
216 logger.info(f"Per-minute rate limit reached, waiting {wait_seconds:.1f}s")
217 await asyncio.sleep(wait_seconds)
219 # Record this request
220 self.minute_window.append(now)
221 self.daily_count += 1
223 def get_stats(self) -> Dict[str, Any]:
224 """Get rate limiter statistics"""
225 now = datetime.now()
227 # Count recent requests
228 cutoff = now - timedelta(seconds=60)
229 recent_count = sum(1 for ts in self.minute_window if ts > cutoff)
231 return {
232 "requests_last_minute": recent_count,
233 "requests_today": self.daily_count,
234 "limit_per_minute": self.requests_per_minute,
235 "limit_per_day": self.requests_per_day,
236 "minute_remaining": self.requests_per_minute - recent_count,
237 "day_remaining": self.requests_per_day - self.daily_count
238 }
241class ProviderFallback:
242 """
243 Automatic fallback to alternative providers/models when primary fails.
244 """
246 def __init__(self):
247 """Initialize fallback manager"""
248 self.fallback_chains: Dict[str, List[str]] = {
249 "openai": ["anthropic", "groq", "deepseek"],
250 "anthropic": ["openai", "groq", "deepseek"],
251 "google": ["openai", "anthropic", "groq"],
252 "groq": ["deepseek", "openai", "anthropic"],
253 "deepseek": ["groq", "anthropic", "openai"],
254 }
256 self.model_fallbacks: Dict[str, str] = {
257 "gpt-4": "gpt-3.5-turbo",
258 "claude-3-5-sonnet": "claude-3-haiku",
259 "gemini-pro": "gemini-1.5-flash",
260 }
262 def get_fallback_provider(self, failed_provider: str) -> Optional[str]:
263 """
264 Get fallback provider when primary fails.
266 Args:
267 failed_provider: Provider that failed
269 Returns:
270 Fallback provider name or None
271 """
272 chain = self.fallback_chains.get(failed_provider, [])
273 if chain:
274 fallback = chain[0]
275 logger.info(f"Falling back from {failed_provider} to {fallback}")
276 return fallback
277 return None
279 def get_fallback_model(self, failed_model: str) -> Optional[str]:
280 """
281 Get fallback model when primary fails.
283 Args:
284 failed_model: Model that failed
286 Returns:
287 Fallback model name or None
288 """
289 fallback = self.model_fallbacks.get(failed_model)
290 if fallback:
291 logger.info(f"Falling back from model {failed_model} to {fallback}")
292 return fallback
295# Global instances
296circuit_breakers: Dict[str, CircuitBreaker] = {}
297rate_limiters: Dict[str, RateLimiter] = {}
298fallback_manager = ProviderFallback()
301def get_circuit_breaker(provider: str) -> CircuitBreaker:
302 """Get or create circuit breaker for provider"""
303 if provider not in circuit_breakers:
304 circuit_breakers[provider] = CircuitBreaker(
305 failure_threshold=5,
306 recovery_timeout=60
307 )
308 return circuit_breakers[provider]
311def get_rate_limiter(provider: str) -> RateLimiter:
312 """Get or create rate limiter for provider"""
313 if provider not in rate_limiters:
314 # Provider-specific limits (can be configured)
315 limits = {
316 "openai": (60, 10000),
317 "anthropic": (50, 5000),
318 "google": (60, 1000),
319 "groq": (30, 14400),
320 "deepseek": (60, 1000),
321 }
322 rpm, rpd = limits.get(provider, (60, 1000))
323 rate_limiters[provider] = RateLimiter(rpm, rpd)
324 return rate_limiters[provider]