import os
import random
import time
from collections import defaultdict
from datetime import datetime
from typing import Optional, Literal, Iterator

import openai
from openai.types.chat import ChatCompletion, ChatCompletionChunk

from suno_utils.utils.bedrock import BedrockAdapter

try:
    import together
except ImportError:
    print("Warning: together not installed locally, will not use together client")
    together = None

MINUTES = 60  # seconds / minute

FAST_TIER_FRACTION = 0.01
MODAL_BASE_URL = "https://suno-ai--lyrics-gen-serve.modal.run"


def _get_openai_compatible_url_from_base_url(base_url: str) -> str:
    return os.path.join(base_url, "v1")


REMI = "ft:gpt-4o-2024-08-06:suno:recent-billboard3:ASqCxlGH"
DEFAULT_GPT_MODEL = "gpt-4o-2024-11-20"
DUSAN_BOT2 = "ft:gpt-4o-2024-08-06:suno:dusan-ascii:APwxFPs4"  # 4o FT with ascii chars fixed
LORENZO_V1_LYRICS_MODEL = "ft:gpt-4o-2024-08-06:suno:lorenzo-gym:AXXM0K5u"
CLAUDE_SONNET_MODEL = "us.anthropic.claude-3-7-sonnet-20250219-v1:0"
CLAUDE_SONNET_4_MODEL = "us.anthropic.claude-sonnet-4-20250514-v1:0"

# Provider types as constants
PROVIDER_ANTHROPIC = "anthropic"
PROVIDER_OPENAI = "openai"
PROVIDER_TOGETHER = "together"
PROVIDER_MODAL = "modal"
PROVIDER_VERTEX = "vertex"
PROVIDER_UNKNOWN = "unknown"

ProviderType = Literal["anthropic", "openai", "together", "modal", "vertex", "unknown"]

LYRICS_PARAMETER_CONFIGS = defaultdict(
    dict,  # if not found, assume no special param values should be passed in
    {
        "my-remi-8b-1": {"top_p": 0.95, "temperature": 0.7},
        "remi-v1-hps1": {"top_p": 0.95, "temperature": 1},
        "remi-v1-hps2": {"top_p": 0.9, "temperature": 1},
        "remi-v1-hps3": {"top_p": 0.8, "temperature": 1},
        "remi-v1-hps4": {"top_p": 0.95, "temperature": 0.9},
        "remi-v1-hps5": {"top_p": 0.9, "temperature": 0.9},
        "remi-v1-hps6": {"top_p": 0.8, "temperature": 0.9},
        "remi-v1-hps7": {"top_p": 0.95, "temperature": 1.1},
        "remi-v1-hps8": {"top_p": 0.9, "temperature": 1.1},
        "remi-v1-hps9": {"top_p": 0.8, "temperature": 1.1},
        "remi-v1-hps10": {"top_p": 0.8, "temperature": 1.2},
        "remi-v1": {"top_p": 0.8, "temperature": 1.1},
        "claude-sonnet": {"topP": 0.8, "temperature": 0.9},
        "claude-sonnet-4": {"topP": 0.8, "temperature": 0.9},
        "default-hps1": {"top_p": 0.95, "temperature": 1},
        "default-hps2": {"top_p": 0.9, "temperature": 1},
        "default-hps3": {"top_p": 0.8, "temperature": 1},
        "default-hps4": {"top_p": 0.95, "temperature": 0.9},
        "default-hps5": {"top_p": 0.9, "temperature": 0.9},
        "default-hps6": {"top_p": 0.8, "temperature": 0.9},
        "default-hps7": {"top_p": 0.95, "temperature": 1.1},
        "default-hps8": {"top_p": 0.9, "temperature": 1.1},
        "default-hps9": {"top_p": 0.8, "temperature": 1.1},
        "default-hps10": {"top_p": 0.8, "temperature": 1.2},
        # A/B hyperparam sweep winner, corresponding to default-hps6
        "default": {"top_p": 0.8, "temperature": 0.9},
    },
)


def _get_backend_lyrics_model(frontend_lyrics_model: str | None) -> str:
    """Translate the lyrics model from the frontend-visible name to the corresponding backend name."""
    # map lyrics model options from FE to real names
    print("getting backend lyrics model for:", frontend_lyrics_model)
    LYRICS_MODEL_LOOKUP = {
        "dusanbot2": DUSAN_BOT2,
        "remi-v1": REMI,
        "lorenzo-v1": LORENZO_V1_LYRICS_MODEL,
        "my-remi-8b-1": "my-remi-8b-1",
        "remi-v1-hps1": REMI,
        "remi-v1-hps2": REMI,
        "remi-v1-hps3": REMI,
        "remi-v1-hps4": REMI,
        "remi-v1-hps5": REMI,
        "remi-v1-hps6": REMI,
        "remi-v1-hps7": REMI,
        "remi-v1-hps8": REMI,
        "remi-v1-hps9": REMI,
        "claude-sonnet": CLAUDE_SONNET_MODEL,
        "claude-sonnet-4": CLAUDE_SONNET_4_MODEL,
        "default-hps1": DEFAULT_GPT_MODEL,
        "default-hps2": DEFAULT_GPT_MODEL,
        "default-hps3": DEFAULT_GPT_MODEL,
        "default-hps4": DEFAULT_GPT_MODEL,
        "default-hps5": DEFAULT_GPT_MODEL,
        "default-hps6": DEFAULT_GPT_MODEL,
        "default-hps7": DEFAULT_GPT_MODEL,
        "default-hps8": DEFAULT_GPT_MODEL,
        "default-hps9": DEFAULT_GPT_MODEL,
        "default-hps10": DEFAULT_GPT_MODEL,
        None: DEFAULT_GPT_MODEL,
        "default": DEFAULT_GPT_MODEL,
        DEFAULT_GPT_MODEL: DEFAULT_GPT_MODEL,
        "gemini-2.0": "google/gemini-2.0-flash-001",
    }
    backend_lyrics_model = LYRICS_MODEL_LOOKUP.get(frontend_lyrics_model)
    if backend_lyrics_model is None:
        valid_choices = list(LYRICS_MODEL_LOOKUP.keys())
        err_msg = f"Lyrics model: {frontend_lyrics_model} unrecognized, provide one of {valid_choices}"
        raise ValueError(err_msg)
    return backend_lyrics_model


class LyricsClient:
    def __init__(self, openai_client, together_client, vertex_client, bedrock_client):
        print("initializing lyrics llm client")
        self.openai_client = openai_client
        self.together_client = together_client
        self.modal_client = openai.OpenAI(
            base_url=_get_openai_compatible_url_from_base_url(MODAL_BASE_URL), api_key="sunosunosuno"
        )
        self.anthropic_client = BedrockAdapter(bedrock_client)
        self.vertex_client = vertex_client

    def _get_provider_type(self, model: str | None = None) -> ProviderType:
        """Determine the provider type based on the model name.

        Args:
            model: Model name to check

        Returns:
            Provider type as a string constant
        """
        if model is None:
            return PROVIDER_OPENAI

        if (
            model == "lorenzo-v1"
            or model == "dusanbot2"
            or model.startswith("gpt-4")
            or model.startswith("remi-v1")
            or model.startswith("default")
        ):
            return PROVIDER_OPENAI
        elif model.startswith("patsuno/Meta-Llama"):
            return PROVIDER_TOGETHER
        elif model.startswith("my-remi-8b-1"):
            return PROVIDER_MODAL
        elif model.startswith("gemini"):
            return PROVIDER_VERTEX
        elif model.startswith("anthropic") or model == "claude-sonnet" or model == "claude-sonnet-4":
            return PROVIDER_ANTHROPIC
        else:
            return PROVIDER_UNKNOWN

    def get_provider_name(self, model: str | None = None) -> str:
        """Return provider name based on model.

        Args:
            model: Model name to check. If None, returns "openai".

        Returns:
            Provider name: "openai", "together", "modal", or "unknown"
        """
        return self._get_provider_type(model)

    def create(
        self,
        model: str,
        system_prompt: str,
        user_prompt: str,
        max_tokens: Optional[int] = None,
        n: int = 1,
        stream: bool = False,
    ) -> ChatCompletion | Iterator[ChatCompletionChunk]:
        print("creating with LyricsLLMClient for:", model, user_prompt)
        client = self._get_client_for_model(model)
        lyrics_param_config = LYRICS_PARAMETER_CONFIGS[model]
        backend_model = _get_backend_lyrics_model(model)
        print("lyrics model:", model, "lyrics_param_config:", lyrics_param_config)
        if model.startswith("default"):
            lyrics_param_config = {
                **lyrics_param_config,
                "timeout": 30,
            }
            fast_tier_fraction = _interpolate_time(time.ctime())
            if random.random() < fast_tier_fraction:
                print(
                    f"upgrading {model} to Fast Tier at: {time.ctime()}, fraction: {fast_tier_fraction}"
                )
                lyrics_param_config["service_tier"] = "fast_tier_temp_pilot"
        try:
            resp = client.chat.completions.create(
                model=backend_model,
                messages=[
                    {"role": "system", "content": system_prompt},
                    {"role": "user", "content": user_prompt},
                ],
                max_tokens=max_tokens,
                n=n,
                stream=stream,
                **lyrics_param_config,
            )
        except Exception as e:
            print("got exception in LyricsClient.create:", e)
            raise e
        return resp

    def _get_client_for_model(self, model: str):
        print("getting model for:", model)
        provider_type = self._get_provider_type(model)

        if provider_type == PROVIDER_OPENAI:
            return self.openai_client
        elif provider_type == PROVIDER_TOGETHER:
            return self.together_client
        elif provider_type == PROVIDER_MODAL:
            print("getting modal client")
            return self.modal_client
        elif provider_type == PROVIDER_VERTEX:
            return self.vertex_client
        elif provider_type == PROVIDER_ANTHROPIC:
            return self.anthropic_client
        else:
            raise ValueError(f"Couldn't find a client for model: {model}")


def _is_openai_client(client):
    """Determine whether client is the openai client."""
    is_openai = "OpenAI" in str(client)
    print(f"checking if {client} is openai: {is_openai}")
    return is_openai


def _interpolate_time(current_time):
    """
    Interpolates between start_time and end_time based on current_time.

    Args:
        current_time (str): Current time in ctime format (e.g., 'Thu Apr  3 12:34:56 2025')
        start_time (str): Start time in ctime format
        end_time (str): End time in ctime format

    Returns:
        float: 0 if current_time < start_time,
               1 if current_time > end_time,
               linear interpolation between 0 and 1 otherwise
    """
    # Convert ctime strings to datetime objects
    START_TIME = "Thu Apr  3 22:00:00 2025"  # NB UTC time in modal land
    END_TIME = "Fri Apr  4 02:00:00 2025"
    TIME_FORMAT = "%a %b %d %H:%M:%S %Y"

    current_dt = datetime.strptime(current_time, TIME_FORMAT)
    start_dt = datetime.strptime(START_TIME, TIME_FORMAT)
    end_dt = datetime.strptime(END_TIME, TIME_FORMAT)

    # Convert to timestamps (seconds since epoch)
    current_ts = current_dt.timestamp()
    start_ts = start_dt.timestamp()
    end_ts = end_dt.timestamp()

    # Check boundary conditions
    if current_ts <= start_ts:
        return 0.0
    if current_ts >= end_ts:
        return 1.0

    # Linear interpolation
    progress = (current_ts - start_ts) / (end_ts - start_ts)
    return progress
