import torch
import numpy as np
import gradio as gr
from librosa import resample
import logging


SAMPLE_RATE = 16000


def _convert_input(arr):
    if "int" in str(arr.dtype):
        max_val = np.iinfo(arr.dtype).max
        arr = (arr.astype(np.float64) / max_val).astype(np.float32)
    if len(arr.shape) == 1:
        pass
    elif len(arr.shape) == 2 and arr.shape[1] in (1, 2):
        arr = arr.mean(axis=1)
    else:
        raise ValueError("wrong input shape")
    return arr


def _convert_output(arr):
    return np.vstack([arr / 2, arr / 2]).T


#########
# Model #
#########


from torch import hub


model = hub.load('JorisCos/asteroid', 'conv_tasnet', 'JorisCos/ConvTasNet_Libri2Mix_sepnoisy_16k')


def apply_model(arr):
    arr = arr.reshape(1, -1)
    sources = model.separate(arr)[0]
    source1 = sources[0, :]
    source2 = sources[1, :]
    noise = arr - sources.sum(axis=0)
    return source1, source2, noise


##########
# Server #
##########


def main(audio, audio_rec):
    if audio is None and audio_rec is None:
        return "no audio defined"
    elif audio is None:
        sr, arr = audio_rec
    else:
        sr, arr = audio
    arr = _convert_input(arr)
    arr = resample(arr, sr, SAMPLE_RATE)
    source1, source2, noise = apply_model(arr)
    source1 = _convert_output(source1)
    source2 = _convert_output(source2)
    noise = _convert_output(noise)
    return (SAMPLE_RATE, source1), (SAMPLE_RATE, source2), (SAMPLE_RATE, noise)


iface = gr.Interface(
    fn=main, inputs=[
        gr.inputs.Audio("upload", label="upload a file", optional=True),
        gr.inputs.Audio("microphone", label="create a recording", optional=True),
    ], outputs=[
        gr.outputs.Audio(type="auto", label="Voice 1"),
        gr.outputs.Audio(type="auto", label="Voice 2"),
        gr.outputs.Audio(type="auto", label="Background Noise"),
    ],
    examples=[["samples/speech_mix.wav", ""]],
    allow_screenshot=False, allow_flagging=True, server_name="0.0.0.0", server_port=7862
)
iface.launch(
    ssl=('../cert/383fff33778762e8.crt', '../cert/383fff33778762e8.key'),
)
