# test_models.py
from transformers import AutoTokenizer, AutoModelForSequenceClassification, pipeline
import pickle
import json
import torch
from pathlib import Path
from typing import List, Dict, Any

# --- Chemins absolus vers vos ressources ---
MODELS_DIR = r"C:\Users\dev29\xamp\htdocs\models"
CIBLES_JSON_PATH = r"C:\Users\dev29\xamp\htdocs\models\cibles_indicateurs.json"

# --- Dictionnaire des noms ODD ---
odd_names = {
    "1": "Pas de pauvreté", "2": "Faim zéro", "3": "Bonnes conditions de vie",
    "4": "Éducation de qualité", "5": "Égalité entre les sexes", "6": "Eau propre et assainissement",
    "7": "Énergie propre et d'un coût abordable", "8": "Travail décent et croissance économique",
    "9": "Industrie, innovation et infrastructure", "10": "Inégalités réduites",
    "11": "Villes et communautés durables", "12": "Consommation et production responsables",
    "13": "Mesures relatives à la lutte contre les changements climatiques",
    "14": "Vie aquatique", "15": "Vie terrestre",
    "16": "Paix, justice et institutions efficaces", "17": "Partenariats pour les objectifs durables"
}

# --- Chargement des modèles et encodeurs ---
print("🔄 Chargement des modèles et encodeurs...")

try:
    # 1. Modèle de nature
    nature_tokenizer = AutoTokenizer.from_pretrained(f"{MODELS_DIR}/nature-classifier", local_files_only=True)
    nature_model = AutoModelForSequenceClassification.from_pretrained(f"{MODELS_DIR}/nature-classifier", local_files_only=True)
    with open(f'{MODELS_DIR}/nature_label_encoder.pkl', 'rb') as f:
        nature_label_encoder = pickle.load(f)

    # 2. Modèle ODD
    odd_tokenizer = AutoTokenizer.from_pretrained(f"{MODELS_DIR}/odd-classifier", local_files_only=True)
    odd_model = AutoModelForSequenceClassification.from_pretrained(f"{MODELS_DIR}/odd-classifier", local_files_only=True)
    with open(f'{MODELS_DIR}/odd_label_encoder.pkl', 'rb') as f:
        odd_label_encoder = pickle.load(f)

    # 3. Modèle Cible
    cible_tokenizer = AutoTokenizer.from_pretrained(f"{MODELS_DIR}/cible-classifier", local_files_only=True)
    cible_model = AutoModelForSequenceClassification.from_pretrained(f"{MODELS_DIR}/cible-classifier", local_files_only=True)
    with open(f'{MODELS_DIR}/cible_label_encoder.pkl', 'rb') as f:
        cible_label_encoder = pickle.load(f)

    # 4. Modèle Indicateur
    indicateur_tokenizer = AutoTokenizer.from_pretrained(f"{MODELS_DIR}/indicateurs-classifier", local_files_only=True)
    indicateur_model = AutoModelForSequenceClassification.from_pretrained(f"{MODELS_DIR}/indicateurs-classifier", local_files_only=True)
    with open(f'{MODELS_DIR}/indicateurs_label_encoder.pkl', 'rb') as f:
        indicateur_label_encoder = pickle.load(f)

    # Création des pipelines
    nature_classifier = pipeline("text-classification", model=nature_model, tokenizer=nature_tokenizer, top_k=None)
    odd_classifier = pipeline("text-classification", model=odd_model, tokenizer=odd_tokenizer, top_k=None)
    cible_classifier = pipeline("text-classification", model=cible_model, tokenizer=cible_tokenizer, top_k=None)
    indicateur_classifier = pipeline("text-classification", model=indicateur_model, tokenizer=indicateur_tokenizer, top_k=None)

    print("✅ Modèles chargés avec succès")

except Exception as e:
    print(f"❌ Erreur lors du chargement des modèles: {e}")
    print(f"Vérifiez que tous les fichiers existent dans {MODELS_DIR}")
    exit(1)

# --- Chargement des données cibles/indicateurs ---
try:
    with open(CIBLES_JSON_PATH, "r", encoding="utf-8") as f:
        cibles_data = json.load(f)
    print(f"✅ Fichier cibles_indicateurs.json chargé ({len(cibles_data)} entrées)")
except Exception as e:
    print(f"❌ Erreur lors du chargement de cibles_indicateurs.json: {e}")
    exit(1)

# --- Fonctions utilitaires (reprises de votre code) ---
def get_first_prediction(preds):
    """Retourne le premier dict d'une liste ou liste imbriquée."""
    if not preds:
        return None
    first = preds[0]
    if isinstance(first, list) and first:
        return first[0]
    return first

def map_model_label_to_nature(label):
    """Mappe un label modèle vers un ID métier pour la nature"""
    try:
        if isinstance(label, str) and label.upper().startswith("LABEL_"):
            model_idx = int(label.split("_", 1)[1])
        else:
            model_idx = int(label)
    except Exception:
        return None, None, None

    # Récupération du nom via nature_label_encoder
    class_name = None
    try:
        if hasattr(nature_label_encoder, "classes_") and model_idx < len(nature_label_encoder.classes_):
            class_name = nature_label_encoder.classes_[model_idx]
    except Exception:
        pass

    return model_idx, None, class_name

# --- Liste d'activités à analyser (à modifier) ---
activities_to_analyze = [
    {"id": 1, "name": "Transport des grumes des fôrets vers la menuiserie de Cotonou"},
    {"id": 2, "name": "Organisation d'une mission de collecte des données dans les zones dans le cadre de l'élaboration du PTAB 2022"},
    {"id": 3, "name": "Activité 4:Organisation trimestrielle des sessions de la Commission Nationale de surveillance, de securité et de sureté des transports fluvio-lagunaires"},
    {"id": 4, "name": "Octroi de secours temporaire (Appui à la mise en œuvre des AGR) au profit de 75 personnes vulnérables (personnes en situation difficile) dans les communes du Plateau"},
    {"id": 5, "name": "Acquisition d'un groupe électrogène de 50 Kva avec système d'inverseur au profit de la Direction Départementale de la Santé de l'Alibori pour alimenter la salle de réunion départementale"},
    {"id": 5, "name": "Appuyer l'organisation du concours inter-Classes Culturelles (Département du COUFFO)"},
    {"id": 1, "name": "Acquisition d'un  logiciel de gestion électronique des courriers et formation des sécrétaires sur son utilisation"},
    {"id": 1, "name": "Fournir une assistance technique et financiere au PNLMT pour rechercher les complications liées à la FL et la PEC des cas de lymphoedemes identifiés dans 9 communes"},
    {"id": 1, "name": "Activité opérationnelle : Elaborer le manuel des procédures de la CCMP"},
    {"id": 1, "name": "Identifier les différentes actions prioritaires d'adaptation des CDN"},
    {"id": 1, "name": "Poursuite des travaux de construction de quatre (04) magasins sur les sites des barrages de Péhunco, Kérou, Nikki et Kandi"},
    {"id": 1, "name": "Acquisition de vivres et de matériels dans le cadre de la coordination des secours  à apporter aux populations sinistrés du département de l'Atlantique"},
    {"id": 1, "name": "Célébration dans le département de l'Atlantique-Littoral de la Journée Mondiale de la Population (JMP) 2020"},
    {"id": 1, "name": "Visite du Chef de l'Etat au Brésil"},
    {"id": 1, "name": "Acquérir et mettre en place six (06) batteuses vanneuses au profit des coopératives de riziculteurs du Pôle3"},
    {"id": 1, "name": "Organiser les évènements culturels et les manifestations officielles (ANECSMO)"},
    {"id": 1, "name": "Renforcement des capacités des associations de jeunesse dans l'identification, la détection et la prévention des comportements déviants"},
    {"id": 1, "name": "Suivi de l'élaboration  des plans communaux de contingence"},
    {"id": 1, "name": "Organisation des visites de prospection par les 12 antennes départementales dans le cadre de l'insertion professionnelle des jeunes diplômés pour le compte d'octobre"},
    {"id": 1, "name": "Réalisation des latrines publiques                                   (Tranche BAD)"},
    {"id": 1, "name": "Réaliser l'étude de référence du projet"},
    {"id": 1, "name": "Réaliser les travaux d'aménagement de 2500 ha d'aires de pâturage"},
   
]

# --- Analyse des activités (code identique à votre version) ---
print("\n" + "="*80)
print(f"🚀 DÉBUT DE L'ANALYSE - {len(activities_to_analyze)} activité(s) à traiter")
print("="*80)

all_results = []

for i, activity in enumerate(activities_to_analyze):
    activity_id = activity["id"]
    activity_name = activity["name"]

    print(f"\n{'='*60}")
    print(f"📝 ACTIVITÉ #{i+1}/{len(activities_to_analyze)}")
    print(f"   ID: {activity_id}")
    print(f"   Nom: '{activity_name}'")
    print(f"{'='*60}")

    # 1. Prédiction ODD
    print("\n🎯 ÉTAPE 1: Prédiction ODD")
    odd_predictions = odd_classifier(activity_name)
    print(f"   Prédictions brutes: {odd_predictions}")

    odd_pred = get_first_prediction(odd_predictions)
    if not odd_pred:
        print("   ⚠️ Aucune prédiction ODD valide")
        continue

    odd_label = odd_pred.get("label", "LABEL_0")
    odd_id = int(odd_label.replace("LABEL_", "")) if isinstance(odd_label, str) and odd_label.startswith("LABEL_") else int(odd_label)
    odd_name = odd_names.get(str(odd_id), f"ODD {odd_id}")
    print(f"   ✅ ODD sélectionné: {odd_name} (ID={odd_id})")

    # 2. Prédiction Nature
    print("\n🌿 ÉTAPE 2: Prédiction Nature")
    nature_predictions = nature_classifier(activity_name)
    nature_pred = get_first_prediction(nature_predictions)

    if not nature_pred:
        nature_id = 4  # ID par défaut pour "Soutien"
        nature_name = "Soutien"
        print(f"   ⚠️ Aucune prédiction Nature - Défaut: {nature_name}")
    else:
        nature_label = nature_pred.get("label", "LABEL_0")
        model_idx, _, nature_name = map_model_label_to_nature(nature_label)
        nature_id = model_idx
        print(f"   ✅ Nature: {nature_name} (ID={nature_id})")

    # 3. Prédiction Cible
    print("\n🎯 ÉTAPE 3: Prédiction Cible")
    possible_cibles = [c for c in cibles_data if str(c["odd_code"]) == str(odd_id)]
    print(f"   {len(possible_cibles)} cibles possibles pour ODD {odd_id}")

    cible_match = None
    if possible_cibles:
        cible_predictions = cible_classifier(activity_name)
        cible_pred = get_first_prediction(cible_predictions)

        if cible_pred:
            cible_label = cible_pred.get("label", "0")
            try:
                cible_id = int(cible_label.replace("LABEL_", "")) if isinstance(cible_label, str) and cible_label.startswith("LABEL_") else int(cible_label)
                cible_match = next((c for c in possible_cibles if c["cible_id"] == cible_id), possible_cibles[0])
            except:
                cible_match = possible_cibles[0]
        else:
            cible_match = possible_cibles[0]

        print(f"   ✅ Cible sélectionnée: {cible_match['cible']} (ID={cible_match['cible_id']})")
    else:
        print("   ❌ Aucune cible trouvée pour cet ODD")

    # 4. Prédiction Indicateur
    print("\n📊 ÉTAPE 4: Prédiction Indicateur")
    indicateur_id = None
    indicateur_name = "Aucun indicateur"

    if cible_match and "indicateurs" in cible_match:
        indicateurs = cible_match["indicateurs"]
        indicateur_predictions = indicateur_classifier(activity_name)
        indicateur_pred = get_first_prediction(indicateur_predictions)

        if indicateur_pred:
            try:
                ind_label = indicateur_pred.get("label", "0")
                ind_id = int(ind_label.replace("LABEL_", "")) if isinstance(ind_label, str) and ind_label.startswith("LABEL_") else int(ind_label)
                ind_id = ind_id % len(indicateurs)
                selected = indicateurs[ind_id]
                indicateur_id = selected.get("id")
                indicateur_name = selected.get("name", "Aucun indicateur")
            except:
                indicateur_id = indicateurs[0].get("id")
                indicateur_name = indicateurs[0].get("name", "Aucun indicateur")

        print(f"   ✅ Indicateur: {indicateur_name} (ID={indicateur_id})")
    else:
        print("   ❌ Aucun indicateur disponible")

    # Stockage du résultat
    result_item = {
        "activity": activity_name,
        "activity_id": activity_id,
        "odd": {"id": odd_id, "name": odd_name},
        "nature": {"id": nature_id, "name": nature_name},
        "cible": {"id": cible_match["cible_id"] if cible_match else None, "name": cible_match["cible"] if cible_match else None},
        "indicateur": {"id": indicateur_id, "name": indicateur_name}
    }
    all_results.append(result_item)

# --- Affichage des résultats ---
print(f"\n{'='*80}")
print(f"🏁 FIN DE L'ANALYSE - {len(all_results)} résultat(s)")
print("="*80)

for i, result in enumerate(all_results, 1):
    print(f"\n{'='*60}")
    print(f"RÉSULTAT #{i}: {result['activity']}")
    print(f"   ODD: {result['odd']['name']} (ID: {result['odd']['id']})")
    print(f"   Nature: {result['nature']['name']} (ID: {result['nature']['id']})")
    print(f"   Cible: {result['cible']['name'] if result['cible']['name'] else 'N/A'}")
    print(f"   Indicateur: {result['indicateur']['name']}")
    print(f"{'='*60}")

# Sauvegarde optionnelle
save = input("\nSauvegarder les résultats? (o/n): ").strip().lower()
if save == 'o':
    with open("resultats_analyse.json", "w", encoding="utf-8") as f:
        json.dump({"results": all_results}, f, indent=2, ensure_ascii=False)
    print("✅ Résultats sauvegardés dans résultats_analyse.json")
