import multiprocessing
from threading import Thread
from queue import Queue
import sys
from contextlib import contextmanager


class _BeatProcessorWorker:
    def __init__(self):
        self.q_new = Queue()
        self.q_ready = Queue()
        self.worker_thd = Thread(target=self._worker_thread_entry, daemon=True)
        self.worker_thd.start()
        self.q_new.put(1)

    def _worker_thread_entry(self):
        from madmom.features import DBNDownBeatTrackingProcessor, RNNDownBeatProcessor

        while True:
            item = self.q_new.get()
            if item is None:
                break
            try:
                beat_processor = RNNDownBeatProcessor(fps=100)
                down_beat_processor = DBNDownBeatTrackingProcessor(
                    beats_per_bar=[3, 4], fps=100
                )
                self.q_ready.put((beat_processor, down_beat_processor))
            except Exception as e:
                print(f"error constructing madmom resources: {e}", file=sys.stderr)
                sys.exit(1)

    @contextmanager
    def processors(self):
        p = self.q_ready.get()
        try:
            yield p
        finally:
            self.q_new.put(1)


_global_beat_processor_worker = None


def _create_worker():
    global _global_beat_processor_worker
    if _global_beat_processor_worker is None:
        _global_beat_processor_worker = _BeatProcessorWorker()


def _detect_beats(wav_path):
    with _global_beat_processor_worker.processors() as (
        beat_processor,
        down_beat_processor,
    ):
        dist = beat_processor.process(str(wav_path))
        return down_beat_processor.process(dist)


class BeatProcessor:
    def __init__(self, num_workers=1):
        self.pool = multiprocessing.Pool(num_workers, initializer=_create_worker)

    def close(self):
        self.pool.close()
        self.pool.join()

    def detect_beats(self, wav_path):
        return self.pool.apply(_detect_beats, (str(wav_path),))
