"""
Addestra LightFM in modalita' IBRIDA (collaborative + content-based).

Usa sia le interazioni user-item (reviews) sia le item features
(categorie prodotto + fasce di prezzo) per generare embedding piu' ricchi.

Input:  artifacts/reviews_cleaned.json
        artifacts/metadata_cleaned.json
Output: artifacts/recommender_model.pkl

Tempo stimato: 2-7 minuti.
"""
import json
import os
import sys
import pickle
from collections import Counter
from lightfm import LightFM
from lightfm.data import Dataset

PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
REVIEWS_FILE = os.path.join(PROJECT_ROOT, "artifacts", "reviews_cleaned.json")
METADATA_FILE = os.path.join(PROJECT_ROOT, "artifacts", "metadata_cleaned.json")
MODEL_FILE = os.path.join(PROJECT_ROOT, "artifacts", "recommender_model.pkl")

# --- Configurazione Feature ---

# Categoria presente sul 99.98% dei prodotti, zero valore discriminante
EXCLUDED_EXACT = {"Clothing, Shoes & Jewelry"}

# Keyword per filtrare categorie promozionali / rumore
NOISE_KEYWORDS = {
    "sale", "off", "save", "deal", "prime", "gift", "clearance",
    "black friday", "holiday", "test", "exclusion", "shopbop",
    "cyber", "top 50", "top rated", "featured", "our brands"
}

# Minimo numero di prodotti per considerare una categoria
MIN_CATEGORY_COUNT = 10

# Fasce di prezzo
PRICE_BUCKETS = [
    ("price:budget",      0,   15),
    ("price:affordable", 15,   30),
    ("price:mid",        30,   50),
    ("price:premium",    50,  100),
    ("price:luxury",    100, 10000),
]


def get_price_feature(price):
    if price is None or price == "N/A":
        return None
    try:
        p = float(price)
        for name, lo, hi in PRICE_BUCKETS:
            if lo <= p < hi:
                return name
        return "price:luxury"
    except (ValueError, TypeError):
        return None


def flatten_categories(categories):
    """Appiattisce le categorie (possono essere stringhe o liste annidate)."""
    result = []
    for c in categories:
        if isinstance(c, str):
            result.append(c)
        elif isinstance(c, list):
            for sub in c:
                if isinstance(sub, str):
                    result.append(sub)
    return result


def main():
    print("--- TRAINING RECOMMENDER IBRIDO (LightFM + Item Features) ---")

    # --- 1. Caricamento reviews ---
    if not os.path.exists(REVIEWS_FILE):
        print(f"Errore: {REVIEWS_FILE} non trovato!")
        sys.exit(1)

    raw_data = []
    users_set = set()
    items_set = set()

    print("-> Lettura interazioni...")
    with open(REVIEWS_FILE, "r", encoding="utf-8") as f:
        for line in f:
            try:
                d = json.loads(line)
                u, p = d["reviewerID"], d["asin"]
                raw_data.append((u, p))
                users_set.add(u)
                items_set.add(p)
            except Exception:
                continue

    print(f"   Interazioni: {len(raw_data):,}")
    print(f"   Utenti unici: {len(users_set):,}")
    print(f"   Prodotti unici: {len(items_set):,}")

    if len(raw_data) == 0:
        print("Errore: nessuna interazione trovata!")
        sys.exit(1)

    # --- 2. Caricamento metadati per le features ---
    if not os.path.exists(METADATA_FILE):
        print(f"Errore: {METADATA_FILE} non trovato!")
        sys.exit(1)

    print("-> Caricamento metadati prodotti...")
    with open(METADATA_FILE, "r", encoding="utf-8") as f:
        products = json.load(f)
    meta_by_asin = {p["asin"]: p for p in products}

    # --- 3. Analisi e filtraggio categorie ---
    print("-> Analisi categorie...")
    cat_counter = Counter()
    for asin in items_set:
        p = meta_by_asin.get(asin)
        if p:
            for c in flatten_categories(p.get("categories", [])):
                if c not in EXCLUDED_EXACT:
                    cat_counter[c] += 1

    allowed_categories = set()          # Set vuoto dove metteremo le categorie "buone"
    excluded_noise = 0
    excluded_rare = 0
    for cat, count in cat_counter.items():
        if count < MIN_CATEGORY_COUNT:  # MIN_CATEGORY_COUNT = 10
            excluded_rare += 1
            continue                    # Salta, non la aggiunge
        lower = cat.lower()
        if any(kw in lower for kw in NOISE_KEYWORDS):
            excluded_noise += 1
            continue
        allowed_categories.add(cat)

    print(f"   Categorie totali trovate: {len(cat_counter)}")
    print(f"   Escluse (rumore/promo): {excluded_noise}")
    print(f"   Escluse (troppo rare <{MIN_CATEGORY_COUNT}): {excluded_rare}")
    print(f"   Categorie utili: {len(allowed_categories)}")

    # --- 4. Preparazione feature names ---
    price_feature_names = [name for name, _, _ in PRICE_BUCKETS]
    all_feature_names = sorted(allowed_categories) + price_feature_names

    print(f"   Feature prezzo: {len(price_feature_names)}")
    print(f"   Feature totali: {len(all_feature_names)}")

    # --- 5. Costruzione Dataset LightFM con features ---
    print("-> Costruzione Dataset LightFM...")
    dataset = Dataset()
    dataset.fit(
        users=list(users_set),
        items=list(items_set),
        item_features=all_feature_names
    )

    (interactions, _) = dataset.build_interactions(raw_data)

    # --- 6. Costruzione matrice item features ---
    print("-> Costruzione matrice item features...")
    item_feature_tuples = []        # Lista di coppie (ASIN, [lista feature])
    items_with_features = 0
    items_with_price = 0

    for asin in items_set:
        p = meta_by_asin.get(asin)  # Cerca i metadati di questo prodotto dentro a "metadata_cleaned.json"
        features = []
        if p:
            for c in flatten_categories(p.get("categories", [])):
                if c in allowed_categories:
                    features.append(c)
            price_feat = get_price_feature(p.get("price"))
            if price_feat:
                features.append(price_feat)
                items_with_price += 1
        if features:
            item_feature_tuples.append((asin, features))
            items_with_features += 1

                                                                            # normalize=True fa sì che i pesi delle feature siano normalizzati (sommano a 1 per prodotto), 
                                                                            # così un prodotto con 8 categorie non ha più "peso" di uno con 2.
    item_features_matrix = dataset.build_item_features(item_feature_tuples, normalize=True)

    print(f"   Prodotti con features: {items_with_features:,} / {len(items_set):,}")
    print(f"   Prodotti con fascia prezzo: {items_with_price:,}")
    print(f"   Shape matrice: {item_features_matrix.shape}")

    # --- 7. Training ibrido ---
    print("-> Addestramento Modello Ibrido (WARP, 30 epochs)...")
    model = LightFM(loss="warp", no_components=10)
    model.fit(
        interactions,                        # matrice sparsa [utenti x prodotti] (131k interazioni)
        item_features=item_features_matrix,  # matrice sparsa [prodotti x (prodotti + 413 feature)]
        epochs=30,                           # passa 30 volte su tutti i dati
        num_threads=2                        # parallelizza su 2 thread CPU
    )

    print(f"   Shape item_embeddings: {model.item_embeddings.shape}")

    # --- 8. Salvataggio ---
    print("-> Salvataggio modello...")
    user_id_map, _, item_id_map, _ = dataset.mapping()
    idx_to_item = {v: k for k, v in item_id_map.items()}

    package = {
        "model": model,
        "item_id_map": item_id_map,
        "idx_to_item": idx_to_item,
        "item_features_matrix": item_features_matrix,
    }

    with open(MODEL_FILE, "wb") as f:
        pickle.dump(package, f)

    print(f"\n--- COMPLETATO ---")
    print(f"Modello salvato in: {MODEL_FILE}")
    print(f"Tipo: Ibrido (collaborative + {len(allowed_categories)} categorie + {len(price_feature_names)} fasce prezzo)")


if __name__ == "__main__":
    main()
