#!/usr/bin/env python3
"""
Script para processar e plotar dados SCANTEC (ERA5)
Autor: Thaisa Lopes
Descrição: Gera gráficos comparativos entre modelos BAM_IBIS e BAM_OPER
Suporte para variáveis de superfície e níveis de pressão (850, 500, 250 hPa)
"""

import os
import re
import pandas as pd
import matplotlib.pyplot as plt
from pathlib import Path
from typing import Dict, Tuple, Optional, List
import numpy as np

# ============================================================================
# CONFIGURAÇÕES GLOBAIS
# ============================================================================

TIPO = "SURF" 
IDENTIFICADOR_MODELO = "AGCM"

DIRETORIO_BASE = f"/mnt/beegfs/thaisa.lopes/SCANTEC/dataout/SMNA_x_NCEP.bkp/{IDENTIFICADOR_MODELO}_{TIPO}/"
ARQUIVO_VARS = "/mnt/beegfs/thaisa.lopes/SCANTEC/tables/scantec.vars.total"


# 🆕 Nome da pasta de saída (alterado de "fig" para "scantec_plot_espacial")
PASTA_SAIDA = "plot_serie_scantec"

# Cores para os modelos
CORES_MODELOS = {
    'BAM_SMNA': '#2E86C1',  # Azul
    'BAM_NCEP': '#E74C3C',  # Vermelho
}

# Marcadores para os modelos
MARCADORES_MODELOS = {
    'BAM_SMNA': 'o',
    'BAM_NCEP': 's',
}

# ============================================================================
# FUNÇÕES AUXILIARES
# ============================================================================

def limpar_nome_variavel(nome_coluna: str) -> str:
    """
    Remove sufixos e mantém níveis de pressão
    
    Args:
        nome_coluna: Nome da coluna (ex: 'psnm:000', 'geop:850')
    
    Returns:
        Nome limpo (ex: 'psnm', 'geop_850')
    
    Exemplos:
        PSNM:000 -> PSNM
        U10:000 -> U10
        geop:850 -> geop_850
        temp:500 -> temp_500
    """
    # Substitui : por _ para manter os níveis
    return re.sub(r':', '_', nome_coluna)


def carregar_unidades() -> Dict[str, str]:
    """
    Carrega as unidades das variáveis do arquivo scantec.vars
    
    Returns:
        Dicionário com variável -> unidade
    """
    unidades = {}
    
    try:
        with open(ARQUIVO_VARS, 'r') as f:
            for linha in f:
                linha = linha.strip()
                if not linha or linha.startswith('#') or linha.startswith('variables:'):
                    continue
                
                # Padrão: PSNM:000 "Pressure reduced to MSL [hPa]"
                match = re.search(r'(\w+):\d+\s+"[^[]*\[([^\]]+)\]"', linha)
                if match:
                    variavel = match.group(1).lower()
                    unidade = match.group(2)
                    unidades[variavel] = unidade
                    unidades[variavel.upper()] = unidade
                else:
                    # Tentar outro padrão sem unidade
                    match = re.search(r'(\w+):\d+\s+"(.*?)"', linha)
                    if match:
                        variavel = match.group(1).lower()
                        unidades[variavel] = ""
                        unidades[variavel.upper()] = ""
    
    except Exception as e:
        print(f"Aviso: Não foi possível ler {ARQUIVO_VARS}: {e}")
        # Unidades padrão
        unidades_padrao = {
            'psnm': 'hPa', 'PSNM': 'hPa',
            'u10': 'm/s', 'U10': 'm/s',
            'v10': 'm/s', 'V10': 'm/s',
            't2m': 'K', 'T2M': 'K',
            'sp': 'Pa', 'SP': 'Pa',
            'tcwv': 'kg/m²', 'TCWV': 'kg/m²',
            'geop': 'm²/s²', 'GEOP': 'm²/s²',
            'temp': 'K', 'TEMP': 'K',
            'umes': 'kg/kg', 'UMES': 'kg/kg',
            'uwnd': 'm/s', 'UWND': 'm/s',
            'vwnd': 'm/s', 'VWND': 'm/s',
            'omeg': 'Pa/s', 'OMEG': 'Pa/s'
        }
        unidades.update(unidades_padrao)
    
    return unidades


def extrair_metadata(nome_arquivo: str) -> Tuple[Optional[str], Optional[str]]:
    """
    Extrai tipo (ACOR, MEAN, RMSE, VIES) e modelo (BAM_IBIS, BAM_OPER)
    
    Args:
        nome_arquivo: Nome do arquivo .scan
    
    Returns:
        Tupla (tipo, modelo) ou (None, None) se não encontrar
    
    Exemplos:
        ACORBAM_IBIS_20251020002025110800T.scan -> ('ACOR', 'BAM_IBIS')
        RMSEBAM_OPER_20251020002025110800F.scan -> ('RMSE', 'BAM_OPER')
    """
    tipos = r"(ACOR|MEAN|RMSE|VIES)"
    padrao = f"{tipos}([A-Z]+_[A-Z]+)(?=_\\d{{14}})"
    
    match = re.search(padrao, nome_arquivo)
    if match:
        return match.group(1), match.group(2)
    
    return None, None


def ler_arquivo_scan(caminho_arquivo: str) -> Optional[pd.DataFrame]:
    """
    Lê o arquivo .scan e retorna um DataFrame
    Mantém os níveis das variáveis (ex: geop_850, geop_500, geop_250)
    
    Args:
        caminho_arquivo: Caminho completo do arquivo .scan
    
    Returns:
        DataFrame com os dados ou None em caso de erro
    """
    try:
        with open(caminho_arquivo, 'r') as f:
            linhas = f.readlines()
        
        dados_linhas = []
        cabecalho = None
        
        for linha in linhas:
            linha = linha.strip()
            if not linha:
                continue
            
            if linha.startswith('%'):
                cabecalho = linha.replace('%', '').split()
                if cabecalho:
                    cabecalho[0] = '%Previsao'
                    # Mantém os níveis (ex: geop:850 -> geop_850)
                    cabecalho = [limpar_nome_variavel(col) for col in cabecalho]
            else:
                valores = linha.split()
                if valores and cabecalho:
                    if len(valores) == len(cabecalho):
                        dados_linhas.append(valores)
                    elif len(valores) > len(cabecalho):
                        dados_linhas.append(valores[:len(cabecalho)])
                    elif len(valores) < len(cabecalho):
                        valores.extend([None] * (len(cabecalho) - len(valores)))
                        dados_linhas.append(valores)
        
        if not cabecalho or not dados_linhas:
            print(f"  Erro: Cabeçalho ou dados não encontrados em {caminho_arquivo}")
            return None
        
        df = pd.DataFrame(dados_linhas, columns=cabecalho)
        
        # Converter colunas para numérico
        for col in df.columns:
            try:
                df[col] = pd.to_numeric(df[col])
            except (ValueError, TypeError):
                continue
        
        return df
    
    except Exception as e:
        print(f"  Erro ao ler {caminho_arquivo}: {e}")
        return None


def obter_rotulo_eixo_y(tipo_variavel: str, nome_variavel: str, 
                        unidades: Dict[str, str]) -> str:
    """
    Retorna o rótulo apropriado para o eixo Y
    Suporte para variáveis com níveis de pressão
    
    Args:
        tipo_variavel: ACOR, RMSE, VIES, ou MEAN
        nome_variavel: Nome da variável (ex: 'geop_850')
        unidades: Dicionário de unidades
    
    Returns:
        Rótulo formatado para o eixo Y
    
    Exemplos:
        geop_850 -> "RMSE - GEOP (Nível 850 hPa) [m²/s²]"
        PSNM -> "RMSE - PSNM [hPa]"
    """
    # Separar variável base e nível (se existir)
    if '_' in nome_variavel:
        partes = nome_variavel.split('_')
        variavel_base = partes[0]
        nivel = partes[1] if len(partes) > 1 else ''
    else:
        variavel_base = nome_variavel
        nivel = ''
    
    nome_var_lower = variavel_base.lower()
    unidade = unidades.get(nome_var_lower, "")
    
    # Construir rótulo com nível
    sufixo_nivel = f" (Nível {nivel} hPa)" if nivel else ""
    
    # Mapeamento de tipos
    if tipo_variavel == "ACOR":
        return f"Correlação de Anomalia - {variavel_base.upper()}{sufixo_nivel}"
    elif tipo_variavel == "RMSE":
        return f"RMSE - {variavel_base.upper()}{sufixo_nivel}{f' [{unidade}]' if unidade else ''}"
    elif tipo_variavel == "VIES":
        return f"VIES - {variavel_base.upper()}{sufixo_nivel}{f' [{unidade}]' if unidade else ''}"
    elif tipo_variavel == "MEAN":
        return f"MÉDIA - {variavel_base.upper()}{sufixo_nivel}{f' [{unidade}]' if unidade else ''}"
    else:
        return f"{variavel_base.upper()}{sufixo_nivel}"


def configurar_estilo_grafico():
    """Configura o estilo global dos gráficos"""
    plt.style.use('seaborn-v0_8-whitegrid')
    plt.rcParams['font.size'] = 11
    plt.rcParams['axes.labelsize'] = 12
    plt.rcParams['axes.titlesize'] = 14
    plt.rcParams['legend.fontsize'] = 10
    plt.rcParams['figure.figsize'] = (14, 7)
    plt.rcParams['lines.linewidth'] = 2.5


# ============================================================================
# FUNÇÕES PRINCIPAIS
# ============================================================================

def gerar_graficos_por_variavel(dfs_por_modelo: Dict[str, pd.DataFrame], 
                                tipo_variavel: str, 
                                dominio: str, 
                                pasta_fig: str, 
                                unidades: Dict[str, str]):
    """
    Gera gráficos para cada variável comparando os modelos
    """
    if not dfs_por_modelo:
        print(f"  Nenhum dado para {tipo_variavel}")
        return
    
    primeiro_modelo = list(dfs_por_modelo.keys())[0]
    df_primeiro = dfs_por_modelo[primeiro_modelo]
    
    if df_primeiro is None or df_primeiro.empty:
        print(f"  DataFrame vazio para {tipo_variavel}")
        return
    
    # Identificar variáveis (excluindo '%Previsao')
    variaveis = [col for col in df_primeiro.columns if col != '%Previsao']
    
    if not variaveis:
        print(f"  Nenhuma variável encontrada para {tipo_variavel}")
        return
    
    print(f"  Variáveis encontradas: {', '.join(variaveis)}")
    
    # Gerar gráfico para cada variável (incluindo níveis)
    for variavel in variaveis:
        try:
            plt.figure(figsize=(12, 7))
            modelos_plotados = 0
            
            for modelo, df in dfs_por_modelo.items():
                if df is None or variavel not in df.columns:
                    continue
                
                dados_validos = df[['%Previsao', variavel]].dropna()
                if dados_validos.empty:
                    continue
                
                # Extrair dados como arrays numpy
                x = dados_validos['%Previsao'].to_numpy().ravel()
                y = dados_validos[variavel].to_numpy().ravel()
                
                # Verificar se as dimensões são compatíveis
                if len(x) != len(y):
                    print(f"    ⚠️  Dimensões incompatíveis para {variavel} - {modelo}: x={len(x)}, y={len(y)}")
                    continue
                
                cor = CORES_MODELOS.get(modelo, '#000000')
                marcador = MARCADORES_MODELOS.get(modelo, 'o')
                
                plt.plot(x, y, 
                        marker=marcador,
                        label=modelo,
                        color=cor,
                        linewidth=2.5,
                        markersize=9,
                        markeredgecolor='black',
                        markeredgewidth=0.8,
                        alpha=0.9)
                
                modelos_plotados += 1
            
            if modelos_plotados == 0:
                print(f"    Nenhum dado válido para {variavel}")
                plt.close()
                continue
            
            # ===== CONFIGURAR GRÁFICO =====
            
            # Eixo X
            plt.xlabel('Tempo de Previsão (horas)', fontsize=13, fontweight='bold')
            
            # Eixo Y
            rotulo_y = obter_rotulo_eixo_y(tipo_variavel, variavel, unidades)
            plt.ylabel(rotulo_y, fontsize=13, fontweight='bold')
            
            # Título
            titulo = f'{dominio} - {tipo_variavel} - {variavel.upper()}'
            plt.title(titulo, fontsize=15, fontweight='bold', pad=20)
            
            # Grade
            plt.grid(True, alpha=0.3, linestyle='--', linewidth=0.8)
            
            # Legenda
            plt.legend(loc='best', fontsize=11, framealpha=0.95)
            
            # Eixo X com os tempos disponíveis
            tempos = df_primeiro['%Previsao'].dropna().unique()
            if len(tempos) > 0:
                plt.xticks(sorted(tempos), fontsize=10)
            
            # Ajustar layout
            plt.tight_layout()
            
            # ===== SALVAR FIGURA =====
            # 🆕 Nome do arquivo com o identificador no início
            nome_arquivo = f"{IDENTIFICADOR_MODELO}_{TIPO}_{dominio}_{tipo_variavel}_{variavel}.png"
            caminho_salvar = os.path.join(pasta_fig, nome_arquivo)
            
            plt.savefig(caminho_salvar, dpi=300, bbox_inches='tight', 
                       facecolor='white', edgecolor='none')
            plt.close()
            
            print(f"    ✅ Gráfico salvo: {nome_arquivo}")
        
        except Exception as e:
            print(f"    ❌ ERRO ao plotar {variavel}: {e}")
            plt.close()
            continue


def processar_dominio(caminho_dominio: str, nome_dominio: str, 
                      unidades: Dict[str, str]):
    """
    Processa um domínio específico
    
    Args:
        caminho_dominio: Caminho do diretório do domínio
        nome_dominio: Nome do domínio
        unidades: Dicionário de unidades
    """
    print(f"\n{'='*70}")
    print(f"📍 Processando domínio: {nome_dominio}")
    print(f"{'='*70}")
    
    # 🆕 Cria a pasta com o novo nome (scantec_plot_espacial)
    pasta_fig = os.path.join(caminho_dominio, PASTA_SAIDA)
    Path(pasta_fig).mkdir(exist_ok=True)
    print(f"📁 Pasta de figuras: {pasta_fig}")
    
    # Encontrar todos os arquivos .scan com T.scan
    arquivos_scan = [f for f in os.listdir(caminho_dominio) 
                     if f.endswith('T.scan')]
    
    print(f"📄 Arquivos T.scan encontrados: {len(arquivos_scan)}")
    
    if not arquivos_scan:
        print("⚠️  Nenhum arquivo .scan encontrado!")
        return
    
    # Organizar arquivos por tipo (ACOR, MEAN, RMSE, VIES)
    arquivos_por_tipo = {}
    
    for arquivo in arquivos_scan:
        tipo, modelo = extrair_metadata(arquivo)
        
        if tipo and modelo:
            print(f"  Processando: {arquivo}")
            print(f"    → Tipo: {tipo}, Modelo: {modelo}")
            
            if tipo not in arquivos_por_tipo:
                arquivos_por_tipo[tipo] = {}
            
            caminho_completo = os.path.join(caminho_dominio, arquivo)
            df = ler_arquivo_scan(caminho_completo)
            
            if df is not None and not df.empty:
                arquivos_por_tipo[tipo][modelo] = df
                print(f"    ✅ {len(df)} linhas carregadas")
            else:
                print(f"    ⚠️  Falha ao carregar dados")
        else:
            print(f"  ⚠️  Não foi possível extrair metadados de: {arquivo}")
    
    # Gerar gráficos para cada tipo
    for tipo, dfs_modelos in arquivos_por_tipo.items():
        print(f"\n🎨 Gerando gráficos para {tipo}...")
        gerar_graficos_por_variavel(dfs_modelos, tipo, nome_dominio, pasta_fig, unidades)
    
    print(f"\n✅ Processamento do domínio {nome_dominio} concluído!")


def criar_grafico_resumo(dominio: str, pasta_fig: str):
    """
    Cria um gráfico de resumo com todas as variáveis para um domínio
    (Função opcional para futuras implementações)
    """
    pass


# ============================================================================
# FUNÇÃO PRINCIPAL
# ============================================================================

def main():
    """Função principal do script"""
    print("\n" + "="*70)
    print("🚀 SCANTEC Data Processor - ERA5")
    print("="*70)
    print(f"📂 Diretório base: {DIRETORIO_BASE}")
    print(f"📁 Pasta de saída: {PASTA_SAIDA}")  # 🆕 Mostra o novo nome
    
    # Configurar estilo dos gráficos
    configurar_estilo_grafico()
    
    # Carregar unidades
    print(f"📋 Carregando unidades do arquivo: {ARQUIVO_VARS}")
    unidades = carregar_unidades()
    print(f"   ✅ {len(unidades)} unidades carregadas")
    
    # Verificar diretório base
    if not os.path.exists(DIRETORIO_BASE):
        print(f"❌ ERRO: Diretório base não encontrado: {DIRETORIO_BASE}")
        return
    
    # Listar domínios
    dominios = [d for d in os.listdir(DIRETORIO_BASE) 
                if os.path.isdir(os.path.join(DIRETORIO_BASE, d)) 
                and d not in [PASTA_SAIDA, 'fig', '__pycache__']]  # 🆕 Ignora as pastas de figura
    
    if not dominios:
        print("❌ Nenhum domínio encontrado!")
        return
    
    print(f"\n📊 Domínios encontrados: {', '.join(dominios)}")
    
    # Processar cada domínio
    for dominio in dominios:
        caminho_dominio = os.path.join(DIRETORIO_BASE, dominio)
        processar_dominio(caminho_dominio, dominio, unidades)
    
    # Resumo final
    print(f"\n{'='*70}")
    print("✅ PROCESSAMENTO FINALIZADO COM SUCESSO!")
    print(f"   {len(dominios)} domínios processados")
    print(f"   Figuras salvas em: {DIRETORIO_BASE}/[dominio]/{PASTA_SAIDA}/")
    print(f"{'='*70}\n")


if __name__ == "__main__":
    main()
