#!/usr/bin/env python3
import argparse
import csv
import os
from datetime import datetime, timedelta
import re
import sys
import unicodedata
from pathlib import Path

import mysql.connector

#BASE_DIR = Path(r"C:\Users\genge\Desktop\Projet\rh-appweb-stage\ExportAccess") # desktop genge 
BASE_DIR = Path(r"C:\Users\leo.gengembre\Desktop\Projet\rh-appweb-stage\ExportAccess") #pc envie2e
#BASE_DIR = Path(r"C:\Users\genge\OneDrive\Desktop\Projet\StageEnvie2e\ExportAccess") 


TABLE_MAP = {
    "joursfériés": "jours_feries",
    "joursferies": "jours_feries",
    "tabsences": "absences",
    "ttypeabsence": "types_absence",
    "tdisciplinaire": "disciplinaire",
    "tsanctions": "sanctions",
    "prov_mutuelle": "mutuelles",
    "tcpam": "cpam",
    "tmotifsdérogationmutuelle": "motifs_derogation_mutuelle",
    "tmotifsderogationmutuelle": "motifs_derogation_mutuelle",
    "tmédecinedutravail": "medecine_travail",
    "tmedecinedutravail": "medecine_travail",
    "tmedecine_du_travail": "medecine_travail",
    "tobjetvisitemédicale": "objet_visite_medicale",
    "tobjetvisitemedicale": "objet_visite_medicale",
    "tobjet_visite_medicale": "objet_visite_medicale",
    "tvisitesmédicales": "visites_medicales",
    "tvisitesmedicales": "visites_medicales",
    "prov_alertes": "alertes",
    "prov_alerteurgence": "alertes_urgence",
    "tbl_alertes1mail": "mails_alertes",
    "tbl_alertes1mail_liste": "alertes_mail_liste",
    "tbl_alertesliste": "listes_alertes",
    "documentauto": "documents_auto",
    "impressions": "impressions",
    "prov_publi": "prov_publi",
    "ref_adressesmailpourautomatiques": "ref_adresses_mail_pour_automatiques",
    "ref_signataires": "ref_signataires",
    "tautomatique": "automatisations",
    "tparamètres": "parametres",
    "tparametres": "parametres",
    "tadministrateurs": "administrateurs",
    "tsecteurs": "secteurs",
    "tsecteurs_salarié": "secteurs_employes",
    "tsecteurs_salarie": "secteurs_employes",
    "tsecteurstravail": "secteurs_travail",
    "tétablissement": "etablissement",
    "tetablissement": "etablissement",
    "code_insee_emploi": "code_insee_emploi",
    "tdemandenouvelles": "demande_nouvelles",
    "tmatériel": "materiel",
    "tmateriel": "materiel",
    "tmotifssortie": "motifs_sortie",
    "tplans": "plans",
    "prov_fichierspaie": "fichiers_paie",
    "prov_equipenuit": "prov_equipenuit",
    "t_equipenuit": "equipenuit",
    "tarh": "arh",
    "tavenants": "avenants_contrat",
    "tcan_convinser": "can_convinser",
    "tcan_metiers": "can_metiers",
    "tcan_ruptbudg": "can_ruptbudg",
    "tcan_typcontr": "can_typcontr",
    "tclassif_salairesbase": "classif_salairesbase",
    "tclassifications": "classifications",
    "tcontrats": "contrats",
    "tdiplomes": "diplomes",
    "tposte": "postes",
    "tsalariés": "employes",
    "tsalaries": "employes",
    "ttypeavenants": "types_avenants",
    "ttypecontratsalarié": "types_contrat_salarie",
    "ttypecontratsalarie": "types_contrat_salarie",
    "ttypecs": "types_cs",
    "ttypesdiplomes": "types_diplome",
}


def strip_accents(value: str) -> str:
    normalized = unicodedata.normalize("NFD", value)
    return "".join(char for char in normalized if unicodedata.category(char) != "Mn")


def normalize_name(value: str) -> str:
    value = strip_accents(value)
    value = value.lower()
    value = re.sub(r"[^a-z0-9]+", "_", value)
    return value.strip("_")


def mysql_identifier(name: str) -> str:
    return "`" + name.replace("`", "``") + "`"


def resolve_table_name(csv_path: Path) -> str | None:
    stem = csv_path.stem
    normalized = normalize_name(stem)
    return TABLE_MAP.get(stem, TABLE_MAP.get(normalized, normalized))


def clean_value(value, column_name: str | None = None):
    if value is None:
        return None
    if isinstance(value, str):
        value = value.strip()
        if value == "":
            return None

        # Fix import dates Access/Excel -> MySQL DATE.
        # Accepte par exemple : 2025-10-20, 20/10/2025, 20/10/2025 00:00:00, 45950.
        normalized_column = normalize_name(column_name or "")
        looks_like_date_column = (
            "date" in normalized_column
            or normalized_column.endswith("du")
            or normalized_column.endswith("au")
            or normalized_column in {"dae", "dsp", "dsr", "dsrr", "dan"}
        )

        if looks_like_date_column:
            parsed_date = parse_date_value(value)
            if parsed_date is not None:
                return parsed_date

        return value
    return value


def parse_date_value(value: str) -> str | None:
    value = value.strip()
    if not value:
        return None

    # Excel serial date, ex: 45950
    if re.fullmatch(r"\d+(?:\.0+)?", value):
        try:
            serial = int(float(value))
            if 20000 <= serial <= 60000:
                return (datetime(1899, 12, 30) + timedelta(days=serial)).strftime("%Y-%m-%d")
        except ValueError:
            pass

    formats = [
        "%Y-%m-%d",
        "%Y/%m/%d",
        "%d/%m/%Y",
        "%d-%m-%Y",
        "%d/%m/%Y %H:%M:%S",
        "%Y-%m-%d %H:%M:%S",
        "%d/%m/%y",
    ]
    for fmt in formats:
        try:
            return datetime.strptime(value, fmt).strftime("%Y-%m-%d")
        except ValueError:
            continue

    return None


def connect_mysql():
    return mysql.connector.connect(
        host=os.getenv("MYSQL_HOST", "127.0.0.1"),
        port=int(os.getenv("MYSQL_PORT", "3306")),
        user=os.getenv("MYSQL_USER", "app_user"),
        password=os.getenv("MYSQL_PASSWORD", "apppass"),
        database=os.getenv("MYSQL_DATABASE", "app_db"),
        autocommit=False,
    )


def iter_csv_files(base_dir: Path, only: str | None = None):
    wanted = normalize_name(only) if only else None
    for folder in sorted(path for path in base_dir.iterdir() if path.is_dir()):
        for csv_path in sorted(folder.glob("*.csv")):
            if wanted and normalize_name(csv_path.stem) != wanted:
                continue
            yield csv_path


def get_table_columns(cursor, table_name: str) -> dict[str, str]:
    cursor.execute(
        "SELECT COLUMN_NAME FROM INFORMATION_SCHEMA.COLUMNS WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = %s",
        (table_name,),
    )
    return {normalize_name(row[0]): row[0] for row in cursor.fetchall()}


def main():
    parser = argparse.ArgumentParser(description="Import all normalized CSV files into MySQL.")
    parser.add_argument("--base-dir", default=str(BASE_DIR), help="Root ExportAccess directory")
    parser.add_argument("--truncate", action="store_true", help="Truncate each table before importing it")
    parser.add_argument("--dry-run", action="store_true", help="Print actions without writing to MySQL")
    parser.add_argument("--only", help="Importer uniquement un CSV précis, ex: tavenants")
    args = parser.parse_args()

    base_dir = Path(args.base_dir)
    if not base_dir.exists():
        print(f"Base directory not found: {base_dir}", file=sys.stderr)
        return 1

    csv_files = list(iter_csv_files(base_dir, args.only))
    if not csv_files:
        print(f"No CSV files found under {base_dir} for filter {args.only!r}")
        return 1

    if args.dry_run:
        print("Dry run mode")

    connection = connect_mysql()
    cursor = connection.cursor()

    # Désactiver les contraintes FK pendant l'import pour éviter les conflits d'auto-référence
    cursor.execute("SET FOREIGN_KEY_CHECKS=0")
    connection.commit()

    try:
        for csv_path in csv_files:
            table_name = resolve_table_name(csv_path)
            if not table_name:
                print(f"SKIP {csv_path.relative_to(base_dir)} -> table not resolved")
                continue

            with csv_path.open("r", encoding="utf-8-sig", newline="") as handle:
                reader = csv.reader(handle, delimiter=",")
                rows = list(reader)

            if not rows:
                print(f"SKIP {csv_path.relative_to(base_dir)} -> empty file")
                continue

            headers = rows[0]
            data_rows = rows[1:]
            if not headers:
                print(f"SKIP {csv_path.relative_to(base_dir)} -> missing headers")
                continue

            table_columns = get_table_columns(cursor, table_name)
            column_pairs = []
            for index, header in enumerate(headers):
                if not header:
                    continue
                normalized_header = normalize_name(header)
                if normalized_header == "id":
                    continue
                actual_column = table_columns.get(normalized_header)
                if actual_column:
                    column_pairs.append((index, actual_column))

            if not column_pairs:
                print(f"SKIP {csv_path.relative_to(base_dir)} -> no matching table columns")
                continue

            column_indexes = [index for index, _ in column_pairs]
            columns = [column for _, column in column_pairs]
            insert_sql = (
                f"INSERT INTO {mysql_identifier(table_name)} "
                f"({', '.join(mysql_identifier(column) for column in columns)}) "
                f"VALUES ({', '.join(['%s'] * len(columns))})"
            )

            payload = []
            for row in data_rows:
                values = [clean_value(row[index], columns[pos]) if index < len(row) else None for pos, index in enumerate(column_indexes)]
                payload.append(values)

            print(f"IMPORT {csv_path.relative_to(base_dir)} -> {table_name} ({len(payload)} rows)")

            if args.dry_run:
                continue

            if args.truncate:
                cursor.execute(f"TRUNCATE TABLE {mysql_identifier(table_name)}")

            if payload:
                cursor.executemany(insert_sql, payload)
                connection.commit()
            else:
                print(f"  No data rows for {table_name}")

        return 0
    except mysql.connector.Error as exc:
        connection.rollback()
        print(f"MySQL error: {exc}", file=sys.stderr)
        return 1
    finally:
        # Réactiver les contraintes FK après l'import
        try:
            cursor.execute("SET FOREIGN_KEY_CHECKS=1")
            connection.commit()
        except:
            pass
        cursor.close()
        connection.close()


if __name__ == "__main__":
    raise SystemExit(main())