import pandas as pd
import xgboost as xgb
from sklearn.model_selection import train_test_split
from sklearn.metrics import classification_report, roc_auc_score, confusion_matrix
from sklearn.preprocessing import LabelEncoder
import matplotlib.pyplot as plt
import numpy as np
import warnings
import pymysql
from dotenv import load_dotenv
import os
warnings.filterwarnings('ignore')

load_dotenv()

print("🔌 Conectando ao banco de dados...")
try:
    connection = pymysql.connect(
        host=os.getenv('host'),
        user=os.getenv('user'),
        password=os.getenv('password'),
        port=int(os.getenv('port')),
        database=os.getenv('database', 'producao'),
        charset='utf8mb4',
        cursorclass=pymysql.cursors.DictCursor
    )
except Exception as e:
    print(f"❌ Erro ao conectar ao banco de dados: {e}")
    print("\n⚠️  Verifique se:")
    print("   1. O servidor de banco está acessível")
    print("   2. As credenciais no arquivo .env estão corretas")
    print("   3. Você tem acesso ao banco a partir desta rede")
    print("\n💡 Como alternativa, você pode:")
    print("   - Executar a query manualmente no banco")
    print("   - Exportar os dados como glosas.csv")
    print("   - Executar o script treina_modelo.py")
    exit(1)

query = """
SELECT
    s.id,
    s.protocolo,
    s.id_usuario,
    s.id_convenio,
    s.tipo_guia,
    s.data_cirurgia,
    s.data_hora as data_cadastro,
    l.data_solicitacao as data_envio_operadora,
    s.id_local,
    IF(DATEDIFF(s.data_cirurgia, s.data_nascimento_paciente) > 28, 'N', 'S') as recem_nascido,
    s.tipo_atendimento,
    s.tipo_acomodacao,
    sp.id_cod,
    c.cod,
    sp.via_acesso,
    COUNT(*) as quantidade,
    IF(sp.recurso = 'N', '0', '1') as glosado
FROM servicos s
LEFT JOIN servicos_procedimentos sp ON sp.id_servico = s.id
LEFT JOIN cbhpm c ON c.id = sp.id_cod
LEFT JOIN lote_guias lg ON lg.id_servico = s.id
LEFT JOIN lotes l ON l.id = lg.id_lote
WHERE s.vigente = 'S' AND s.data_hora >= '2024-12-01 00:00:00'
GROUP BY s.id, sp.id_cod, sp.via_acesso
"""

print("📊 Executando query...")
df = pd.read_sql(query, connection)
connection.close()
print(f"✅ Dados extraídos: {df.shape[0]} registros, {df.shape[1]} colunas")

df.to_csv('glosas.csv', index=False)
print("💾 Dados salvos em 'glosas.csv'")

print("\n🔧 Pré-processamento dos dados...")
df['data_cadastro'] = pd.to_datetime(df['data_cadastro'])
df['data_envio_operadora'] = pd.to_datetime(df['data_envio_operadora'])
df['data_cirurgia'] = pd.to_datetime(df['data_cirurgia'])

df['dias_envio'] = (df['data_envio_operadora'] - df['data_cadastro']).dt.days
df['dias_ate_cirurgia'] = (df['data_cirurgia'] - df['data_cadastro']).dt.days
df['mes_cirurgia'] = df['data_cirurgia'].dt.month
df['dia_semana_cirurgia'] = df['data_cirurgia'].dt.dayofweek

df = df.dropna(subset=['dias_envio'])
df['dias_envio'] = df['dias_envio'].fillna(0)
df['dias_ate_cirurgia'] = df['dias_ate_cirurgia'].fillna(0)

cat_cols = ['tipo_guia', 'recem_nascido', 'tipo_atendimento', 'tipo_acomodacao', 'via_acesso']
encoders = {}

for col in cat_cols:
    if col in df.columns:
        le = LabelEncoder()
        df[col] = le.fit_transform(df[col].astype(str))
        encoders[col] = le
        print(f"  ✓ Codificada coluna: {col}")

feature_cols = [col for col in df.columns if col not in ['glosado', 'id', 'protocolo', 'data_cadastro', 
                                                         'data_envio_operadora', 'data_cirurgia']]
X = df[feature_cols]
y = df['glosado'].astype(int)

print(f"\n📊 Distribuição das classes:")
print(f"  • Não glosado (0): {(y == 0).sum()} ({(y == 0).sum() / len(y) * 100:.1f}%)")
print(f"  • Glosado (1): {(y == 1).sum()} ({(y == 1).sum() / len(y) * 100:.1f}%)")

X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42, stratify=y)
print(f"\n✂️ Divisão treino/teste:")
print(f"  • Treino: {X_train.shape[0]} registros")
print(f"  • Teste: {X_test.shape[0]} registros")

print("\n🚀 Treinando modelo XGBoost...")
model = xgb.XGBClassifier(
    objective='binary:logistic',
    eval_metric='auc',
    tree_method='hist',
    max_depth=6,
    n_estimators=100,
    learning_rate=0.1,
    subsample=0.8,
    colsample_bytree=0.8,
    scale_pos_weight=(y == 0).sum() / (y == 1).sum(),
    n_jobs=-1,
    random_state=42,
    verbosity=1
)

eval_set = [(X_train, y_train), (X_test, y_test)]
model.fit(
    X_train, 
    y_train,
    eval_set=eval_set,
    early_stopping_rounds=10,
    verbose=False
)

print("\n📈 Avaliação do modelo:")
y_pred = model.predict(X_test)
y_proba = model.predict_proba(X_test)[:, 1]

print("\n" + "="*50)
print("RELATÓRIO DE CLASSIFICAÇÃO:")
print("="*50)
print(classification_report(y_test, y_pred, target_names=['Não Glosado', 'Glosado']))

auc_score = roc_auc_score(y_test, y_proba)
print(f"AUC-ROC Score: {auc_score:.4f}")

cm = confusion_matrix(y_test, y_pred)
print("\nMatriz de Confusão:")
print(f"  Verdadeiro Negativo: {cm[0,0]}")
print(f"  Falso Positivo: {cm[0,1]}")
print(f"  Falso Negativo: {cm[1,0]}")
print(f"  Verdadeiro Positivo: {cm[1,1]}")

plt.figure(figsize=(12, 8))
xgb.plot_importance(model, max_num_features=15, importance_type='gain')
plt.title('Top 15 Features Mais Importantes', fontsize=14)
plt.xlabel('Ganho', fontsize=12)
plt.tight_layout()
plt.savefig('feature_importance.png', dpi=300, bbox_inches='tight')
print("\n💾 Gráfico de importância das features salvo como 'feature_importance.png'")

import joblib
joblib.dump(model, 'modelo_glosas_xgboost.pkl')
print("💾 Modelo salvo como 'modelo_glosas_xgboost.pkl'")

joblib.dump(encoders, 'encoders.pkl')
print("💾 Encoders salvos como 'encoders.pkl'")

print("\n✅ Treinamento concluído com sucesso!")
print(f"\n📋 Resumo final:")
print(f"  • Acurácia: {(y_pred == y_test).mean():.2%}")
print(f"  • AUC-ROC: {auc_score:.4f}")
print(f"  • Melhor iteração: {model.best_iteration}")