#!/usr/bin/env python3
"""mcp-telemetry.py — Trazabilidad OTLP-lite para sesiones MCP en swl-ses.

Registra llamadas a herramientas MCP en formato compatible con el sistema
de trazas de swl-ses (.planning/traces/mcp-YYYY-MM-DD.jsonl).

El formato de cada entrada es el mismo OTLP-lite que usa otlp-exporter.js:
  { traceId, spanId, nombre, inicio, fin, duracionMs, estado, atributos }

Modo modulo (importable):
  from scripts.mcp_telemetry import MCPTelemetrySession, registrar_traza

Modo CLI:
  python scripts/mcp-telemetry.py report          -- trazas de hoy
  python scripts/mcp-telemetry.py report --days 3 -- trazas de los ultimos 3 dias
  python scripts/mcp-telemetry.py stats            -- estadisticas por servidor
"""
from __future__ import annotations

import argparse
import asyncio
import json
import os
import secrets
import sys
import time
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from pathlib import Path
from typing import Any

# ---------------------------------------------------------------------------
# Dependencias — raw mcp SDK
# ---------------------------------------------------------------------------
try:
    from mcp import ClientSession
    from mcp.client.stdio import stdio_client, StdioServerParameters
    from mcp.client.sse import sse_client
    HAS_MCP = True
except ImportError:
    HAS_MCP = False

# ---------------------------------------------------------------------------
# Constantes
# ---------------------------------------------------------------------------

TRACES_DIR   = Path('.planning') / 'traces'
MCP_PREFIX   = 'mcp'                          # prefijo de archivos de traza MCP
TIMEOUT_S    = 15

# ---------------------------------------------------------------------------
# Utilidades de traza OTLP-lite
# ---------------------------------------------------------------------------


def _trace_id() -> str:
    return secrets.token_hex(16)   # 128 bits, 32 hex chars


def _span_id() -> str:
    return secrets.token_hex(8)    # 64 bits, 16 hex chars


def _iso_now() -> str:
    return datetime.now(timezone.utc).isoformat()


def _ruta_hoy(cwd: Path) -> Path:
    """Ruta del JSONL de trazas MCP del dia actual."""
    directorio = cwd / TRACES_DIR
    directorio.mkdir(parents=True, exist_ok=True)
    fecha = datetime.now(timezone.utc).strftime('%Y-%m-%d')
    return directorio / f'{MCP_PREFIX}-{fecha}.jsonl'


def registrar_traza(cwd: Path, nombre: str, atributos: dict,
                    estado: str = 'OK',
                    inicio_iso: str | None = None,
                    fin_iso: str | None = None,
                    duracion_ms: int = 0) -> dict:
    """Escribe una entrada OTLP-lite al JSONL de trazas del dia.

    Parametros:
        cwd:         Directorio raiz del proyecto.
        nombre:      Nombre de la operacion (ej. 'mcp:call_tool').
        atributos:   Metadatos adicionales de la traza.
        estado:      'OK' o 'ERROR'.
        inicio_iso:  Timestamp ISO de inicio (se genera si no se provee).
        fin_iso:     Timestamp ISO de fin (se genera si no se provee).
        duracion_ms: Duracion en milisegundos.

    Retorna el dict de la traza escrita.
    """
    ahora = _iso_now()
    traza = {
        'traceId':    _trace_id(),
        'spanId':     _span_id(),
        'nombre':     nombre,
        'inicio':     inicio_iso or ahora,
        'fin':        fin_iso or ahora,
        'duracionMs': duracion_ms,
        'estado':     estado,
        'atributos':  atributos,
    }
    try:
        ruta = _ruta_hoy(cwd)
        with open(ruta, 'a', encoding='utf-8') as fh:
            fh.write(json.dumps(traza, ensure_ascii=False) + '\n')
    except Exception as exc:
        sys.stderr.write(f'[mcp-telemetry] No se pudo escribir traza: {exc}\n')
    return traza

# ---------------------------------------------------------------------------
# Context manager: MCPTelemetrySession
# ---------------------------------------------------------------------------


class MCPTelemetrySession:
    """Envuelve una sesion MCP con trazabilidad automatica.

    Cada llamada a call_tool() se registra como una traza OTLP-lite.
    La sesion se conecta y desconecta al entrar/salir del contexto.

    Ejemplo:
        async with MCPTelemetrySession(cwd, server_name, server_cfg) as sess:
            herramientas = await sess.list_tools()
            resultado    = await sess.call_tool('mi_tool', {'arg': 'valor'})
    """

    def __init__(self, cwd: Path, server_name: str, server_cfg: dict) -> None:
        self.cwd         = cwd
        self.server_name = server_name
        self.server_cfg  = server_cfg
        self._session: ClientSession | None = None
        self._ctx_stack: list = []

    async def __aenter__(self) -> 'MCPTelemetrySession':
        if not HAS_MCP:
            raise RuntimeError('La libreria mcp no esta instalada. Ejecuta: pip install mcp')

        cfg = self.server_cfg
        if 'url' in cfg:
            transport = sse_client(cfg['url'])
        else:
            env: dict | None = None
            if extra := cfg.get('env'):
                env = {**os.environ, **extra}
            params = StdioServerParameters(
                command=cfg['command'],
                args=cfg.get('args', []),
                env=env,
                cwd=cfg.get('cwd'),
            )
            transport = stdio_client(params)

        t_ctx = transport
        rw    = await t_ctx.__aenter__()
        self._ctx_stack.append(t_ctx)

        sess_ctx = ClientSession(rw[0], rw[1])
        self._session = await sess_ctx.__aenter__()
        self._ctx_stack.append(sess_ctx)

        await asyncio.wait_for(self._session.initialize(), timeout=TIMEOUT_S)
        return self

    async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:
        for ctx in reversed(self._ctx_stack):
            try:
                await ctx.__aexit__(exc_type, exc_val, exc_tb)
            except Exception:
                pass
        self._ctx_stack.clear()
        self._session = None

    async def list_tools(self) -> list:
        """Lista herramientas disponibles en el servidor."""
        if not self._session:
            raise RuntimeError('Sesion no inicializada — usar como context manager')
        resp = await asyncio.wait_for(self._session.list_tools(), timeout=TIMEOUT_S)
        return resp.tools or []

    async def call_tool(self, tool: str, arguments: dict | None = None) -> dict:
        """Llama una herramienta MCP y registra la traza automaticamente.

        Retorna dict con: result (list[str]), is_error (bool), duration_ms (int).
        """
        if not self._session:
            raise RuntimeError('Sesion no inicializada — usar como context manager')

        inicio_ts = time.time()
        inicio_iso = _iso_now()
        estado    = 'OK'
        resultado: dict = {'result': [], 'is_error': False, 'duration_ms': 0}

        try:
            resp = await asyncio.wait_for(
                self._session.call_tool(tool, arguments or {}),
                timeout=TIMEOUT_S,
            )
            resultado['is_error'] = bool(getattr(resp, 'isError', False))
            resultado['result']   = [
                c.text if hasattr(c, 'text') else str(c)
                for c in (resp.content or [])
            ]
            if resultado['is_error']:
                estado = 'ERROR'
        except Exception as exc:
            resultado['error'] = str(exc)
            estado = 'ERROR'
        finally:
            fin_ts  = time.time()
            dur_ms  = int((fin_ts - inicio_ts) * 1000)
            fin_iso = _iso_now()
            resultado['duration_ms'] = dur_ms

            registrar_traza(
                cwd       = self.cwd,
                nombre    = f'mcp:call_tool',
                atributos = {
                    'server':      self.server_name,
                    'tool':        tool,
                    'args_keys':   list((arguments or {}).keys()),
                    'duration_ms': dur_ms,
                },
                estado      = estado,
                inicio_iso  = inicio_iso,
                fin_iso     = fin_iso,
                duracion_ms = dur_ms,
            )

        return resultado

# ---------------------------------------------------------------------------
# CLI helpers
# ---------------------------------------------------------------------------


def _leer_trazas_mcp(cwd: Path, dias: int = 1) -> list:
    """Lee todas las trazas MCP de los ultimos N dias (archivos mcp-*.jsonl)."""
    directorio = cwd / TRACES_DIR
    if not directorio.exists():
        return []

    trazas: list = []
    archivos = sorted(directorio.glob(f'{MCP_PREFIX}-*.jsonl'), reverse=True)[:dias]
    for archivo in archivos:
        try:
            for linea in archivo.read_text(encoding='utf-8').splitlines():
                linea = linea.strip()
                if linea:
                    try:
                        trazas.append(json.loads(linea))
                    except json.JSONDecodeError:
                        continue
        except Exception:
            continue
    return trazas


def _cmd_report(cwd: Path, dias: int, as_json: bool) -> None:
    out    = sys.stdout.write
    trazas = _leer_trazas_mcp(cwd, dias)

    if as_json:
        out(json.dumps(trazas, ensure_ascii=False, indent=2) + '\n')
        return

    if not trazas:
        out(f'Sin trazas MCP en los ultimos {dias} dia(s).\n')
        out(f'Directorio: {cwd / TRACES_DIR}\n')
        return

    out(f'\nTrazas MCP ({len(trazas)} entradas, ultimos {dias} dia(s)):\n')
    out('-' * 72 + '\n')
    for t in trazas[-50:]:        # mostrar ultimas 50
        estado = t.get('estado', '?')
        nombre = t.get('nombre', '?')
        dur    = t.get('duracionMs', 0)
        attrs  = t.get('atributos', {})
        server = attrs.get('server', '')
        tool   = attrs.get('tool', '')
        ts     = t.get('inicio', '')[:19]
        out(f'  {ts}  {estado:<5}  {nombre:<20}  {server}/{tool}  {dur}ms\n')


def _cmd_stats(cwd: Path, dias: int, as_json: bool) -> None:
    out    = sys.stdout.write
    trazas = _leer_trazas_mcp(cwd, dias)

    # Agregar por servidor y herramienta
    stats: dict = {}   # server -> {tool -> {count, ok, errors, total_ms}}
    for t in trazas:
        attrs  = t.get('atributos', {})
        server = attrs.get('server', 'desconocido')
        tool   = attrs.get('tool', 'desconocido')
        estado = t.get('estado', 'OK')
        dur    = t.get('duracionMs', 0)

        if server not in stats:
            stats[server] = {}
        if tool not in stats[server]:
            stats[server][tool] = {'count': 0, 'ok': 0, 'errors': 0, 'total_ms': 0}

        s = stats[server][tool]
        s['count']    += 1
        s['total_ms'] += dur
        if estado == 'OK':
            s['ok'] += 1
        else:
            s['errors'] += 1

    # Calcular promedio
    for server in stats:
        for tool in stats[server]:
            s = stats[server][tool]
            s['avg_ms'] = int(s['total_ms'] / s['count']) if s['count'] else 0

    if as_json:
        out(json.dumps(stats, ensure_ascii=False, indent=2) + '\n')
        return

    if not stats:
        out(f'Sin datos de uso MCP en los ultimos {dias} dia(s).\n')
        return

    out(f'\nEstadisticas MCP — ultimos {dias} dia(s) ({len(trazas)} trazas):\n')
    out('-' * 72 + '\n')
    out(f'  {"Servidor":<22} {"Herramienta":<24} {"Calls":>6} {"OK":>4} {"Err":>4} {"Avg ms":>7}\n')
    out('-' * 72 + '\n')
    for server, tools in sorted(stats.items()):
        for tool, s in sorted(tools.items()):
            out(
                f'  {server:<22} {tool:<24}'
                f' {s["count"]:>6} {s["ok"]:>4} {s["errors"]:>4} {s["avg_ms"]:>7}\n'
            )


# ---------------------------------------------------------------------------
# main()
# ---------------------------------------------------------------------------


def main() -> None:
    if hasattr(sys.stdout, 'reconfigure'):
        sys.stdout.reconfigure(encoding='utf-8', errors='replace')

    parser = argparse.ArgumentParser(
        description='Trazabilidad OTLP-lite para sesiones MCP en swl-ses'
    )
    parser.add_argument('--json', action='store_true', help='Salida en formato JSON')
    parser.add_argument('--cwd',  default='.', help='Directorio raiz del proyecto')

    sub = parser.add_subparsers(dest='cmd')

    p_rep = sub.add_parser('report', help='Muestra trazas MCP recientes')
    p_rep.add_argument('--days', type=int, default=1,
                       help='Numero de dias hacia atras (default: 1)')

    p_stats = sub.add_parser('stats', help='Estadisticas de uso por servidor/herramienta')
    p_stats.add_argument('--days', type=int, default=7,
                         help='Numero de dias hacia atras (default: 7)')

    args = parser.parse_args()
    cwd  = Path(args.cwd).resolve()

    as_json = bool(getattr(args, 'json', False))

    if args.cmd == 'report':
        _cmd_report(cwd, args.days, as_json)
    elif args.cmd == 'stats':
        _cmd_stats(cwd, args.days, as_json)
    else:
        parser.print_help()


if __name__ == '__main__':
    main()
