"""Objeto en memoria CallSession: métricas de una llamada activa."""

from __future__ import annotations

from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, Literal


def utc_now() -> datetime:
    return datetime.now(timezone.utc)


TurnRole = Literal["user", "bot"]


@dataclass
class TurnTiming:
    """Duración de una intervención del usuario o del bot."""

    role: TurnRole
    duration_ms: int
    started_at: datetime
    ended_at: datetime
    meta: dict[str, Any] = field(default_factory=dict)

    def to_dict(self) -> dict[str, Any]:
        return {
            "role": self.role,
            "duration_ms": self.duration_ms,
            "started_at": self.started_at.isoformat(),
            "ended_at": self.ended_at.isoformat(),
            **({"meta": self.meta} if self.meta else {}),
        }


@dataclass
class CallSession:
    """Métricas acumuladas en memoria durante una llamada.

    Se crea al iniciar la llamada y se persiste una sola vez al finalizar
    (ChannelDestroyed / StasisEnd / hangup).
    """

    call_id: str
    start_time: datetime = field(default_factory=utc_now)
    answer_time: datetime | None = None
    first_bot_audio_time: datetime | None = None
    end_time: datetime | None = None

    document_id: str | None = None
    number: str | None = None
    direction: str | None = None

    user_turns: int = 0
    bot_turns: int = 0

    stt_requests: int = 0
    llm_requests: int = 0
    tts_requests: int = 0

    total_prompt_tokens: int = 0
    total_completion_tokens: int = 0
    stt_audio_seconds: float = 0.0
    tts_characters: int = 0

    # Sumas para promedios
    _llm_latency_sum_ms: float = 0.0
    _tts_latency_sum_ms: float = 0.0
    _stt_latency_sum_ms: float = 0.0
    _response_latency_sum_ms: float = 0.0
    _response_latency_count: int = 0

    interruptions: int = 0
    silence_seconds: float = 0.0

    disconnect_reason: str | None = None
    turn_timings: list[TurnTiming] = field(default_factory=list)

    stt_provider: str | None = None
    tts_provider: str | None = None
    llm_provider: str | None = None
    llm_model: str | None = None

    _persisted: bool = False
    # Marca el instante en que terminó STT del turno actual (para latencia e2e)
    _pending_user_ready_at: datetime | None = None

    # --- derived ---

    @property
    def duration_seconds(self) -> float | None:
        if not self.end_time:
            return None
        return max(0.0, (self.end_time - self.start_time).total_seconds())

    @property
    def time_to_first_audio_ms(self) -> int | None:
        """Ms desde answer (o start) hasta el primer audio TTS del bot."""
        if not self.first_bot_audio_time:
            return None
        base = self.answer_time or self.start_time
        return max(0, int((self.first_bot_audio_time - base).total_seconds() * 1000))

    @property
    def avg_llm_latency_ms(self) -> float | None:
        if self.llm_requests <= 0:
            return None
        return self._llm_latency_sum_ms / self.llm_requests

    @property
    def avg_tts_latency_ms(self) -> float | None:
        if self.tts_requests <= 0:
            return None
        return self._tts_latency_sum_ms / self.tts_requests

    @property
    def avg_stt_latency_ms(self) -> float | None:
        if self.stt_requests <= 0:
            return None
        return self._stt_latency_sum_ms / self.stt_requests

    @property
    def avg_response_latency_ms(self) -> float | None:
        """Latencia promedio STT-listo → audio TTS listo (cuello de botella e2e)."""
        if self._response_latency_count <= 0:
            return None
        return self._response_latency_sum_ms / self._response_latency_count

    @property
    def silence_percent(self) -> float | None:
        dur = self.duration_seconds
        if dur is None or dur <= 0:
            return None
        return min(100.0, max(0.0, (self.silence_seconds / dur) * 100.0))

    def mark_answered(self, when: datetime | None = None) -> None:
        if self.answer_time is None:
            self.answer_time = when or utc_now()

    def mark_first_bot_audio(self, when: datetime | None = None) -> None:
        if self.first_bot_audio_time is None:
            self.first_bot_audio_time = when or utc_now()

    def record_user_turn(
        self,
        duration_ms: int,
        *,
        started_at: datetime | None = None,
        ended_at: datetime | None = None,
        silence_ms: int = 0,
    ) -> None:
        end = ended_at or utc_now()
        start = started_at or (
            end
            if duration_ms <= 0
            else datetime.fromtimestamp(
                end.timestamp() - duration_ms / 1000.0, tz=timezone.utc
            )
        )
        self.user_turns += 1
        if silence_ms > 0:
            self.silence_seconds += silence_ms / 1000.0
        self.turn_timings.append(
            TurnTiming(
                role="user",
                duration_ms=max(0, duration_ms),
                started_at=start,
                ended_at=end,
            )
        )
        self._pending_user_ready_at = end

    def record_bot_turn(
        self,
        duration_ms: int,
        *,
        started_at: datetime | None = None,
        ended_at: datetime | None = None,
    ) -> None:
        end = ended_at or utc_now()
        start = started_at or (
            end
            if duration_ms <= 0
            else datetime.fromtimestamp(
                end.timestamp() - duration_ms / 1000.0, tz=timezone.utc
            )
        )
        self.bot_turns += 1
        self.turn_timings.append(
            TurnTiming(
                role="bot",
                duration_ms=max(0, duration_ms),
                started_at=start,
                ended_at=end,
            )
        )

    def record_stt(
        self,
        latency_ms: float,
        *,
        audio_seconds: float = 0.0,
    ) -> None:
        self.stt_requests += 1
        self._stt_latency_sum_ms += max(0.0, latency_ms)
        if audio_seconds > 0:
            self.stt_audio_seconds += audio_seconds

    def record_llm(
        self,
        latency_ms: float,
        *,
        prompt_tokens: int = 0,
        completion_tokens: int = 0,
    ) -> None:
        self.llm_requests += 1
        self._llm_latency_sum_ms += max(0.0, latency_ms)
        self.total_prompt_tokens += max(0, prompt_tokens)
        self.total_completion_tokens += max(0, completion_tokens)

    def record_tts(
        self,
        latency_ms: float,
        *,
        characters: int = 0,
        audio_ready_at: datetime | None = None,
    ) -> None:
        self.tts_requests += 1
        self._tts_latency_sum_ms += max(0.0, latency_ms)
        if characters > 0:
            self.tts_characters += characters
        ready = audio_ready_at or utc_now()
        self.mark_first_bot_audio(ready)
        if self._pending_user_ready_at is not None:
            e2e_ms = (ready - self._pending_user_ready_at).total_seconds() * 1000.0
            if e2e_ms >= 0:
                self._response_latency_sum_ms += e2e_ms
                self._response_latency_count += 1
            self._pending_user_ready_at = None

    def add_silence_seconds(self, seconds: float) -> None:
        if seconds > 0:
            self.silence_seconds += seconds

    def record_interruption(self) -> None:
        self.interruptions += 1

    def finalize(self, reason: str, *, when: datetime | None = None) -> None:
        if self.end_time is None:
            self.end_time = when or utc_now()
        if reason and not self.disconnect_reason:
            self.disconnect_reason = reason

    def to_row(self, estimated_cost: float) -> dict[str, Any]:
        """Dict listo para INSERT MySQL."""
        return {
            "call_id": self.call_id,
            "document_id": self.document_id,
            "number": self.number,
            "direction": self.direction,
            "start_time": self.start_time,
            "answer_time": self.answer_time,
            "first_bot_audio_time": self.first_bot_audio_time,
            "end_time": self.end_time,
            "duration_seconds": self.duration_seconds,
            "time_to_first_audio_ms": self.time_to_first_audio_ms,
            "user_turns": self.user_turns,
            "bot_turns": self.bot_turns,
            "stt_requests": self.stt_requests,
            "llm_requests": self.llm_requests,
            "tts_requests": self.tts_requests,
            "total_prompt_tokens": self.total_prompt_tokens,
            "total_completion_tokens": self.total_completion_tokens,
            "stt_audio_seconds": round(self.stt_audio_seconds, 3),
            "tts_characters": self.tts_characters,
            "estimated_cost": round(estimated_cost, 6),
            "avg_llm_latency_ms": (
                round(self.avg_llm_latency_ms, 2)
                if self.avg_llm_latency_ms is not None
                else None
            ),
            "avg_tts_latency_ms": (
                round(self.avg_tts_latency_ms, 2)
                if self.avg_tts_latency_ms is not None
                else None
            ),
            "avg_stt_latency_ms": (
                round(self.avg_stt_latency_ms, 2)
                if self.avg_stt_latency_ms is not None
                else None
            ),
            "avg_response_latency_ms": (
                round(self.avg_response_latency_ms, 2)
                if self.avg_response_latency_ms is not None
                else None
            ),
            "interruptions": self.interruptions,
            "silence_seconds": round(self.silence_seconds, 3),
            "silence_percent": (
                round(self.silence_percent, 2)
                if self.silence_percent is not None
                else None
            ),
            "disconnect_reason": self.disconnect_reason,
            "turn_timings_json": [t.to_dict() for t in self.turn_timings],
            "providers_json": {
                "stt": self.stt_provider,
                "tts": self.tts_provider,
                "llm": self.llm_provider,
                "llm_model": self.llm_model,
            },
        }
