#!/usr/bin/env python3
"""mcp-orchestrator.py — Orquestador multi-servidor MCP para swl-ses.

Descubrimiento dinamico de herramientas MCP disponibles en todos los
servidores configurados. Health check, busqueda de tools y reportes
estructurados usados por el comando /swl:mcp-status.

Combina mcp-pool-manager.py (conexion) con mcp-telemetry.py (trazas).

Uso:
  python scripts/mcp-orchestrator.py status           -- health de todos los servidores
  python scripts/mcp-orchestrator.py discover         -- herramientas completas por servidor
  python scripts/mcp-orchestrator.py find-tool <kw>   -- buscar herramienta por nombre/desc
  python scripts/mcp-orchestrator.py summary          -- resumen compacto para /swl:mcp-status

Flags:
  --json   Salida en formato JSON (default: texto legible)
  --cwd    Directorio raiz del proyecto (default: .)
  --trace  Registrar resultado en .planning/traces/ (default: False)
"""
from __future__ import annotations

import argparse
import asyncio
import json
import os
import sys
import time
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
# ---------------------------------------------------------------------------

SETTINGS_CANDIDATES = [
    '.claude/settings.local.json',
    '.claude/settings.json',
    'mcp-servers.json',
]

TIMEOUT_S           = 12
SNAPSHOT_FILE       = Path('.planning') / 'mcp-snapshot.json'

# ---------------------------------------------------------------------------
# Carga de config (identica a mcp-pool-manager para evitar dependencia circular)
# ---------------------------------------------------------------------------


def _cargar_config(cwd: Path, config_path: str | None = None) -> dict:
    candidates = [Path(config_path)] if config_path else [cwd / p for p in SETTINGS_CANDIDATES]
    for p in candidates:
        if p.exists():
            try:
                data = json.loads(p.read_text(encoding='utf-8'))
                servers = data.get('mcpServers', {})
                if servers:
                    return servers
            except Exception:
                continue
    return {}


def _build_env(cfg: dict) -> dict | None:
    extra = cfg.get('env') or {}
    if not extra:
        return None
    merged = dict(os.environ)
    merged.update(extra)
    return merged

# ---------------------------------------------------------------------------
# Conexion async (igual que pool-manager, duplicada para zero-imports externos)
# ---------------------------------------------------------------------------


async def _probe_server(nombre: str, cfg: dict) -> dict:
    """Conecta a un servidor, devuelve su inventario completo de herramientas."""
    resultado: dict = {
        'server':    nombre,
        'transport': 'http' if 'url' in cfg else 'stdio',
        'tools':     [],
        'error':     None,
        'estado':    'OK',
        'duration_ms': 0,
    }
    t0 = time.time()
    try:
        if 'url' in cfg:
            transport = sse_client(cfg['url'])
        else:
            env = _build_env(cfg)
            params = StdioServerParameters(
                command=cfg['command'],
                args=cfg.get('args', []),
                env=env,
                cwd=cfg.get('cwd'),
            )
            transport = stdio_client(params)

        async with transport as (read, write):
            async with ClientSession(read, write) as session:
                await asyncio.wait_for(session.initialize(), timeout=TIMEOUT_S)
                resp = await asyncio.wait_for(session.list_tools(), timeout=TIMEOUT_S)
                resultado['tools'] = [
                    {
                        'name':        t.name,
                        'description': t.description or '',
                    }
                    for t in (resp.tools or [])
                ]
    except asyncio.TimeoutError:
        resultado['error'] = f'Timeout ({TIMEOUT_S}s)'
        resultado['estado'] = 'ERROR'
    except Exception as exc:
        resultado['error'] = str(exc)
        resultado['estado'] = 'ERROR'
    finally:
        resultado['duration_ms'] = int((time.time() - t0) * 1000)
    return resultado


async def _probe_all(servers: dict) -> list:
    """Prueba todos los servidores en paralelo."""
    tareas = [_probe_server(n, c) for n, c in servers.items()]
    return list(await asyncio.gather(*tareas))

# ---------------------------------------------------------------------------
# Snapshot — persiste el ultimo estado conocido
# ---------------------------------------------------------------------------


def _guardar_snapshot(cwd: Path, resultados: list) -> None:
    """Guarda el ultimo estado de discovery en .planning/mcp-snapshot.json."""
    snap_file = cwd / SNAPSHOT_FILE
    snap_file.parent.mkdir(parents=True, exist_ok=True)
    snapshot = {
        'generado': datetime.now(timezone.utc).isoformat(),
        'servidores': resultados,
        'resumen': {
            'total':     len(resultados),
            'activos':   sum(1 for r in resultados if r['estado'] == 'OK'),
            'errores':   sum(1 for r in resultados if r['estado'] == 'ERROR'),
            'tools':     sum(len(r['tools']) for r in resultados),
        },
    }
    try:
        tmp = snap_file.with_suffix('.tmp')
        tmp.write_text(json.dumps(snapshot, ensure_ascii=False, indent=2), encoding='utf-8')
        os.replace(tmp, snap_file)
    except Exception as exc:
        sys.stderr.write(f'[mcp-orchestrator] No se pudo guardar snapshot: {exc}\n')


def _cargar_snapshot(cwd: Path) -> dict | None:
    snap_file = cwd / SNAPSHOT_FILE
    if not snap_file.exists():
        return None
    try:
        return json.loads(snap_file.read_text(encoding='utf-8'))
    except Exception:
        return None

# ---------------------------------------------------------------------------
# Telemetria opcional (importa mcp-telemetry si esta disponible)
# ---------------------------------------------------------------------------


def _registrar_traza_discovery(cwd: Path, resumen: dict) -> None:
    """Registra una traza del discovery si mcp-telemetry esta disponible."""
    try:
        sys.path.insert(0, str(cwd / 'scripts'))
        from mcp_telemetry import registrar_traza  # type: ignore
        registrar_traza(
            cwd       = cwd,
            nombre    = 'mcp:discovery',
            atributos = resumen,
            estado    = 'OK' if resumen.get('errores', 0) == 0 else 'ERROR',
        )
    except ImportError:
        pass   # mcp-telemetry no disponible — silencioso

# ---------------------------------------------------------------------------
# Subcomandos
# ---------------------------------------------------------------------------


async def cmd_status(servers: dict, cwd: Path, as_json: bool, con_traza: bool) -> None:
    out        = sys.stdout.write
    resultados = await _probe_all(servers)
    _guardar_snapshot(cwd, resultados)

    resumen = {
        'total':   len(resultados),
        'activos': sum(1 for r in resultados if r['estado'] == 'OK'),
        'errores': sum(1 for r in resultados if r['estado'] == 'ERROR'),
        'tools':   sum(len(r['tools']) for r in resultados),
    }

    if con_traza:
        _registrar_traza_discovery(cwd, resumen)

    if as_json:
        out(json.dumps({'resumen': resumen, 'servidores': resultados},
                       ensure_ascii=False, indent=2) + '\n')
        return

    out('\nEstado de servidores MCP:\n')
    out('-' * 68 + '\n')
    for r in resultados:
        estado = r['estado']
        tools  = len(r['tools'])
        ms     = r['duration_ms']
        trans  = r['transport']
        err    = f'  -- {r["error"]}' if r['error'] else ''
        out(f'  {r["server"]:<28} {estado:<5} {tools:>3} tools  {ms:>5}ms  [{trans}]{err}\n')
    out('\n')
    out(f'  Total: {resumen["total"]} servidores, {resumen["activos"]} activos, '
        f'{resumen["tools"]} herramientas disponibles\n')


async def cmd_discover(servers: dict, cwd: Path, as_json: bool, con_traza: bool) -> None:
    out        = sys.stdout.write
    resultados = await _probe_all(servers)
    _guardar_snapshot(cwd, resultados)

    if con_traza:
        resumen = {
            'total':   len(resultados),
            'activos': sum(1 for r in resultados if r['estado'] == 'OK'),
            'errores': sum(1 for r in resultados if r['estado'] == 'ERROR'),
            'tools':   sum(len(r['tools']) for r in resultados),
        }
        _registrar_traza_discovery(cwd, resumen)

    if as_json:
        out(json.dumps(resultados, ensure_ascii=False, indent=2) + '\n')
        return

    for r in resultados:
        if r['error']:
            out(f'\n[{r["server"]}] ERROR: {r["error"]}\n')
            continue
        out(f'\n[{r["server"]}] {len(r["tools"])} herramientas:\n')
        for t in r['tools']:
            desc = t['description']
            if len(desc) > 72:
                desc = desc[:69] + '...'
            out(f'  - {t["name"]:<32} {desc}\n')


async def cmd_find_tool(servers: dict, keyword: str, as_json: bool) -> None:
    out        = sys.stdout.write
    resultados = await _probe_all(servers)
    kw_lower   = keyword.lower()

    encontradas: list = []
    for r in resultados:
        for t in r.get('tools', []):
            if (kw_lower in t['name'].lower() or
                    kw_lower in t['description'].lower()):
                encontradas.append({
                    'server': r['server'],
                    'name':   t['name'],
                    'description': t['description'],
                })

    if as_json:
        out(json.dumps(encontradas, ensure_ascii=False, indent=2) + '\n')
        return

    if not encontradas:
        out(f'Ninguna herramienta encontrada con keyword: "{keyword}"\n')
        return

    out(f'\nHerramientas que coinciden con "{keyword}" ({len(encontradas)}):\n')
    out('-' * 68 + '\n')
    for e in encontradas:
        desc = e['description']
        if len(desc) > 60:
            desc = desc[:57] + '...'
        out(f'  [{e["server"]}]  {e["name"]:<28}  {desc}\n')


def cmd_summary(cwd: Path, as_json: bool) -> None:
    """Muestra el ultimo snapshot guardado sin re-conectar a los servidores."""
    out      = sys.stdout.write
    snapshot = _cargar_snapshot(cwd)

    if not snapshot:
        out('Sin snapshot disponible. Ejecuta primero: mcp-orchestrator.py status\n')
        return

    if as_json:
        out(json.dumps(snapshot, ensure_ascii=False, indent=2) + '\n')
        return

    generado = snapshot.get('generado', '?')[:19]
    resumen  = snapshot.get('resumen', {})
    out(f'\nResumen MCP (snapshot del {generado} UTC):\n')
    out(f'  Servidores: {resumen.get("total", 0)} '
        f'({resumen.get("activos", 0)} activos, {resumen.get("errores", 0)} con error)\n')
    out(f'  Herramientas disponibles: {resumen.get("tools", 0)}\n')
    out('\nDetalle por servidor:\n')
    for r in snapshot.get('servidores', []):
        estado = r.get('estado', '?')
        tools  = len(r.get('tools', []))
        ms     = r.get('duration_ms', 0)
        out(f'  {r["server"]:<28} {estado:<5} {tools:>3} tools  {ms:>5}ms\n')

# ---------------------------------------------------------------------------
# main()
# ---------------------------------------------------------------------------


def main() -> None:
    if hasattr(sys.stdout, 'reconfigure'):
        sys.stdout.reconfigure(encoding='utf-8', errors='replace')

    if not HAS_MCP:
        sys.stderr.write(
            '[mcp-orchestrator] ERROR: la libreria mcp no esta instalada.\n'
            '[mcp-orchestrator] Ejecuta: pip install mcp\n'
        )
        sys.exit(1)

    parser = argparse.ArgumentParser(
        description='Orquestador multi-servidor MCP para swl-ses'
    )
    parser.add_argument('--json',  action='store_true', help='Salida en formato JSON')
    parser.add_argument('--cwd',   default='.', help='Directorio raiz del proyecto')
    parser.add_argument('--trace', action='store_true',
                        help='Registrar resultado en .planning/traces/')
    parser.add_argument('--config', default=None, help='Ruta al archivo de config MCP')

    sub = parser.add_subparsers(dest='cmd')

    sub.add_parser('status',   help='Health check de todos los servidores MCP')
    sub.add_parser('discover', help='Lista completa de herramientas por servidor')

    p_find = sub.add_parser('find-tool', help='Busca herramienta por nombre o descripcion')
    p_find.add_argument('keyword', help='Texto a buscar en nombre o descripcion')

    sub.add_parser('summary',  help='Muestra el ultimo snapshot sin reconectar')

    args    = parser.parse_args()
    cwd     = Path(args.cwd).resolve()
    as_json = bool(getattr(args, 'json', False))
    traza   = bool(getattr(args, 'trace', False))

    if args.cmd == 'summary':
        cmd_summary(cwd, as_json)
        return

    servers = _cargar_config(cwd, args.config)
    if not servers:
        sys.stderr.write(
            '[mcp-orchestrator] No se encontraron servidores MCP.\n'
            '[mcp-orchestrator] Verifica la clave "mcpServers" en .claude/settings.json\n'
        )
        sys.exit(1)

    if args.cmd == 'status':
        asyncio.run(cmd_status(servers, cwd, as_json, traza))
    elif args.cmd == 'discover':
        asyncio.run(cmd_discover(servers, cwd, as_json, traza))
    elif args.cmd == 'find-tool':
        asyncio.run(cmd_find_tool(servers, args.keyword, as_json))
    else:
        parser.print_help()


if __name__ == '__main__':
    main()
