"""Persistencia MySQL: una sola INSERT al finalizar cada llamada."""

from __future__ import annotations

import json
import logging
from typing import Any

from call_metrics.session import CallSession

logger = logging.getLogger(__name__)

# Tabla propia del módulo. Nunca UPDATE de filas previas: cada llamada = 1 INSERT.
_CREATE_TABLE_SQL = """
CREATE TABLE IF NOT EXISTS call_sessions (
  id BIGINT AUTO_INCREMENT PRIMARY KEY,
  call_id VARCHAR(64) NOT NULL,
  document_id VARCHAR(64) NULL,
  number VARCHAR(64) NULL,
  direction VARCHAR(16) NULL,
  start_time DATETIME(3) NULL,
  answer_time DATETIME(3) NULL,
  first_bot_audio_time DATETIME(3) NULL,
  end_time DATETIME(3) NULL,
  duration_seconds DOUBLE NULL,
  time_to_first_audio_ms INT NULL,
  user_turns INT NOT NULL DEFAULT 0,
  bot_turns INT NOT NULL DEFAULT 0,
  stt_requests INT NOT NULL DEFAULT 0,
  llm_requests INT NOT NULL DEFAULT 0,
  tts_requests INT NOT NULL DEFAULT 0,
  total_prompt_tokens INT NOT NULL DEFAULT 0,
  total_completion_tokens INT NOT NULL DEFAULT 0,
  stt_audio_seconds DOUBLE NOT NULL DEFAULT 0,
  tts_characters INT NOT NULL DEFAULT 0,
  estimated_cost DECIMAL(12,6) NOT NULL DEFAULT 0,
  avg_llm_latency_ms DOUBLE NULL,
  avg_tts_latency_ms DOUBLE NULL,
  avg_stt_latency_ms DOUBLE NULL,
  avg_response_latency_ms DOUBLE NULL,
  interruptions INT NOT NULL DEFAULT 0,
  silence_seconds DOUBLE NOT NULL DEFAULT 0,
  silence_percent DOUBLE NULL,
  disconnect_reason VARCHAR(255) NULL,
  turn_timings_json JSON NULL,
  providers_json JSON NULL,
  created_at DATETIME(3) NOT NULL DEFAULT CURRENT_TIMESTAMP(3),
  KEY idx_call_sessions_call_id (call_id),
  KEY idx_call_sessions_start_time (start_time),
  KEY idx_call_sessions_document (document_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4
"""

# Si la tabla ya existía con UNIQUE(call_id), el upsert pisaba costos anteriores.
_DROP_LEGACY_UNIQUE_SQL = (
    "ALTER TABLE call_sessions DROP INDEX uq_call_sessions_call_id"
)

_INSERT_SQL = """
INSERT INTO call_sessions (
  call_id, document_id, number, direction,
  start_time, answer_time, first_bot_audio_time, end_time,
  duration_seconds, time_to_first_audio_ms,
  user_turns, bot_turns,
  stt_requests, llm_requests, tts_requests,
  total_prompt_tokens, total_completion_tokens,
  stt_audio_seconds, tts_characters, estimated_cost,
  avg_llm_latency_ms, avg_tts_latency_ms, avg_stt_latency_ms,
  avg_response_latency_ms,
  interruptions, silence_seconds, silence_percent,
  disconnect_reason, turn_timings_json, providers_json
) VALUES (
  %(call_id)s, %(document_id)s, %(number)s, %(direction)s,
  %(start_time)s, %(answer_time)s, %(first_bot_audio_time)s, %(end_time)s,
  %(duration_seconds)s, %(time_to_first_audio_ms)s,
  %(user_turns)s, %(bot_turns)s,
  %(stt_requests)s, %(llm_requests)s, %(tts_requests)s,
  %(total_prompt_tokens)s, %(total_completion_tokens)s,
  %(stt_audio_seconds)s, %(tts_characters)s, %(estimated_cost)s,
  %(avg_llm_latency_ms)s, %(avg_tts_latency_ms)s, %(avg_stt_latency_ms)s,
  %(avg_response_latency_ms)s,
  %(interruptions)s, %(silence_seconds)s, %(silence_percent)s,
  %(disconnect_reason)s, %(turn_timings_json)s, %(providers_json)s
)
"""


class CallMetricsMySqlStore:
    """Pool aiomysql + ensure schema + insert atómico por call_id."""

    def __init__(
        self,
        *,
        host: str,
        port: int,
        user: str,
        password: str,
        database: str,
        autocommit: bool = True,
    ) -> None:
        self._host = host
        self._port = port
        self._user = user
        self._password = password
        self._database = database
        self._autocommit = autocommit
        self._pool: Any = None
        self._schema_ready = False

    @property
    def ready(self) -> bool:
        return self._pool is not None

    async def connect(self) -> None:
        if self._pool is not None:
            return
        try:
            import aiomysql
        except ImportError as exc:
            raise RuntimeError(
                "CALL_METRICS_ENABLED=true requiere aiomysql. "
                "Instalá: pip install aiomysql PyMySQL"
            ) from exc
        self._pool = await aiomysql.create_pool(
            host=self._host,
            port=self._port,
            user=self._user,
            password=self._password,
            db=self._database,
            autocommit=self._autocommit,
            minsize=1,
            maxsize=5,
            charset="utf8mb4",
        )
        await self.ensure_schema()
        logger.info(
            "Call metrics MySQL conectado %s:%s/%s",
            self._host,
            self._port,
            self._database,
        )

    async def ensure_schema(self) -> None:
        if self._schema_ready or self._pool is None:
            return
        async with self._pool.acquire() as conn:
            async with conn.cursor() as cur:
                await cur.execute(_CREATE_TABLE_SQL)
                # Migración: tablas creadas con UNIQUE(call_id) + upsert
                # reescribían la misma fila y “borraban” costos previos.
                try:
                    await cur.execute(_DROP_LEGACY_UNIQUE_SQL)
                    logger.info(
                        "Call metrics: removido UNIQUE(call_id) legacy "
                        "(cada llamada queda como fila independiente)"
                    )
                except Exception as exc:
                    # 1091 = Can't DROP; index already gone / never existed
                    errno = getattr(exc, "args", [None])[0]
                    if errno not in (1091,):
                        logger.debug(
                            "Call metrics: no se eliminó UNIQUE(call_id): %s",
                            exc,
                        )
        self._schema_ready = True

    async def insert_session(self, row: dict[str, Any]) -> None:
        """Siempre INSERT (append). Nunca UPDATE de filas anteriores."""
        if self._pool is None:
            raise RuntimeError("Call metrics MySQL no conectado")
        payload = dict(row)
        payload["turn_timings_json"] = json.dumps(
            payload.get("turn_timings_json") or [],
            ensure_ascii=False,
            default=str,
        )
        payload["providers_json"] = json.dumps(
            payload.get("providers_json") or {},
            ensure_ascii=False,
            default=str,
        )
        async with self._pool.acquire() as conn:
            async with conn.cursor() as cur:
                await cur.execute(_INSERT_SQL, payload)
                inserted_id = cur.lastrowid
        logger.info(
            "Call metrics persistido id=%s call=%s cost≈$%.6f turns=%d/%d",
            inserted_id,
            row.get("call_id"),
            float(row.get("estimated_cost") or 0),
            int(row.get("user_turns") or 0),
            int(row.get("bot_turns") or 0),
        )

    async def persist(self, session: CallSession, estimated_cost: float) -> None:
        await self.insert_session(session.to_row(estimated_cost))

    async def close(self) -> None:
        if self._pool is None:
            return
        self._pool.close()
        await self._pool.wait_closed()
        self._pool = None
        self._schema_ready = False
