import logging
import sys
import threading
import time

import numpy as np
import sounddevice as sd
import ctranslate2
from faster_whisper import WhisperModel
from huggingface_hub import snapshot_download
from huggingface_hub.utils import enable_progress_bars

_real_stdout = sys.stdout
_active_bar = None  # barre LoadingBar active, pour que les logs s'écrivent proprement sur leur propre ligne
_at_line_start = True  # le curseur est-il en début de ligne ? (les barres tqdm de HF écrivent sans \n final)


def _write_raw(text):
	"""Écrit brut sur la sortie réelle en suivant la position du curseur."""
	global _at_line_start
	_real_stdout.write(text)
	_real_stdout.flush()
	if text:
		_at_line_start = text.endswith("\n")


class LineStartStdout:
	"""Passe-plat stdout/stderr : centralise l'écriture pour suivre la position du curseur."""

	def write(self, text):
		_write_raw(text)
		return len(text)

	def flush(self):
		_real_stdout.flush()

	def __getattr__(self, name):
		# délègue tout le reste (fileno, isatty, encoding...) au vrai stdout
		return getattr(_real_stdout, name)


sys.stdout = LineStartStdout()
sys.stderr = LineStartStdout()


class LoadingBar(threading.Thread):
	"""Barre indéterminée (bloc rebondissant) ; les logs passent au-dessus d'elle sans l'écraser."""

	def __init__(self, label, width=28, block=6):
		super().__init__(daemon=True)
		self.label = label
		self.width = width
		self.block = block
		self._stop_event = threading.Event()
		self._t0 = None
		self._line_len = 0

	def render_line(self, line):
		"""Dessine une ligne de barre à la position du curseur (et mémorise sa longueur)."""
		_write_raw("\r" + line)
		self._line_len = len(line)

	def pause_for_log(self):
		"""Efface la barre le temps d'écrire un log ; elle se redessine au tick suivant."""
		if self._line_len:
			_write_raw("\r" + (" " * self._line_len) + "\r")
			self._line_len = 0

	def run(self):
		global _active_bar
		_active_bar = self
		self._t0 = time.time()
		pos, direction = 0, 1
		while not self._stop_event.wait(0.1):
			cells = [" "] * self.width
			for i in range(self.block):
				if pos + i < self.width:
					cells[pos + i] = "="
			line = "[%s] %s (%3.0fs)" % ("".join(cells), self.label, time.time() - self._t0)
			self.render_line(line)
			pos += direction
			if pos + self.block >= self.width:
				direction = -1
			elif pos <= 0:
				direction = 1

	def stop(self):
		global _active_bar
		self._stop_event.set()
		self.join()
		self.pause_for_log()
		if _active_bar is self:
			_active_bar = None


class BarAwareHandler(logging.StreamHandler):
	"""Handler de log qui garantit la date en début de ligne, même après une barre de progression."""

	def emit(self, record):
		try:
			msg = self.format(record)
			bar = _active_bar
			if bar is not None and bar.is_alive():
				bar.pause_for_log()
			# une barre externe (tqdm HF...) peut avoir laissé le curseur en milieu de ligne
			if not _at_line_start:
				_write_raw("\n")
			_write_raw(msg + "\n")
		except Exception:
			self.handleError(record)


log = logging.getLogger("transcript-audio")
log.setLevel(logging.INFO)
_handler = BarAwareHandler(sys.stdout)
_handler.setFormatter(logging.Formatter("%(asctime)s [%(levelname)s] %(message)s"))
log.addHandler(_handler)

SAMPLE_RATE = 16000
BLOCK_SECONDS = 5  # transcrire toutes les 5 secondes d'audio
BLOCKSIZE = int(SAMPLE_RATE * BLOCK_SECONDS)

MODEL_REPO = "Systran/faster-whisper-large-v3"
MODEL_PATTERNS = [
	"config.json",
	"preprocessor_config.json",
	"model.bin",
	"tokenizer.json",
	"vocabulary.*",
]


log.info("Périphériques d'entrée disponibles :")
for i, dev in enumerate(sd.query_devices()):
	if dev["max_input_channels"] > 0:
		log.info("  [%d] %s (canaux=%d, taux=%d)", i, dev["name"], dev["max_input_channels"], dev["default_samplerate"])

default_in = sd.default.device[0]
if default_in < 0:
	log.error("Aucun périphérique d'entrée par défaut défini !")
else:
	log.info("Entrée par défaut : [%d] %s", default_in, sd.query_devices(default_in)["name"])

t0 = time.time()
log.info("Vérification du modèle %s (téléchargement avec barre de progression si incomplet)...", MODEL_REPO)
enable_progress_bars()
try:
	model_path = snapshot_download(repo_id=MODEL_REPO, allow_patterns=MODEL_PATTERNS)
except Exception:
	log.warning("Téléchargement en ligne impossible, tentative depuis le cache local...")
	try:
		model_path = snapshot_download(
			repo_id=MODEL_REPO,
			allow_patterns=MODEL_PATTERNS,
			local_files_only=True,
		)
	except Exception:
		log.exception("Modèle introuvable dans le cache et téléchargement impossible")
		raise
log.info("Fichiers du modèle prêts : %s", model_path)


def pick_device_and_compute_type():
	"""Choisit CUDA si un GPU NVIDIA est réellement utilisable, sinon retombe sur CPU."""
	if ctranslate2.get_cuda_device_count() > 0:
		try:
			probe = ctranslate2.models.Whisper(
				"tiny", device="CUDA", compute_type="float16"
			)
			del probe
			return "cuda", "float16"
		except Exception as e:
			log.warning("CUDA détecté mais inutilisable (%s) -> bascule sur CPU", e)
	log.info("Aucun GPU CUDA utilisable -> exécution sur CPU (int8)")
	return "cpu", "int8"


bar = LoadingBar("Chargement du modèle large-v3")
bar.start()
try:
	device, compute_type = pick_device_and_compute_type()
	log.info("Device sélectionné : %s (%s)", device, compute_type)
	model = WhisperModel(model_path, device=device, compute_type=compute_type)
finally:
	bar.stop()
log.info("Modèle chargé en %.1fs au total", time.time() - t0)


def callback(indata, frames, time_info, status):
	try:
		if status:
			log.warning("Statut du flux : %s", status)
		audio = indata.copy().astype(np.float32).flatten()
		segments, _ = model.transcribe(
			audio,
			language="fr",
			vad_filter=True,
			beam_size=5,
		)
		for segment in segments:
			_write_raw(segment.text.strip() + "\n")
	except Exception:
		log.exception("Erreur dans le callback audio")


try:
	with sd.InputStream(
		samplerate=SAMPLE_RATE,
		channels=1,
		dtype="float32",
		blocksize=BLOCKSIZE,
		callback=callback,
	) as stream:
		log.info(
			"Flux ouvert : samplerate=%s, blocksize=%s, canaux=%s, device=%s",
			stream.samplerate,
			stream.blocksize,
			stream.channels,
			stream.device,
		)
		print("Listening... appuyez sur Ctrl+C pour arrêter.")
		try:
			# boucle d'attente interruptible : le Ctrl+C est absorbé ici,
			# pas pendant la fermeture C du flux (source du traceback Pa_StopStream)
			while True:
				time.sleep(0.5)
		except KeyboardInterrupt:
			pass
	log.info("Transcription arrêtée. À bientôt !")
except KeyboardInterrupt:
	# Ctrl+C en rafale pendant la fermeture : on reste silencieux
	log.info("Transcription arrêtée. À bientôt !")
except Exception:
	log.exception("Erreur lors de l'ouverture du flux audio")