FreshPixels commited on
Commit
649caf1
·
verified ·
1 Parent(s): 4ee5306

Create llm/manager.py

Browse files
Files changed (1) hide show
  1. llm/manager.py +395 -0
llm/manager.py ADDED
@@ -0,0 +1,395 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ from dataclasses import dataclass, field
5
+ from typing import Any
6
+
7
+ from core.logging.logger import get_logger
8
+ from llm.providers.base_provider import BaseProvider, LLMProviderError, LLMRequest, LLMResponse
9
+
10
+ logger = get_logger(__name__)
11
+
12
+
13
+ # ---------------------------------------------------------------------------
14
+ # Exceptions
15
+ # ---------------------------------------------------------------------------
16
+
17
+ class LLMManagerError(Exception):
18
+ """Base exception for LLMManager errors."""
19
+
20
+
21
+ class ConfigurationError(LLMManagerError):
22
+ """Raised when LLMManager is misconfigured."""
23
+
24
+
25
+ class ProviderNotFoundError(LLMManagerError):
26
+ """Raised when a requested provider is not registered."""
27
+
28
+
29
+ class ProviderAlreadyRegisteredError(LLMManagerError):
30
+ """Raised when a provider with the same name is already registered."""
31
+
32
+
33
+ class AllProvidersFailedError(LLMManagerError):
34
+ """Raised when all fallback providers have been exhausted."""
35
+
36
+ def __init__(self, attempts: list[dict[str, Any]]) -> None:
37
+ self.attempts = attempts
38
+ provider_names = [a["provider"] for a in attempts]
39
+ errors = [a["error"] for a in attempts]
40
+ super().__init__(
41
+ f"All providers failed: {provider_names}. Errors: {errors}"
42
+ )
43
+
44
+
45
+ # ---------------------------------------------------------------------------
46
+ # Fallback Policy
47
+ # ---------------------------------------------------------------------------
48
+
49
+ @dataclass(frozen=True)
50
+ class FallbackPolicy:
51
+ """Determines which providers to try and in what order.
52
+
53
+ If ``providers`` is empty, the manager falls back to iterating over all
54
+ registered providers in insertion order.
55
+
56
+ Attributes:
57
+ providers: Ordered list of provider names to try.
58
+ max_retries_per_provider: How many times to retry each provider
59
+ before moving to the next one.
60
+ base_delay_seconds: Base delay for exponential backoff between
61
+ retries (seconds).
62
+ max_delay_seconds: Maximum delay between retries (seconds).
63
+ """
64
+
65
+ providers: list[str] = field(default_factory=list)
66
+ max_retries_per_provider: int = 2
67
+ base_delay_seconds: float = 1.0
68
+ max_delay_seconds: float = 30.0
69
+
70
+
71
+ # ---------------------------------------------------------------------------
72
+ # LLMManager
73
+ # ---------------------------------------------------------------------------
74
+
75
+ class LLMManager:
76
+ """Registry and dispatcher for LLM providers.
77
+
78
+ Implements the Strategy pattern with retry and fallback:
79
+
80
+ - ``generate()`` delegates to the specified provider (or default).
81
+ - On ``LLMProviderError``, retries with exponential backoff.
82
+ - When retries are exhausted, falls back to the next provider in the
83
+ ``FallbackPolicy``.
84
+ - Thread-safe provider registration via ``asyncio.Lock``.
85
+
86
+ Usage::
87
+
88
+ manager = LLMManager(default_provider="openai")
89
+ manager.register_provider("openai", openai_provider)
90
+ manager.register_provider("claude", claude_provider)
91
+
92
+ response = await manager.generate(request, provider_name="openai")
93
+ """
94
+
95
+ def __init__(
96
+ self,
97
+ default_provider: str = "openai",
98
+ fallback_policy: FallbackPolicy | None = None,
99
+ ) -> None:
100
+ self._providers: dict[str, BaseProvider] = {}
101
+ self._default_provider: str = default_provider
102
+ self._fallback_policy: FallbackPolicy = fallback_policy or FallbackPolicy()
103
+ self._lock: asyncio.Lock = asyncio.Lock()
104
+
105
+ async def register_provider(
106
+ self, name: str, provider: BaseProvider, overwrite: bool = False
107
+ ) -> None:
108
+ """Register an LLM provider under the given name.
109
+
110
+ Args:
111
+ name: Unique provider identifier (e.g. "openai", "claude").
112
+ provider: Provider instance implementing BaseProvider.
113
+ overwrite: If True, silently replaces an existing provider.
114
+ If False, raises ProviderAlreadyRegisteredError.
115
+
116
+ Raises:
117
+ ProviderAlreadyRegisteredError: If name is already registered
118
+ and overwrite is False.
119
+ """
120
+ async with self._lock:
121
+ if name in self._providers and not overwrite:
122
+ raise ProviderAlreadyRegisteredError(
123
+ f"Provider '{name}' is already registered. "
124
+ f"Use overwrite=True to replace it."
125
+ )
126
+
127
+ self._providers[name] = provider
128
+
129
+ logger.info(
130
+ "provider_registered",
131
+ provider_name=name,
132
+ provider_class=provider.__class__.__name__,
133
+ )
134
+
135
+ def get_provider(self, name: str | None = None) -> BaseProvider:
136
+ """Retrieve a registered provider by name.
137
+
138
+ If name is None, returns the default provider.
139
+
140
+ Args:
141
+ name: Provider identifier, or None for the default.
142
+
143
+ Returns:
144
+ The registered BaseProvider instance.
145
+
146
+ Raises:
147
+ ProviderNotFoundError: If the requested provider is not found.
148
+ """
149
+ lookup_name = name or self._default_provider
150
+
151
+ provider = self._providers.get(lookup_name)
152
+ if provider is None:
153
+ available = list(self._providers.keys())
154
+ raise ProviderNotFoundError(
155
+ f"Provider '{lookup_name}' not found. "
156
+ f"Available providers: {available}"
157
+ )
158
+
159
+ return provider
160
+
161
+ def _resolve_fallback_order(self, explicit_name: str | None = None) -> list[str]:
162
+ """Determine the ordered list of provider names to attempt.
163
+
164
+ If an explicit provider name is given, start with it and then
165
+ fall back through the policy (or all registered providers).
166
+
167
+ If no explicit name, start with the default provider.
168
+ """
169
+ policy_names = (
170
+ list(self._fallback_policy.providers)
171
+ if self._fallback_policy.providers
172
+ else list(self._providers.keys())
173
+ )
174
+
175
+ primary = explicit_name or self._default_provider
176
+
177
+ # Build ordered list: primary first, then policy order (excluding primary)
178
+ order: list[str] = []
179
+ if primary in self._providers:
180
+ order.append(primary)
181
+
182
+ for name in policy_names:
183
+ if name not in order and name in self._providers:
184
+ order.append(name)
185
+
186
+ return order
187
+
188
+ async def _attempt_with_retries(
189
+ self,
190
+ provider: BaseProvider,
191
+ request: LLMRequest,
192
+ max_retries: int,
193
+ base_delay: float,
194
+ max_delay: float,
195
+ ) -> LLMResponse:
196
+ """Call provider.generate() with exponential backoff retries.
197
+
198
+ Args:
199
+ provider: The provider instance.
200
+ request: The LLM request.
201
+ max_retries: Maximum number of retries after the first attempt.
202
+ base_delay: Base delay for exponential backoff (seconds).
203
+ max_delay: Maximum delay between retries (seconds).
204
+
205
+ Returns:
206
+ LLMResponse on success.
207
+
208
+ Raises:
209
+ LLMProviderError: If all retries are exhausted.
210
+ """
211
+ last_error: LLMProviderError | None = None
212
+ total_attempts = max_retries + 1
213
+
214
+ for attempt in range(total_attempts):
215
+ try:
216
+ response = await provider.generate(request)
217
+
218
+ if attempt > 0:
219
+ logger.info(
220
+ "llm_generate_retry_succeeded",
221
+ provider=provider.name,
222
+ attempt=attempt + 1,
223
+ total_attempts=total_attempts,
224
+ )
225
+
226
+ return response
227
+
228
+ except LLMProviderError as exc:
229
+ last_error = exc
230
+ if attempt < max_retries:
231
+ delay = min(base_delay * (2 ** attempt), max_delay)
232
+ logger.warning(
233
+ "llm_generate_retry",
234
+ provider=provider.name,
235
+ attempt=attempt + 1,
236
+ max_retries=max_retries + 1,
237
+ retry_in_seconds=delay,
238
+ error=str(exc),
239
+ )
240
+ await asyncio.sleep(delay)
241
+ else:
242
+ logger.error(
243
+ "llm_generate_exhausted_retries",
244
+ provider=provider.name,
245
+ total_attempts=total_attempts,
246
+ error=str(exc),
247
+ )
248
+
249
+ # Should never reach here, but for type safety
250
+ raise last_error or LLMProviderError( # type: ignore[misc]
251
+ "Unexpected retry loop exit",
252
+ provider_name=provider.name,
253
+ )
254
+
255
+ async def generate(
256
+ self,
257
+ request: LLMRequest,
258
+ provider_name: str | None = None,
259
+ ) -> LLMResponse:
260
+ """Generate a response using the specified or default provider.
261
+
262
+ Implements retry with exponential backoff per provider, then
263
+ fallback to the next provider in the FallbackPolicy order.
264
+
265
+ Args:
266
+ request: The LLM request to send.
267
+ provider_name: Optional provider to use; falls back to default.
268
+
269
+ Returns:
270
+ The LLMResponse from the provider.
271
+
272
+ Raises:
273
+ ProviderNotFoundError: If no providers are registered
274
+ at all.
275
+ AllProvidersFailedError: If all providers fail after retries.
276
+ """
277
+ fallback_order = self._resolve_fallback_order(provider_name)
278
+
279
+ if not fallback_order:
280
+ raise ProviderNotFoundError(
281
+ f"No providers registered. "
282
+ f"Requested: '{provider_name or self._default_provider}'"
283
+ )
284
+
285
+ policy = self._fallback_policy
286
+ attempts: list[dict[str, Any]] = []
287
+
288
+ for name in fallback_order:
289
+ provider = self._providers.get(name)
290
+ if provider is None:
291
+ logger.debug(
292
+ "llm_generate_skip_missing_provider",
293
+ provider_name=name,
294
+ )
295
+ continue
296
+
297
+ logger.debug(
298
+ "llm_generate_attempt",
299
+ provider=provider.name,
300
+ model=request.model,
301
+ messages_count=len(request.messages),
302
+ attempt=len(attempts) + 1,
303
+ total_providers=len(fallback_order),
304
+ )
305
+
306
+ try:
307
+ response = await self._attempt_with_retries(
308
+ provider=provider,
309
+ request=request,
310
+ max_retries=policy.max_retries_per_provider,
311
+ base_delay=policy.base_delay_seconds,
312
+ max_delay=policy.max_delay_seconds,
313
+ )
314
+
315
+ logger.debug(
316
+ "llm_generate_complete",
317
+ provider=provider.name,
318
+ model=response.model,
319
+ content_length=len(response.content),
320
+ )
321
+
322
+ return response
323
+
324
+ except LLMProviderError as exc:
325
+ attempts.append({
326
+ "provider": name,
327
+ "error": str(exc),
328
+ })
329
+ logger.warning(
330
+ "llm_generate_provider_failed",
331
+ provider=name,
332
+ error=str(exc),
333
+ remaining_providers=[
334
+ n for n in fallback_order
335
+ if n not in [a["provider"] for a in attempts]
336
+ ],
337
+ )
338
+
339
+ raise AllProvidersFailedError(attempts)
340
+
341
+ def list_providers(self) -> list[str]:
342
+ """Return a list of all registered provider names."""
343
+ return list(self._providers.keys())
344
+
345
+ @property
346
+ def default_provider(self) -> str:
347
+ """Return the name of the default provider."""
348
+ return self._default_provider
349
+
350
+ @default_provider.setter
351
+ def default_provider(self, name: str, strict: bool = False) -> None:
352
+ """Set the default provider name.
353
+
354
+ Args:
355
+ name: Provider name to set as default.
356
+ strict: If True, raises ValueError when the provider is not
357
+ registered. If False, logs a warning but allows the assignment.
358
+
359
+ Raises:
360
+ ValueError: In strict mode, when provider is not registered.
361
+ """
362
+ if name not in self._providers:
363
+ if strict:
364
+ raise ValueError(
365
+ f"Cannot set default_provider to '{name}': "
366
+ f"provider not registered. "
367
+ f"Available: {self.list_providers()}"
368
+ )
369
+ logger.warning(
370
+ "default_provider_set_to_unregistered",
371
+ provider_name=name,
372
+ available=list(self._providers.keys()),
373
+ )
374
+ self._default_provider = name
375
+
376
+ def has_providers(self) -> bool:
377
+ """Return True if at least one provider is registered."""
378
+ return len(self._providers) > 0
379
+
380
+ @property
381
+ def fallback_policy(self) -> FallbackPolicy:
382
+ """Return the current fallback policy."""
383
+ return self._fallback_policy
384
+
385
+ @fallback_policy.setter
386
+ def fallback_policy(self, policy: FallbackPolicy) -> None:
387
+ """Set a new fallback policy."""
388
+ self._fallback_policy = policy
389
+
390
+ def __repr__(self) -> str:
391
+ return (
392
+ f"<LLMManager providers={self.list_providers()} "
393
+ f"default='{self._default_provider}'>"
394
+ )
395
+