from fastapi import FastAPI, HTTPException
from fastapi.responses import JSONResponse
import os
import sys
import uuid
import io
import base64
import logging
import re
from bs4 import BeautifulSoup
import html
from typing import List
from pydantic import BaseModel, Field

import torch
import numpy as np
import soundfile as sf

# ===================== LOGGING =====================
logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(message)s"
)
logger = logging.getLogger(__name__)

# ===================== ENV ACTIVATION =====================
ENV_PATH = "/var/www/eduai.edurigo.com/doc_train/edurigo_ai/Puru/ai4bharat/ai4bharat_env"

def activate_environment():
    site_packages = os.path.join(ENV_PATH, "lib", "python3.10", "site-packages")
    if site_packages not in sys.path:
        sys.path.insert(0, site_packages)

    os.environ["VIRTUAL_ENV"] = ENV_PATH
    os.environ["PATH"] = os.path.join(ENV_PATH, "bin") + os.pathsep + os.environ.get("PATH", "")

activate_environment()

# ===================== HF IMPORTS =====================
from transformers import VitsModel, AutoTokenizer

# ===================== FASTAPI =====================
app = FastAPI()

# ===================== MODEL CACHE =====================
MODEL_CACHE = {}

def load_tts_model(language_code: str):
    language_code = language_code.strip().lower()

    if not language_code:
        raise ValueError("Language code is empty")

    if language_code not in MODEL_CACHE:
        model_name = f"facebook/mms-tts-{language_code}"
        logger.info("Loading TTS model: %s", model_name)

        try:
            model = VitsModel.from_pretrained(model_name)
            tokenizer = AutoTokenizer.from_pretrained(model_name)
        except Exception:
            raise ValueError(f"TTS model not found for language code '{language_code}'")

        model.eval()
        MODEL_CACHE[language_code] = (model, tokenizer)

    return MODEL_CACHE[language_code]

# ===================== TEXT CLEANER =====================
def clean_text(text: str) -> str:

    # Decode HTML entities
    text = html.unescape(text)

    # Remove HTML tags safely
    soup = BeautifulSoup(text, "html.parser")
    text = soup.get_text(separator=" ")

    # Remove invalid attributes remnants
    text = re.sub(r"xss\s*=\s*removed", "", text, flags=re.IGNORECASE)

    # Replace abbreviations for better Hindi pronunciation
    replacements = {
        "P-M-S": "पी एम एस",
        "U-P-S-I": "यू पी एस आई",
        "C-R-M": "सी आर एम",
    }

    for k, v in replacements.items():
        text = text.replace(k, v)

    # Remove emojis
    emoji_pattern = re.compile(
        "["
        "\U0001F600-\U0001F64F"
        "\U0001F300-\U0001F5FF"
        "\U0001F680-\U0001F6FF"
        "\U0001F1E0-\U0001F1FF"
        "]+",
        flags=re.UNICODE
    )
    text = emoji_pattern.sub("", text)

    # Remove weird unicode spaces
    text = text.replace("\xa0", " ")

    # Normalize spaces
    text = re.sub(r"\s+", " ", text).strip()

    return text
# ===================== REQUEST MODELS =====================
class TTSItem(BaseModel):
    id: int = Field(alias="contentId")
    contentLanguageId: int
    name: str
    backName: str = ""   # optional — use name as fallback if empty
    type: str

    class Config:
        populate_by_name = True


class TTSBatchRequest(BaseModel):
    clientId: int
    storigoId: int
    host: str
    hostType: str
    storigoLanguage: str   # expects MMS language code (dynamic)
    storigoLanguageName: str
    languageData: List[TTSItem]

# ===================== API =====================
@app.post("/generate-storigo-content-hindi-file")
async def generate_tts_batch(request: TTSBatchRequest):

    if not request.languageData:
        raise HTTPException(status_code=400, detail="languageData cannot be empty")

    try:
        language_code = request.storigoLanguage
        model, tokenizer = load_tts_model(language_code)

        # ── helper: text → base64 MP3 ──────────────────────────────────────
        def generate_audio_base64(text: str) -> tuple[str, str, int]:
            """Returns (audio_base64, filename, token_count)"""
            cleaned = clean_text(text)
            if not cleaned:
                raise ValueError("Text empty after cleaning")

            logger.info("CLEANED TEXT: %s", cleaned)

            inputs = tokenizer(cleaned, return_tensors="pt")
            if inputs["input_ids"].shape[1] == 0:
                raise ValueError("Tokenization produced empty input")

            inputs["input_ids"] = inputs["input_ids"].long()
            if "attention_mask" in inputs:
                inputs["attention_mask"] = inputs["attention_mask"].long()

            n_tokens = inputs["input_ids"].shape[1]

            with torch.no_grad():
                waveform = model(**inputs).waveform

            audio_np = waveform.squeeze().cpu().numpy().astype(np.float32)
            buf = io.BytesIO()
            sf.write(buf, audio_np, samplerate=16000, format="MP3")
            b64 = base64.b64encode(buf.getvalue()).decode("utf-8")

            fname = (
                f"techno_{item.contentLanguageId}_"
                f"{request.storigoLanguageName}_"
                f"{uuid.uuid4().hex[:8]}.mp3"
            )
            return b64, fname, n_tokens

        results = []
        token_usage = {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}

        for item in request.languageData:
            try:
                has_name     = bool(item.name and item.name.strip())
                has_backname = bool(item.backName and item.backName.strip())

                if not has_name and not has_backname:
                    raise ValueError("Both name and backName are empty")

                # ── CASE: both name and backName present → two audios ──────
                if has_name and has_backname:
                    name_b64, name_file, name_tokens = generate_audio_base64(item.name)
                    back_b64, back_file, back_tokens = generate_audio_base64(item.backName)

                    token_usage["input_tokens"] += name_tokens + back_tokens

                    results.append({
                        "id": item.id,
                        "success": True,
                        "contentLanguageId": item.contentLanguageId,
                        "storigoLanguageName": request.storigoLanguageName,
                        # name audio
                        "name_content": clean_text(item.name),
                        "filename": name_file,
                        "audio_base64": name_b64,
                        # backName audio
                        "back_name_content": clean_text(item.backName),
                        "back_name_filename": back_file,
                        "back_name_audio_base64": back_b64,
                    })

                # ── CASE: only one present → single audio ──────────────────
                else:
                    raw_text = item.name if has_name else item.backName
                    audio_b64, fname, n_tokens = generate_audio_base64(raw_text)
                    token_usage["input_tokens"] += n_tokens

                    results.append({
                        "id": item.id,
                        "success": True,
                        "filename": fname,
                        "audio_base64": audio_b64,
                        "content": clean_text(raw_text),
                        "contentLanguageId": item.contentLanguageId,
                        "storigoLanguageName": request.storigoLanguageName,
                    })

            except Exception as e:
                logger.error("TTS failed for ID %s: %s", item.id, str(e))
                results.append({
                    "id": item.id,
                    "success": False,
                    "error": str(e),
                    "contentType": item.type
                })

        token_usage["total_tokens"] = token_usage["input_tokens"]

        return JSONResponse(
            status_code=200,
            content={
                "success": True,
                "message": f"Processed {len(results)} items",
                "token_usage": token_usage,
                "results": results
            }
        )

    except Exception as e:
        logger.exception("Batch TTS generation failed")
        raise HTTPException(status_code=500, detail=str(e))
