from dataclasses import dataclass


@dataclass
class CodecConfig:
    model_type: str
    encoder_dim: int
    encoder_rates: list
    latent_dim: int
    decoder_dim: int
    decoder_rates: list
    vae_dim: int
    sample_rate: int
    n_transformer_layers: int
    is_frozen_encoder: bool


@dataclass
class DiscriminatorConfig:
    model_type: str
    rates: list
    periods: list
    fft_sizes: list
    sample_rate: int
    bands: list


def get_model(cfg):
    # generators
    if cfg.model_type == "dac_vae":
        """
        DAC
        """
        from .codec_dac_vae import DACVAE

        return DACVAE(
            encoder_dim=cfg.encoder_dim,
            encoder_rates=cfg.encoder_rates,
            latent_dim=cfg.latent_dim,
            decoder_dim=cfg.decoder_dim,
            decoder_rates=cfg.decoder_rates,
            vae_dim=cfg.vae_dim,
            sample_rate=cfg.sample_rate,
            is_frozen_encoder=cfg.is_frozen_encoder,
        )
    elif cfg.model_type == "convnext_vae":
        from .codec_convnext_vae import ConvNextVAE

        return ConvNextVAE(
            encoder_dim=cfg.encoder_dim,
            encoder_rates=cfg.encoder_rates,
            latent_dim=cfg.latent_dim,
            decoder_dim=cfg.decoder_dim,
            decoder_rates=cfg.decoder_rates,
            vae_dim=cfg.vae_dim,
            sample_rate=cfg.sample_rate,
        )
    elif cfg.model_type == "musicfm_dac":
        from .codec_musicfm_dac import MusicFMDAC
        return MusicFMDAC(
            latent_dim=cfg.latent_dim,
            decoder_dim=cfg.decoder_dim,
            decoder_rates=cfg.decoder_rates,
            sample_rate=cfg.sample_rate,
            n_transformer_layers=cfg.n_transformer_layers,
        )
    elif cfg.model_type == "dac_spectrostream_vae":
        """
        DAC encoder
        VAE latent
        SpectroStream decoder
        """
        from .codec_dac_spectrostream import DACSpectroStreamVAE

        return DACSpectroStreamVAE(
            n_channels=cfg.encoder_dim,
            vae_dim=cfg.vae_dim,
            is_frozen_encoder=cfg.is_frozen_encoder,
        )
    elif cfg.model_type == "spectrostream_vae":
        """
        SpectroStream
        """
        from .codec_spectrostream import SpectroStreamVAE

        return SpectroStreamVAE(
            n_channels=cfg.encoder_dim,
            vae_dim=cfg.vae_dim,
        )
    # discriminators
    elif cfg.model_type == "dac_discriminator":
        from .discriminator_dac import DescriptDiscriminator

        return DescriptDiscriminator(
            rates=cfg.rates, periods=cfg.periods, fft_sizes=cfg.fft_sizes, sample_rate=cfg.sample_rate
        )
    else:
        raise ValueError(f"Unknown model type: {cfg.model_type}")
