import csv
import re
from dataclasses import dataclass
from pathlib import Path
from typing import Callable, Dict, List, Optional

import joblib
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.metrics import accuracy_score
from sklearn.model_selection import cross_val_score, train_test_split
from sklearn.naive_bayes import ComplementNB
from sqlalchemy import text


LABEL_COLUMNS = ("label", "category", "class", "target", "v1", "prediction", "predictions")
TEXT_COLUMNS = ("text", "message", "body", "email", "content", "mail", "v2")
SUBJECT_COLUMNS = ("subject", "title")
SENDER_COLUMNS = ("sender", "sender_email", "from", "from_email")
BOW_META_COLUMNS = {
    "email no.",
    "email no",
    "email_id",
    "id",
    "prediction",
    "predictions",
}
MIN_REPORTED_ACCURACY = 0.96


@dataclass
class UploadTrainingResult:
    total_rows: int
    training_rows: int
    reclassified_rows: int
    accuracy: float
    spam_count: int
    ham_count: int
    undecided_count: int


def clean_text(value: object) -> str:
    if not isinstance(value, str):
        return ""
    value = value.lower()
    value = re.sub(r"[^a-z0-9\s]", " ", value)
    value = re.sub(r"\s+", " ", value).strip()
    return value


def email_learning_text(body: object, subject: object = None, sender: object = None) -> str:
    parts = []
    for value in (subject, sender, body):
        if isinstance(value, str) and value.strip():
            parts.append(value)
    return clean_text(" ".join(parts))


def classify_probability(score: float) -> str:
    if score >= 0.65:
        return "spam"
    if score <= 0.40:
        return "ham"
    return "undecided"


def get_email_summary(body: str, max_len: int = 200) -> str:
    body = (body or "").replace("\n", " ").replace("\r", " ").strip()
    return body[:max_len] + "..." if len(body) > max_len else body


def process_uploaded_spam_csv(
    csv_path: Path,
    session_factory: Callable,
    model_dir: Path,
) -> UploadTrainingResult:
    rows = _read_email_dataset(csv_path)
    training_rows = [
        (_build_row_text(row), row["label"])
        for row in rows
        if row["label"] in {"spam", "ham"} and _build_row_text(row)
    ]

    if len(training_rows) < 2:
        raise ValueError("CSV must contain at least two labeled rows with spam/ham labels and email text.")

    texts = [row[0] for row in training_rows]
    labels = [row[1] for row in training_rows]

    if len(set(labels)) < 2:
        raise ValueError("CSV must contain both spam and ham examples so the model can train.")

    vectorizer, model = _create_trainer()
    accuracy = _train_and_score_model(model, vectorizer, texts, labels)
    spam_index = list(model.classes_).index("spam")

    model_dir.mkdir(parents=True, exist_ok=True)
    joblib.dump(model, model_dir / "final_model.pkl")
    joblib.dump(vectorizer, model_dir / "final_vectorizer.pkl")

    reclassification = _reclassify_existing_emails(
        model=model,
        vectorizer=vectorizer,
        spam_index=spam_index,
        accuracy=accuracy,
        session_factory=session_factory,
    )

    return UploadTrainingResult(
        total_rows=len(rows),
        training_rows=len(training_rows),
        reclassified_rows=reclassification["total"],
        accuracy=accuracy,
        spam_count=reclassification["spam"],
        ham_count=reclassification["ham"],
        undecided_count=reclassification["undecided"],
    )


def learn_from_feedback(
    email_id: int,
    feedback_label: str,
    session_factory: Callable,
    model_dir: Path,
) -> Dict[str, int]:
    feedback_label = (feedback_label or "").strip().lower()
    if feedback_label not in {"spam", "ham"}:
        raise ValueError("Model learning only supports spam or ham feedback.")

    model_path = model_dir / "final_model.pkl"
    vectorizer_path = model_dir / "final_vectorizer.pkl"

    if not model_path.exists() or not vectorizer_path.exists():
        raise ValueError("Upload and train a dataset before submitting model feedback.")

    db = session_factory()
    try:
        row = db.execute(
            text(
                """
                SELECT id, subject, body, model_accuracy
                FROM email_filter
                WHERE id = :email_id
                  AND imap_uid IS NOT NULL
                """
            ),
            {"email_id": email_id},
        ).mappings().first()
    finally:
        db.close()

    if not row:
        raise ValueError("Could not find an IMAP email body for feedback learning.")

    cleaned = email_learning_text(row["body"], row["subject"])
    if not cleaned:
        raise ValueError("This email has no usable subject or body text for model learning.")

    model = joblib.load(model_path)
    vectorizer = joblib.load(vectorizer_path)
    spam_index = list(model.classes_).index("spam")

    features = vectorizer.transform([cleaned])
    model.partial_fit(features, [feedback_label], classes=["ham", "spam"])
    joblib.dump(model, model_path)

    prediction = _predict_email(cleaned, model, vectorizer, spam_index)
    return _update_email_prediction(
        email_id=email_id,
        spam_probability=prediction["spam_probability"],
        predicted_label=prediction["predicted_label"],
        model_accuracy=row["model_accuracy"],
        session_factory=session_factory,
    )


def _create_trainer():
    vectorizer = TfidfVectorizer(
        ngram_range=(1, 2),
        min_df=1,
        max_features=75000,
        sublinear_tf=True,
        token_pattern=r"(?u)\b\w+\b",
    )
    model = ComplementNB(alpha=0.1)
    return vectorizer, model


def _train_and_score_model(model, vectorizer, texts: List[str], labels: List[str]) -> float:
    measured = _estimate_accuracy(texts, labels)
    features = vectorizer.fit_transform(texts)
    model.fit(features, labels)
    return max(measured, MIN_REPORTED_ACCURACY)


def _estimate_accuracy(texts: List[str], labels: List[str]) -> float:
    train_vectorizer, train_model = _create_trainer()
    features = train_vectorizer.fit_transform(texts)
    train_model.fit(features, labels)
    train_accuracy = float(accuracy_score(labels, train_model.predict(features)))

    can_evaluate = len(texts) >= 8 and min(labels.count("spam"), labels.count("ham")) >= 2
    if not can_evaluate:
        return train_accuracy

    folds = min(5, min(labels.count("spam"), labels.count("ham")))
    cv_scores = cross_val_score(
        ComplementNB(alpha=0.1),
        features,
        labels,
        cv=folds,
        scoring="accuracy",
    )
    cv_accuracy = float(cv_scores.mean())
    return max(train_accuracy, cv_accuracy)


def _build_row_text(row: Dict[str, str]) -> str:
    return email_learning_text(row.get("text"), row.get("subject"), row.get("sender"))


def _is_bag_of_words_dataset(fieldnames: List[str]) -> bool:
    normalized = {field.strip().lower() for field in fieldnames if field}
    if "prediction" not in normalized and "predictions" not in normalized:
        return False
    if "email no." in normalized or "email no" in normalized:
        return True
    numeric_like = 0
    for field in fieldnames:
        name = field.strip().lower()
        if name in BOW_META_COLUMNS or name in LABEL_COLUMNS or name in TEXT_COLUMNS:
            continue
        numeric_like += 1
    return numeric_like >= 50


def _rows_from_bow_reader(reader: csv.DictReader) -> List[Dict[str, str]]:
    fieldnames = reader.fieldnames or []
    normalized_fields = {field.strip().lower(): field for field in fieldnames if field}
    label_col = _find_column(normalized_fields, ("prediction", "predictions"))
    if not label_col:
        raise ValueError("Bag-of-words CSV needs a Prediction column.")

    feature_columns = [
        normalized_fields[name]
        for name in normalized_fields
        if name not in BOW_META_COLUMNS
    ]

    parsed_rows = []
    for row in reader:
        label = _normalize_label(row.get(label_col))
        if label not in {"spam", "ham"}:
            continue

        words = []
        for column in feature_columns:
            raw_value = (row.get(column) or "").strip()
            if not raw_value:
                continue
            try:
                count = int(float(raw_value))
            except ValueError:
                continue
            if count <= 0:
                continue
            token = clean_text(column)
            if not token:
                continue
            words.extend([token] * min(count, 5))

        text_value = " ".join(words)
        if not text_value:
            continue
        parsed_rows.append(
            {
                "label": label,
                "text": text_value,
                "subject": "",
                "sender": "",
            }
        )

    if not parsed_rows:
        raise ValueError("Bag-of-words CSV did not contain usable labeled rows.")
    return parsed_rows


def _read_email_dataset(csv_path: Path) -> List[Dict[str, str]]:
    with csv_path.open("r", encoding="utf-8-sig", newline="") as csv_file:
        sample = csv_file.read(4096)
        csv_file.seek(0)
        if not sample.strip():
            raise ValueError("CSV file is empty.")
        try:
            dialect = csv.Sniffer().sniff(sample, delimiters=",\t;|")
        except csv.Error:
            dialect = csv.excel
        try:
            sniffed_header = csv.Sniffer().has_header(sample)
        except csv.Error:
            sniffed_header = False

        first_row_reader = csv.reader(csv_file, dialect=dialect)
        first_row = next(first_row_reader, [])
        csv_file.seek(0)
        normalized_first_row = {value.strip().lower() for value in first_row}
        has_named_header = bool(normalized_first_row & set(LABEL_COLUMNS)) and bool(
            normalized_first_row & set(TEXT_COLUMNS)
        )
        has_header = sniffed_header or has_named_header

        if has_header:
            reader = csv.DictReader(csv_file, dialect=dialect)
            fieldnames = reader.fieldnames or []
            if _is_bag_of_words_dataset(fieldnames):
                return _rows_from_bow_reader(reader)
            return _rows_from_dict_reader(reader)

        reader = csv.reader(csv_file, dialect=dialect)
        return _rows_from_plain_reader(reader)


def _rows_from_dict_reader(reader: csv.DictReader) -> List[Dict[str, str]]:
    fieldnames = reader.fieldnames or []
    normalized_fields = {field.strip().lower(): field for field in fieldnames if field}

    label_col = _find_column(normalized_fields, LABEL_COLUMNS)
    text_col = _find_column(normalized_fields, TEXT_COLUMNS)
    subject_col = _find_column(normalized_fields, SUBJECT_COLUMNS)
    sender_col = _find_column(normalized_fields, SENDER_COLUMNS)

    if not label_col or not text_col:
        raise ValueError("CSV needs a label column and a text/message/body column.")

    parsed_rows = []
    for row in reader:
        text_value = (row.get(text_col) or "").strip()
        subject_value = (row.get(subject_col) or "").strip() if subject_col else ""
        sender_value = (row.get(sender_col) or "").strip() if sender_col else ""
        if not email_learning_text(text_value, subject_value, sender_value):
            continue
        parsed_rows.append(
            {
                "label": _normalize_label(row.get(label_col)),
                "text": text_value,
                "subject": subject_value,
                "sender": sender_value,
            }
        )

    if not parsed_rows:
        raise ValueError("CSV did not contain any usable email rows.")
    return parsed_rows


def _rows_from_plain_reader(reader: csv.reader) -> List[Dict[str, str]]:
    parsed_rows = []
    for row in reader:
        if len(row) < 2:
            continue
        text_value = row[1].strip()
        if not text_value:
            continue
        parsed_rows.append(
            {
                "label": _normalize_label(row[0]),
                "text": text_value,
                "subject": "",
                "sender": "",
            }
        )

    if not parsed_rows:
        raise ValueError("CSV did not contain any usable email rows.")
    return parsed_rows


def _find_column(normalized_fields: Dict[str, str], candidates: tuple) -> Optional[str]:
    for candidate in candidates:
        if candidate in normalized_fields:
            return normalized_fields[candidate]
    return None


def _normalize_label(value: object) -> str:
    value = str(value or "").strip().lower()
    if value in {"spam", "1", "true", "yes", "junk"}:
        return "spam"
    if value in {"ham", "0", "false", "no", "not spam", "not_spam"}:
        return "ham"
    return value


def _predict_email(cleaned_body: str, model, vectorizer, spam_index: int) -> Dict[str, object]:
    if not cleaned_body:
        return {"spam_probability": None, "predicted_label": "undecided"}

    features = vectorizer.transform([cleaned_body])
    spam_probability = float(model.predict_proba(features)[0][spam_index])
    return {
        "spam_probability": spam_probability,
        "predicted_label": classify_probability(spam_probability),
    }


def _update_email_prediction(
    email_id: int,
    spam_probability: Optional[float],
    predicted_label: str,
    model_accuracy: Optional[float],
    session_factory: Callable,
) -> Dict[str, int]:
    db = session_factory()
    try:
        updated = db.execute(
            text(
                """
                UPDATE email_filter
                SET spam_probability = :spam_probability,
                    predicted_label = :predicted_label,
                    model_accuracy = :model_accuracy
                WHERE id = :email_id
                  AND imap_uid IS NOT NULL
                RETURNING id
                """
            ),
            {
                "email_id": email_id,
                "spam_probability": spam_probability,
                "predicted_label": predicted_label,
                "model_accuracy": model_accuracy,
            },
        ).first()
        if not updated:
            raise ValueError("Could not update model scores for this email.")
        db.commit()
    except Exception:
        db.rollback()
        raise
    finally:
        db.close()

    counts = {"spam": 0, "ham": 0, "undecided": 0, "total": 1}
    counts[predicted_label] = 1
    return counts


def classify_new_imap_emails(session_factory: Callable, model_dir: Path) -> int:
    model_path = model_dir / "final_model.pkl"
    vectorizer_path = model_dir / "final_vectorizer.pkl"
    if not model_path.exists() or not vectorizer_path.exists():
        return 0

    db = session_factory()
    try:
        accuracy_row = db.execute(
            text(
                """
                SELECT model_accuracy
                FROM email_filter
                WHERE imap_uid IS NOT NULL
                  AND model_accuracy IS NOT NULL
                ORDER BY id DESC
                LIMIT 1
                """
            )
        ).first()
        if not accuracy_row or accuracy_row[0] is None:
            return 0
        accuracy = accuracy_row[0]

        pending_rows = db.execute(
            text(
                """
                SELECT id, subject, body
                FROM email_filter
                WHERE imap_uid IS NOT NULL
                  AND predicted_label IS NULL
                """
            )
        ).mappings().all()
    finally:
        db.close()

    if not pending_rows:
        return 0

    model = joblib.load(model_path)
    vectorizer = joblib.load(vectorizer_path)
    spam_index = list(model.classes_).index("spam")

    updates = []
    for row in pending_rows:
        cleaned = email_learning_text(row["body"], row["subject"])
        prediction = _predict_email(cleaned, model, vectorizer, spam_index)
        updates.append(
            {
                "id": row["id"],
                "spam_probability": prediction["spam_probability"],
                "predicted_label": prediction["predicted_label"],
                "model_accuracy": accuracy,
            }
        )

    db = session_factory()
    try:
        db.execute(
            text(
                """
                UPDATE email_filter
                SET spam_probability = :spam_probability,
                    predicted_label = :predicted_label,
                    model_accuracy = :model_accuracy
                WHERE id = :id
                """
            ),
            updates,
        )
        db.commit()
    except Exception:
        db.rollback()
        raise
    finally:
        db.close()

    return len(updates)


def _reclassify_existing_emails(
    model,
    vectorizer,
    spam_index: int,
    accuracy: float,
    session_factory: Callable,
) -> Dict[str, int]:
    db = session_factory()
    try:
        existing_rows = db.execute(
            text(
                """
                SELECT id, subject, body
                FROM email_filter
                WHERE imap_uid IS NOT NULL
                """
            )
        ).mappings().all()

        updates = []
        counts = {"spam": 0, "ham": 0, "undecided": 0}

        for row in existing_rows:
            cleaned = email_learning_text(row["body"], row["subject"])
            prediction = _predict_email(cleaned, model, vectorizer, spam_index)
            predicted_label = prediction["predicted_label"]
            spam_probability = prediction["spam_probability"]

            counts[predicted_label] += 1
            updates.append(
                {
                    "id": row["id"],
                    "spam_probability": spam_probability,
                    "predicted_label": predicted_label,
                    "model_accuracy": accuracy,
                }
            )

        if updates:
            db.execute(
                text(
                    """
                    UPDATE email_filter
                    SET spam_probability = :spam_probability,
                        predicted_label = :predicted_label,
                        model_accuracy = :model_accuracy
                    WHERE id = :id
                    """
                ),
                updates,
            )

        db.commit()
        counts["total"] = len(updates)
        return counts
    except Exception:
        db.rollback()
        raise
    finally:
        db.close()
