/**
 * ChainedLlmProvider — tries providers in order, falling through on transient/rate-limit errors.
 */

import { inArray, sql } from 'drizzle-orm';
import type { DrizzleClient } from '@/server/db/client';
import { llmProviders } from '@/server/db/schema.js';
import type { AiCredentials } from '@/server/ai/credentials';
import { geminiProvider } from './gemini';
import { createOpenAICompatProvider } from './openai';
import { mockProvider } from './mock';
import type {
  ImageRequest,
  LlmProvider,
  NormalizedTextResult,
  ProviderId,
  TextRequest,
} from './types';
import {
  ProviderRateLimitError,
  ProviderTransientError,
  ProviderQuotaExhaustedError,
  ProviderAuthError,
} from './types';

export interface LlmChainConfig {
  chain: Array<{ llmProviderId: string; model: string }>;
}

interface ChainEntry {
  provider: LlmProvider;
  creds: AiCredentials;
  model: string;
}

function isFallbackError(err: unknown): boolean {
  return (
    err instanceof ProviderTransientError ||
    err instanceof ProviderRateLimitError ||
    err instanceof ProviderQuotaExhaustedError ||
    err instanceof ProviderAuthError
  );
}

export class ChainedLlmProvider implements LlmProvider {
  readonly id: ProviderId;
  readonly supportsImage: boolean;

  constructor(private readonly chain: ChainEntry[]) {
    if (chain.length === 0) {
      throw new Error('ChainedLlmProvider requires at least one configured provider');
    }
    this.id = chain[0]!.provider.id;
    this.supportsImage = chain.some((entry) => entry.provider.supportsImage);
  }

  async generateText(creds: AiCredentials, req: TextRequest): Promise<NormalizedTextResult> {
    let lastError: unknown;
    let retryableError: ProviderTransientError | ProviderRateLimitError | undefined;
    for (const entry of this.chain) {
      try {
        const result = await entry.provider.generateText(entry.creds ?? creds, {
          ...req,
          model: entry.model,
        });
        return { ...result, model: entry.model };
      } catch (err) {
        if (isFallbackError(err)) {
          console.warn(
            `[ChainedLlmProvider] falling through from ${entry.model} (${entry.provider.id}):`,
            err instanceof Error ? err.message : String(err),
          );
          lastError = err;
          if (err instanceof ProviderTransientError || err instanceof ProviderRateLimitError) {
            retryableError = err;
          }
          continue;
        }
        throw err;
      }
    }
    throw retryableError ?? lastError;
  }

  async generateTextWithImage(
    creds: AiCredentials,
    req: ImageRequest,
  ): Promise<NormalizedTextResult> {
    let lastError: unknown;
    let retryableError: ProviderTransientError | ProviderRateLimitError | undefined;
    for (const entry of this.chain) {
      try {
        const result = await entry.provider.generateTextWithImage(entry.creds ?? creds, {
          ...req,
          model: entry.model,
        });
        return { ...result, model: entry.model };
      } catch (err) {
        if (isFallbackError(err)) {
          console.warn(
            `[ChainedLlmProvider] falling through from ${entry.model} (${entry.provider.id}):`,
            err instanceof Error ? err.message : String(err),
          );
          lastError = err;
          if (err instanceof ProviderTransientError || err instanceof ProviderRateLimitError) {
            retryableError = err;
          }
          continue;
        }
        throw err;
      }
    }
    throw retryableError ?? lastError;
  }
}

export async function buildChainedProvider(
  db: DrizzleClient,
  piiKey: string,
  config: LlmChainConfig,
): Promise<ChainedLlmProvider> {
  const ids = config.chain.map((entry) => entry.llmProviderId);
  if (ids.length === 0) {
    throw new Error('ChainedLlmProvider requires at least one configured provider');
  }

  const rows = await db
    .select({
      id: llmProviders.id,
      slug: llmProviders.slug,
      type: llmProviders.type,
      base_url: llmProviders.baseUrl,
      api_key: sql<string | null>`pgp_sym_decrypt(${llmProviders.apiKeyEnc}, ${piiKey})::text`,
    })
    .from(llmProviders)
    .where(inArray(llmProviders.id, ids));

  const rowById = new Map(rows.map((row) => [row.id, row]));
  const entries: ChainEntry[] = [];

  for (const link of config.chain) {
    const row = rowById.get(link.llmProviderId);
    if (!row) {
      console.warn(`buildChainedProvider: llm provider not found: ${link.llmProviderId}`);
      continue;
    }
    if (!row.api_key) {
      console.warn(`buildChainedProvider: missing api_key for provider ${row.slug} (${row.id})`);
      continue;
    }

    const creds: AiCredentials = { apiKey: row.api_key };
    let provider: LlmProvider;

    switch (row.type) {
      case 'google':
        provider = geminiProvider;
        break;
      case 'openai-compat':
        if (!row.base_url) {
          console.warn(`buildChainedProvider: openai-compat provider ${row.slug} missing base_url`);
          continue;
        }
        provider = createOpenAICompatProvider(row.base_url);
        break;
      case 'mock':
        provider = mockProvider;
        break;
      default:
        console.warn(
          `buildChainedProvider: unsupported provider type "${row.type}" for ${row.slug}`,
        );
        continue;
    }

    entries.push({ provider, creds, model: link.model });
  }

  return new ChainedLlmProvider(entries);
}
