#!/usr/bin/env python3
"""
API simplificada para predição de glosas
Retorna apenas JSON sem outputs adicionais
"""

import json
import pandas as pd
import numpy as np
import joblib
import os
import sys
import warnings
warnings.filterwarnings('ignore')

class PreditorAPI:
    def __init__(self, modo_factoring=True):
        self.tipos_guia = {
            1: 'Internação',
            2: 'SADT',
            3: 'Consulta',
            4: 'Honorários',
            5: 'Odonto'
        }
        self.modelos = {}
        self.modelos_arquivos = {}  # Mapeia modelo -> arquivo
        self.modo_factoring = modo_factoring  # Modo ultra-conservador para recebíveis
        self.THRESHOLD_FACTORING = 0.15  # Threshold fixo para factoring
        self._carregar_modelos_silencioso()
        
    def _carregar_modelos_silencioso(self):
        """Carrega modelos sem output - prioriza modelos OOV do diretório modelos/"""
        import glob

        # Diretórios onde buscar modelos (em ordem de prioridade)
        diretorios = ['modelos/', './']

        # Primeiro tenta carregar modelos OOV (prioridade máxima)
        for tipo in self.tipos_guia.keys():
            modelo_carregado = False

            # Procura primeiro por modelos OOV em todos os diretórios
            for diretorio in diretorios:
                if modelo_carregado:
                    break

                arquivos_oov = glob.glob(f'{diretorio}modelo_*_tipo{tipo}_oov.pkl')
                if arquivos_oov:
                    # Ordena por data de modificação (mais recente primeiro)
                    arquivo_oov = sorted(arquivos_oov, key=os.path.getmtime, reverse=True)[0]
                    self.modelos[f'tipo_{tipo}'] = joblib.load(arquivo_oov)
                    self.modelos_arquivos[f'tipo_{tipo}'] = arquivo_oov
                    modelo_carregado = True
                    break

            # Se não encontrou OOV, tenta modelo normal
            if not modelo_carregado:
                for diretorio in diretorios:
                    arquivos = glob.glob(f'{diretorio}modelo_*_tipo{tipo}.pkl')
                    if arquivos:
                        # Evita carregar arquivos que terminam com _oov.pkl
                        arquivos_normais = [a for a in arquivos if not a.endswith('_oov.pkl')]
                        if arquivos_normais:
                            # Ordena por data de modificação (mais recente primeiro)
                            arquivo_recente = sorted(arquivos_normais, key=os.path.getmtime, reverse=True)[0]
                            self.modelos[f'tipo_{tipo}'] = joblib.load(arquivo_recente)
                            self.modelos_arquivos[f'tipo_{tipo}'] = arquivo_recente
                            break

        # Modelo geral OOV
        modelo_geral_carregado = False
        for diretorio in diretorios:
            if modelo_geral_carregado:
                break

            modelos_gerais_oov = glob.glob(f'{diretorio}modelo_*_oov.pkl')
            modelos_gerais_oov = [m for m in modelos_gerais_oov if '_tipo' not in m]
            if modelos_gerais_oov:
                # Ordena por data de modificação (mais recente primeiro)
                arquivo_recente = sorted(modelos_gerais_oov, key=os.path.getmtime, reverse=True)[0]
                self.modelos['geral'] = joblib.load(arquivo_recente)
                self.modelos_arquivos['geral'] = arquivo_recente
                modelo_geral_carregado = True

        # Fallback para modelo legado específico
        if not modelo_geral_carregado:
            for diretorio in diretorios:
                arquivo_legado = f'{diretorio}modelo_export_guias_2025-07-05_23-04-07.pkl'
                if os.path.exists(arquivo_legado):
                    self.modelos['geral'] = joblib.load(arquivo_legado)
                    self.modelos_arquivos['geral'] = arquivo_legado
                    break

        # Modelo avançado de consulta (se existir)
        for diretorio in diretorios:
            arquivo_avancado = f'{diretorio}modelo_consulta_avancado.pkl'
            if os.path.exists(arquivo_avancado):
                self.modelos['consulta_avancado'] = joblib.load(arquivo_avancado)
                self.modelos_arquivos['consulta_avancado'] = arquivo_avancado
                break
    
    def preparar_dados_com_oov(self, df, modelo_info):
        """Prepara dados com features OOV se o modelo suportar"""
        # Verifica se é modelo OOV
        if 'codigos_conhecidos' in modelo_info and 'oov_features' in modelo_info:
            codigos_conhecidos = set(modelo_info['codigos_conhecidos'])
            
            # Adiciona features OOV
            df['codigo_desconhecido'] = 0
            df['codigo_suspeito'] = 0
            df['frequencia_codigo'] = 0
            df['categoria_codigo'] = 0
            
            if 'cod' in df.columns:
                # Marca códigos desconhecidos
                df['codigo_desconhecido'] = (~df['cod'].isin(codigos_conhecidos)).astype(int)
                
                # Códigos suspeitos
                df.loc[df['cod'] > 99999999, 'codigo_suspeito'] = 1
                df.loc[df['cod'] < 10000000, 'codigo_suspeito'] = 1
                df.loc[df['codigo_desconhecido'] == 1, 'codigo_suspeito'] = 1
                
                # Categoria do código
                df['categoria_codigo'] = df['cod'].fillna(0).astype(str).str[:2].apply(
                    lambda x: float(x) if x.isdigit() else 999
                )
                df.loc[df['codigo_desconhecido'] == 1, 'categoria_codigo'] = 999
                
                # Frequência (0 para desconhecidos)
                df.loc[df['codigo_desconhecido'] == 1, 'frequencia_codigo'] = 0
                df.loc[df['codigo_desconhecido'] == 0, 'frequencia_codigo'] = 0.5
        
        return self.preparar_dados(df)
    
    def preparar_dados(self, df):
        """Prepara dados do DataFrame"""
        # Converte datas
        for col in ['data_cadastro', 'data_cirurgia', 'data_envio_operadora']:
            if col in df.columns:
                df[col] = pd.to_datetime(df[col], errors='coerce')
        
        # Features temporais
        if 'data_envio_operadora' in df.columns and 'data_cadastro' in df.columns:
            df['dias_envio'] = (df['data_envio_operadora'] - df['data_cadastro']).dt.days
            df['dias_envio'] = df['dias_envio'].fillna(5)
        else:
            df['dias_envio'] = 5
        
        if 'data_cirurgia' in df.columns and 'data_cadastro' in df.columns:
            df['dias_ate_cirurgia'] = (df['data_cirurgia'] - df['data_cadastro']).dt.days
            df['dias_ate_cirurgia'] = df['dias_ate_cirurgia'].fillna(0)
        else:
            df['dias_ate_cirurgia'] = 0
        
        if 'data_cirurgia' in df.columns:
            df['mes_cirurgia'] = df['data_cirurgia'].dt.month
            df['dia_semana_cirurgia'] = df['data_cirurgia'].dt.dayofweek
            df['mes_cirurgia'] = df['mes_cirurgia'].fillna(6)
            df['dia_semana_cirurgia'] = df['dia_semana_cirurgia'].fillna(3)
        else:
            df['mes_cirurgia'] = 6
            df['dia_semana_cirurgia'] = 3
        
        # Procedimentos combinados
        if 'procedimentos_combinados' in df.columns:
            df['procedimentos_combinados'] = df['procedimentos_combinados'].fillna('').astype(str).str.count(',') + 1
            df.loc[df['procedimentos_combinados'] == 1, 'procedimentos_combinados'] = 0
        
        # IMPORTANTE: Remove colunas de data originais após criar features
        # XGBoost não aceita datetime
        for col in ['data_cadastro', 'data_cirurgia', 'data_envio_operadora']:
            if col in df.columns:
                df = df.drop(col, axis=1)
        
        return df
    
    def predizer_guia(self, guia_data, modelo_info):
        """Prediz uma única guia com suporte a OOV"""
        model = modelo_info['model']
        encoders = modelo_info['encoders']
        best_threshold = modelo_info['best_threshold']
        feature_names = modelo_info['feature_names']
        
        # Converte para DataFrame
        df = pd.DataFrame([guia_data])
        
        # Detecta se código é desconhecido ANTES da predição
        codigo = guia_data.get('cod', guia_data.get('id_cod'))
        is_unknown = False
        suspicion_score = 0.0
        
        if 'codigos_conhecidos' in modelo_info and codigo:
            codigos_conhecidos = set(modelo_info['codigos_conhecidos'])
            try:
                codigo_num = int(str(codigo).replace('.', '').replace('-', ''))
                is_unknown = codigo_num not in codigos_conhecidos
                if is_unknown:
                    suspicion_score = 0.8
                    if codigo_num > 88000000:
                        suspicion_score = 1.0
            except:
                is_unknown = True
                suspicion_score = 1.0
        
        # Prepara dados com OOV se suportado
        if 'oov_features' in modelo_info:
            df = self.preparar_dados_com_oov(df, modelo_info)
        else:
            df = self.preparar_dados(df)
        
        # Codifica categóricas
        for col, encoder in encoders.items():
            if col in df.columns:
                df[col] = df[col].fillna('Unknown').astype(str)
                try:
                    df[col] = encoder.transform(df[col])
                except:
                    df[col] = 0
        
        # Features específicas para consulta avançada
        if 'convenios_risco' in modelo_info:
            df['convenio_risco'] = df['id_convenio'].isin(modelo_info['convenios_risco']).astype(int)
            if 'id_cod_agrupado' not in df.columns:
                df['id_cod_agrupado'] = df['id_cod']
        
        # Remove glosado
        if 'glosado' in df.columns:
            df = df.drop('glosado', axis=1)
        
        # Prepara features
        X = pd.DataFrame()
        for feat in feature_names:
            if feat in df.columns:
                X[feat] = df[feat]
            else:
                X[feat] = 0
        
        # Garante tipos numéricos
        for col in X.columns:
            if X[col].dtype == 'object':
                X[col] = pd.to_numeric(X[col], errors='coerce').fillna(0)
        
        # Predição
        if 'model_focal' in modelo_info:  # Modelo avançado
            proba_focal = modelo_info['model_focal'].predict_proba(X)[:, 1][0]
            proba_cv = 0
            for cv_model in modelo_info['cv_models']:
                proba_cv += cv_model.predict_proba(X)[:, 1][0]
            proba_cv /= len(modelo_info['cv_models'])
            probabilidade = modelo_info['weights']['focal'] * proba_focal + modelo_info['weights']['cv'] * proba_cv
        else:
            probabilidade = model.predict_proba(X)[:, 1][0]
        
        # Ajusta score se código é desconhecido e modelo tem threshold OOV
        if is_unknown and 'oov_threshold' in modelo_info:
            oov_threshold = modelo_info['oov_threshold']
            # Penalidade para códigos desconhecidos
            penalty = max(oov_threshold, suspicion_score * 0.8)
            probabilidade_ajustada = max(probabilidade, penalty)
            return float(np.float64(probabilidade_ajustada))
        
        return float(np.float64(probabilidade))
    
    def analisar(self, json_data):
        """Analisa dados JSON e retorna resultado simplificado"""
        if not self.modelos:
            return {"erro": "Nenhum modelo disponível"}
        
        # Converte para lista se for dict único
        if isinstance(json_data, dict):
            json_data = [json_data]
        
        # Agrupa por protocolo
        protocolos = {}
        for item in json_data:
            protocolo = item.get('protocolo', 'sem_protocolo')
            if protocolo not in protocolos:
                protocolos[protocolo] = []
            protocolos[protocolo].append(item)
        
        resultados = []
        
        # Analisa cada protocolo
        for protocolo, itens in protocolos.items():
            # Identifica tipo
            tipo_guia = itens[0].get('tipo_guia', 0)
            tipo_nome = self.tipos_guia.get(tipo_guia, 'Desconhecido')
            
            # Seleciona modelo - prioriza modelos OOV por tipo
            modelo_key = None
            if f'tipo_{tipo_guia}' in self.modelos:
                modelo_info = self.modelos[f'tipo_{tipo_guia}']
                modelo_key = f'tipo_{tipo_guia}'
            elif tipo_guia == 3 and 'consulta_avancado' in self.modelos:
                modelo_info = self.modelos['consulta_avancado']
                modelo_key = 'consulta_avancado'
            else:
                modelo_info = self.modelos.get('geral')
                modelo_key = 'geral'
            
            if not modelo_info:
                continue
            
            # Pega informações do modelo
            arquivo_modelo = self.modelos_arquivos.get(modelo_key, 'modelo_desconhecido.pkl')
            versao_modelo = modelo_info.get('versao', 'v1')
            is_oov = 'oov' in versao_modelo.lower() or '_oov' in arquivo_modelo
            
            # Analisa procedimentos
            analise_procedimentos = []
            scores = []
            
            for item in itens:
                score = self.predizer_guia(item, modelo_info)
                scores.append(score)
                
                # Adiciona ao resultado com novo formato
                proc_resultado = {
                    "cod": str(item.get('cod', 'sem_cod')),
                    "risco": round(score, 2)
                }
                
                # Se modelo tem OOV, verifica se código é desconhecido
                if 'codigos_conhecidos' in modelo_info:
                    codigo = item.get('cod', item.get('id_cod'))
                    if codigo:
                        codigos_conhecidos = set(modelo_info['codigos_conhecidos'])
                        try:
                            codigo_num = int(str(codigo).replace('.', '').replace('-', ''))
                            if codigo_num not in codigos_conhecidos:
                                proc_resultado["codigo_desconhecido"] = True
                                proc_resultado["alerta"] = "Código nunca visto no treinamento"
                        except:
                            proc_resultado["codigo_desconhecido"] = True
                            proc_resultado["alerta"] = "Código com formato inválido"
                
                # Adiciona id_procedimento se existir
                if 'id_procedimento' in item:
                    proc_resultado["id_procedimento"] = str(item['id_procedimento'])
                
                analise_procedimentos.append(proc_resultado)
            
            # Calcula métricas
            score_maximo = float(max(scores)) if scores else 0
            score_medio = float(np.mean(scores)) if scores else 0
            
            # Verifica se há códigos desconhecidos
            tem_codigo_desconhecido = any(
                p.get('codigo_desconhecido', False) for p in analise_procedimentos
            )
            
            # MODO FACTORING: Threshold ultra-conservador fixo
            if self.modo_factoring:
                threshold = self.THRESHOLD_FACTORING  # Sempre 0.15 para factoring
                # Adiciona penalidades extras no modo factoring
                factoring_penalty = 0.0
                
                # Penalidade se houver código desconhecido
                if tem_codigo_desconhecido:
                    factoring_penalty += 0.85  # Força rejeição (1.0 total)
                
                # Penalidade para múltiplos procedimentos
                if len(analise_procedimentos) > 2:
                    factoring_penalty += 0.10
                
                # Aplica penalidade em todos os scores
                for i, proc in enumerate(analise_procedimentos):
                    score_original = scores[i]
                    scores[i] = min(1.0, scores[i] + factoring_penalty)
                    proc['risco'] = round(scores[i], 2)
                    proc['risco_sem_factoring'] = round(score_original, 2)
                    if factoring_penalty > 0:
                        proc['penalidade_factoring'] = round(factoring_penalty, 2)
                
                # Recalcula métricas
                score_maximo = float(max(scores)) if scores else 0
                score_medio = float(np.mean(scores)) if scores else 0
            else:
                # Modo normal: usa threshold do modelo
                if tem_codigo_desconhecido and 'oov_threshold' in modelo_info:
                    threshold = float(modelo_info['oov_threshold'])
                else:
                    threshold = float(modelo_info['best_threshold'])
            
            # Nível de risco
            if self.modo_factoring:
                # Faixas ultra-conservadoras para factoring
                if score_maximo < 0.10:
                    risco = 'baixo'  # 0-10%: Liberar
                elif score_maximo < 0.15:
                    risco = 'medio'  # 10-15%: Revisar
                else:
                    risco = 'muito_alto'  # >15%: Rejeitar
            else:
                # Faixas normais
                if score_maximo < 0.3:
                    risco = 'baixo'
                elif score_maximo < 0.5:
                    risco = 'medio'
                elif score_maximo < 0.7:
                    risco = 'alto'
                else:
                    risco = 'muito_alto'
            
            # Resultado do protocolo
            resultado_protocolo = {
                'protocolo': protocolo,
                'tipo_guia': tipo_nome,
                'analise_procedimentos': analise_procedimentos,
                'resultado': {
                    'media': round(score_medio, 2),
                    'maximo': round(score_maximo, 2),
                    'risco': risco,
                    'predicao': 'glosa_provavel' if score_maximo >= threshold else 'aprovado',
                    'threshold': round(threshold, 2)
                },
                'modelo_usado': {
                    'arquivo': arquivo_modelo,
                    'versao': versao_modelo,
                    'tipo': 'OOV' if is_oov else 'Normal',
                    'tem_deteccao_oov': is_oov
                }
            }
            
            # Adiciona informações do modo factoring
            if self.modo_factoring:
                resultado_protocolo['modo_analise'] = 'FACTORING_ULTRA_CONSERVADOR'
                resultado_protocolo['factoring'] = {
                    'aprovado_para_antecipacao': score_maximo < self.THRESHOLD_FACTORING,
                    'margem_seguranca': round(self.THRESHOLD_FACTORING - score_maximo, 2),
                    'confianca': 'ALTA' if score_maximo < 0.10 else ('MEDIA' if score_maximo < 0.15 else 'BAIXA'),
                    'recomendacao': self._gerar_recomendacao_factoring(score_maximo, tem_codigo_desconhecido)
                }
            
            # Adiciona alertas se houver códigos desconhecidos
            if tem_codigo_desconhecido:
                num_desconhecidos = sum(1 for p in analise_procedimentos if p.get('codigo_desconhecido'))
                resultado_protocolo['alertas'] = [
                    f"⚠️ {num_desconhecidos} código(s) nunca visto(s) no treinamento",
                    "📊 Score ajustado com penalidade OOV"
                ]
            
            resultados.append(resultado_protocolo)
        
        # Se apenas um protocolo, retorna direto
        if len(resultados) == 1:
            return resultados[0]
        
        return resultados
    
    def _gerar_recomendacao_factoring(self, score_maximo, tem_codigo_desconhecido):
        """Gera recomendação específica para factoring"""
        if tem_codigo_desconhecido:
            return {
                'decisao': 'REJEITAR',
                'motivo': 'Código nunca visto no treinamento',
                'acao': 'Não aprovar para antecipação - risco extremo'
            }
        elif score_maximo >= 0.15:
            return {
                'decisao': 'REJEITAR',
                'motivo': f'Risco {score_maximo:.1%} acima do limite de {self.THRESHOLD_FACTORING:.1%}',
                'acao': 'Não aprovar para antecipação - risco inaceitável'
            }
        elif score_maximo >= 0.10:
            return {
                'decisao': 'REVISAR',
                'motivo': f'Risco {score_maximo:.1%} próximo ao limite',
                'acao': 'Revisão manual obrigatória antes de aprovar'
            }
        else:
            return {
                'decisao': 'APROVAR',
                'motivo': f'Risco {score_maximo:.1%} dentro da margem segura',
                'acao': 'Pode aprovar para antecipação com segurança'
            }


def main():
    if len(sys.argv) < 2:
        print(json.dumps({
            "erro": "Uso: python api_predicao.py <arquivo.json> ou python api_predicao.py '<json_string>'"
        }))
        sys.exit(1)
    
    entrada = sys.argv[1]
    preditor = PreditorAPI()
    
    try:
        # Tenta como arquivo
        if os.path.exists(entrada):
            with open(entrada, 'r', encoding='utf-8') as f:
                json_data = json.load(f)
        else:
            # Tenta como string JSON
            json_data = json.loads(entrada)
        
        resultado = preditor.analisar(json_data)
        print(json.dumps(resultado, indent=2, ensure_ascii=False))
        
    except Exception as e:
        print(json.dumps({"erro": str(e)}))
        sys.exit(1)


if __name__ == "__main__":
    main()