import { and, eq, inArray, sql } from 'drizzle-orm';
import type { DrizzleClient } from '@/server/db/client.js';
import { llmProviders } from '@/server/db/schema.js';
import { getAllSystemConfig } from '@/server/db/queries/system-config.js';
import { getLlmQueue } from '@/server/db/queries/llm-queues.js';
import { env } from '@/server/env.js';
import {
  GeminiTranslationProvider,
  DEFAULT_INPUT_USD_PER_MTOK,
  DEFAULT_OUTPUT_USD_PER_MTOK,
} from './gemini.js';
import { TranslationProviderError } from './types.js';
import type { TranslationProvider, TranslationRequest, TranslationResponse } from './types.js';
import { requireConfiguredLlmModel } from '@/server/ai/model-config.js';

class ChainedTranslationProvider implements TranslationProvider {
  readonly id = 'chained';
  constructor(private readonly entries: TranslationProvider[]) {}
  async translateBatch(req: TranslationRequest): Promise<TranslationResponse> {
    let lastErr: unknown;
    let retryableError: TranslationProviderError | undefined;
    for (const entry of this.entries) {
      try {
        return await entry.translateBatch(req);
      } catch (err) {
        if (
          err instanceof TranslationProviderError &&
          (err.kind === 'auth' || err.kind === 'rate_limit' || err.kind === 'transport')
        ) {
          lastErr = err;
          if (err.kind === 'rate_limit' || err.kind === 'transport') {
            retryableError = err;
          }
          continue;
        }
        throw err;
      }
    }
    throw retryableError ?? lastErr;
  }
}

let cached: TranslationProvider | null = null;

export async function getProvider(db: DrizzleClient): Promise<TranslationProvider> {
  if (cached) return cached;
  const queue = await getLlmQueue(db, 'translation');
  const cfg = await getAllSystemConfig(db);
  const inputUsdPerMTok = cfg['translation_model_input_usd_per_mtok']
    ? parseFloat(cfg['translation_model_input_usd_per_mtok'])
    : DEFAULT_INPUT_USD_PER_MTOK;
  const outputUsdPerMTok = cfg['translation_model_output_usd_per_mtok']
    ? parseFloat(cfg['translation_model_output_usd_per_mtok'])
    : DEFAULT_OUTPUT_USD_PER_MTOK;
  const chain = queue?.chain ?? [];
  if (chain.length === 0) {
    const modelId = requireConfiguredLlmModel(cfg['translation_model_id'], 'translation');
    const apiKey = cfg['google_api_key'] || env.GOOGLE_API_KEY || '';
    cached = new GeminiTranslationProvider(apiKey, modelId, inputUsdPerMTok, outputUsdPerMTok);
    return cached;
  }
  const piiKey = env.PII_KEY;
  if (!piiKey) throw new Error('PII_KEY not configured');
  const ids = chain.map((e: { llmProviderId: string }) => e.llmProviderId);
  const providerRows = await db
    .select({
      id: sql<string>`${llmProviders.id}::text`,
      type: llmProviders.type,
      api_key: sql<string | null>`pgp_sym_decrypt(${llmProviders.apiKeyEnc}, ${piiKey})::text`,
    })
    .from(llmProviders)
    .where(and(inArray(llmProviders.id, ids), eq(llmProviders.isActive, true)));
  const rowById = new Map(providerRows.map((r) => [r.id, r]));
  const entries: TranslationProvider[] = [];
  for (const link of chain) {
    const row = rowById.get(link.llmProviderId);
    if (!row || !row.api_key) continue;
    if (row.type === 'google')
      entries.push(
        new GeminiTranslationProvider(row.api_key, link.model, inputUsdPerMTok, outputUsdPerMTok),
      );
  }
  if (entries.length === 0) throw new Error('No active translation providers in queue chain');
  cached = entries.length === 1 ? entries[0]! : new ChainedTranslationProvider(entries);
  return cached;
}

export function resetProvider(): void {
  cached = null;
}
