import OpenAI from 'openai';
import { ChatCompletionStream } from 'openai/resources/beta/chat/completions.js';
import { encoding_for_model } from 'tiktoken';
import Sentry from '../external-services/sentry.js';
import { getAsyncStorage, logger, runWithAsyncStorage } from '../server-utils/logger.js';
import query from '../server-utils/query.js';

export type OpenAIUsage = {
  total: number;
  system: number;
  user: number;
  assistant: number;
  streamed: number;
};

export const addOpenAIUsages = (a: { [key: string]: number }, b: { [key: string]: number }): OpenAIUsage => {
  return Object.fromEntries(Object.keys(a).map((k) => [k, a[k] + b[k]])) as unknown as OpenAIUsage;
};

export type OpenAIResponse = {
  streamResult: string;
  tokensUsed: OpenAIUsage;
};

type OpenAIFunctionCallResponse = OpenAIResponse & {
  toolCalls: OpenAI.Chat.Completions.ChatCompletionChunk.Choice.Delta.ToolCall[];
};

const progressTimeout = 20000; // ms

export class OpenAIRateLimitError extends Error {}
export class TokenQuotaError extends Error {}

const quotaUpsert = (tokensUsed) => {
  try {
    query(
      `INSERT INTO conductor_token_quota (minute, tokens_used) VALUES (floor(extract(epoch from now()) / 60)::integer, $1) ON CONFLICT (minute) DO UPDATE SET tokens_used = conductor_token_quota.tokens_used + $1`,
      [tokensUsed]
    );
  } catch (e) {
    Sentry.captureException(e);
    logger.warning(`Failed to async update conductor_token_quota: ${e.message}`);
  }
};

export const streamOpenAICall = async (
  openai: OpenAI,
  onStreamMessage: (content: string) => void,
  onCancellable: (cancel: () => void) => void,
  request: OpenAI.Chat.Completions.ChatCompletionCreateParams,
  maxInitialTokens?: number
): Promise<OpenAIFunctionCallResponse> => {
  // ref https://github.com/openai/openai-cookbook/blob/main/examples/How_to_count_tokens_with_tiktoken.ipynb
  const enc = encoding_for_model('gpt-4-0613'); //request.model as TiktokenModel);
  const usage: OpenAIUsage = {
    total: 3,
    system: 3,
    user: 0,
    assistant: 0,
    streamed: 0,
  };
  for (const message of request.messages) {
    const tokenCount = 3 + ((message as any).name ? 1 : 0) + enc.encode((message as any).content || '').length;
    usage.total += tokenCount;
    usage[message.role] += tokenCount;
  }

  // Could not find any documentation explaining how to do this. Assuming this is roughly correct.
  request.tools?.forEach((tool) => {
    const tokenCount =
      enc.encode(tool.function.name).length +
      enc.encode(JSON.stringify(tool.function.parameters)).length +
      enc.encode(tool.function.description).length;
    usage.total += tokenCount;
    usage.system += tokenCount;
  });
  enc.free();

  if (typeof maxInitialTokens === 'number' && usage.total > maxInitialTokens) {
    logger.info('Token quota exceeded - refusing to make OpenAI request');
    throw new TokenQuotaError(`usage.total=${usage.total} > maxInitialTokens=${maxInitialTokens}`);
  }

  let stream: ChatCompletionStream;

  try {
    stream = await openai.beta.chat.completions.stream({
      ...request,
      stream: true,
    });
  } catch (err) {
    if (err instanceof OpenAI.APIError) {
      const message = `OpenAI streaming request failed with status ${err.status}`;
      if (err.status === 429) {
        throw new OpenAIRateLimitError(message);
      } else {
        throw new Error(message);
      }
    } else {
      throw err;
    }
  }

  const result = await new Promise<OpenAIFunctionCallResponse>(async (resolve, reject) => {
    let done = false;
    const finish = (result) => {
      if (done) {
        return;
      }

      done = true;

      stream.controller.abort();

      if (usage.streamed > 0) {
        quotaUpsert(usage.streamed);
      }
      if (result?.error) {
        return reject(result.error);
      }

      resolve({
        streamResult: result.message,
        toolCalls: result.toolCalls,
        tokensUsed: usage,
      });
    };

    try {
      const asyncStorage = getAsyncStorage();
      onCancellable(() =>
        runWithAsyncStorage(asyncStorage, () => {
          if (!done) {
            logger.info('user cancelled conductor request');
            finish({ message: 'cancelled' });
          }
        })
      );
    } catch (error) {
      finish({ error });
    }

    let progressTimeoutHandle = null;
    let partialMessage = '';

    const consume = (chunk: OpenAI.Chat.Completions.ChatCompletionChunk) => {
      if (progressTimeoutHandle) {
        clearTimeout(progressTimeoutHandle);
      }
      progressTimeoutHandle = setTimeout(() => {
        finish({ error: new Error('OpenAI streaming response went ' + progressTimeout + ' ms without progress') });
      }, progressTimeout);

      const choices = chunk.choices;
      if (!Array.isArray(choices) || choices.length < 1) {
        logger.info('OpenAI streaming resposne ignoring missing nonempty choices array:', chunk);
        return;
      }

      const content = choices[0].delta?.content;

      if (content) {
        partialMessage = partialMessage + content;
        usage.streamed += 1;
        usage.assistant += 1;
        usage.total += 1;
        onStreamMessage(content);
      }

      if (choices[0].delta.tool_calls) {
        usage.streamed += 1;
        usage.system += 1;
        usage.total += 1;
      }
    };

    try {
      for await (const chunk of stream) {
        if (done) {
          return;
        }
        consume(chunk);
      }
      const finalCompletion = await stream.finalChatCompletion();
      finish({
        message: finalCompletion.choices[0].message.content,
        toolCalls: finalCompletion.choices[0].message.tool_calls,
      });
    } catch (error) {
      finish({ error });
    } finally {
      if (progressTimeoutHandle) {
        clearTimeout(progressTimeoutHandle);
      }
      if (!done) {
        finish({ error: new Error('OpenAI streaming request connection closed before end of response') });
      }
    }
  });

  quotaUpsert(usage.total);

  return result;
};
