from __future__ import annotations

import hashlib
import json
import logging
import os
import pickle
import tempfile
import unicodedata
from collections import Counter
from datetime import datetime, timezone
from pathlib import Path
from typing import Any

import sklearn
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import accuracy_score, classification_report
from sklearn.model_selection import (
    StratifiedKFold,
    cross_val_predict,
    cross_val_score,
)
from sklearn.pipeline import Pipeline

logger = logging.getLogger(__name__)


class IntentTrainer:
    """Melatih, mengevaluasi, menyimpan, dan memuat model klasifikasi intent."""

    PROJECT_ROOT = Path(__file__).resolve().parents[2]

    def __init__(
        self,
        data_path: str | os.PathLike[str] = "app/data/intents.json",
        model_path: str | os.PathLike[str] = "app/model",
        *,
        confidence_threshold: float = 0.5,
        random_state: int = 42,
        cv_folds: int = 5,
        max_features: int = 5000,
    ) -> None:
        self.data_path = self._resolve_path(data_path)
        self.model_path = self._resolve_path(model_path)

        if not 0.0 <= confidence_threshold <= 1.0:
            raise ValueError("confidence_threshold harus berada di antara 0.0 dan 1.0")
        if cv_folds < 2:
            raise ValueError("cv_folds minimal 2")
        if max_features < 100:
            raise ValueError("max_features minimal 100")

        self.confidence_threshold = confidence_threshold
        self.random_state = random_state
        self.cv_folds = cv_folds
        self.max_features = max_features

        self.vectorizer: TfidfVectorizer | None = None
        self.model: LogisticRegression | None = None
        self.intents_data: dict[str, dict[str, Any]] | None = None
        self.label_to_intent: dict[int, str] = {}
        self.intent_to_label: dict[str, int] = {}
        self.training_metrics: dict[str, Any] = {}

    @classmethod
    def _resolve_path(cls, path: str | os.PathLike[str]) -> Path:
        """Resolve path relatif terhadap root project, bukan current working directory."""
        resolved = Path(path).expanduser()
        if resolved.is_absolute():
            return resolved
        return (cls.PROJECT_ROOT / resolved).resolve()

    @staticmethod
    def normalize_text(text: str) -> str:
        """Normalisasi ringan tanpa menghapus karakter penting dari pertanyaan."""
        if not isinstance(text, str):
            raise TypeError("Text harus berupa string")

        normalized = unicodedata.normalize("NFKC", text)
        normalized = " ".join(normalized.strip().lower().split())
        return normalized

    def load_intents(self) -> dict[str, dict[str, Any]]:
        """Muat dan validasi struktur dasar file intents.json."""
        if not self.data_path.exists():
            raise FileNotFoundError(f"File intents tidak ditemukan: {self.data_path}")

        try:
            with self.data_path.open("r", encoding="utf-8") as file:
                data = json.load(file)
        except json.JSONDecodeError as exc:
            raise ValueError(
                f"Format JSON tidak valid pada {self.data_path}, "
                f"baris {exc.lineno}, kolom {exc.colno}: {exc.msg}"
            ) from exc

        if not isinstance(data, dict) or not data:
            raise ValueError("Root intents.json harus berupa object JSON yang tidak kosong")

        for intent_name, intent_data in data.items():
            if not isinstance(intent_name, str) or not intent_name.strip():
                raise ValueError("Setiap nama intent harus berupa string yang tidak kosong")
            if not isinstance(intent_data, dict):
                raise ValueError(f"Data intent '{intent_name}' harus berupa object JSON")

            patterns = intent_data.get("patterns")
            if not isinstance(patterns, list) or not patterns:
                raise ValueError(
                    f"Intent '{intent_name}' harus memiliki list 'patterns' yang tidak kosong"
                )

            responses = intent_data.get("responses")
            if not isinstance(responses, list) or not responses:
                logger.warning(
                    "Intent '%s' tidak memiliki responses yang valid. "
                    "Training tetap berjalan, tetapi inference dapat memakai fallback.",
                    intent_name,
                )

        self.intents_data = data
        logger.info("✓ Intents loaded: %d intent(s) dari %s", len(data), self.data_path)
        return data

    def prepare_training_data(self) -> tuple[list[str], list[int], list[str]]:
        """Siapkan dataset, mapping label, dan deteksi pattern duplikat/konflik."""
        if self.intents_data is None:
            self.load_intents()

        assert self.intents_data is not None

        self.label_to_intent = {}
        self.intent_to_label = {}

        patterns: list[str] = []
        labels: list[int] = []
        intent_names = list(self.intents_data.keys())
        seen_patterns: dict[str, str] = {}

        if len(intent_names) < 2:
            raise ValueError("Training memerlukan minimal 2 intent berbeda")

        for label, intent_name in enumerate(intent_names):
            self.intent_to_label[intent_name] = label
            self.label_to_intent[label] = intent_name

            raw_patterns = self.intents_data[intent_name].get("patterns", [])
            valid_count = 0

            for raw_pattern in raw_patterns:
                if not isinstance(raw_pattern, str):
                    logger.warning(
                        "Pattern non-string di intent '%s' dilewati: %r",
                        intent_name,
                        raw_pattern,
                    )
                    continue

                pattern = self.normalize_text(raw_pattern)
                if not pattern:
                    logger.warning("Pattern kosong di intent '%s' dilewati", intent_name)
                    continue

                previous_intent = seen_patterns.get(pattern)
                if previous_intent is not None:
                    if previous_intent != intent_name:
                        raise ValueError(
                            "Pattern konflik ditemukan: "
                            f"'{pattern}' digunakan oleh intent '{previous_intent}' "
                            f"dan '{intent_name}'"
                        )

                    logger.warning(
                        "Pattern duplikat dalam intent '%s' dilewati: '%s'",
                        intent_name,
                        pattern,
                    )
                    continue

                seen_patterns[pattern] = intent_name
                patterns.append(pattern)
                labels.append(label)
                valid_count += 1

            if valid_count == 0:
                raise ValueError(f"Intent '{intent_name}' tidak memiliki pattern valid")
            if valid_count < 5:
                logger.warning(
                    "Intent '%s' hanya memiliki %d pattern. "
                    "Disarankan minimal 20 pattern untuk kualitas produksi.",
                    intent_name,
                    valid_count,
                )

        class_counts = Counter(labels)
        logger.info(
            "✓ Training data prepared: %d pattern, %d intent, min/max pattern per intent: %d/%d",
            len(patterns),
            len(intent_names),
            min(class_counts.values()),
            max(class_counts.values()),
        )

        return patterns, labels, intent_names

    def _build_vectorizer(self) -> TfidfVectorizer:
        return TfidfVectorizer(
            lowercase=False,  # Sudah dinormalisasi oleh normalize_text().
            strip_accents="unicode",
            analyzer="word",
            ngram_range=(1, 2),
            min_df=1,
            max_df=1.0,
            max_features=self.max_features,
            sublinear_tf=True,
            norm="l2",
            token_pattern=r"(?u)\b\w+\b",
        )

    def _build_model(self) -> LogisticRegression:
        # Parameter multi_class sengaja tidak digunakan karena deprecated/dihapus
        # pada versi scikit-learn terbaru. Solver lbfgs menangani multiclass.
        return LogisticRegression(
            solver="lbfgs",
            max_iter=2000,
            class_weight="balanced",
            random_state=self.random_state,
            C=1.0,
        )

    def _evaluate_with_cross_validation(
        self,
        patterns: list[str],
        labels: list[int],
        intent_names: list[str],
    ) -> dict[str, Any]:
        """Evaluasi tanpa data leakage menggunakan Pipeline dan Stratified K-Fold."""
        class_counts = Counter(labels)
        minimum_class_count = min(class_counts.values())
        n_splits = min(self.cv_folds, minimum_class_count)

        if n_splits < 2:
            logger.warning(
                "Cross-validation dilewati karena ada intent dengan kurang dari 2 pattern."
            )
            return {
                "cv_enabled": False,
                "cv_folds": 0,
                "cv_accuracy_mean": None,
                "cv_accuracy_std": None,
                "cv_fold_scores": [],
                "cv_classification_report": None,
            }

        pipeline = Pipeline(
            steps=[
                ("tfidf", self._build_vectorizer()),
                ("classifier", self._build_model()),
            ]
        )
        splitter = StratifiedKFold(
            n_splits=n_splits,
            shuffle=True,
            random_state=self.random_state,
        )

        logger.info("📊 Running %d-fold stratified cross-validation...", n_splits)

        fold_scores = cross_val_score(
            pipeline,
            patterns,
            labels,
            scoring="accuracy",
            cv=splitter,
            n_jobs=None,
            error_score="raise",
        )
        cv_predictions = cross_val_predict(
            pipeline,
            patterns,
            labels,
            cv=splitter,
            n_jobs=None,
            method="predict",
        )

        report = classification_report(
            labels,
            cv_predictions,
            labels=list(range(len(intent_names))),
            target_names=intent_names,
            digits=3,
            zero_division=0,
        )

        accuracy_mean = float(fold_scores.mean())
        accuracy_std = float(fold_scores.std())

        logger.info(
            "✓ Cross-validation accuracy: %.2f%% ± %.2f%%",
            accuracy_mean * 100,
            accuracy_std * 100,
        )
        logger.info("\n📈 Cross-validation Classification Report:\n%s", report)

        return {
            "cv_enabled": True,
            "cv_folds": n_splits,
            "cv_accuracy_mean": accuracy_mean,
            "cv_accuracy_std": accuracy_std,
            "cv_fold_scores": [float(score) for score in fold_scores],
            "cv_classification_report": report,
        }

    def train(self) -> dict[str, Any]:
        """Evaluasi model, lalu latih model final menggunakan seluruh dataset."""
        logger.info("🚀 Starting intent model training...")

        patterns, labels, intent_names = self.prepare_training_data()
        evaluation = self._evaluate_with_cross_validation(
            patterns,
            labels,
            intent_names,
        )

        logger.info("📊 Fitting final TF-IDF vectorizer on full dataset...")
        self.vectorizer = self._build_vectorizer()
        features = self.vectorizer.fit_transform(patterns)

        logger.info("🤖 Fitting final Logistic Regression model...")
        self.model = self._build_model()
        self.model.fit(features, labels)

        # Training accuracy hanya diagnostik, bukan ukuran generalisasi.
        training_predictions = self.model.predict(features)
        training_accuracy = float(accuracy_score(labels, training_predictions))

        class_counts = Counter(labels)
        dataset_hash = hashlib.sha256(self.data_path.read_bytes()).hexdigest()

        self.training_metrics = {
            "trained_at": datetime.now(timezone.utc).isoformat(),
            "sklearn_version": sklearn.__version__,
            "dataset_path": str(self.data_path),
            "dataset_sha256": dataset_hash,
            "total_patterns": len(patterns),
            "total_intents": len(intent_names),
            "patterns_per_intent": {
                self.label_to_intent[label]: count
                for label, count in sorted(class_counts.items())
            },
            "vocabulary_size": len(self.vectorizer.vocabulary_),
            "training_accuracy": training_accuracy,
            "confidence_threshold": self.confidence_threshold,
            "model_type": "logistic_regression_tfidf",
            **evaluation,
        }

        logger.info(
            "✓ Final model trained. Training accuracy: %.2f%% "
            "(diagnostic only, gunakan CV untuk evaluasi)",
            training_accuracy * 100,
        )

        return self.training_metrics

    @staticmethod
    def _atomic_pickle_dump(value: Any, destination: Path) -> None:
        """Simpan pickle secara atomik agar file lama tidak rusak jika proses gagal."""
        destination.parent.mkdir(parents=True, exist_ok=True)
        temporary_path: str | None = None

        try:
            with tempfile.NamedTemporaryFile(
                mode="wb",
                dir=destination.parent,
                prefix=f".{destination.name}.",
                suffix=".tmp",
                delete=False,
            ) as temporary_file:
                temporary_path = temporary_file.name
                pickle.dump(value, temporary_file, protocol=pickle.HIGHEST_PROTOCOL)
                temporary_file.flush()
                os.fsync(temporary_file.fileno())

            os.replace(temporary_path, destination)
        except Exception:
            if temporary_path and os.path.exists(temporary_path):
                os.unlink(temporary_path)
            raise

    @staticmethod
    def _atomic_json_dump(value: Any, destination: Path) -> None:
        destination.parent.mkdir(parents=True, exist_ok=True)
        temporary_path: str | None = None

        try:
            with tempfile.NamedTemporaryFile(
                mode="w",
                encoding="utf-8",
                dir=destination.parent,
                prefix=f".{destination.name}.",
                suffix=".tmp",
                delete=False,
            ) as temporary_file:
                temporary_path = temporary_file.name
                json.dump(value, temporary_file, ensure_ascii=False, indent=2)
                temporary_file.write("\n")
                temporary_file.flush()
                os.fsync(temporary_file.fileno())

            os.replace(temporary_path, destination)
        except Exception:
            if temporary_path and os.path.exists(temporary_path):
                os.unlink(temporary_path)
            raise

    def save_model(self) -> dict[str, str]:
        """Simpan vectorizer, model, label mappings, dan metadata training."""
        if self.model is None or self.vectorizer is None:
            raise RuntimeError("Model belum dilatih. Jalankan train() sebelum save_model().")
        if not self.label_to_intent or not self.intent_to_label:
            raise RuntimeError("Label mappings belum tersedia")

        self.model_path.mkdir(parents=True, exist_ok=True)

        vectorizer_path = self.model_path / "vectorizer.pkl"
        model_file_path = self.model_path / "intent_model.pkl"
        mappings_path = self.model_path / "label_mappings.pkl"
        metadata_path = self.model_path / "model_metadata.json"

        self._atomic_pickle_dump(self.vectorizer, vectorizer_path)
        self._atomic_pickle_dump(self.model, model_file_path)
        self._atomic_pickle_dump(
            {
                "intent_to_label": self.intent_to_label,
                "label_to_intent": self.label_to_intent,
            },
            mappings_path,
        )
        self._atomic_json_dump(self.training_metrics, metadata_path)

        saved_files = {
            "vectorizer": str(vectorizer_path),
            "model": str(model_file_path),
            "mappings": str(mappings_path),
            "metadata": str(metadata_path),
        }

        for file_type, file_path in saved_files.items():
            logger.info("✓ %s saved: %s", file_type.capitalize(), file_path)

        return saved_files

    def load_model(self) -> bool:
        """Muat model hasil training. Hanya muat pickle dari sumber terpercaya."""
        model_file_path = self.model_path / "intent_model.pkl"
        vectorizer_path = self.model_path / "vectorizer.pkl"
        mappings_path = self.model_path / "label_mappings.pkl"

        required_files = [model_file_path, vectorizer_path, mappings_path]
        missing_files = [str(path) for path in required_files if not path.exists()]
        if missing_files:
            raise FileNotFoundError(
                "Model files belum lengkap. Jalankan training terlebih dahulu. "
                f"File yang belum ada: {', '.join(missing_files)}"
            )

        try:
            with model_file_path.open("rb") as file:
                model = pickle.load(file)
            with vectorizer_path.open("rb") as file:
                vectorizer = pickle.load(file)
            with mappings_path.open("rb") as file:
                mappings = pickle.load(file)
        except (pickle.UnpicklingError, EOFError, AttributeError, ValueError) as exc:
            raise RuntimeError(f"Gagal memuat model pickle: {exc}") from exc

        if not hasattr(model, "predict") or not hasattr(model, "predict_proba"):
            raise TypeError("intent_model.pkl bukan classifier yang valid")
        if not hasattr(vectorizer, "transform"):
            raise TypeError("vectorizer.pkl bukan vectorizer yang valid")
        if not isinstance(mappings, dict):
            raise TypeError("label_mappings.pkl harus berisi dictionary")

        label_to_intent = mappings.get("label_to_intent")
        intent_to_label = mappings.get("intent_to_label")
        if not isinstance(label_to_intent, dict) or not isinstance(intent_to_label, dict):
            raise ValueError("Label mappings tidak lengkap atau tidak valid")

        self.model = model
        self.vectorizer = vectorizer
        self.label_to_intent = label_to_intent
        self.intent_to_label = intent_to_label

        logger.info("✓ Model loaded successfully dari %s", self.model_path)
        return True

    def predict(
        self,
        text: str,
        threshold: float | None = None,
        *,
        top_k: int = 3,
    ) -> dict[str, Any]:
        """Prediksi intent dan tampilkan kandidat intent dengan confidence tertinggi."""
        if self.model is None or self.vectorizer is None:
            raise RuntimeError("Model belum tersedia. Jalankan train() atau load_model().")

        active_threshold = (
            self.confidence_threshold if threshold is None else float(threshold)
        )
        if not 0.0 <= active_threshold <= 1.0:
            raise ValueError("threshold harus berada di antara 0.0 dan 1.0")
        if top_k < 1:
            raise ValueError("top_k minimal 1")

        normalized_text = self.normalize_text(text)
        if not normalized_text:
            raise ValueError("Text tidak boleh kosong")

        features = self.vectorizer.transform([normalized_text])
        probabilities = self.model.predict_proba(features)[0]
        classes = list(self.model.classes_)

        ranked_indices = probabilities.argsort()[::-1][: min(top_k, len(classes))]
        top_predictions: list[dict[str, Any]] = []

        for probability_index in ranked_indices:
            label = int(classes[int(probability_index)])
            intent = self.label_to_intent.get(label, "unknown")
            top_predictions.append(
                {
                    "intent": intent,
                    "confidence": float(probabilities[int(probability_index)]),
                }
            )

        best_prediction = top_predictions[0]
        confidence = float(best_prediction["confidence"])

        return {
            "intent": best_prediction["intent"],
            "confidence": confidence,
            "is_confident": confidence >= active_threshold,
            "threshold": active_threshold,
            "normalized_text": normalized_text,
            "top_predictions": top_predictions,
        }


def main() -> None:
    """Entry point ketika file trainer.py dijalankan langsung."""
    logging.basicConfig(
        level=logging.INFO,
        format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
    )

    trainer = IntentTrainer()
    metrics = trainer.train()
    saved_files = trainer.save_model()

    logger.info(
        "Training selesai. CV accuracy: %s, model: %s",
        (
            f"{metrics['cv_accuracy_mean']:.2%}"
            if metrics.get("cv_accuracy_mean") is not None
            else "tidak tersedia"
        ),
        saved_files["model"],
    )

    result = trainer.predict("apa itu sertifikasi k3", top_k=3)
    logger.info("Test prediction: %s", json.dumps(result, ensure_ascii=False))


if __name__ == "__main__":
    main()