import numpy as np
import matplotlib.pyplot as plt


def draw_pitch_waveform_transcription(y, sr, pitch, title,
                                      transcription_data=None):
    """
    Plot the audio waveform, pitch contour, and optional transcription
    on the same figure.
    """
    # Reset matplotlib settings
    plt.rcParams.update(plt.rcParamsDefault)

    # Configure matplotlib with multiple fallback fonts
    plt.rcParams.update({
        'font.family': 'DejaVu Sans',
        'text.usetex': False,
    })

    times = np.linspace(0, len(y) / sr, len(y))
    pitch_values = pitch.selected_array['frequency']
    pitch_values[pitch_values == 0] = np.nan

    # Normalize waveform to bring it to a similar scale as pitch
    y_normalized = y / np.max(np.abs(y)) * np.nanmax(pitch_values)

    fig, ax = plt.subplots(figsize=(14, 6))

    # Plot waveform
    ax.plot(times, y_normalized, color='gray', alpha=0.5, label='Waveform')

    # Plot pitch
    ax.scatter(pitch.xs(), pitch_values, color='r', s=10, label='Pitch')

    if transcription_data:
        # Add transcription data
        transcription = transcription_data['transcription']
        start_times = np.array(transcription_data[
                                   'start_timestamps']) / 1000  #
        # Convert to seconds
        end_times = np.array(transcription_data[
                                 'end_timestamps']) / 1000  # Convert to
        # seconds
        probabilities = transcription_data['probabilities']

        # Iterate over each transcription unit and its corresponding
        # timestamps
        for char, start, end, prob in zip(transcription, start_times,
                                          end_times, probabilities):
            # Plot a rectangle over the corresponding waveform region
            # with color based on probability
            ax.axvspan(start, end, color=plt.cm.viridis(prob), alpha=0.3)
            # Annotate the transcription character above the waveform
            mid_time = (start + end) / 2
            ax.text(mid_time, 0, char, ha='center', va='center',
                    fontsize=12, alpha=0.9, weight='bold')

    ax.set_title(title)
    # ax.set_xlabel("Time (s)")
    # ax.set_ylabel("Normalized Amplitude / Fundamental Frequency [Hz]")
    # ax.legend()
    return fig


# First, ensure we have the necessary fonts installed
def setup_fonts():
    import subprocess

    """Install required fonts if not present"""
    try:
        # Install fonts if they're not already installed
        subprocess.run(['apt-get', 'update'], check=True)
        subprocess.run(
            ['apt-get', 'install', '-y', 'fonts-noto-cjk', 'fonts-noto'],
            check=True)

        # Clear matplotlib's font cache
        import matplotlib.font_manager as fm
        fm._get_font.cache_clear()
    except:
        print(
            "Could not install fonts automatically. Please ensure you have "
            "the required fonts installed.")


if __name__ == '__main__':
    setup_fonts()
