from __future__ import annotations

import io
import logging
import math
import os
import time
from pathlib import Path
from typing import Any

import numpy as np
import soundfile as sf

try:
    import torch
except ImportError:  # pragma: no cover - optional until a model is installed
    torch = None

LOGGER = logging.getLogger("kaonic.tts")


class KaonicTTS:
    def __init__(self, model_path: str | None = None) -> None:
        self.model_path = Path(model_path or os.getenv("KAONIC_TTS_MODEL", "models/kaonic_voice_1.pt"))
        self.sample_rate = int(os.getenv("KAONIC_TTS_SAMPLE_RATE", "22050"))
        self.device = "cuda" if torch is not None and torch.cuda.is_available() else "cpu"
        self.model: Any = None
        self.model_status = "fallback"
        if self.model_path.is_file() and torch is not None:
            try:
                self.model = torch.jit.load(str(self.model_path), map_location=self.device)
                self.model.eval()
                self.model_status = "loaded"
                LOGGER.info("Loaded Kaonic TTS model path=%s device=%s", self.model_path, self.device)
            except Exception:
                LOGGER.exception("Failed to load Kaonic TTS model path=%s", self.model_path)
        else:
            LOGGER.warning("Kaonic TTS model unavailable; using deterministic local fallback path=%s", self.model_path)

    def synthesize(self, text: str, voice: str, duration: int) -> np.ndarray:
        if self.model is not None:
            try:
                with torch.no_grad():
                    output = self.model.generate(text=text, voice=voice, duration=duration)
                audio = output.detach().float().cpu().numpy()
                return self._normalize(audio)
            except Exception:
                LOGGER.exception("Model inference failed; falling back to local synthesis")
        return self._fallback(text, duration)

    def encode_wav(self, audio: np.ndarray) -> bytes:
        output = io.BytesIO()
        sf.write(output, self._normalize(audio), self.sample_rate, format="WAV", subtype="PCM_16")
        return output.getvalue()

    def _normalize(self, audio: np.ndarray) -> np.ndarray:
        audio = np.asarray(audio, dtype=np.float32).reshape(-1)
        peak = float(np.max(np.abs(audio))) if audio.size else 0
        if peak == 0:
            return audio
        return np.clip(audio * (10 ** (-1 / 20)) / peak, -1, 1)

    def _fallback(self, text: str, duration: int) -> np.ndarray:
        samples = self.sample_rate * max(1, min(30, int(duration or 8)))
        seed = sum((i + 1) * ord(char) for i, char in enumerate(text)) & 0xFFFFFFFF
        base = 150 + seed % 180
        t = np.arange(samples, dtype=np.float32) / self.sample_rate
        modulation = 0.5 + 0.5 * np.sin(t * (2 + seed % 3) * math.pi)
        envelope = np.minimum(np.minimum(t * 8, (samples / self.sample_rate - t) * 8), 1)
        return (0.42 * np.sin(2 * math.pi * (base + modulation * 35) * t) +
                0.32 * np.sin(2 * math.pi * base * 2 * t)) * np.maximum(envelope, 0)

    def health(self) -> dict[str, str]:
        return {"status": "ok", "model": self.model_status, "device": self.device}
