import asyncio
import base64
import json
import logging
import struct

import httpx
from sarvamai import AsyncSarvamAI
from sarvamai.speech_to_text_realtime_streaming.socket_client import AsyncSpeechToTextRealtimeStreamingSocketClient
from sarvamai.types import RealtimeAudioInput, RealtimeFlush

from app.core.config import settings

logger = logging.getLogger(__name__)

SARVAM_STT_URL = "https://api.sarvam.ai/speech-to-text"
SARVAM_BATCH_BASE_URL = "https://api.sarvam.ai/speech-to-text/job/v1"
SARVAM_TRANSLATE_URL = "https://api.sarvam.ai/translate"
SARVAM_TTS_URL = "https://api.sarvam.ai/text-to-speech"

SUPPORTED_LANGUAGES = (
    "te-IN", "hi-IN", "ta-IN", "en-IN",
    "bn-IN", "gu-IN", "kn-IN", "ml-IN", "mr-IN", "pa-IN", "od-IN",
)

_LANGUAGE_SPEAKER = {
    "te-IN": "anushka",
    "hi-IN": "manisha",
    "ta-IN": "vidya",
    "en-IN": "anushka",
}

_http_client = httpx.AsyncClient()

MAX_POLL_ATTEMPTS = 60       # 30s total (0.5s × 60) — plain transcription (structured mode)
POLL_INTERVAL_S = 0.5
MAX_PENDING_POLLS = 60       # bail out after 30s stuck in Pending — Sarvam's diarized queue
                             # has been seen backlogged well past the original 15s ceiling

# Diarization is noticeably slower server-side than plain transcription — 30s was routinely
# too short and the job was still "Running" (not stuck in Pending) when it hit that ceiling.
# 90s gives real jobs room to finish while staying under the frontend's 120s override for
# this endpoint (see submitConversationSegment in frontend/src/api/sessions.ts).
DIARIZED_MAX_POLL_ATTEMPTS = 180   # 90s total (0.5s × 180)

# Upper bound, not a requirement — covers doctor + patient, plus an accompanying family
# member who sometimes joins in. Diarization still resolves to fewer speaker ids when
# only two people are actually talking.
CONVERSATION_NUM_SPEAKERS = 3

STREAM_SAMPLE_RATE = 16000
STREAM_RECONNECT_ATTEMPTS = 3
STREAM_RECONNECT_BACKOFF_S = (1.0, 2.0, 4.0)

# One shared SDK client, mirroring the module-level _http_client above — the SDK's
# httpx client underneath is reused across every SarvamStreamSession connection.
_sarvam_client = AsyncSarvamAI(api_subscription_key=settings.SARVAM_API_KEY)


class SarvamStreamSession:
    """One persistent Sarvam speech-to-text-realtime WebSocket connection (saaras:v3-realtime),
    carrying one conversation-mode session's live audio. Unlike the batch path above, this
    endpoint does not support diarization — every entry comes back as a single speaker
    stream, so callers treat this as speakerId "0" throughout.

    Runs with endpointing="vad" (the default) so Sarvam's own VAD marks utterance
    boundaries — flush() and speech_start/speech_end are no-ops in that mode; they only
    take effect under endpointing="manual", which this session does not use."""

    def __init__(self, language: str = "te-IN", sample_rate: int = STREAM_SAMPLE_RATE):
        self._language = language
        self._sample_rate = sample_rate
        # speech_to_text_realtime_streaming.connect() is an @asynccontextmanager — entered
        # and exited manually (rather than via `async with`) so the connection can outlive
        # a single call and be torn down independently on reconnect/close.
        self._ctx = None
        self._ws: AsyncSpeechToTextRealtimeStreamingSocketClient | None = None

    async def connect(self) -> None:
        self._ctx = _sarvam_client.speech_to_text_realtime_streaming.connect(
            language_code=self._language,
            model="saaras:v3-realtime",
            mode="translate",
            sample_rate=str(self._sample_rate),
            encoding="linear16",
            endpointing="vad",
        )
        self._ws = await self._ctx.__aenter__()

    async def send_audio(self, pcm_bytes: bytes) -> None:
        if self._ws is None:
            raise RuntimeError("SarvamStreamSession.send_audio called before connect()")
        await self._ws.send_realtime_audio_input(
            RealtimeAudioInput(audio=base64.b64encode(pcm_bytes).decode("ascii"))
        )

    async def flush(self) -> None:
        # No-op under endpointing="vad" — kept so the periodic stall-nudge call site in
        # sessions.py doesn't need special-casing; Sarvam's own VAD is what actually
        # closes out a stalled continuous-talker utterance on this endpoint.
        if self._ws is not None:
            await self._ws.send_realtime_flush(RealtimeFlush())

    async def events(self):
        if self._ws is None:
            raise RuntimeError("SarvamStreamSession.events called before connect()")
        async for response in self._ws:
            yield response.dict()

    async def close(self) -> None:
        if self._ctx is not None:
            await self._ctx.__aexit__(None, None, None)
            self._ctx = None
            self._ws = None


def pcm_to_wav_bytes(pcm_bytes: bytes, sample_rate: int = STREAM_SAMPLE_RATE) -> bytes:
    """Wraps raw 16-bit mono PCM in a WAV header — used to submit audio buffered during a
    Sarvam streaming outage through the batch STT path above, which expects a file."""
    channels = 1
    bits_per_sample = 16
    byte_rate = sample_rate * channels * bits_per_sample // 8
    block_align = channels * bits_per_sample // 8
    header = struct.pack(
        "<4sI4s4sIHHIIHH4sI",
        b"RIFF", 36 + len(pcm_bytes), b"WAVE",
        b"fmt ", 16, 1, channels, sample_rate, byte_rate, block_align, bits_per_sample,
        b"data", len(pcm_bytes),
    )
    return header + pcm_bytes


async def _run_batch_stt_job(
    audio_bytes: bytes, filename: str, mime_type: str, job_parameters: dict,
    max_poll_attempts: int = MAX_POLL_ATTEMPTS,
) -> dict:
    """Runs the full Sarvam batch STT job lifecycle (create → upload → start → poll →
    download → fetch) and returns the parsed file_result (dict or list, depending on
    whether diarization was requested). Shared by transcribe_audio and
    transcribe_audio_diarized — only job_parameters, max_poll_attempts, and the caller's
    result parsing differ."""
    headers = {"api-subscription-key": settings.SARVAM_API_KEY, "Content-Type": "application/json"}

    # Step 1 — create batch job
    create_resp = await _http_client.post(
        SARVAM_BATCH_BASE_URL,
        headers=headers,
        json={"job_parameters": job_parameters},
    )
    logger.warning("STT step1 create %d: %s", create_resp.status_code, create_resp.text[:300])
    create_resp.raise_for_status()
    create_data = create_resp.json()
    job_id = create_data.get("job_id") or create_data.get("id")
    if not job_id:
        raise RuntimeError(f"No job_id in create response: {create_data}")
    logger.warning("STT job created: %s", job_id)

    # Step 2 — get presigned upload URL
    upload_resp = await _http_client.post(
        f"{SARVAM_BATCH_BASE_URL}/upload-files",
        headers=headers,
        json={"job_id": job_id, "files": [filename]},
    )
    logger.warning("STT step2 upload-files %d: %s", upload_resp.status_code, upload_resp.text[:300])
    upload_resp.raise_for_status()
    upload_data = upload_resp.json()
    upload_urls = upload_data.get("upload_urls", upload_data)
    url_entry = upload_urls.get(filename) or next(iter(upload_urls.values()), None)
    if not url_entry:
        raise RuntimeError(f"No presigned upload URL for {filename}: {upload_data}")
    presigned_url = url_entry["file_url"] if isinstance(url_entry, dict) else url_entry

    # Step 3 — upload audio to Azure presigned URL
    base_mime = mime_type.split(";")[0].strip()
    put_resp = await _http_client.put(
        presigned_url,
        content=audio_bytes,
        headers={"Content-Type": base_mime, "x-ms-blob-type": "BlockBlob"},
        timeout=30.0,
    )
    logger.warning("STT step3 upload audio %d", put_resp.status_code)
    put_resp.raise_for_status()

    # Step 4 — start the job
    start_resp = await _http_client.post(
        f"{SARVAM_BATCH_BASE_URL}/{job_id}/start",
        headers=headers,
        json={},
    )
    logger.warning("STT step4 start %d: %s", start_resp.status_code, start_resp.text[:200])
    start_resp.raise_for_status()

    # Step 5 — poll until completed
    status_data: dict = {}
    pending_polls = 0
    for attempt in range(max_poll_attempts):
        await asyncio.sleep(POLL_INTERVAL_S)
        status_resp = await _http_client.get(
            f"{SARVAM_BATCH_BASE_URL}/{job_id}/status",
            headers={"api-subscription-key": settings.SARVAM_API_KEY},
        )
        status_resp.raise_for_status()
        status_data = status_resp.json()
        job_state = status_data.get("job_state", status_data.get("status", ""))
        if attempt % 20 == 0:
            logger.warning("STT poll #%d state=%s", attempt, job_state)
        if job_state in ("Completed", "completed", "SUCCESS", "success"):
            logger.warning("STT step5 completed after %d polls: %s", attempt, str(status_data)[:300])
            break
        if job_state in ("Failed", "failed", "ERROR", "error"):
            raise RuntimeError(f"Batch STT job failed: {status_data}")
        if job_state in ("Pending", "pending"):
            pending_polls += 1
            if pending_polls >= MAX_PENDING_POLLS:
                raise RuntimeError(f"Batch STT job stuck in Pending for {pending_polls} polls — Sarvam queue may be backlogged")
        else:
            pending_polls = 0
    else:
        raise RuntimeError(f"Batch STT timed out after {max_poll_attempts * POLL_INTERVAL_S:.0f}s, last state: {status_data}")

    # Step 6 — extract output filename from status
    logger.warning("STT step6 status_data: %s", str(status_data)[:500])
    job_details = status_data.get("job_details") or status_data.get("details") or []
    if job_details and isinstance(job_details, list):
        outputs = job_details[0].get("outputs") or []
        output_filename = outputs[0].get("file_name") or outputs[0].get("filename") if outputs else None
    else:
        output_filename = None

    if not output_filename:
        # fallback: derive output filename from input filename (common pattern)
        base = filename.rsplit(".", 1)[0]
        output_filename = f"{base}.json"
        logger.warning("STT step6 no output filename in status, guessing: %s", output_filename)

    # Step 7 — get presigned download URL
    download_resp = await _http_client.post(
        f"{SARVAM_BATCH_BASE_URL}/download-files",
        headers=headers,
        json={"job_id": job_id, "files": [output_filename]},
    )
    logger.warning("STT step7 download-files %d: %s", download_resp.status_code, download_resp.text[:300])
    download_resp.raise_for_status()
    download_result = download_resp.json()
    download_urls = download_result.get("download_urls", download_result)
    url_entry = download_urls.get(output_filename) or next(iter(download_urls.values()), None)
    if not url_entry:
        raise RuntimeError(f"No download URL for {output_filename}: {download_result}")
    if isinstance(url_entry, dict):
        presigned_download_url = url_entry.get("file_url") or url_entry.get("url") or ""
        if not presigned_download_url:
            raise RuntimeError(f"download_urls entry has no file_url/url key: {url_entry}")
    else:
        presigned_download_url = str(url_entry)
    if not presigned_download_url.startswith("http"):
        raise RuntimeError(f"presigned_download_url is not a valid URL: {presigned_download_url!r}")

    # Step 8 — fetch transcript content from presigned URL
    fetch_resp = await _http_client.get(presigned_download_url, timeout=30.0)
    logger.warning("STT step8 fetch %d content-type=%s body=%s", fetch_resp.status_code, fetch_resp.headers.get("content-type"), fetch_resp.text[:300])
    fetch_resp.raise_for_status()

    raw = fetch_resp.text.strip()
    try:
        file_result = json.loads(raw)
    except json.JSONDecodeError:
        file_result = raw

    if isinstance(file_result, list):
        file_result = file_result[0] if file_result else {}

    logger.warning("STT job %s done — file_result: %s", job_id, str(file_result)[:500])
    return file_result if isinstance(file_result, dict) else {"transcript": str(file_result).strip()}


async def transcribe_audio(audio_bytes: bytes, filename: str, mime_type: str, language: str = "te-IN") -> tuple[str, str]:
    file_result = await _run_batch_stt_job(
        audio_bytes, filename, mime_type,
        {"language_code": "unknown", "model": "saaras:v3", "mode": "transcribe"},
    )
    transcript = file_result.get("transcript", "")
    detected_lang = file_result.get("language_code") or language
    logger.warning("STT done — detected: %s, transcript: %r", detected_lang, transcript)
    return transcript, detected_lang


async def transcribe_audio_diarized(
    audio_bytes: bytes, filename: str, mime_type: str, language: str = "te-IN", num_speakers: int = CONVERSATION_NUM_SPEAKERS,
) -> tuple[list[dict], str]:
    """Same batch pipeline as transcribe_audio, but with Sarvam speaker diarization enabled —
    used by free-conversation intake mode, where one audio segment can contain more than one
    speaker's turn (doctor + patient, or patient + an accompanying family member). Returns a
    list of {"speakerId": ..., "text": ...} entries in chronological order, plus the detected
    language for the segment as a whole (diarization does not give a language per entry)."""
    file_result = await _run_batch_stt_job(
        audio_bytes, filename, mime_type,
        {
            "language_code": "unknown", "model": "saaras:v3", "mode": "transcribe",
            "with_diarization": True, "num_speakers": num_speakers,
        },
        max_poll_attempts=DIARIZED_MAX_POLL_ATTEMPTS,
    )
    raw_entries = (file_result.get("diarized_transcript") or {}).get("entries") or []
    entries = [
        {"speakerId": str(e.get("speaker_id", "")), "text": (e.get("transcript") or "").strip()}
        for e in raw_entries
        if (e.get("transcript") or "").strip()
    ]
    detected_lang = file_result.get("language_code") or language
    logger.warning("Diarized STT done — detected: %s, entries: %r", detected_lang, entries)
    return entries, detected_lang


async def translate_to_english(text: str, source_language_code: str = "te-IN") -> str:
    if not text.strip():
        return text
    response = await _http_client.post(
        SARVAM_TRANSLATE_URL,
        headers={
            "api-subscription-key": settings.SARVAM_API_KEY,
            "Content-Type": "application/json",
        },
        json={
            "input": text,
            "source_language_code": source_language_code,
            "target_language_code": "en-IN",
            "speaker_gender": "Female",
            "mode": "formal",
            "model": "mayura:v1",
            "enable_preprocessing": False,
        },
    )
    response.raise_for_status()
    data = response.json()
    translated = data.get("translated_text", text)
    logger.warning("Translate %s → en: %r → %r", source_language_code, text, translated)
    return translated


async def translate_to_language(text: str, target_language_code: str) -> str:
    if not text.strip() or target_language_code == "en-IN":
        return text
    response = await _http_client.post(
        SARVAM_TRANSLATE_URL,
        headers={
            "api-subscription-key": settings.SARVAM_API_KEY,
            "Content-Type": "application/json",
        },
        json={
            "input": text,
            "source_language_code": "en-IN",
            "target_language_code": target_language_code,
            "speaker_gender": "Female",
            "mode": "formal",
            "model": "mayura:v1",
            "enable_preprocessing": False,
        },
    )
    response.raise_for_status()
    data = response.json()
    translated = data.get("translated_text", text)
    logger.warning("Translate en → %s: %r → %r", target_language_code, text, translated)
    return translated


async def text_to_speech(text: str, language_code: str) -> bytes:
    speaker = _LANGUAGE_SPEAKER.get(language_code, "anushka")
    response = await _http_client.post(
        SARVAM_TTS_URL,
        headers={
            "api-subscription-key": settings.SARVAM_API_KEY,
            "Content-Type": "application/json",
        },
        json={
            "inputs": [text],
            "target_language_code": language_code,
            "speaker": speaker,
            "model": "bulbul:v2",
        },
        timeout=15.0,
    )
    if response.status_code != 200:
        logger.error("Sarvam TTS error %d: %s", response.status_code, response.text)
    response.raise_for_status()
    data = response.json()
    audio_b64 = data["audios"][0]
    return base64.b64decode(audio_b64)
