"""
Lead data cleaning engine.

Design goals:
- Handle millions of rows without loading everything into memory unnecessarily
  (Polars lazy/streaming where practical).
- Never silently corrupt data: every rejected row is logged with a reason.
- Be delimiter-agnostic for multi-value fields by extracting values via
  pattern matching rather than assuming a fixed separator character.
"""

import re
import polars as pl
from pathlib import Path
from dataclasses import dataclass, field

# ---------------------------------------------------------------------------
# Constants / blocklists
# ---------------------------------------------------------------------------

EMAIL_REGEX = r"[a-zA-Z0-9._%+\-]+@[a-zA-Z0-9.\-]+\.[a-zA-Z]{2,}"

# Loose phone matcher: sequences with at least 7 digits, allowing common
# separators/punctuation used inside a single phone number.
PHONE_SPLIT_REGEX = r"\s*[:;|]\s*|\s*,\s*(?=\+|\d{3,})"
PHONE_MIN_DIGITS = 7

# Educational domain patterns (international academic TLD/SLD conventions)
EDU_PATTERNS = [
    r"\.edu$",
    r"\.edu\.[a-z]{2,3}$",     # edu.in, edu.au, edu.cn ...
    r"\.ac\.[a-z]{2,3}$",      # ac.in, ac.uk, ac.jp ...
    r"\.ac$",
]

# Military / defense domain patterns + a small starter blocklist of known
# defense-sector companies. This list is intentionally conservative (exact
# domain matches, not keyword matches) to avoid false positives like
# "armystrong.com". Extend backend/blocklists/military_domains.txt as needed.
MIL_PATTERNS = [
    r"\.mil$",
    r"\.mil\.[a-z]{2,3}$",
]

DEFAULT_MILITARY_BLOCKLIST_PATH = Path(__file__).parent / "blocklists" / "military_domains.txt"

PLACEHOLDER_STORE_PATTERNS = [
    r"\.myshopify\.com$",
    r"\.myshopify\.io$",
]


def _load_blocklist(path: Path) -> set[str]:
    if not path.exists():
        return set()
    with open(path, "r", encoding="utf-8") as f:
        return {
            line.strip().lower()
            for line in f
            if line.strip() and not line.strip().startswith("#")
        }


# ---------------------------------------------------------------------------
# Result container
# ---------------------------------------------------------------------------

@dataclass
class CleaningResult:
    clean_df: pl.DataFrame
    rejected_df: pl.DataFrame
    stats: dict = field(default_factory=dict)
    category_column: str | None = None


# ---------------------------------------------------------------------------
# Column detection
# ---------------------------------------------------------------------------

def _detect_column(columns: list[str], keywords: list[str]) -> str | None:
    lower_map = {c.lower(): c for c in columns}
    for kw in keywords:
        for lc, orig in lower_map.items():
            if kw in lc:
                return orig
    return None


# ---------------------------------------------------------------------------
# Value extraction helpers (row-level, applied via map_elements)
# ---------------------------------------------------------------------------

def _extract_emails(raw: str | None) -> list[str]:
    if not raw:
        return []
    found = re.findall(EMAIL_REGEX, raw)
    seen = []
    for e in found:
        e = e.strip().lower()
        if e not in seen:
            seen.append(e)
    return seen


def _extract_phones(raw: str | None) -> list[str]:
    if not raw:
        return []
    parts = re.split(PHONE_SPLIT_REGEX, raw)
    out = []
    for p in parts:
        p = p.strip()
        digits = re.sub(r"\D", "", p)
        if len(digits) >= PHONE_MIN_DIGITS and p not in out:
            out.append(p)
    return out


def _pair_multivalue(emails: list[str], phones: list[str]) -> list[tuple[str | None, str | None]]:
    """Pair emails/phones by position. If one side has exactly one value and
    the other has several, repeat the single value across all rows. If both
    have multiple values of different lengths, pair by index and leave the
    shorter side's extra slots as None."""
    n_e, n_p = len(emails), len(phones)

    if n_e == 0 and n_p == 0:
        return [(None, None)]

    if n_e <= 1 and n_p > 1:
        e_val = emails[0] if emails else None
        return [(e_val, p) for p in phones]

    if n_p <= 1 and n_e > 1:
        p_val = phones[0] if phones else None
        return [(e, p_val) for e in emails]

    n = max(n_e, n_p, 1)
    pairs = []
    for i in range(n):
        e = emails[i] if i < n_e else None
        p = phones[i] if i < n_p else None
        pairs.append((e, p))
    return pairs


# ---------------------------------------------------------------------------
# Main pipeline
# ---------------------------------------------------------------------------

def clean_leads(
    input_path: str | Path | list[str | Path],
    email_column: str | None = None,
    phone_column: str | None = None,
    explode_phones: bool = True,
    dedupe_key: str = "email",  # "email" | "domain" | "row"
    exclude_placeholder_stores: bool = False,  # False -> flag only (default per spec)
    military_blocklist_path: Path = DEFAULT_MILITARY_BLOCKLIST_PATH,
    category_column: str | None = None,
) -> CleaningResult:
    input_paths = [input_path] if isinstance(input_path, (str, Path)) else list(input_path)

    if len(input_paths) == 1:
        df = pl.read_csv(input_paths[0], infer_schema_length=10000, ignore_errors=True)
    else:
        # Bulk merge: files from the same kind of export usually share most
        # columns but may differ slightly (an extra field here, a missing
        # one there). "diagonal_relaxed" unions the column sets — rows from
        # a file missing a given column just get null there — rather than
        # erroring or silently dropping mismatched columns.
        frames = [pl.read_csv(p, infer_schema_length=10000, ignore_errors=True) for p in input_paths]
        df = pl.concat(frames, how="diagonal_relaxed")

    original_columns = df.columns
    n_input = df.height

    email_col = email_column or _detect_column(original_columns, ["email"])
    phone_col = phone_column or _detect_column(original_columns, ["phone"])

    if email_col is None:
        raise ValueError("Could not detect an email column. Pass email_column explicitly.")

    # Category/platform column detection: prefer an exact "platform" column,
    # then exact "category", but never the plural "categories" (that's a
    # different, multi-value taxonomy field in typical lead exports).
    cat_col = category_column
    if cat_col is None:
        exact_lower = {c.lower(): c for c in original_columns}
        cat_col = exact_lower.get("platform") or exact_lower.get("category")

    mil_blocklist = _load_blocklist(military_blocklist_path)

    # --- Step 1: drop exact full-row duplicates -----------------------------
    df = df.with_row_index("_orig_row_id")
    # with_row_index()'s dtype (u32 vs i64) has varied across Polars
    # versions; pin it explicitly rather than relying on it matching
    # whatever dtype gets inferred for exp_df below — a silent mismatch
    # here throws "datatypes of join keys don't match" at the join.
    df = df.with_columns(pl.col("_orig_row_id").cast(pl.Int64))
    n_before_exact_dedupe = df.height
    df = df.unique(subset=[c for c in original_columns], keep="first", maintain_order=True)
    n_exact_dupes_removed = n_before_exact_dedupe - df.height

    # --- Step 2: extract & pair emails/phones per row -----------------------
    email_series = df[email_col].to_list()
    phone_series = df[phone_col].to_list() if (phone_col and explode_phones) else [None] * df.height

    exploded_rows = []
    for i in range(df.height):
        emails = _extract_emails(email_series[i])
        phones = _extract_phones(phone_series[i]) if explode_phones else (
            [phone_series[i]] if phone_series[i] else []
        )
        pairs = _pair_multivalue(emails, phones)
        for e, p in pairs:
            exploded_rows.append((i, e, p))

    exp_df = pl.DataFrame(
        exploded_rows,
        schema={"_orig_row_id": pl.Int64, "_clean_email": pl.String, "_clean_phone": pl.String},
        orient="row",
    )

    base = df.drop([email_col] + ([phone_col] if phone_col else []))
    result = exp_df.join(base, on="_orig_row_id", how="left")
    result = result.rename({"_clean_email": email_col})
    if phone_col:
        result = result.rename({"_clean_phone": phone_col})
    else:
        result = result.drop("_clean_phone")

    n_after_explode = result.height

    # --- Step 3: classify each row (military / edu / placeholder / ok) ------
    def classify(email: str | None) -> str:
        if not email or "@" not in email:
            return "no_email"
        domain = email.rsplit("@", 1)[-1].lower()
        for pat in MIL_PATTERNS:
            if re.search(pat, domain):
                return "military_domain"
        if domain in mil_blocklist:
            return "military_domain"
        for pat in EDU_PATTERNS:
            if re.search(pat, domain):
                return "educational_domain"
        return "ok"

    result = result.with_columns(
        pl.col(email_col)
        .map_elements(classify, return_dtype=pl.String, skip_nulls=False)
        .alias("_reject_reason")
    )

    # --- Step 4: placeholder store flag (e.g. *.myshopify.com dev stores) ---
    domain_col = _detect_column(original_columns, ["domain"])
    if domain_col and domain_col in result.columns:
        def is_placeholder(d: str | None) -> bool:
            if not d:
                return False
            d = d.lower()
            return any(re.search(pat, d) for pat in PLACEHOLDER_STORE_PATTERNS)

        result = result.with_columns(
            pl.col(domain_col)
            .map_elements(is_placeholder, return_dtype=pl.Boolean, skip_nulls=False)
            .alias("is_placeholder_store")
        )
        if exclude_placeholder_stores:
            n_placeholder = result.filter(pl.col("is_placeholder_store")).height
            result = result.filter(~pl.col("is_placeholder_store"))
        else:
            n_placeholder = result.filter(pl.col("is_placeholder_store")).height
    else:
        n_placeholder = 0

    # --- Step 5: split clean vs rejected -------------------------------------
    rejected_df = result.filter(pl.col("_reject_reason").is_in(["military_domain", "educational_domain"]))
    clean_df = result.filter(~pl.col("_reject_reason").is_in(["military_domain", "educational_domain"]))

    n_military = rejected_df.filter(pl.col("_reject_reason") == "military_domain").height
    n_edu = rejected_df.filter(pl.col("_reject_reason") == "educational_domain").height
    n_no_email = clean_df.filter(pl.col("_reject_reason") == "no_email").height

    clean_df = clean_df.drop("_reject_reason")

    # --- Step 6: dedupe -------------------------------------------------------
    n_before_final_dedupe = clean_df.height
    if dedupe_key == "email":
        # keep first occurrence per non-null email; keep all no-email rows as-is
        has_email = clean_df.filter(pl.col(email_col).is_not_null())
        no_email_rows = clean_df.filter(pl.col(email_col).is_null())
        has_email = has_email.unique(subset=[email_col], keep="first", maintain_order=True)
        clean_df = pl.concat([has_email, no_email_rows])
    elif dedupe_key == "domain" and domain_col:
        clean_df = clean_df.unique(subset=[domain_col], keep="first", maintain_order=True)
    # "row" -> already exact-deduped in step 1, nothing further

    n_final_dupes_removed = n_before_final_dedupe - clean_df.height

    clean_df = clean_df.drop("_orig_row_id", strict=False)
    rejected_df = rejected_df.drop("_orig_row_id", strict=False)

    stats = {
        "rows_in_original_file": n_input,
        "exact_duplicate_rows_removed": n_exact_dupes_removed,
        "rows_after_email_phone_explode": n_after_explode,
        "rejected_military_domain": n_military,
        "rejected_educational_domain": n_edu,
        "rows_with_no_email": n_no_email,
        "duplicate_emails_removed": n_final_dupes_removed,
        "placeholder_stores_flagged": n_placeholder,
        "final_clean_row_count": clean_df.height,
    }

    if cat_col and cat_col in clean_df.columns:
        counts = (
            clean_df.group_by(cat_col)
            .agg(pl.len().alias("count"))
            .sort("count", descending=True)
        )
        stats["category_counts"] = {
            (row[cat_col] if row[cat_col] is not None else "(uncategorized)"): row["count"]
            for row in counts.to_dicts()
        }

    return CleaningResult(clean_df=clean_df, rejected_df=rejected_df, stats=stats, category_column=cat_col)


if __name__ == "__main__":
    import sys, json
    path = sys.argv[1] if len(sys.argv) > 1 else "sample.csv"
    res = clean_leads(path)
    print(json.dumps(res.stats, indent=2))
    print("\n--- clean sample ---")
    print(res.clean_df.select(["domain", "emails", "phones", "platform"]).head(10))
    print("\n--- rejected sample ---")
    if res.rejected_df.height:
        print(res.rejected_df.select(["domain", "emails", "_reject_reason"]).head(10))
    else:
        print("(none in this file)")
