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

1""" 

2Resilience and retry logic for Osiris CLI 

3 

4Implements exponential backoff, circuit breaker pattern, 

5and request rate limiting for reliable API calls. 

6""" 

7 

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 

15 

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) 

24 

25from .errors import ( 

26 APITimeoutError, 

27 APIRateLimitError, 

28 APIServerError, 

29 APIError 

30) 

31from .logger import get_logger 

32 

33logger = get_logger() 

34 

35 

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. 

45  

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 

51  

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 

82 

83 return wrapper 

84 return decorator 

85 

86 

87class CircuitBreaker: 

88 """ 

89 Circuit breaker pattern for failing services. 

90  

91 States: 

92 - CLOSED: Normal operation, requests pass through 

93 - OPEN: Too many failures, requests blocked 

94 - HALF_OPEN: Testing if service recovered 

95 """ 

96 

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 

112 

113 self.failure_count = 0 

114 self.last_failure_time: Optional[datetime] = None 

115 self.state = "CLOSED" # CLOSED, OPEN, HALF_OPEN 

116 

117 def call(self, func: Callable, *args, **kwargs) -> Any: 

118 """ 

119 Execute function through circuit breaker. 

120  

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

131 

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 

139 

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

146 

147 def _on_failure(self): 

148 """Handle failed call""" 

149 self.failure_count += 1 

150 self.last_failure_time = datetime.now() 

151 

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 ) 

158 

159 

160class RateLimiter: 

161 """ 

162 Token bucket rate limiter for API calls. 

163  

164 Enforces rate limits like: 

165 - 60 requests per minute 

166 - 1000 requests per day 

167 """ 

168 

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 

177 

178 # Sliding window for minute tracking 

179 self.minute_window: deque = deque(maxlen=requests_per_minute) 

180 

181 # Counter for daily tracking 

182 self.day_start = datetime.now() 

183 self.daily_count = 0 

184 

185 async def acquire(self): 

186 """ 

187 Acquire permission to make request. 

188 Blocks if rate limit exceeded. 

189 """ 

190 now = datetime.now() 

191 

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 

196 

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 ) 

205 

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

211 

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) 

218 

219 # Record this request 

220 self.minute_window.append(now) 

221 self.daily_count += 1 

222 

223 def get_stats(self) -> Dict[str, Any]: 

224 """Get rate limiter statistics""" 

225 now = datetime.now() 

226 

227 # Count recent requests 

228 cutoff = now - timedelta(seconds=60) 

229 recent_count = sum(1 for ts in self.minute_window if ts > cutoff) 

230 

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 } 

239 

240 

241class ProviderFallback: 

242 """ 

243 Automatic fallback to alternative providers/models when primary fails. 

244 """ 

245 

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 } 

255 

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 } 

261 

262 def get_fallback_provider(self, failed_provider: str) -> Optional[str]: 

263 """ 

264 Get fallback provider when primary fails. 

265  

266 Args: 

267 failed_provider: Provider that failed 

268  

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 

278 

279 def get_fallback_model(self, failed_model: str) -> Optional[str]: 

280 """ 

281 Get fallback model when primary fails. 

282  

283 Args: 

284 failed_model: Model that failed 

285  

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 

293 

294 

295# Global instances 

296circuit_breakers: Dict[str, CircuitBreaker] = {} 

297rate_limiters: Dict[str, RateLimiter] = {} 

298fallback_manager = ProviderFallback() 

299 

300 

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] 

309 

310 

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]