"""
routers/websocket.py
─────────────────────────────────────────────────────────────────
WebSocket endpoint: /ws/{session_id}?token=<jwt>

Protocol (JSON frames):
  CLIENT → SERVER:
    { "type": "text",        "content": "..." }
    { "type": "audio_b64",   "audio_b64": "<base64 webm/opus>" }
    { "type": "end_session" }
    { "type": "ping" }

  SERVER → CLIENT:
    { "type": "text_delta",  "content": "token" }
    { "type": "text_done",   "content": "full response", "turn": N }
    { "type": "audio_chunk", "audio_b64": "<base64 PCM16 @ 24kHz>" }
    { "type": "audio_done" }        ← signals end of audio stream for this turn
    { "type": "sentiment",   "data": {...} }
    { "type": "transcription","content": "..." }
    { "type": "error",       "content": "message" }
    { "type": "pong" }
    { "type": "session_end", "turn": N }

AUDIO FORMAT NOTE:
  audio_chunk payloads contain raw signed-16-bit little-endian PCM @ 24 kHz
  mono (no WAV header). The frontend AudioEngine decodes each chunk directly
  into an AudioBuffer — zero buffering, zero lag.
  Backend: openai_service.stream_tts_chunks must use response_format="pcm".
"""
from __future__ import annotations
import asyncio
import base64
import json
import logging
from datetime import datetime, timezone

from fastapi import APIRouter, Query, WebSocket, WebSocketDisconnect
from starlette.websockets import WebSocketState

from models.schemas import ChatMessage
from services import store
from services.auth import decode_token
from services.openai_service import (
    analyse_sentiment,
    stream_ai_response,
    stream_tts_chunks,
    transcribe_audio,
)
# video_engine is imported lazily inside the video branch — heavy imports
# (asyncio.subprocess, hashlib, file I/O) are fine, but keeping it lazy
# means a misconfigured Wav2Lip install can't take the WS router down at
# import time and silently 404 the whole role-play feature.

log = logging.getLogger(__name__)
router = APIRouter()


async def _send(ws: WebSocket, obj: dict) -> None:
    """Best-effort send. Skips if the socket has already disconnected so we
    don't trip the "WebSocket is not connected" RuntimeError storm Starlette
    raises after the client closes."""
    if ws.client_state != WebSocketState.CONNECTED:
        return
    try:
        await ws.send_text(json.dumps(obj))
    except Exception:
        pass


@router.websocket("/ws/{session_id}")
async def roleplay_ws(
    ws: WebSocket,
    session_id: str,
    token: str = Query(default=""),
):
    await ws.accept()
    log.info("[WS] connected  session=%s", session_id)

    # ── Auth: validate JWT token from query param ────────────────
    # The frontend passes ?token=<jwt> when opening the WebSocket.
    # Reject connections that carry an invalid or missing token.
    # Auth: accept either a JWT (legacy/internal flow) or an access_key
    # (new external People Hub login flow). Token may be passed either as
    # ?token=<jwt> or ?token=ak:<access_key>.
    auth_ok = False
    if token:
        try:
            if token.startswith("ak:"):
                # External access-key auth — validate non-empty for now;
                # future: re-fetch /login/by-accesskey to confirm validity.
                if len(token) > 3:
                    auth_ok = True
            else:
                decoded = decode_token(token)  # returns None if invalid
                auth_ok = bool(decoded)
        except Exception as exc:  # noqa: BLE001
            log.warning("[WS] auth error session=%s reason=%s", session_id, exc)

    if not auth_ok:
        # Allow unauthenticated connections only in dev mode
        from config import settings
        if not settings.APP_DEBUG:
            await _send(ws, {"type": "error", "content": "Authentication failed"})
            await ws.close(code=4401)
            return
        log.warning("[WS] no/invalid token — allowing in APP_DEBUG mode")

    # ── Load session & scenario ──────────────────────────────────
    session = await store.get_session(session_id)
    if not session:
        await _send(ws, {"type": "error", "content": "Session not found"})
        await ws.close()
        return

    scenario = await store.get_scenario(session.scenario_id)
    if not scenario:
        await _send(ws, {"type": "error", "content": "Scenario not found"})
        await ws.close()
        return

    # ── Per-connection state (mode + chosen avatar) ──────────────
    #
    # `mode` mirrors the client-side tab selection: 'chat' | 'voice' |
    # 'video'. The client sends `{type: "mode", mode: "..."}` whenever
    # the candidate flips tabs, plus once at startup. We default to
    # 'voice' so an old client (no mode message) keeps the existing
    # voice behaviour.
    #
    # `video_avatar` is the filename of the MP4 used for lip-sync. It
    # is picked once per session (deterministically from session_id) so
    # the candidate doesn't see the face change between turns.
    # `video_voice` is the OpenAI TTS voice ID paired with that avatar's
    # gender (see services/video_engine.voice_for_avatar). Pinning both
    # at session start guarantees the candidate never hears a female
    # voice over a male avatar (or vice versa) — the bug that this
    # state-pinning was added to fix. The voice is also threaded through
    # the voice-mode path below so the gender stays consistent if the
    # candidate flips between voice and video mid-session.
    conn_state: dict = {
        "mode":         "voice",
        "video_avatar": None,
        "video_voice":  None,
    }

    def _video_mode() -> bool:
        return conn_state.get("mode") == "video"

    def _tts_voice() -> str | None:
        """Voice to use for any TTS call on this connection. Returns
        the avatar-paired voice once it's been picked, else None so
        stream_tts_chunks() falls back to the env default."""
        return conn_state.get("video_voice")

    async def _ensure_avatar() -> str:
        """Pick an avatar lazily, on first use of video mode. Stored on
        the connection so subsequent turns reuse the same face — AND
        the same matched voice.

        Selection precedence:
          1. session.selected_avatar — set by the avatar-picker modal
             when IS_VIDEO_ROLEPLAY_SELECTION=true.
          2. seeded random pick keyed off session_id — the existing
             behaviour when the picker is disabled or the candidate
             didn't make a choice.
        """
        if conn_state["video_avatar"]:
            return conn_state["video_avatar"]
        from services import video_engine

        avatar = ""
        # 1) Explicit candidate choice from the picker. We still
        #    re-validate against the on-disk list so a stale session
        #    referencing a removed file falls through to random.
        picked = (getattr(session, "selected_avatar", None) or "").strip()
        if picked and picked in set(video_engine.list_available_avatars()):
            avatar = picked
            log.info("[WS] avatar from picker session=%s avatar=%s", session_id, avatar)
        # 2) Seeded random fallback.
        if not avatar:
            avatar = video_engine.pick_avatar(seed=session_id)
            log.info("[WS] avatar random session=%s avatar=%s", session_id, avatar)

        conn_state["video_avatar"] = avatar
        # Pin the matched voice on the same line as the avatar so the
        # two can never drift apart on subsequent turns.
        conn_state["video_voice"]  = video_engine.voice_for_avatar(avatar)
        # Tell the client which avatar to display while video renders —
        # the static MP4 plays muted as a placeholder, so the face on
        # screen during 'rendering…' is the same one in the final video.
        await _send(ws, {
            "type":       "video_avatar",
            "avatar":     avatar,
            "avatar_url": video_engine.avatar_url(avatar),
        })
        return avatar

    async def _stream_tts_response(full_response: str) -> None:
        """Emit TTS for the AI's reply according to the active mode.

        chat:  no TTS, no video — just the audio_done sentinel so the
               frontend state machine keeps cycling.
        voice: stream PCM chunks live as before; frontend plays them
               through AudioEngine.
        video: collect the full PCM into a buffer, run Wav2Lip, and
               send a video_ready frame with the cached/generated MP4
               URL. Audio chunks are NOT forwarded to the client in
               this mode — the video has its own audio track baked in,
               so playing both would double the voice.
        """
        if not scenario.voice_on:
            await _send(ws, {"type": "audio_done"})
            return

        if not _video_mode():
            # voice mode (or chat-but-voice_on, which acts like voice):
            # stream chunks immediately so the user hears speech start
            # within ~100 ms. _tts_voice() returns the avatar-paired
            # voice if one was pinned (e.g. the candidate was in video
            # mode earlier and switched back), else None → env default.
            async for audio_chunk in stream_tts_chunks(full_response, voice=_tts_voice()):
                await _send(ws, {"type": "audio_chunk", "audio_b64": audio_chunk})
            await _send(ws, {"type": "audio_done"})
            return

        # ── video mode ──────────────────────────────────────────────
        # 1. Tell the client we're rendering so it can show a state.
        await _send(ws, {"type": "video_rendering"})
        # 1.a Pick the avatar (and pair its voice) NOW, before we run
        #     the TTS — so the speech we synthesize is in the voice
        #     that matches the avatar the client will show. If we
        #     deferred this to the post-TTS render step the speech
        #     would already have been synthesized with a possibly-
        #     mismatched voice; the lip-sync render would then put
        #     that audio on the avatar's mouth → exactly the bug we
        #     are fixing.
        await _ensure_avatar()
        # 2. Collect the full PCM. The render needs the entire utterance.
        pcm_buf = bytearray()
        async for audio_chunk_b64 in stream_tts_chunks(full_response, voice=_tts_voice()):
            try:
                pcm_buf.extend(base64.b64decode(audio_chunk_b64))
            except Exception as exc:
                log.warning("[WS][video] bad TTS chunk: %s", exc)
        # 3. Pick avatar + render. video_engine.get_or_generate has its
        #    own two-tier strategy (Wav2Lip then ffmpeg dub), so the
        #    only way this raises is if BOTH fail — usually a missing
        #    avatar file or a broken ffmpeg install. In that case we
        #    fall back one more level to raw audio chunks so the
        #    candidate at least hears the AI.
        try:
            from services import video_engine
            avatar = await _ensure_avatar()
            result = await video_engine.get_or_generate(bytes(pcm_buf), avatar)
            await _send(ws, {
                "type":     "video_ready",
                "url":      result["url"],
                "avatar":   result["avatar"],
                "cached":   result["cached"],
                "fallback": result.get("fallback", False),
            })
        except Exception as exc:  # noqa: BLE001
            log.error("[WS][video] both render paths failed: %s", exc, exc_info=True)
            # Last-ditch fallback — stream the buffered PCM so the
            # candidate at least hears the AI rather than going silent.
            await _send(ws, {"type": "video_error", "content": str(exc)})
            if pcm_buf:
                CHUNK = 32 * 1024
                mv = memoryview(pcm_buf)
                for i in range(0, len(mv), CHUNK):
                    seg = bytes(mv[i:i + CHUNK])
                    await _send(ws, {
                        "type":     "audio_chunk",
                        "audio_b64": base64.b64encode(seg).decode("ascii"),
                    })
        # Always emit audio_done so the frontend's "AI turn finished"
        # state machine cycles (your-turn banner, mic re-enable, etc.).
        await _send(ws, {"type": "audio_done"})

    # ── Helpers ──────────────────────────────────────────────────

    # Words that count as "yes, I'm ready to start the interview" in
    # the readiness round. We match a permissive substring set rather
    # than a strict regex because the candidate may say "okay let's
    # do it", "sure thing", "I'm ready", etc., and STT outputs vary.
    _READY_WORDS = {
        "yes", "yeah", "yep", "yup", "sure", "ok", "okay",
        "ready", "let's", "lets", "go", "start", "begin",
        "alright", "all right", "absolutely", "of course", "sounds good",
        "हाँ", "हां", "ठीक", "तैयार",  # Hindi
    }

    def _is_ready_response(text: str) -> bool:
        t = (text or "").strip().lower()
        if not t:
            return False
        # If the candidate says ANY of these (even buried in a longer
        # sentence) we treat it as confirmation. False positives like
        # "Sorry, can you start over?" are tolerable — the AI just
        # begins the interview, no harm done.
        return any(w in t for w in _READY_WORDS)

    async def _send_greeting() -> None:
        """First AI turn: friendly greeting that doubles as a mic test.

        We deliberately don't run this through the LLM — keeping the
        wording deterministic means the cache key for the Wav2Lip /
        ffmpeg-dub render is stable across sessions, so every candidate
        after the first one gets an instant cached video instead of a
        30-90 s render.
        """
        name = (session.candidate_name or "").strip().split()[0] if session.candidate_name else ""
        if name:
            greeting = f"Hi {name}! How are you today? Shall we start the interview?"
        else:
            greeting = "Hi there! How are you today? Shall we start the interview?"

        # Persist the greeting as the first assistant message so the
        # rest of the flow (sentiment, transcript, report) sees a
        # complete conversation history.
        session.messages.append(ChatMessage(
            role="assistant",
            content=greeting,
            timestamp=datetime.now(timezone.utc).isoformat(),
        ))
        await store.save_session(session)

        # Stream as if it came from the model so the client's text
        # streaming UI animates naturally.
        for tok in greeting.split(" "):
            await _send(ws, {"type": "text_delta", "content": tok + " "})
        await _send(ws, {"type": "text_done", "content": greeting, "turn": 0})
        await _stream_tts_response(greeting)

    async def _process_learner_text(text: str) -> None:
        """Handle one learner turn: stream AI reply + optional TTS."""
        text = text.strip()
        if not text:
            return

        # ── Greeting / mic-check round ────────────────────────────
        # Until the candidate has confirmed they're ready, we don't
        # spend interview turns. This also gives us a cheap mic test:
        # if they can say "yes", their hardware is working.
        if not session.greeting_done:
            session.messages.append(ChatMessage(
                role="user", content=text,
                timestamp=datetime.now(timezone.utc).isoformat(),
            ))
            if _is_ready_response(text):
                session.greeting_done = True
                await store.save_session(session)
                # Generate the actual scenario opening from the LLM now
                # that we know the candidate is ready.
                opening_prompt = (
                    f"You are {scenario.ai_character}. The candidate has just confirmed they're "
                    "ready to begin. Start the interview now — open with your first scripted line "
                    "in character."
                )
                session.messages.append(ChatMessage(
                    role="user", content=opening_prompt,
                    timestamp=datetime.now(timezone.utc).isoformat(),
                ))
                full_opening = ""
                async for delta in stream_ai_response(session.messages, scenario):
                    full_opening += delta
                    await _send(ws, {"type": "text_delta", "content": delta})
                # Replace the prompt with a neutral marker and add the
                # actual assistant turn.
                session.messages[-1] = ChatMessage(role="user", content="[ready]")
                session.messages.append(ChatMessage(
                    role="assistant", content=full_opening,
                    timestamp=datetime.now(timezone.utc).isoformat(),
                ))
                session.turn = 1
                await store.save_session(session)
                await _send(ws, {"type": "text_done", "content": full_opening, "turn": session.turn})
                await _stream_tts_response(full_opening)
                return
            # Not ready yet — gentle nudge and re-greet. We don't bump
            # turn counter on these, they're free.
            nudge = ("No problem — just say 'yes' or 'I'm ready' "
                     "whenever you'd like to start the interview.")
            session.messages.append(ChatMessage(
                role="assistant", content=nudge,
                timestamp=datetime.now(timezone.utc).isoformat(),
            ))
            await store.save_session(session)
            for tok in nudge.split(" "):
                await _send(ws, {"type": "text_delta", "content": tok + " "})
            await _send(ws, {"type": "text_done", "content": nudge, "turn": 0})
            await _stream_tts_response(nudge)
            return

        # ── Normal interview turn ────────────────────────────────
        session.messages.append(ChatMessage(
            role="user",
            content=text,
            timestamp=datetime.now(timezone.utc).isoformat(),
        ))

        # ── Stream AI text ────────────────────────────────────────
        full_response = ""
        async for delta in stream_ai_response(session.messages, scenario):
            full_response += delta
            await _send(ws, {"type": "text_delta", "content": delta})

        session.turn += 1
        session.messages.append(ChatMessage(
            role="assistant",
            content=full_response,
            timestamp=datetime.now(timezone.utc).isoformat(),
        ))
        await store.save_session(session)

        await _send(ws, {
            "type": "text_done",
            "content": full_response,
            "turn": session.turn,
        })

        # ── Sentiment analysis (parallel, non-blocking) ───────────
        sentiment_task = asyncio.create_task(analyse_sentiment(text))

        # ── Stream TTS / render video per active mode ─────────────
        # Mode dispatch lives in _stream_tts_response — it picks chat /
        # voice / video and emits the right frames. Always sends an
        # audio_done sentinel so the frontend state machine resets.
        await _stream_tts_response(full_response)

        # ── Collect sentiment result ──────────────────────────────
        try:
            sentiment_result = await sentiment_task
            if isinstance(sentiment_result, dict):
                await _send(ws, {"type": "sentiment", "data": sentiment_result})
        except Exception as exc:
            log.warning("[WS] sentiment skipped: %s", exc)

        # ── Check turn limit ──────────────────────────────────────
        if session.turn >= scenario.turns:
            await _send(ws, {"type": "session_end", "turn": session.turn})

    async def _process_audio(audio_b64: str) -> None:
        """Transcribe audio blob then process as text."""
        try:
            raw  = base64.b64decode(audio_b64)
            text = await transcribe_audio(raw, session.language_code)
            if text:
                await _send(ws, {"type": "transcription", "content": text})
                await _process_learner_text(text)
            else:
                # Empty transcription — re-enable mic on frontend
                await _send(ws, {"type": "transcription", "content": ""})
        except Exception as exc:
            log.error("[WS] audio error  session=%s: %s", session_id, exc)
            await _send(ws, {"type": "error", "content": f"Audio processing failed: {exc}"})

    # ── Wait briefly for the client's first `mode` frame ────────────
    #
    # Why: the opening AI turn fires immediately after `accept`. If we
    # process it before the client tells us its mode, we default to
    # 'voice' and the greeting goes out as raw audio_chunks — which is
    # bad in two ways for video-mode users:
    #   1. Audio plays through AudioEngine WHILE the init loader is up,
    #      so the candidate hears the AI before they see the avatar.
    #   2. The opener is never lip-synced (no Wav2Lip render), so the
    #      first thing the candidate sees is a static face talking.
    #
    # We block here for up to 1.5 s waiting for ONE frame. If the first
    # frame is `mode`, we apply it and proceed. If it's anything else
    # (e.g. an old client that never sends `mode`), we put it back into
    # the receive queue by handling it inline (so it isn't lost).
    async def _drain_first_frame_or_timeout(timeout_s: float = 1.5):
        try:
            raw = await asyncio.wait_for(ws.receive_text(), timeout=timeout_s)
        except (asyncio.TimeoutError, Exception):
            return None
        try:
            return json.loads(raw)
        except json.JSONDecodeError:
            return None

    pending_first_msg = await _drain_first_frame_or_timeout()
    if pending_first_msg and pending_first_msg.get("type") == "mode":
        new_mode = (pending_first_msg.get("mode") or "voice").strip()
        if new_mode in ("chat", "voice", "video"):
            conn_state["mode"] = new_mode
            log.info("[WS] mode (pre-open)=%s session=%s", new_mode, session_id)
            if new_mode == "video":
                try: await _ensure_avatar()
                except Exception as exc:
                    log.warning("[WS] avatar pick failed: %s", exc)
        # Consumed — don't put back.
        pending_first_msg = None

    # ── Opening AI turn (first connection) ───────────────────────
    # The first thing the candidate hears is a friendly greeting that
    # also doubles as a mic check. The actual interview only begins
    # after they confirm with "yes / sure / ok / ready". This lives in
    # _send_greeting / the greeting_done branch in _process_learner_text.
    try:
        if not session.messages:
            await _send_greeting()

        # ── Main receive loop ─────────────────────────────────────
        # If a non-mode frame arrived during the pre-open drain, dispatch
        # it now so it isn't dropped.
        async def _dispatch(msg: dict) -> bool:
            mtype = msg.get("type", "")
            if mtype == "ping":
                await _send(ws, {"type": "pong"})
            elif mtype == "text":
                await _process_learner_text(msg.get("content", ""))
            elif mtype == "audio_b64":
                await _process_audio(msg.get("audio_b64", ""))
            elif mtype == "mode":
                new_mode = (msg.get("mode") or "voice").strip()
                if new_mode in ("chat", "voice", "video"):
                    conn_state["mode"] = new_mode
                    log.info("[WS] mode=%s session=%s", new_mode, session_id)
                    if new_mode == "video":
                        try: await _ensure_avatar()
                        except Exception as exc:
                            log.warning("[WS] avatar pick failed: %s", exc)
            elif mtype == "end_session":
                await _send(ws, {"type": "session_end", "turn": session.turn})
                return False  # stop loop
            else:
                log.debug("[WS] unknown message type=%s  session=%s", mtype, session_id)
            return True

        if pending_first_msg:
            keep_going = await _dispatch(pending_first_msg)
            if not keep_going:
                raise WebSocketDisconnect()

        while True:
            raw_msg = await ws.receive_text()
            try:
                msg = json.loads(raw_msg)
            except json.JSONDecodeError:
                log.warning("[WS] invalid JSON  session=%s", session_id)
                continue
            keep_going = await _dispatch(msg)
            if not keep_going:
                break

    except WebSocketDisconnect:
        log.info("[WS] disconnected  session=%s", session_id)
    except RuntimeError as exc:
        # Starlette raises plain RuntimeError ("WebSocket is not connected.
        # Need to call \"accept\" first.") when the peer closes mid-receive
        # — treat the same as a clean disconnect, no stack trace.
        if "not connected" in str(exc).lower() or "accept" in str(exc).lower():
            log.info("[WS] client closed mid-stream  session=%s", session_id)
        else:
            log.error("[WS] runtime error  session=%s: %s", session_id, exc, exc_info=True)
            await _send(ws, {"type": "error", "content": "Session ended unexpectedly."})
    except Exception as exc:  # noqa: BLE001
        log.error("[WS] error  session=%s: %s", session_id, exc, exc_info=True)
        # Never echo the raw exception text — it may contain provider URLs
        # or stack hints. Frontend just needs to know to recover.
        await _send(ws, {"type": "error", "content": "Session ended unexpectedly."})
    finally:
        try:
            await store.save_session(session)
        except Exception as exc:  # noqa: BLE001
            log.warning("[WS] save_session at close failed: %s", exc)
        # Only close if the socket is still connected; calling close on a
        # disconnected socket also raises RuntimeError under Starlette.
        if ws.client_state == WebSocketState.CONNECTED:
            try:
                await ws.close()
            except Exception:  # noqa: BLE001
                pass
