import computeCosineSimilarity from 'compute-cosine-similarity';
import openai from '../external-services/openai.js';
import { logger } from '../server-utils/logger.js';
import query from '../server-utils/query.js';

export const createEmbedding = async (content: string) => {
  const result = await openai.embeddings.create({
    model: 'text-embedding-ada-002',
    input: content,
  });
  logger.info('Created embedding for', content);
  return result.data[0].embedding as number[];
};

const getEmbedding = async (content: string) => {
  const queryResult = await query(`SELECT * FROM embedded_strings WHERE content = $1`, [content]);
  if (queryResult.rows.length > 0) {
    return queryResult.rows[0].embedding as number[];
  } else {
    const embedding = await createEmbedding(content);
    await query(`INSERT INTO embedded_strings (content, embedding) VALUES ($1, $2) ON CONFLICT (content) DO NOTHING`, [
      content,
      JSON.stringify(embedding),
    ]);
    return embedding;
  }
};

const semanticRank = async (query: string, options: string[]) => {
  const normalizedQuery = query.toLowerCase().trim();
  const normalizedOptions = options.map((option) => option.toLowerCase().trim());

  const existingEmbedding = await getEmbedding(normalizedQuery);
  const existingOptionEmbeddings = await Promise.all(normalizedOptions.map((option) => getEmbedding(option)));

  const similarities = existingOptionEmbeddings.map((optionEmbedding) => {
    return computeCosineSimilarity(existingEmbedding, optionEmbedding);
  });

  const sortedOptions = options
    .map((option, index) => {
      return {
        option,
        similarity: similarities[index],
      };
    })
    .sort((a, b) => b.similarity - a.similarity);

  return sortedOptions;
};

export default semanticRank;
