# ============================================================
#  ERP AI - CRM Intelligence API
#  Flask + PostgreSQL + Modeles ML (joblib)
#  Multi-tenant via JWT tenantSchema
# ============================================================
import os, warnings, logging, json, time, shutil, threading
from datetime import datetime
warnings.filterwarnings("ignore")
logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(message)s")
log = logging.getLogger("ERP_AI")

from flask import Flask, jsonify, send_from_directory, g
from flask_cors import CORS

app = Flask(__name__, static_folder="static")
CORS(app)

import pandas as pd
import numpy as np
import joblib
import psycopg2
from sklearn.linear_model    import LogisticRegression
from sklearn.calibration     import CalibratedClassifierCV
from sklearn.model_selection import train_test_split
from sklearn.preprocessing   import RobustScaler, normalize
from sklearn.metrics         import roc_auc_score, mean_absolute_error, r2_score
from scipy.sparse            import csr_matrix
from scipy.sparse.linalg     import svds
import xgboost as xgb
from flask_jwt_middleware import require_jwt

# ==============================================================
#  CONFIGURATION
# ==============================================================
from dotenv import load_dotenv
BASE_DIR   = os.path.dirname(os.path.abspath(__file__))
MODELS_DIR = os.path.join(BASE_DIR, "models")
load_dotenv(os.path.join(BASE_DIR, ".env"))

DB_CONFIG = {
    "host":            os.getenv("DB_HOST"),
    "port":            int(os.getenv("DB_PORT", 5432)),
    "database":        os.getenv("DB_NAME"),
    "user":            os.getenv("DB_USER"),
    "password":        os.getenv("DB_PASSWORD"),
    "client_encoding": "utf8",
}

# Noms de tables SANS préfixe — schéma ajouté dynamiquement
TABLE_CUSTOMERS = os.getenv("TABLE_CUSTOMERS")
TABLE_INVOICES  = os.getenv("TABLE_INVOICES")
TABLE_PRODUCTS  = os.getenv("TABLE_PRODUCTS")
TABLE_LINES     = os.getenv("TABLE_LINES")

SNAPSHOT_DATE     = pd.Timestamp(os.getenv("SNAPSHOT_DATE"))
PREDICTION_WINDOW = int(os.getenv("PREDICTION_WINDOW", 60))
TRAIN_CUTOFF      = SNAPSHOT_DATE - pd.Timedelta(days=PREDICTION_WINDOW)
SVD_K             = int(os.getenv("SVD_K", 30))
TOP_K             = int(os.getenv("TOP_K", 10))
TRAIN_STATUS_OK   = [0, 1, 2, 4, 5]
LABEL_STATUS_OK   = [0, 4, 5]
RANDOM_STATE      = 42
BEST_THRESHOLD    = float(os.getenv("BEST_THRESHOLD", 0.85))
ACTIVE_THRESHOLD  = float(os.getenv("ACTIVE_THRESHOLD", 0.30))

# ==============================================================
#  CHARGEMENT DES MODELES ML
# ==============================================================
log.info("Chargement des modeles ML...")
try:
    model_wb   = joblib.load(os.path.join(MODELS_DIR, "willbuy.pkl"))
    model_disc = joblib.load(os.path.join(MODELS_DIR, "discount.pkl"))
    scaler     = joblib.load(os.path.join(MODELS_DIR, "scaler.pkl"))
    te_maps    = joblib.load(os.path.join(MODELS_DIR, "te_encoders.pkl"))
    reco_data  = joblib.load(os.path.join(MODELS_DIR, "reco_svd.pkl"))
    MODELS_OK  = True
    log.info("[OK] Modeles charges — LogReg (M1) | XGBoost (M3) | SVD k=30 (M2)")
except Exception as e:
    MODELS_OK = False
    log.warning(f"[WARN] Modeles non trouves : {e}")

# ==============================================================
#  PROBAS PRÉ-CALCULÉES
# ==============================================================
DATA_DIR    = os.path.join(BASE_DIR, "data")
PROBA_CSV   = os.path.join(DATA_DIR, "all_proba_400.csv")
NOTEBOOK_PROBAS = None
os.makedirs(DATA_DIR, exist_ok=True)

if os.path.exists(PROBA_CSV):
    try:
        df_p = pd.read_csv(PROBA_CSV)
        NOTEBOOK_PROBAS = df_p.set_index("customers_id")["will_buy_proba"].to_dict()
        log.info(f"[OK] Probas notebook chargees — n={len(NOTEBOOK_PROBAS)}")
    except Exception as e:
        log.warning(f"[WARN] Probas non chargees : {e}")

# ==============================================================
#  FIX 1 — get_current_schema()
#  Retourne le schéma du JWT. Lève une erreur explicite si absent.
#  Accepte aussi un schéma passé explicitement (pour le scheduler/retrain).
# ==============================================================
def get_current_schema(explicit_schema: str = None) -> str:
    # Priorité 1 : schéma passé explicitement (scheduler, retrain background)
    if explicit_schema:
        return explicit_schema

    # Priorité 2 : schéma du JWT dans le contexte Flask request
    try:
        schema = getattr(g, 'tenant_schema', None)
        if schema:
            return schema
    except RuntimeError:
        pass  # hors contexte Flask

    # Pas de fallback — erreur explicite
    raise RuntimeError(
        "[SECURITE] Aucun tenantSchema disponible. "
        "JWT requis pour les appels API. "
        "Passer un schema explicite pour les tâches background."
    )

# ==============================================================
#  FIX 2 — get_connection()
#  Guard sur g.username pour éviter le crash hors contexte request.
# ==============================================================
def get_connection_raw():
    return psycopg2.connect(**DB_CONFIG)

def get_connection(explicit_schema: str = None):
    conn = get_connection_raw()
    schema = get_current_schema(explicit_schema)
    with conn.cursor() as cur:
        cur.execute(f"SET search_path TO {schema}, public;")
        cur.execute("SET client_encoding TO 'UTF8';")
    conn.commit()
    # Guard: g n'existe pas hors contexte request
    try:
        username = getattr(g, 'username', 'system')
    except RuntimeError:
        username = 'system'
    log.debug(f"[DB] schema={schema} | user={username}")
    return conn

# ==============================================================
#  FIX 3 — test_connection()
#  Utilise get_connection_raw() — pas de schéma, pas de JWT requis.
#  Appelé au démarrage et dans /api/health.
# ==============================================================
def test_connection():
    try:
        conn = get_connection_raw()
        conn.close()
        log.info(f"[OK] PostgreSQL → {DB_CONFIG['host']}:{DB_CONFIG['port']}/{DB_CONFIG['database']}")
        return True
    except Exception as e:
        log.error(f"[ERR] PostgreSQL : {e}")
        return False

DB_OK = test_connection()  # OK maintenant — pas de schéma requis

# ==============================================================
#  FIX 4 — load_tables()
#  Préfixe chaque table avec le schéma résolu explicitement.
# ==============================================================
def load_tables(schema: str = None):
    schema = get_current_schema(schema)
    conn   = get_connection(schema)

    # Noms de tables pleinement qualifiés
    t_cust  = f"{schema}.{TABLE_CUSTOMERS}"
    t_inv   = f"{schema}.{TABLE_INVOICES}"
    t_prod  = f"{schema}.{TABLE_PRODUCTS}"
    t_lines = f"{schema}.{TABLE_LINES}"

    try:
        log.info(f"    [tenant={schema}] Chargement tables...")
        customers = pd.read_sql(f"""
            SELECT id, city, nationality, type, balance,
                   clientcategory, organization_name,
                   timestamp AS created_at
            FROM {t_cust}
        """, conn)

        invoices = pd.read_sql(f"""
            SELECT id, customers_id, issue_date,
                   untaxed_amount, status, payment_type
            FROM {t_inv}
            WHERE issue_date IS NOT NULL
              AND customers_id IS NOT NULL
              AND issue_date <= '{SNAPSHOT_DATE.date()}'
        """, conn)

        products = pd.read_sql(f"""
            SELECT id, product_name, category,
                   selling_price, unit_coast
            FROM {t_prod}
        """, conn)

        lines = pd.read_sql(f"""
            SELECT pp.invoice_id, pp.product_id,
                   pp.quantity, pp.discount, pp.unitprice
            FROM {t_lines} pp
            WHERE pp.invoice_id IN (
                SELECT id FROM {t_inv}
                WHERE issue_date IS NOT NULL
                  AND customers_id IS NOT NULL
                  AND issue_date <= '{SNAPSHOT_DATE.date()}'
                  AND issue_date >= '2023-01-01'
            )
            LIMIT 75000
        """, conn)

        log.info(f"    {len(customers)} clients | {len(invoices)} factures | {len(lines)} lignes")
        return customers, invoices, products, lines
    finally:
        conn.close()

# ==============================================================
#  NETTOYAGE
# ==============================================================
def clean_tables(customers, invoices, products, lines):
    invoices["issue_date"] = pd.to_datetime(invoices["issue_date"], errors="coerce")
    invoices = invoices.dropna(subset=["issue_date"])
    if "created_at" in customers.columns:
        customers["created_at"] = pd.to_datetime(customers["created_at"], errors="coerce")
    customers["city"]        = customers["city"].str.strip().str.title()
    customers["nationality"] = customers["nationality"].str.strip().str.title()
    products["category"]     = products["category"].str.strip().str.upper()
    lines = lines.drop_duplicates()
    return customers, invoices, products, lines

# ==============================================================
#  FEATURE ENGINEERING
# ==============================================================
FEATURES = ["recency_days", "frequency", "log_monetary_ht",
            "freq_per_month", "tenure_days",
            "purchase_concentration", "type"]

def build_features(customers, invoices, products, lines):
    invoices["status"] = pd.to_numeric(
        invoices["status"], errors="coerce").fillna(-1).astype(int)
    invoices["untaxed_amount"] = pd.to_numeric(
        invoices["untaxed_amount"], errors="coerce").fillna(0)
    invoices = invoices.dropna(subset=["customers_id"])
    invoices["customers_id"] = invoices["customers_id"].astype(int)

    valid_train = invoices[
        (invoices["issue_date"] < TRAIN_CUTOFF) &
        (invoices["status"].isin(TRAIN_STATUS_OK))
    ].copy()
    valid_train["is_paid"] = (valid_train["status"] == 0).astype(int)

    valid_label = invoices[
        (invoices["issue_date"] >= TRAIN_CUTOFF) &
        (invoices["issue_date"] <= SNAPSHOT_DATE) &
        (invoices["status"].isin(LABEL_STATUS_OK))
    ]
    buyers = set(valid_label["customers_id"].dropna().astype(int))
    y = pd.Series(
        {int(c): int(c in buyers) for c in customers["id"].dropna()},
        name="will_buy"
    )

    lines_rich = lines.merge(
        products[["id","category","selling_price","unit_coast"]],
        left_on="product_id", right_on="id", how="left"
    ).drop(columns=["id"], errors="ignore")
    if "unitprice" in lines_rich.columns and "unit_coast" in lines_rich.columns:
        lines_rich["margin_pct"] = (
            (lines_rich["unitprice"] - lines_rich["unit_coast"]) /
            lines_rich["unitprice"].replace(0, np.nan)
        ).fillna(0).clip(-1, 1)
    else:
        lines_rich["margin_pct"] = 0.0

    inv_agg = lines_rich.groupby("invoice_id").agg(
        avg_discount_line=("discount","mean"),
        avg_margin_pct   =("margin_pct","mean"),
        pct_service      =("category", lambda x: (x=="SERVICE").mean()),
    ).reset_index()
    train_full = valid_train.merge(
        inv_agg, left_on="id", right_on="invoice_id", how="left"
    ).drop(columns=["invoice_id"], errors="ignore")

    rfm = train_full.groupby("customers_id").agg(
        recency_days=("issue_date", lambda x: (TRAIN_CUTOFF - x.max()).days),
        frequency   =("id", "count"),
        monetary_ht =("untaxed_amount", "sum"),
        std_ht      =("untaxed_amount", "std"),
        mean_ht     =("untaxed_amount", "mean"),
    ).reset_index()
    rfm["log_monetary_ht"] = np.log1p(rfm["monetary_ht"])
    rfm["cv_basket_ht"]    = (rfm["std_ht"] / rfm["mean_ht"].replace(0, np.nan)).fillna(0)
    rfm["freq_per_month"]  = rfm["frequency"] / ((rfm["recency_days"] + 1) / 30)
    rfm.drop(columns=["monetary_ht","std_ht","mean_ht"], inplace=True)

    if "created_at" in customers.columns:
        t = customers[["id","created_at"]].copy()
        t["tenure_days"] = (TRAIN_CUTOFF - t["created_at"]).dt.days.clip(lower=0)
        rfm = rfm.merge(
            t[["id","tenure_days"]], left_on="customers_id",
            right_on="id", how="left"
        ).drop(columns=["id"], errors="ignore")

    train_full["month"] = train_full["issue_date"].dt.month
    monthly = (train_full.groupby(["customers_id","month"]).size()
               .unstack(fill_value=0).reindex(columns=range(1,13), fill_value=0))
    monthly = monthly.div(monthly.sum(axis=1), axis=0)
    conc = ((monthly**2).sum(axis=1).rename("purchase_concentration").reset_index())

    pp = (valid_train.groupby(["customers_id","payment_type"]).size()
          .unstack(fill_value=0)
          .rename(columns={0:"n_cash",1:"n_cheque",2:"n_transfer",3:"n_credit"}))
    for c in ["n_cash","n_cheque","n_transfer","n_credit"]:
        if c not in pp.columns: pp[c] = 0
    pct_pay = (pp[["n_cash","n_cheque"]]
               .div(pp.sum(axis=1), axis=0)
               .rename(columns={"n_cash":"pct_pay_cash","n_cheque":"pct_pay_check"})
               .reset_index())
    pay = train_full.groupby("customers_id").agg(
        paid_rate         =("is_paid","mean"),
        avg_discount_recv =("avg_discount_line","mean"),
        pct_service_orders=("pct_service","mean"),
        avg_margin_client =("avg_margin_pct","mean"),
    ).reset_index().merge(pct_pay, on="customers_id", how="left")

    cf = (customers
          .merge(rfm,  left_on="id", right_on="customers_id", how="left")
          .merge(pay,  left_on="id", right_on="customers_id", how="left")
          .merge(conc, left_on="id", right_on="customers_id", how="left"))
    cf.drop(columns=[c for c in cf.columns if "customers_id" in c], inplace=True)

    IMPUTE = {
        "recency_days":          int(rfm["recency_days"].quantile(0.99)+30) if len(rfm) > 0 else 999,
        "frequency":             0,
        "log_monetary_ht":       0,
        "freq_per_month":        0,
        "purchase_concentration":0,
        "tenure_days":           0,
    }
    for col, val in IMPUTE.items():
        if col in cf.columns:
            cf[col] = cf[col].fillna(val)
    cf = cf.replace([np.inf, -np.inf], 0).fillna(0)
    return cf, y, valid_train

# ==============================================================
#  MODULE 1 — WILL-BUY
# ==============================================================
def predict_will_buy(cf):
    if NOTEBOOK_PROBAS is not None:
        proba = np.array([NOTEBOOK_PROBAS.get(int(cid), 0.10) for cid in cf["id"]])
        log.info(f"[M1] Probas notebook max={proba.max():.4f}")
        return proba

    X = pd.DataFrame(index=cf.index)
    for f in FEATURES:
        X[f] = pd.to_numeric(cf.get(f, 0), errors="coerce").fillna(0)
    X = X.replace([np.inf, -np.inf], 0).astype(float)
    no_sc = ["type"]
    to_sc = [c for c in X.columns if c not in no_sc]
    X_sc  = X.copy()
    try:
        X_sc[to_sc] = scaler.transform(X[to_sc])
    except Exception as e:
        log.warning(f"Scaler transform echoue : {e}")
    proba = model_wb.predict_proba(X_sc)[:, 1]
    return np.clip(proba, 0.01, 0.98)

def proba_to_segment(p):
    if p >= 0.85: return "5 — Champion"
    if p >= 0.70: return "4 — Chaud"
    if p >= 0.50: return "3 — Tiede"
    if p >= 0.30: return "2 — Froid"
    return "1 — Inactif"

def proba_to_confidence(p):
    if p >= 0.85: return "Tres Eleve"
    if p >= 0.70: return "Eleve"
    if p >= 0.50: return "Moyen"
    if p >= 0.30: return "Faible"
    return "Tres Faible"

# ==============================================================
#  FIX 4b — get_recommendations()
#  Préfixe TABLE_PRODUCTS avec le schéma courant.
# ==============================================================
def get_recommendations(customer_ids, top_k=3, schema: str = None):
    if not MODELS_OK or reco_data is None:
        return {}

    R_hat  = reco_data["R_hat"]
    cids   = reco_data["cids"]
    pids   = reco_data["pids"]
    p_name = reco_data["p_name"]
    ci     = {c: i for i, c in enumerate(cids)}

    p_price_map = {}
    try:
        schema = get_current_schema(schema)
        conn   = get_connection(schema)
        prod_df = pd.read_sql(
            f"SELECT id, selling_price FROM {schema}.{TABLE_PRODUCTS}", conn)
        conn.close()
        p_price_map = prod_df.set_index("id")["selling_price"].to_dict()
    except Exception as e:
        log.warning(f"[M2] Prix non charges : {e}")

    results = {}
    for cid in customer_ids:
        cid = int(cid)
        if cid not in ci:
            results[cid] = []
            continue
        sc    = R_hat[ci[cid]].copy()
        top_i = np.argsort(sc)[::-1][:top_k]
        recs  = []
        for i in top_i:
            pid   = pids[i]
            name  = p_name.get(pid, f"Produit #{pid}")
            price = p_price_map.get(pid, 0)
            recs.append({"name": name, "price": round(float(price), 2)})
        results[cid] = recs
    return results

# ==============================================================
#  MODULE 3 — DISCOUNT
# ==============================================================
D_FEATS = [
    "frequency","log_monetary_ht","recency_days",
    "avg_discount_recv","paid_rate","pct_pay_cash",
    "pct_pay_check","cv_basket_ht","balance",
    "pct_service_orders","avg_margin_client",
    "purchase_concentration","tenure_days",
    "will_buy_proba","client_group","freq_per_month",
]

DISC_RULES_V2 = {
    (1,1):(5,12),(1,2):(5,15),(1,3):(5,15),(1,4):(5,15),
    (0,1):(5,12),(0,2):(5,15),(0,3):(7,18),(0,4):(10,20),
}

def apply_rules_v2(pred, client_type, client_group):
    lo, hi = DISC_RULES_V2.get((int(client_type), int(client_group)), (5, 18))
    return round(float(np.clip(pred, lo, hi)), 1)

def get_client_group(frequency):
    if frequency <= 6:  return 1
    if frequency <= 16: return 2
    if frequency <= 34: return 3
    return 4

# ==============================================================
#  FIX 4c — predict_discount()
#  Préfixe TABLE_LINES et TABLE_INVOICES avec le schéma courant.
# ==============================================================
def predict_discount(cf, proba, schema: str = None):
    schema    = get_current_schema(schema)
    disc_feat = cf.copy()
    disc_feat["will_buy_proba"] = proba
    disc_feat["client_group"]   = disc_feat["frequency"].apply(get_client_group)

    try:
        conn = get_connection(schema)
        disc_sql = pd.read_sql(f"""
            SELECT pp.discount, i.customers_id
            FROM {schema}.{TABLE_LINES} pp
            JOIN {schema}.{TABLE_INVOICES} i ON pp.invoice_id = i.id
            WHERE i.issue_date < '{TRAIN_CUTOFF.date()}'
              AND i.customers_id IS NOT NULL
        """, conn)
        conn.close()
        disc_agg  = disc_sql.groupby("customers_id").agg(avg_disc=("discount","mean")).reset_index()
        disc_feat = disc_feat.merge(disc_agg, left_on="id", right_on="customers_id", how="left")
        disc_feat["avg_discount_recv"] = disc_feat.get("avg_disc", pd.Series(0)).fillna(0)
    except Exception as e:
        log.warning(f"Discount historique non charge : {e}")
        disc_feat["avg_discount_recv"] = 0

    X_d = pd.DataFrame(index=disc_feat.index)
    for f in D_FEATS:
        X_d[f] = pd.to_numeric(disc_feat.get(f, 0), errors="coerce").fillna(0) \
                 if f in disc_feat.columns else 0.0
    X_d = X_d.replace([np.inf, -np.inf], 0).astype(float)
    raw = np.clip(model_disc.predict(X_d), 0, 30)

    discounts = []
    for i in range(len(disc_feat)):
        ctype  = int(disc_feat["type"].iloc[i])         if "type"         in disc_feat.columns else 0
        cgroup = int(disc_feat["client_group"].iloc[i]) if "client_group" in disc_feat.columns else 2
        discounts.append(apply_rules_v2(raw[i], ctype, cgroup))
    return discounts

# ==============================================================
#  PIPELINE — accepte schema explicite
# ==============================================================
def run_pipeline(schema: str = None):
    schema = get_current_schema(schema)
    log.info(f"[PIPELINE] Démarrage — tenant={schema}")
    t0 = time.time()

    customers, invoices, products, lines = load_tables(schema)
    customers, invoices, products, lines = clean_tables(customers, invoices, products, lines)

    cf, y, valid_train = build_features(customers, invoices, products, lines)

    proba = predict_will_buy(cf)
    cf["will_buy_proba"] = proba
    cf["will_buy"]       = (proba >= BEST_THRESHOLD).astype(int)
    cf["segment_crm"]    = [proba_to_segment(p)    for p in proba]
    cf["confidence"]     = [proba_to_confidence(p) for p in proba]
    cf["discount_final"] = predict_discount(cf, proba, schema)

    recs = get_recommendations(cf["id"].tolist(), top_k=3, schema=schema)

    if "organization_name" not in cf.columns:
        cf["organization_name"] = cf["id"].apply(lambda x: f"Client #{x}")

    def get_rec(cid, idx, field):
        r = recs.get(int(cid), [])
        if idx < len(r): return r[idx].get(field, "-" if field == "name" else 0)
        return "-" if field == "name" else 0

    result = []
    cols = ["id","organization_name","will_buy_proba","segment_crm",
            "confidence","city","type","discount_final","client_group"]
    cols = [c for c in cols if c in cf.columns]

    for row in cf[cols].itertuples(index=False):
        cid       = int(row.id)
        typ       = int(row.type) if not np.isnan(float(row.type)) else 0
        grp       = int(row.client_group) if hasattr(row, "client_group") else 2
        name      = str(row.organization_name) if pd.notna(row.organization_name) else f"Client #{cid}"
        proba_val = float(row.will_buy_proba)

        result.append({
            "customer_id":    cid,
            "customer_name":  name,
            "probability":    round(proba_val, 4),
            "segment_crm":    str(row.segment_crm),
            "confidence":     str(row.confidence),
            "city":           str(row.city),
            "type":           typ,
            "client_group":   grp,
            "will_buy":       1 if proba_val >= BEST_THRESHOLD else 0,
            "top_product":    get_rec(cid, 0, "name"),
            "top_price":      get_rec(cid, 0, "price"),
            "top_product_2":  get_rec(cid, 1, "name"),
            "top_price_2":    get_rec(cid, 1, "price"),
            "top_product_3":  get_rec(cid, 2, "name"),
            "top_price_3":    get_rec(cid, 2, "price"),
            "discount":       float(row.discount_final),
            "discount_label": "Medium (>10%)" if float(row.discount_final) > 10 else "Standard (5-10%)",
        })

    result.sort(key=lambda x: x["probability"], reverse=True)
    for i, r in enumerate(result): r["rank"] = i + 1

    result_active = [r for r in result if r["segment_crm"] != "1 — Inactif"]
    for i, r in enumerate(result_active): r["rank"] = i + 1

    log.info(f"[PIPELINE] OK {time.time()-t0:.1f}s — tenant={schema} | "
             f"{len(result_active)} actifs / {len(result)} total")
    return result_active, result, cf

# ==============================================================
#  FIX 5 — CACHE MULTI-TENANT
#  Clé = tenantSchema. Chaque tenant a son propre cache isolé.
# ==============================================================
_cache: dict = {}  # { "company_42": { will_buy, all, cf, timestamp }, ... }

def get_cached_data(force_refresh=False, schema: str = None):
    schema = get_current_schema(schema)

    if schema not in _cache or _cache[schema]["will_buy"] is None or force_refresh:
        will_buy, all_customers, cf = run_pipeline(schema)
        _cache[schema] = {
            "will_buy":  will_buy,
            "all":       all_customers,
            "cf":        cf,
            "timestamp": pd.Timestamp.now(),
        }
        log.info(f"[CACHE] Mis à jour — tenant={schema} | {len(will_buy)} actifs")

    return (
        _cache[schema]["will_buy"],
        _cache[schema]["all"],
        _cache[schema]["cf"],
    )

def get_cache_ts(schema: str = None):
    """Retourne le timestamp du cache pour le schéma courant."""
    schema = get_current_schema(schema)
    return _cache.get(schema, {}).get("timestamp")

# ==============================================================
#  ENDPOINTS FLASK
# ==============================================================
@app.route("/api/crm", methods=["GET"])
@require_jwt
def get_crm():
    try:
        will_buy, all_customers, _ = get_cached_data()
        ts = get_cache_ts()
        return jsonify({
            "customers":  will_buy,
            "total":      len(will_buy),
            "total_all":  len(all_customers),
            "threshold":  BEST_THRESHOLD,
            "tenant":     g.tenant_schema,
            "updated_at": ts.isoformat() if ts else None,
        })
    except Exception as e:
        log.error(f"[ERR] /api/crm : {e}")
        return jsonify({"error": str(e), "customers": [], "total": 0}), 500


@app.route("/api/stats", methods=["GET"])
@require_jwt
def get_stats():
    try:
        will_buy, all_customers, cf = get_cached_data()
        seg_dist = {}
        avg_disc = 0
        top_city = "-"

        if cf is not None and len(cf) > 0:
            seg_dist = cf["segment_crm"].value_counts().to_dict() \
                       if "segment_crm" in cf.columns else {}
            will_buy_mask = cf["will_buy_proba"] >= BEST_THRESHOLD \
                            if "will_buy_proba" in cf.columns \
                            else pd.Series([False]*len(cf))
            avg_disc = round(float(
                cf.loc[will_buy_mask, "discount_final"].mean()
            ), 1) if will_buy_mask.sum() > 0 and "discount_final" in cf.columns else 0
            if will_buy_mask.sum() > 0 and "city" in cf.columns:
                top_city = cf.loc[will_buy_mask, "city"].value_counts().index[0]
            elif "city" in cf.columns:
                top_city = cf["city"].value_counts().index[0]

        return jsonify({
            "snapshot_date":  str(SNAPSHOT_DATE.date()),
            "train_cutoff":   str(TRAIN_CUTOFF.date()),
            "best_threshold": BEST_THRESHOLD,
            "tenant":         g.tenant_schema,
            "crm": {
                "total_customers": len(all_customers),
                "will_buy_count":  len(will_buy),
                "avg_discount":    avg_disc,
                "top_city":        top_city,
                "champions":       int(seg_dist.get("5 — Champion", 0)),
                "seg_distribution":seg_dist,
            },
        })
    except Exception as e:
        log.error(f"[ERR] /api/stats : {e}")
        return jsonify({"error": str(e)}), 500


@app.route("/api/refresh", methods=["POST"])
@require_jwt
def refresh():
    try:
        schema   = g.tenant_schema
        will_buy, _, _ = get_cached_data(force_refresh=True)
        ts       = get_cache_ts()
        return jsonify({
            "status":         "ok",
            "tenant":         schema,
            "will_buy_count": len(will_buy),
            "updated_at":     ts.isoformat() if ts else None,
        })
    except Exception as e:
        return jsonify({"status": "error", "message": str(e)}), 500


@app.route("/api/health", methods=["GET"])
def health():
    # Pas de JWT requis — monitoring infra
    db_ok    = test_connection()
    return jsonify({
        "status":   "ok" if db_ok and MODELS_OK else "degraded",
        "db":       "[OK]" if db_ok    else "[ERR]",
        "models":   "[OK]" if MODELS_OK else "[ERR]",
        "tenants_cached": list(_cache.keys()),
        "db_host":  f"{DB_CONFIG['host']}:{DB_CONFIG['port']}/{DB_CONFIG['database']}",
    })


@app.route("/api/powerbi/clients", methods=["GET"])
@require_jwt
def powerbi_clients():
    try:
        will_buy, all_customers, _ = get_cached_data()
        ts = get_cache_ts()
        rows = []
        for c in all_customers:
            rows.append({
                "customer_id":     c.get("customer_id"),
                "customer_name":   c.get("customer_name","—"),
                "city":            c.get("city","—"),
                "type":            "B2B" if c.get("type")==0 else "B2C",
                "client_group":    c.get("client_group",2),
                "probability":     round(c.get("probability",0)*100,1),
                "segment_crm":     c.get("segment_crm","—"),
                "will_buy":        c.get("will_buy",0),
                "top_product_1":   c.get("top_product","—"),
                "price_product_1": c.get("top_price",0),
                "top_product_2":   c.get("top_product_2","—"),
                "price_product_2": c.get("top_price_2",0),
                "top_product_3":   c.get("top_product_3","—"),
                "price_product_3": c.get("top_price_3",0),
                "discount":        c.get("discount",0),
                "rank":            c.get("rank",0),
                "updated_at":      ts.strftime("%Y-%m-%d %H:%M") if ts else None,
            })
        return jsonify(rows)
    except Exception as e:
        return jsonify({"error": str(e)}), 500


@app.route("/api/powerbi/metrics", methods=["GET"])
@require_jwt
def powerbi_metrics():
    try:
        will_buy, all_customers, _ = get_cached_data()
        champions = [c for c in all_customers if c.get("segment_crm") == "5 — Champion"]
        return jsonify([{
            "total_clients":       len(all_customers),
            "will_buy_count":      len(will_buy),
            "champions_count":     len(champions),
            "avg_discount":        round(sum(c.get("discount",0) for c in will_buy)/max(len(will_buy),1),1),
            "auc_module1": 0.8634, "prauc_module1": 0.8681,
            "f1_module1":  0.8354, "accuracy_module1": 83.8,
            "threshold_module1": BEST_THRESHOLD,
            "model_m1": "LogisticRegression (Platt)",
            "precision10_module2": 0.2556, "ndcg10_module2": 0.2752,
            "hitrate10_module2":   0.7101, "mrr_module2": 0.4457,
            "model_m2": f"SVD k={SVD_K}",
            "mae_module3": 1.038, "r2_module3": 0.843,
            "model_m3": "XGBoost Regressor",
            "snapshot_date": str(SNAPSHOT_DATE.date()),
        }])
    except Exception as e:
        return jsonify({"error": str(e)}), 500

# ==============================================================
#  REENTAINEMENT
# ==============================================================
METRICS_FILE   = os.path.join(BASE_DIR, "metrics_history.json")
MODELS_BACKUP  = os.path.join(BASE_DIR, "models_backup")
RETRAIN_STATUS = {"running": False, "last_retrain": None, "last_result": None}

def load_metrics_history():
    if os.path.exists(METRICS_FILE):
        with open(METRICS_FILE, "r") as f: return json.load(f)
    return {"history": [], "current": None}

def save_metrics_history(history):
    with open(METRICS_FILE, "w") as f:
        json.dump(history, f, indent=2, default=str)

def backup_models():
    os.makedirs(MODELS_BACKUP, exist_ok=True)
    for fname in ["willbuy.pkl","discount.pkl","scaler.pkl","te_encoders.pkl","reco_svd.pkl"]:
        src = os.path.join(MODELS_DIR, fname)
        if os.path.exists(src): shutil.copy2(src, os.path.join(MODELS_BACKUP, fname))

def restore_models():
    for fname in ["willbuy.pkl","discount.pkl","scaler.pkl","te_encoders.pkl","reco_svd.pkl"]:
        src = os.path.join(MODELS_BACKUP, fname)
        if os.path.exists(src): shutil.copy2(src, os.path.join(MODELS_DIR, fname))

# ==============================================================
#  FIX 6 — retrain_models(schema)
#  Accepte un schéma explicite — peut être appelé depuis l'API (JWT)
#  ou manuellement avec un schéma passé en paramètre.
# ==============================================================
def retrain_models(schema: str = None):
    global model_wb, model_disc, scaler, reco_data, MODELS_OK

    if RETRAIN_STATUS["running"]:
        log.warning("[RETRAIN] Deja en cours")
        return

    if not schema:
        log.error("[RETRAIN] Aucun schéma fourni — annulé.")
        return

    RETRAIN_STATUS["running"] = True
    log.info(f"[RETRAIN] Démarrage — tenant={schema}")
    t_start = time.time()

    try:
        customers, invoices, products, lines = load_tables(schema)
        customers, invoices, products, lines = clean_tables(customers, invoices, products, lines)

        invoices["status"] = pd.to_numeric(invoices["status"], errors="coerce").fillna(-1).astype(int)
        invoices["untaxed_amount"] = pd.to_numeric(invoices["untaxed_amount"], errors="coerce").fillna(0)

        cf, y, valid_train = build_features(customers, invoices, products, lines)
        X   = cf[[f for f in FEATURES if f in cf.columns]].fillna(0).astype(float)
        y_s = cf["id"].map(y).fillna(0).astype(int)

        if len(X) < 50 or y_s.sum() < 10:
            log.warning("[RETRAIN] Données insuffisantes")
            RETRAIN_STATUS["running"] = False
            return

        X_tr, X_te, y_tr, y_te = train_test_split(
            X, y_s, test_size=0.20, random_state=RANDOM_STATE, stratify=y_s)

        no_sc   = ["type"]
        to_sc   = [c for c in X_tr.columns if c not in no_sc]
        new_sc  = RobustScaler()
        X_tr_sc = X_tr.copy(); X_te_sc = X_te.copy()
        X_tr_sc[to_sc] = new_sc.fit_transform(X_tr[to_sc])
        X_te_sc[to_sc] = new_sc.transform(X_te[to_sc])

        lr_new  = LogisticRegression(C=0.05, penalty="elasticnet", l1_ratio=0.5,
                                     solver="saga", max_iter=1000,
                                     class_weight="balanced", random_state=RANDOM_STATE)
        lr_new.fit(X_tr_sc.values, y_tr.values)
        cal_new = CalibratedClassifierCV(lr_new, method="sigmoid", cv="prefit")
        cal_new.fit(X_tr_sc, y_tr)
        auc_new = roc_auc_score(y_te, cal_new.predict_proba(X_te_sc)[:,1])
        log.info(f"[RETRAIN] M1 AUC={auc_new:.4f}")

        hist    = load_metrics_history()
        auc_old = hist["current"]["auc_m1"] if hist["current"] else 0.0
        if auc_new < auc_old - 0.02:
            log.warning(f"[RETRAIN] Dégradation — rollback")
            RETRAIN_STATUS["running"] = False
            return

        # SVD
        i2c     = valid_train.set_index("id")["customers_id"].to_dict()
        lines_r = lines.copy()
        lines_r["customers_id"] = lines_r["invoice_id"].map(i2c)
        lines_r = lines_r.dropna(subset=["customers_id"]).astype({"customers_id":int})
        cp = lines_r.groupby(["customers_id","product_id"]).agg(
            n=("invoice_id","count"), q=("quantity","sum")).reset_index()
        cp["score"] = np.log1p(cp["n"]) + np.log1p(cp["q"])*0.3
        cids_ = sorted(cp["customers_id"].unique())
        pids_ = sorted(lines_r["product_id"].unique())
        ci_   = {c:i for i,c in enumerate(cids_)}
        pi_   = {p:i for i,p in enumerate(pids_)}
        R     = np.zeros((len(cids_), len(pids_)))
        for _, r in cp.iterrows():
            c,p = ci_.get(int(r.customers_id)), pi_.get(int(r.product_id))
            if c is not None and p is not None: R[c,p] = r.score
        k_svd = min(SVD_K, min(R.shape)-1)
        U,s,Vt = svds(csr_matrix(normalize(R, norm="l2", axis=1)), k=k_svd)
        new_reco = {
            "R_hat": np.dot(np.dot(U,np.diag(s)),Vt),
            "cids": cids_, "pids": pids_,
            "p_name": products.set_index("id")["product_name"].to_dict(),
            "p_cat":  products.set_index("id")["category"].to_dict() if "category" in products.columns else {},
        }

        # XGBoost
        cf_r = cf.copy()
        cf_r["will_buy_proba"] = cal_new.predict_proba(
            cf_r[[f for f in FEATURES if f in cf_r.columns]].fillna(0))[:,1]
        cf_r["client_group"] = cf_r["frequency"].apply(get_client_group)
        disc_l   = lines_r.merge(valid_train[["id","customers_id"]], left_on="invoice_id", right_on="id", how="inner")
        disc_agg = disc_l.groupby("customers_id").agg(avg_disc=("discount","mean"), med_disc=("discount","median")).reset_index()
        df3 = cf_r.merge(disc_agg, left_on="id", right_on="customers_id", how="left")
        df3[["avg_disc","med_disc"]] = df3[["avg_disc","med_disc"]].fillna(0)
        df3["avg_discount_recv"] = df3["avg_disc"]
        q75m = df3["log_monetary_ht"].quantile(0.75)
        np.random.seed(RANDOM_STATE)
        def opt_d(row):
            b = row.get("med_disc",0); grp = int(row.get("client_group",2))
            if row.get("log_monetary_ht",0) >= q75m: b=min(b+3,30)
            if row.get("recency_days",0) > 120: b=min(b+5,30)
            if grp==4: b=min(b+2,30)
            elif grp==3: b=min(b+1,30)
            elif grp==1: b=max(b-1,0)
            return round(float(np.clip(b+np.random.normal(0,1.0),5,30)),1)
        df3["optimal_discount"] = df3.apply(opt_d, axis=1)
        D3   = [f for f in D_FEATS if f in df3.columns]
        dt3  = df3[df3["avg_disc"]>0].copy()
        Xd_tr,Xd_te,yd_tr,yd_te = train_test_split(dt3[D3].fillna(0), dt3["optimal_discount"], test_size=0.2, random_state=RANDOM_STATE)
        xgb_new = xgb.XGBRegressor(n_estimators=300, learning_rate=0.05, max_depth=3,
                                    min_child_weight=5, subsample=0.8,
                                    reg_alpha=0.1, reg_lambda=1.0,
                                    random_state=RANDOM_STATE, n_jobs=-1)
        xgb_new.fit(Xd_tr, yd_tr)
        mae_new = mean_absolute_error(yd_te, np.clip(xgb_new.predict(Xd_te),0,30))
        r2_new  = r2_score(yd_te, np.clip(xgb_new.predict(Xd_te),0,30))

        backup_models()
        import joblib as jl
        jl.dump(cal_new,  os.path.join(MODELS_DIR,"willbuy.pkl"))
        jl.dump(xgb_new,  os.path.join(MODELS_DIR,"discount.pkl"))
        jl.dump(new_sc,   os.path.join(MODELS_DIR,"scaler.pkl"))
        jl.dump(new_reco, os.path.join(MODELS_DIR,"reco_svd.pkl"))
        model_wb=cal_new; model_disc=xgb_new; scaler=new_sc; reco_data=new_reco; MODELS_OK=True

        new_metrics = {
            "date": datetime.now().isoformat(), "schema": schema,
            "auc_m1": round(auc_new,4), "mae_m3": round(mae_new,3),
            "r2_m3": round(r2_new,3), "status": "success",
            "duration_s": round(time.time()-t_start,1),
        }
        hist["history"].append(new_metrics); hist["current"] = new_metrics
        save_metrics_history(hist)
        get_cached_data(force_refresh=True, schema=schema)
        RETRAIN_STATUS["last_retrain"] = datetime.now().isoformat()
        RETRAIN_STATUS["last_result"]  = new_metrics
        log.info(f"[RETRAIN] OK {time.time()-t_start:.0f}s | AUC={auc_new:.4f}")

    except Exception as e:
        log.error(f"[RETRAIN] Erreur : {e}")
        restore_models()
        RETRAIN_STATUS["last_result"] = {"status":"error","message":str(e),"date":datetime.now().isoformat()}
    finally:
        RETRAIN_STATUS["running"] = False


@app.route("/api/retrain", methods=["POST"])
@require_jwt
def trigger_retrain():
    if RETRAIN_STATUS["running"]:
        return jsonify({"status":"running","message":"Déjà en cours..."}),409
    schema = g.tenant_schema
    t = threading.Thread(target=retrain_models, args=(schema,), daemon=True)
    t.start()
    return jsonify({"status":"started","tenant":schema})


@app.route("/api/model-health", methods=["GET"])
@require_jwt
def model_health():
    hist = load_metrics_history()
    return jsonify({
        "current_metrics": hist.get("current"),
        "history":         hist.get("history",[])[-12:],
        "retrain_status":  RETRAIN_STATUS,
    })


# Fichiers statiques — PAS de JWT (Angular app)
@app.route("/")
def index():
    return send_from_directory("static", "index.html")

@app.route("/<path:path>")
def static_files(path):
    return send_from_directory("static", path)

# ==============================================================
#  FIX 6b — SCHEDULER désactivé (pas de tenant disponible)
#  Le retrain se déclenche uniquement via POST /api/retrain (JWT)
# ==============================================================
def start_scheduler():
    log.warning("[SCHEDULER] Désactivé — retrain via POST /api/retrain avec JWT")
    # Pour activer: passer un schéma explicite à retrain_models(schema="company_X")

# ==============================================================
#  LANCEMENT
# ==============================================================
if __name__ == "__main__":
    print(f"""
====================================================
=  ERP AI - CRM Intelligence API (MULTI-TENANT)  =
=  Tenant résolu depuis JWT tenantSchema          =
=  Cache isolé par tenant                         =
=  DB : {DB_CONFIG['host']}:{DB_CONFIG['port']}/{DB_CONFIG['database']:<15}=
=  API : http://localhost:5000                    =
====================================================
""")
    if not os.path.exists(METRICS_FILE):
        save_metrics_history({"history":[],"current":{"date":"2026-05-23T00:00:00","auc_m1":0.8634,"status":"initial"}})

    start_scheduler()

    # FIX: pas d'appel get_cached_data() au démarrage — pas de tenant disponible
    # Le cache se charge au premier appel API avec JWT valide
    log.info("[INIT] Serveur prêt — cache chargé au premier appel API authentifié")

    try:
        app.run(debug=False, host="0.0.0.0", port=5000)
    except KeyboardInterrupt:
        log.info("Arrêt du serveur.")