import { randomUUID } from 'node:crypto';
import { Socket } from 'socket.io';
import { enqueueWithTimeout } from '../external-services/redis.js';
import query from '../server-utils/query.js';
import { SocketEvent } from '../server-utils/socketMessageHandler.js';
import interpretComposerPrompt from './interpretComposerPrompt.js';

/**
 * This file contains route handlers and other helper functions
 * for the composer service.
 *
 * Composer is the in-house service that generates midi notes from scratch,
 * with a prompt, or with accompanying midi notes.
 */

// TODO: share type with frontend
enum ComposerAction {
  GetChoices = 'getChoices',
  ChooseChoice = 'chooseChoice',
}

export const getComposerConfig = async () => {
  return {
    queueName: process.env.COMPOSER_QUEUE_NAME || 'composer_requests',
    timeout: Number(process.env.COMPOSER_TIMEOUT || 10.0),
  };
};

const chooseComposerChoice = async (requestUUID: string, choice: any) => {
  await query(
    `INSERT INTO composer_choices (request_id, choice) VALUES ((SELECT id FROM composer_requests WHERE uuid = $1), $2);`,
    [requestUUID, JSON.stringify(choice)]
  );
};

const userComposerPrompts = new Map<
  number,
  { lastInputPrompt: string; outputPrompt: string; biasOverrides: { [key: string]: any } }
>();

const mapComposerRequest = async (userId: number, request: any) => {
  if (!request.textPrompt) return request;

  const prompt = request.textPrompt as string;
  let outputPrompt = '';
  let biasOverrides = {};
  let cachedInterpretation = userComposerPrompts.get(userId);

  if (cachedInterpretation?.lastInputPrompt === prompt) {
    outputPrompt = userComposerPrompts.get(userId)!.outputPrompt;
    biasOverrides = userComposerPrompts.get(userId)!.biasOverrides;
  } else {
    const { bias, textPrompt } = await interpretComposerPrompt(prompt);
    outputPrompt = textPrompt;
    biasOverrides = bias;
    userComposerPrompts.set(userId, { lastInputPrompt: prompt, outputPrompt, biasOverrides });
  }

  return {
    ...request,
    textPrompt: outputPrompt,
    requests: request.requests.map((r) => {
      const combinedBias = {
        ...r.bias,
        ...biasOverrides,
      };

      return {
        ...r,
        cfgScale: 1.3,
        bias: combinedBias,
      };
    }),
  };
};

const getComposerChoices = async (userId: number, request: any) => {
  const startTime = Date.now();
  const interpretedRequest = await mapComposerRequest(userId, request);

  const { queueName, timeout } = await getComposerConfig();
  let composerResponse;

  try {
    composerResponse = await enqueueWithTimeout(queueName, interpretedRequest, timeout);
  } catch (e) {
    if (e.message.match(/Timeout/)) {
      console.error(e);
      composerResponse = { error: 'Request timed out' };
    } else {
      throw e;
    }
  }

  const uuid = randomUUID();

  // intentionally not awaiting this
  query(
    `INSERT INTO composer_requests (user_id, queue_name, request, response, response_time_ms, uuid) VALUES ($1, $2, $3, $4, $5, $6)`,
    [userId, queueName, request, composerResponse, Date.now() - startTime, uuid]
  );

  return { ...composerResponse, uuid };
};

export const composerMessageHandler = async (socket: Socket, message: any) => {
  const userId = (socket.request as any).user.id;
  const action = message.action as ComposerAction;

  if (action === ComposerAction.GetChoices) {
    socket.emit(SocketEvent.Final, await getComposerChoices(userId, message.request));
    return;
  }

  if (action === ComposerAction.ChooseChoice) {
    await chooseComposerChoice(message.payload.requestUUID, message.payload.choice);
    return;
  }

  socket.emit(SocketEvent.Final, { error: true, message: 'Unknown action' });
};
