Team Ai
Apppublic

LetsTryGPT/agent4-implementation

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
fallback.ts464 linesDownload Raw Back to llm
1import { LLMProvider } from './providers/base';2import { HuggingFaceProvider } from './providers/huggingface';3import { MistralProvider } from './providers/mistral';4import { DeepSeekProvider } from './providers/deepseek';5import { OpenRouterProvider } from './providers/openrouter';6import { CodestralProvider } from './providers/codestral';7import { config } from '../config/index';8import { logger, ErrorHandler, llmCache } from '../utils';9 10// Constants11const HEALTH_CHECK_INTERVAL_MS = 5 * 60 * 1000; // 5 minutes12const PROVIDER_TIMEOUT_MS = 30000; // 30 seconds timeout per provider13const MAX_RETRY_ATTEMPTS = 2; // Retry failed requests 2 times14const CIRCUIT_BREAKER_THRESHOLD = 5; // Open circuit after 5 consecutive failures15const CIRCUIT_BREAKER_RESET_MS = 60000; // Reset circuit breaker after 1 minute16 17// Provider configurations with all available models18const PROVIDER_CONFIG = {19  huggingface: {20    model: 'mistralai/Mistral-7B-Instruct-v0.1',21    apiUrl: 'https://api-inference.huggingface.co/models',22    envVar: 'HF_TOKEN',23  },24  mistral: {25    model: 'mistral-small-latest',26    apiUrl: 'https://api.mistral.ai/v1/chat/completions',27    envVar: 'MISTRAL_API_KEY',28  },29  deepseek: {30    model: 'deepseek-coder-33b-instruct',31    apiUrl: 'https://api.deepseek.com/v1/chat/completions',32    envVar: 'DEEPSEEK_API_KEY',33  },34  openrouter: {35    model: 'mistralai/mistral-7b-instruct',36    apiUrl: 'https://openrouter.ai/api/v1/chat/completions',37    envVar: 'OPENROUTER_API_KEY',38    headers: {39      'HTTP-Referer': 'https://github.com/NovusAevum/agent4-implementation',40      'X-Title': 'Agent4 Implementation',41    },42  },43  codestral: {44    model: 'codestral-latest',45    apiUrl: 'https://api.mistral.ai/v1/chat/completions',46    envVar: 'MISTRAL_API_KEY', // Reusing Mistral's API key as per their documentation47  },48} as const;49 50type ProviderName = keyof typeof PROVIDER_CONFIG;51 52interface ProviderInfo {53  name: ProviderName;54  provider: LLMProvider;55  priority: number;56  isHealthy: boolean;57  lastError: Error | null;58  errorCount: number;59  lastUsed: number;60  totalRequests: number;61  failedRequests: number;62  circuitBreakerOpen: boolean;63  circuitBreakerOpenedAt: number | null;64  consecutiveFailures: number;65}66 67export class FallbackLLM {68  private providers: ProviderInfo[] = [];69  private initialized = false;70  private initializationPromise: Promise<void> | null = null;71  private lastError: Error | null = null;72  private healthCheckInterval: ReturnType<typeof setInterval> | null = null;73 74  constructor() {75    this.initialize().catch((error) => {76      const formattedError = ErrorHandler.format(error);77      logger.error('Failed to initialize FallbackLLM', formattedError);78      this.lastError = formattedError;79      // Ensure no health check interval is running if initialization failed80      this.destroy();81    });82  }83 84  private async initialize(): Promise<void> {85    if (this.initialized) return;86 87    if (!this.initializationPromise) {88      this.initializationPromise = this.initializeProviders()89        .then(() => {90          this.initialized = true;91          this.startHealthChecks();92        })93        .catch((error) => {94          logger.error('Failed to initialize providers', ErrorHandler.format(error));95          throw error;96        });97    }98 99    await this.initializationPromise;100  }101 102  private startHealthChecks(): void {103    // Run health check every 5 minutes104    // Use unref() to allow process to exit gracefully even if interval is active105    this.healthCheckInterval = setInterval(async () => {106      await this.checkAllProvidersHealth();107    }, HEALTH_CHECK_INTERVAL_MS);108    this.healthCheckInterval.unref();109  }110 111  /**112   * Stop health checks and clean up resources113   * Call this before destroying the FallbackLLM instance to prevent memory leaks114   */115  destroy(): void {116    if (this.healthCheckInterval) {117      clearInterval(this.healthCheckInterval);118      this.healthCheckInterval = null;119    }120  }121 122  private async checkAllProvidersHealth(): Promise<void> {123    await Promise.all(124      this.providers.map(async (providerInfo) => {125        try {126          const isHealthy = await providerInfo.provider.checkHealth();127          providerInfo.isHealthy = isHealthy;128          if (isHealthy) {129            providerInfo.lastError = null;130            logger.debug('Provider health check passed', { provider: providerInfo.name });131          }132        } catch (error) {133          providerInfo.isHealthy = false;134          providerInfo.lastError = ErrorHandler.format(error);135          logger.warn('Provider health check failed', {136            provider: providerInfo.name,137            error: ErrorHandler.getMessage(error),138          });139        }140      })141    );142  }143 144  private async initializeProviders(): Promise<void> {145    try {146      // Get the list of providers in order of preference147      const providerOrder = (148        Array.isArray(config.FALLBACK_ORDER)149          ? config.FALLBACK_ORDER150          : String(config.FALLBACK_ORDER)151              .split(',')152              .map((p: string) => p.trim().toLowerCase() as ProviderName)153      ).filter((name): name is ProviderName => name in PROVIDER_CONFIG);154 155      // If no valid providers are specified, use a default fallback order156      const effectiveProviderOrder = providerOrder.length > 0 ? providerOrder : ['huggingface'];157 158      // Initialize providers with their respective configurations159      const providers = await Promise.all(160        effectiveProviderOrder.map(async (providerName, index) => {161          try {162            const providerConfig = PROVIDER_CONFIG[providerName as keyof typeof PROVIDER_CONFIG];163            const apiKey = config[providerConfig.envVar as keyof typeof config] as string;164 165            // In test/development, allow test keys; in production, require real keys166            if (!apiKey) {167              logger.warn('No API key found for provider', { provider: providerName });168              return null;169            }170            if (config.NODE_ENV === 'production' && apiKey.startsWith('test-')) {171              logger.warn('Test API key not allowed in production', { provider: providerName });172              return null;173            }174 175            let provider: LLMProvider;176 177            switch (providerName) {178              case 'huggingface':179                const hfConfig = providerConfig as (typeof PROVIDER_CONFIG)['huggingface'];180                provider = new HuggingFaceProvider(apiKey, hfConfig.model, hfConfig.apiUrl);181                break;182 183              case 'mistral':184                const mistralConfig = providerConfig as (typeof PROVIDER_CONFIG)['mistral'];185                provider = new MistralProvider(apiKey, mistralConfig.model, mistralConfig.apiUrl);186                break;187 188              case 'deepseek':189                const deepseekConfig = providerConfig as (typeof PROVIDER_CONFIG)['deepseek'];190                provider = new DeepSeekProvider(191                  apiKey,192                  deepseekConfig.model,193                  deepseekConfig.apiUrl194                );195                break;196 197              case 'openrouter':198                const openrouterConfig = providerConfig as (typeof PROVIDER_CONFIG)['openrouter'];199                provider = new OpenRouterProvider(200                  apiKey,201                  openrouterConfig.model,202                  openrouterConfig.apiUrl203                );204                break;205 206              case 'codestral':207                const codestralConfig = providerConfig as (typeof PROVIDER_CONFIG)['codestral'];208                provider = new CodestralProvider(209                  apiKey,210                  codestralConfig.model,211                  codestralConfig.apiUrl212                );213                break;214 215              default:216                logger.warn('Provider not yet implemented', { provider: providerName });217                return null;218            }219 220            const isHealthy = await provider.checkHealth().catch(() => false);221 222            return {223              name: providerName,224              provider,225              priority: index,226              isHealthy,227              lastError: null,228              errorCount: 0,229              lastUsed: 0,230              totalRequests: 0,231              failedRequests: 0,232              circuitBreakerOpen: false,233              circuitBreakerOpenedAt: null,234              consecutiveFailures: 0,235            };236          } catch (error) {237            logger.error('Failed to initialize provider', ErrorHandler.format(error), {238              provider: providerName,239            });240            return null;241          }242        })243      );244 245      // Filter out any null providers (failed to initialize)246      this.providers = providers.filter((p): p is NonNullable<typeof p> => p !== null);247 248      if (this.providers.length === 0) {249        throw new Error('No valid LLM providers could be initialized');250      }251 252      logger.info('LLM providers initialized', {253        count: this.providers.length,254        providers: this.providers.map((p) => ({255          name: p.name,256          healthy: p.isHealthy,257        })),258      });259    } catch (error) {260      logger.error('Error initializing providers', ErrorHandler.format(error));261      throw new Error(`Failed to initialize providers: ${ErrorHandler.getMessage(error)}`);262    }263  }264 265  /**266   * Generate response with timeout protection267   * @private268   */269  private async generateWithTimeout(270    provider: LLMProvider,271    prompt: string,272    options: Record<string, unknown>,273    timeoutMs: number274  ): Promise<string> {275    return Promise.race([276      provider.generate(prompt, options),277      new Promise<never>((_, reject) =>278        setTimeout(() => reject(new Error(`Provider timeout after ${timeoutMs}ms`)), timeoutMs)279      ),280    ]);281  }282 283  /**284   * Check and reset circuit breaker if cooldown period elapsed285   * @private286   */287  private checkCircuitBreaker(providerInfo: ProviderInfo): boolean {288    if (!providerInfo.circuitBreakerOpen) return false;289 290    const now = Date.now();291    const timeSinceOpen = now - (providerInfo.circuitBreakerOpenedAt || 0);292 293    // Reset circuit breaker after cooldown period294    if (timeSinceOpen >= CIRCUIT_BREAKER_RESET_MS) {295      logger.info('Circuit breaker reset', { provider: providerInfo.name });296      providerInfo.circuitBreakerOpen = false;297      providerInfo.circuitBreakerOpenedAt = null;298      providerInfo.consecutiveFailures = 0;299      return false;300    }301 302    return true;303  }304 305  /**306   * Try provider with retry logic and exponential backoff307   * @private308   */309  private async tryProviderWithRetry(310    providerInfo: ProviderInfo,311    prompt: string,312    options: Record<string, unknown>313  ): Promise<string> {314    let lastError: Error | null = null;315 316    for (let attempt = 0; attempt <= MAX_RETRY_ATTEMPTS; attempt++) {317      try {318        // Add exponential backoff delay for retries319        if (attempt > 0) {320          const backoffMs = Math.min(1000 * Math.pow(2, attempt - 1), 5000);321          logger.debug('Retrying with backoff', {322            provider: providerInfo.name,323            attempt,324            backoffMs,325          });326          await new Promise((resolve) => setTimeout(resolve, backoffMs));327        }328 329        const result = await this.generateWithTimeout(330          providerInfo.provider,331          prompt,332          options,333          PROVIDER_TIMEOUT_MS334        );335 336        return result;337      } catch (error) {338        lastError = ErrorHandler.format(error);339        logger.warn('Provider attempt failed', {340          provider: providerInfo.name,341          attempt: attempt + 1,342          maxAttempts: MAX_RETRY_ATTEMPTS + 1,343          error: ErrorHandler.getMessage(error),344        });345 346        // Don't retry if it's not a retryable error347        if (attempt < MAX_RETRY_ATTEMPTS && !ErrorHandler.isRetryable(error)) {348          logger.debug('Error not retryable, skipping further attempts', {349            provider: providerInfo.name,350          });351          break;352        }353      }354    }355 356    throw lastError || new Error('All retry attempts failed');357  }358 359  async generate(prompt: string, options: Record<string, unknown> = {}): Promise<string> {360    await this.initialize();361 362    if (this.providers.length === 0) {363      throw new Error('No LLM providers available. Check API keys and network connectivity.');364    }365 366    // Check cache first (unless explicitly disabled)367    const useCache = options.cache !== false;368    if (useCache) {369      const cached = llmCache.get(prompt, options);370      if (cached) {371        logger.debug('Returning cached response', { prompt: prompt.substring(0, 50) });372        return cached;373      }374    }375 376    let lastError: Error | null = null;377    const startTime = Date.now();378 379    // Try each provider in order380    for (const providerInfo of this.providers) {381      // Check circuit breaker382      if (this.checkCircuitBreaker(providerInfo)) {383        logger.debug('Circuit breaker open, skipping provider', {384          provider: providerInfo.name,385          openedAt: providerInfo.circuitBreakerOpenedAt,386        });387        continue;388      }389 390      try {391        providerInfo.lastUsed = Date.now();392        providerInfo.totalRequests++;393 394        const result = await this.tryProviderWithRetry(providerInfo, prompt, options);395 396        // Success - mark provider as healthy and update statistics397        providerInfo.isHealthy = true;398        providerInfo.lastError = null;399        providerInfo.consecutiveFailures = 0;400        this.lastError = null; // Clear stale errors on success401 402        const duration = Date.now() - startTime;403        logger.info('Provider generate succeeded', {404          provider: providerInfo.name,405          duration,406          cached: false,407        });408 409        // Cache the result (unless explicitly disabled)410        if (useCache) {411          llmCache.set(prompt, result, options);412        }413 414        return result;415      } catch (error) {416        const formattedError = ErrorHandler.format(error);417        logger.error('Provider generate failed after retries', formattedError, {418          provider: providerInfo.name,419          errorCount: providerInfo.errorCount + 1,420          consecutiveFailures: providerInfo.consecutiveFailures + 1,421        });422 423        // Update failure statistics424        providerInfo.failedRequests++;425        providerInfo.errorCount++;426        providerInfo.consecutiveFailures++;427        providerInfo.isHealthy = false;428        providerInfo.lastError = formattedError;429 430        // Open circuit breaker if threshold reached431        if (providerInfo.consecutiveFailures >= CIRCUIT_BREAKER_THRESHOLD) {432          providerInfo.circuitBreakerOpen = true;433          providerInfo.circuitBreakerOpenedAt = Date.now();434          logger.warn('Circuit breaker opened', {435            provider: providerInfo.name,436            consecutiveFailures: providerInfo.consecutiveFailures,437          });438        }439 440        lastError = formattedError;441      }442    }443 444    this.lastError = lastError;445    const errorMessage =446      lastError?.message ||447      'All LLM providers failed to generate a response. Check logs for details.';448    throw new Error(`FallbackLLM Error: ${errorMessage}`);449  }450 451  getActiveProviderName(): string {452    // Return the first healthy provider, or the first provider if none are healthy453    const healthyProvider = this.providers.find((p) => p.isHealthy);454    return healthyProvider?.name || this.providers[0]?.name || 'none';455  }456 457  getActiveProvider(): LLMProvider | null {458    return this.providers[0]?.provider || null;459  }460  getLastError(): Error | null {461    return this.lastError;462  }463}464