100 lines
3.1 KiB
Python
100 lines
3.1 KiB
Python
"""Fallback chain for AI providers.
|
|
|
|
Tries providers in order with retries per provider.
|
|
YandexGPT -> GigaChat -> Error
|
|
"""
|
|
|
|
import asyncio
|
|
import time
|
|
from typing import Any
|
|
|
|
from app.core.config import get_settings
|
|
from app.integrations.ai.base import AIProvider, AIResult
|
|
from app.integrations.ai.yandex_gpt import YandexGPTProvider
|
|
from app.integrations.ai.gigachat import GigaChatProvider
|
|
|
|
settings = get_settings()
|
|
|
|
|
|
class FallbackChain:
|
|
"""Chain multiple AI providers with retry logic.
|
|
|
|
Tries providers in order (YandexGPT -> GigaChat).
|
|
Each provider gets up to `max_retries` attempts before falling through.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
providers: list[AIProvider] | None = None,
|
|
max_retries: int = 2,
|
|
timeout: int = 30,
|
|
):
|
|
self.providers = providers or [
|
|
YandexGPTProvider(),
|
|
GigaChatProvider(),
|
|
]
|
|
self.max_retries = max_retries
|
|
self.timeout = timeout
|
|
|
|
async def analyze(self, prompt: str, **kwargs: Any) -> AIResult:
|
|
"""Try each provider in order with retries.
|
|
|
|
Args:
|
|
prompt: The formatted prompt to send
|
|
**kwargs: Additional parameters passed to each provider
|
|
|
|
Returns:
|
|
AIResult from the first successful provider,
|
|
or the last error result if all fail.
|
|
"""
|
|
last_error: AIResult | None = None
|
|
|
|
for provider in self.providers:
|
|
for attempt in range(self.max_retries + 1):
|
|
try:
|
|
result = await asyncio.wait_for(
|
|
provider.analyze(prompt, **kwargs),
|
|
timeout=self.timeout,
|
|
)
|
|
|
|
if result.success:
|
|
await self._cleanup_providers(provider)
|
|
return result
|
|
|
|
last_error = result
|
|
|
|
except asyncio.TimeoutError:
|
|
last_error = AIResult(
|
|
success=False,
|
|
error=f"{provider.name} timeout after {self.timeout}s",
|
|
provider=provider.name,
|
|
duration_ms=int(self.timeout * 1000),
|
|
)
|
|
except Exception as e:
|
|
last_error = AIResult(
|
|
success=False,
|
|
error=f"{provider.name} error: {str(e)}",
|
|
provider=provider.name,
|
|
)
|
|
|
|
if attempt < self.max_retries:
|
|
await asyncio.sleep(1 * (attempt + 1))
|
|
|
|
await self._cleanup_providers()
|
|
return last_error or AIResult(
|
|
success=False,
|
|
error="All providers failed",
|
|
provider="fallback_chain",
|
|
)
|
|
|
|
async def _cleanup_providers(self, skip: AIProvider | None = None) -> None:
|
|
"""Close all provider clients except the one to keep."""
|
|
for provider in self.providers:
|
|
if provider is skip:
|
|
continue
|
|
if hasattr(provider, "close"):
|
|
try:
|
|
await provider.close()
|
|
except Exception:
|
|
pass
|