from fastapi import FastAPI, Query
from fastapi.responses import JSONResponse
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, VitsModel
import torch
import base64
import io
import soundfile as sf

import os
from huggingface_hub import login

HF_TOKEN = os.getenv("HF_TOKEN")
login(token=HF_TOKEN)

app = FastAPI(title="Madiak API (Bol)")

# Global device - define here before all functions
device = "cuda" if torch.cuda.is_available() else "cpu"

# Model registry
models = {
    "dik": {
        "id": "facebook/mms-tts-dik",
        "tokenizer": None,
        "model": None
    },
    "dip": {
        "id": "facebook/mms-tts-dip",
        "tokenizer": None,
        "model": None
    }
}
translation_model = {
    "id": "Madiak/Bol",
    "tokenizer": None,
    "Model": None
}

# Load models at startup
@app.on_event("startup")
def load_models():
    for key, cfg in models.items():
        print(f"Loading {cfg['id']}...")
        cfg["tokenizer"] = AutoTokenizer.from_pretrained(cfg["id"])
        cfg["model"] = VitsModel.from_pretrained(cfg["id"])
        cfg["model"].eval()
        print(f"Loaded {cfg['id']} successfully.")

    # Load translation model
    print(f"Loading translation model {translation_model['id']}...")
    translation_model["tokenizer"] = AutoTokenizer.from_pretrained(translation_model["id"], token=HF_TOKEN)
    translation_model["model"] = AutoModelForSeq2SeqLM.from_pretrained(translation_model["id"], token=HF_TOKEN)
    translation_model["model"].eval()
    print("Translation model loaded successfully.")

def generate_audio(model_key: str, text: str):
    cfg = models[model_key]
    tokenizer = cfg["tokenizer"]
    model = cfg["model"]

    # Tokenize
    inputs = tokenizer(text, return_tensors="pt")

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

    audio_np = output.squeeze().cpu().numpy()

    # Convert to WAV bytes
    buffer = io.BytesIO()
    sf.write(buffer, audio_np, 16000, format="WAV")
    wav_bytes = buffer.getvalue()

    # Base64 encode
    b64_audio = base64.b64encode(wav_bytes).decode("utf-8")

    return b64_audio


@app.get("/synthesize/{voice}")
def synthesize(voice: str, text: str = Query(..., description="Text to synthesize")):
    if voice not in models:
        return JSONResponse({"error": "Invalid voice. Use 'dik' or 'dip'."}, status_code=400)

    try:
        audio_b64 = generate_audio(voice, text)
        return {
            "voice": voice,
            "text": text,
            "audio_base64": audio_b64,
            "sampling_rate": 16000
        }
    except Exception as e:
        return JSONResponse({"error": str(e)}, status_code=500)

# ---------- Translation ----------

# Update LANG_MAP if needed (but we'll change the prompt)
LANG_MAP = {
    "en": "english",
    "dik": "dinka",
}

def build_prompt(text: str, source: str, target: str):
    source = source.lower().strip()
    target = target.lower().strip()

    if source == "en" and target == "dik":
        return f"Translate English to Dinka:\n{text}"
    if source == "dik" and target == "en":
        return f"Translate Dinka to English:\n{text}"

    raise ValueError("Unsupported language pair; only en↔dik supported.")

# Update translate_text to strip the prompt from output
def translate_text(text: str, source: str, target: str):
    tokenizer = translation_model["tokenizer"]
    model = translation_model["model"]

    prompt = build_prompt(text, source, target)

    inputs = tokenizer(prompt, return_tensors="pt")
    
    # Move inputs to correct device
    inputs = {k: v.to(device) for k, v in inputs.items()}

    with torch.no_grad():
        outputs = model.generate(**inputs, max_length=256)

    result = tokenizer.decode(outputs[0], skip_special_tokens=True)
    
    # Strip the prompt if present
    if result.startswith(prompt):
        result = result[len(prompt):].strip()
    
    return result


@app.get("/translate")
def translate(
    text: str = Query(...),
    source: str = Query("en"),
    target: str = Query("dik")
):
    try:
        # Validate and normalize language codes
        source = source.lower().strip()
        target = target.lower().strip()
        
        # Map variations to standard codes
        if source not in {"en", "dik"} or target not in {"en", "dik"}:
            return JSONResponse({"error": "Only 'en' and 'dik' are supported."}, status_code=400)
        if source == target:
            return JSONResponse({"error": "Source and target must differ."}, status_code=400)
        
        result = translate_text(text, source, target)

        return {
            "input": text,
            "source": source,
            "target": target,
            "translation": result
        }

    except Exception as e:
        return JSONResponse({"error": str(e)}, 500)
        
@app.get("/")
def root():
    return {
        "status": "ok",
        "voices": ["dik", "dip"],
        "example_endpoints": {
            "dik": "/synthesize/dik?text=Hello",
            "dip": "/synthesize/dip?text=Hello",
            "translate": "/translate?text=Hello"
        }
    }
