"""Customer lifetime value and retention decision engine.

This synthetic portfolio project estimates contribution-based customer value,
tests the ranking against a time-based holdout and produces an intervention
queue.  It is intentionally transparent: the code separates observed value,
forecast assumptions, acquisition economics and treatment eligibility.
"""

from __future__ import annotations

import argparse
import json
from dataclasses import asdict, dataclass
from pathlib import Path

import numpy as np
import pandas as pd
from sklearn.compose import ColumnTransformer
from sklearn.ensemble import GradientBoostingRegressor
from sklearn.metrics import mean_absolute_error
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import OneHotEncoder, StandardScaler


SNAPSHOT = pd.Timestamp("2026-06-30")
HOLDOUT_START = pd.Timestamp("2026-01-01")
REQUIRED = {
    "transaction_id",
    "customer_id",
    "order_date",
    "join_date",
    "segment",
    "acquisition_channel",
    "revenue_gbp",
    "gross_margin_gbp",
    "acquisition_cost_gbp",
    "support_contacts_90d",
}


@dataclass(frozen=True)
class Assumptions:
    forecast_months: int = 12
    annual_discount_rate: float = 0.10
    minimum_contactable_clv_gbp: float = 120.0
    inactivity_trigger_days: int = 120
    treatment_cost_gbp: float = 12.0
    expected_reactivation_rate: float = 0.11


@dataclass(frozen=True)
class DataQuality:
    rows: int
    customers: int
    duplicate_transactions: int
    missing_required_values: int
    negative_revenue_rows: int
    margin_above_revenue_rows: int
    order_before_join_rows: int


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input", type=Path)
    parser.add_argument("--output-dir", type=Path, default=Path("clv_outputs"))
    return parser.parse_args()


def load_transactions(path: Path) -> pd.DataFrame:
    frame = pd.read_csv(path)
    missing = REQUIRED.difference(frame.columns)
    if missing:
        raise ValueError(f"Missing required columns: {sorted(missing)}")
    for column in ("order_date", "join_date"):
        frame[column] = pd.to_datetime(frame[column], errors="raise")
    numeric = [
        "revenue_gbp",
        "gross_margin_gbp",
        "acquisition_cost_gbp",
        "support_contacts_90d",
    ]
    frame[numeric] = frame[numeric].apply(pd.to_numeric, errors="raise")
    return frame.sort_values(["customer_id", "order_date", "transaction_id"])


def assess_quality(frame: pd.DataFrame) -> DataQuality:
    return DataQuality(
        rows=len(frame),
        customers=frame["customer_id"].nunique(),
        duplicate_transactions=int(frame["transaction_id"].duplicated().sum()),
        missing_required_values=int(frame[list(REQUIRED)].isna().sum().sum()),
        negative_revenue_rows=int((frame["revenue_gbp"] < 0).sum()),
        margin_above_revenue_rows=int((frame["gross_margin_gbp"] > frame["revenue_gbp"]).sum()),
        order_before_join_rows=int((frame["order_date"] < frame["join_date"]).sum()),
    )


def fail_on_quality(quality: DataQuality) -> None:
    failures = {
        key: value
        for key, value in asdict(quality).items()
        if key not in {"rows", "customers"} and value
    }
    if failures:
        raise ValueError(f"Data quality controls failed: {failures}")


def customer_features(
    frame: pd.DataFrame,
    observation_end: pd.Timestamp,
) -> pd.DataFrame:
    observed = frame.loc[frame["order_date"] <= observation_end].copy()
    observed["customer_age_days"] = (observation_end - observed["join_date"]).dt.days.clip(lower=1)
    customer = (
        observed.groupby("customer_id", as_index=False)
        .agg(
            segment=("segment", "first"),
            acquisition_channel=("acquisition_channel", "first"),
            join_date=("join_date", "min"),
            first_order_date=("order_date", "min"),
            last_order_date=("order_date", "max"),
            order_count=("transaction_id", "nunique"),
            observed_revenue_gbp=("revenue_gbp", "sum"),
            observed_margin_gbp=("gross_margin_gbp", "sum"),
            acquisition_cost_gbp=("acquisition_cost_gbp", "sum"),
            average_order_value_gbp=("revenue_gbp", "mean"),
            average_order_margin_gbp=("gross_margin_gbp", "mean"),
            support_contacts_90d=("support_contacts_90d", "max"),
        )
    )
    customer["recency_days"] = (observation_end - customer["last_order_date"]).dt.days
    customer["tenure_days"] = (observation_end - customer["join_date"]).dt.days.clip(lower=1)
    customer["active_span_days"] = (
        customer["last_order_date"] - customer["first_order_date"]
    ).dt.days.clip(lower=1)
    customer["orders_per_30d"] = customer["order_count"] / customer["tenure_days"] * 30
    customer["margin_per_30d"] = customer["observed_margin_gbp"] / customer["tenure_days"] * 30
    customer["repeat_customer"] = customer["order_count"].ge(2).astype(int)
    customer["observed_net_value_gbp"] = (
        customer["observed_margin_gbp"] - customer["acquisition_cost_gbp"]
    )
    return customer


def actual_holdout_value(frame: pd.DataFrame) -> pd.Series:
    holdout = frame.loc[
        frame["order_date"].between(HOLDOUT_START, SNAPSHOT, inclusive="both")
    ]
    return holdout.groupby("customer_id")["gross_margin_gbp"].sum().rename("actual_holdout_margin_gbp")


def model_pipeline() -> Pipeline:
    numeric = [
        "order_count",
        "observed_revenue_gbp",
        "observed_margin_gbp",
        "average_order_value_gbp",
        "average_order_margin_gbp",
        "recency_days",
        "tenure_days",
        "active_span_days",
        "orders_per_30d",
        "margin_per_30d",
        "support_contacts_90d",
        "repeat_customer",
    ]
    categorical = ["segment", "acquisition_channel"]
    transformer = ColumnTransformer(
        [
            ("numeric", StandardScaler(), numeric),
            ("categorical", OneHotEncoder(handle_unknown="ignore"), categorical),
        ]
    )
    model = GradientBoostingRegressor(
        n_estimators=180,
        learning_rate=0.035,
        max_depth=2,
        min_samples_leaf=16,
        loss="huber",
        random_state=1210,
    )
    return Pipeline([("features", transformer), ("model", model)])


def train_and_validate(frame: pd.DataFrame) -> tuple[Pipeline, pd.DataFrame, dict]:
    training_features = customer_features(frame, HOLDOUT_START - pd.Timedelta(days=1))
    outcomes = actual_holdout_value(frame)
    training = training_features.merge(outcomes, on="customer_id", how="left")
    training["actual_holdout_margin_gbp"] = training["actual_holdout_margin_gbp"].fillna(0)

    excluded = {
        "customer_id",
        "join_date",
        "first_order_date",
        "last_order_date",
        "actual_holdout_margin_gbp",
        "acquisition_cost_gbp",
        "observed_net_value_gbp",
    }
    feature_columns = [column for column in training.columns if column not in excluded]
    pipeline = model_pipeline()
    pipeline.fit(training[feature_columns], training["actual_holdout_margin_gbp"])
    predictions = np.clip(pipeline.predict(training[feature_columns]), 0, None)

    baseline = np.repeat(training["actual_holdout_margin_gbp"].median(), len(training))
    model_mae = mean_absolute_error(training["actual_holdout_margin_gbp"], predictions)
    baseline_mae = mean_absolute_error(training["actual_holdout_margin_gbp"], baseline)
    ranking = training.assign(prediction=predictions).sort_values("prediction", ascending=False)
    top_decile = max(1, int(np.ceil(len(ranking) * 0.10)))
    captured = ranking.head(top_decile)["actual_holdout_margin_gbp"].sum()
    total = ranking["actual_holdout_margin_gbp"].sum()
    diagnostics = {
        "training_customers": int(len(training)),
        "model_mae_gbp": round(float(model_mae), 2),
        "baseline_mae_gbp": round(float(baseline_mae), 2),
        "mae_improvement_pct": round(float(100 * (baseline_mae - model_mae) / baseline_mae), 1),
        "top_decile_holdout_value_capture_pct": round(float(100 * captured / total), 1),
    }
    return pipeline, training, diagnostics


def score_current_customers(
    frame: pd.DataFrame,
    pipeline: Pipeline,
    assumptions: Assumptions,
) -> pd.DataFrame:
    current = customer_features(frame, SNAPSHOT)
    excluded = {
        "customer_id",
        "join_date",
        "first_order_date",
        "last_order_date",
        "acquisition_cost_gbp",
        "observed_net_value_gbp",
    }
    feature_columns = [column for column in current.columns if column not in excluded]
    six_month_margin = np.clip(pipeline.predict(current[feature_columns]), 0, None)
    annual_discount_factor = 1 / (1 + assumptions.annual_discount_rate)
    current["forecast_12m_margin_gbp"] = six_month_margin * 2
    current["discounted_12m_clv_gbp"] = (
        current["forecast_12m_margin_gbp"] * annual_discount_factor
        - current["acquisition_cost_gbp"]
    )
    current["value_band"] = pd.qcut(
        current["discounted_12m_clv_gbp"].rank(method="first"),
        q=[0, 0.40, 0.75, 0.90, 1.0],
        labels=["Develop", "Core", "High", "Strategic"],
    ).astype(str)
    current["inactivity_risk"] = np.select(
        [
            current["recency_days"] >= 210,
            current["recency_days"] >= assumptions.inactivity_trigger_days,
            current["recency_days"] >= 75,
        ],
        ["Critical", "High", "Watch"],
        default="Current",
    )
    current["contact_priority"] = (
        current["discounted_12m_clv_gbp"].clip(lower=0)
        * (1 + current["recency_days"].clip(upper=365) / 365)
        * (1 + current["support_contacts_90d"].clip(upper=5) * 0.06)
    )
    current["eligible_for_retention"] = (
        current["discounted_12m_clv_gbp"].ge(assumptions.minimum_contactable_clv_gbp)
        & current["recency_days"].ge(assumptions.inactivity_trigger_days)
    )
    current["expected_incremental_margin_gbp"] = np.where(
        current["eligible_for_retention"],
        current["forecast_12m_margin_gbp"] * assumptions.expected_reactivation_rate
        - assumptions.treatment_cost_gbp,
        0,
    )
    current["priority_rank"] = current["contact_priority"].rank(
        method="first", ascending=False
    ).astype(int)
    return current.sort_values("priority_rank")


def cohort_summary(frame: pd.DataFrame) -> pd.DataFrame:
    cohort = frame.copy()
    cohort["join_cohort"] = cohort["join_date"].dt.to_period("Q").astype(str)
    customer_cohort = (
        cohort.groupby("customer_id", as_index=False)
        .agg(
            join_cohort=("join_cohort", "first"),
            acquisition_channel=("acquisition_channel", "first"),
            margin_gbp=("gross_margin_gbp", "sum"),
            acquisition_cost_gbp=("acquisition_cost_gbp", "sum"),
            orders=("transaction_id", "nunique"),
        )
    )
    return (
        customer_cohort.groupby(["join_cohort", "acquisition_channel"], as_index=False)
        .agg(
            customers=("customer_id", "nunique"),
            margin_gbp=("margin_gbp", "sum"),
            acquisition_cost_gbp=("acquisition_cost_gbp", "sum"),
            average_orders=("orders", "mean"),
        )
        .assign(
            net_margin_gbp=lambda x: x["margin_gbp"] - x["acquisition_cost_gbp"],
            margin_to_cac=lambda x: x["margin_gbp"] / x["acquisition_cost_gbp"].replace(0, np.nan),
        )
    )


def executive_summary(
    scored: pd.DataFrame,
    diagnostics: dict,
    quality: DataQuality,
    assumptions: Assumptions,
) -> dict:
    top_decile = max(1, int(np.ceil(len(scored) * 0.10)))
    total_value = scored["discounted_12m_clv_gbp"].clip(lower=0).sum()
    top_value = scored.head(top_decile)["discounted_12m_clv_gbp"].clip(lower=0).sum()
    eligible = scored.loc[scored["eligible_for_retention"]]
    positive = eligible.loc[eligible["expected_incremental_margin_gbp"] > 0]
    return {
        "quality": asdict(quality),
        "assumptions": asdict(assumptions),
        "validation": diagnostics,
        "portfolio": {
            "discounted_12m_clv_gbp": round(float(total_value), 2),
            "top_decile_value_share_pct": round(float(100 * top_value / total_value), 1),
            "retention_eligible_customers": int(len(eligible)),
            "positive_case_customers": int(len(positive)),
            "expected_incremental_margin_gbp": round(
                float(positive["expected_incremental_margin_gbp"].sum()), 2
            ),
            "median_clv_gbp": round(float(scored["discounted_12m_clv_gbp"].median()), 2),
        },
    }


def save_outputs(
    output_dir: Path,
    scored: pd.DataFrame,
    cohorts: pd.DataFrame,
    summary: dict,
) -> None:
    output_dir.mkdir(parents=True, exist_ok=True)
    scored.to_csv(output_dir / "customer_value_scores.csv", index=False)
    scored.loc[scored["eligible_for_retention"]].head(300).to_csv(
        output_dir / "retention_action_queue.csv", index=False
    )
    cohorts.to_csv(output_dir / "cohort_channel_economics.csv", index=False)
    (output_dir / "executive_summary.json").write_text(
        json.dumps(summary, indent=2), encoding="utf-8"
    )


def main() -> None:
    args = parse_args()
    assumptions = Assumptions()
    transactions = load_transactions(args.input)
    quality = assess_quality(transactions)
    fail_on_quality(quality)
    pipeline, _, diagnostics = train_and_validate(transactions)
    scored = score_current_customers(transactions, pipeline, assumptions)
    cohorts = cohort_summary(transactions)
    summary = executive_summary(scored, diagnostics, quality, assumptions)
    save_outputs(args.output_dir, scored, cohorts, summary)
    print(json.dumps(summary, indent=2))


if __name__ == "__main__":
    main()
