#!/usr/bin/env python3
"""MSC-P-031 customer segmentation. MIT License.

Builds the declared segmentation, then submits it to four checks declared before
the file is read: separation, balance, stability under an independent rebuild,
and usefulness against an outcome that never enters the clustering. Nothing is
random here: the initial centroids are the observations at declared ranks of a
declared composite score, so every implementation reproduces the same digits.
The verdict is fail-closed: it authorizes acting on the segments only when the
four checks pass.
"""
import csv
import math
import sys
from pathlib import Path

REQUIRED = ["customer_id", "recency_days", "frequency_12m", "avg_basket_eur", "category_breadth", "digital_share", "next_quarter_revenue_eur"]
FEATURES = ["log_recency", "frequency", "log_basket", "breadth", "digital"]
SEGMENTS = 4
MAX_ITERATIONS = 100
SILHOUETTE_STEP = 3
SILHOUETTE_MIN = 0.25
SMALLEST_SHARE_MIN = 0.05
STABILITY_MIN = 0.60
USEFULNESS_MIN = 1.50


def load(path: Path):
    with path.open(newline="", encoding="utf-8") as handle:
        reader = csv.DictReader(handle)
        if reader.fieldnames != REQUIRED:
            raise ValueError("exact schema required")
        rows = list(reader)
    if len(rows) < 100:
        raise ValueError("too few customers for the declared segmentation")
    parsed = []
    for index, row in enumerate(rows, start=1):
        if set(row) != set(REQUIRED) or any(row[key] is None for key in REQUIRED):
            raise ValueError("exactly seven cells required per row")
        if row["customer_id"].strip() != "S%04d" % index:
            raise ValueError("customers must be ordered S0001, S0002, ... without gaps")
        values = {}
        for key in ["recency_days", "frequency_12m", "avg_basket_eur", "category_breadth"]:
            value = float(row[key])
            if not math.isfinite(value) or value <= 0:
                raise ValueError("recency, frequency, basket and breadth must be finite and strictly positive")
            values[key] = value
        digital = float(row["digital_share"])
        if not math.isfinite(digital) or digital <= 0.0 or digital >= 1.0:
            raise ValueError("the digital share must lie strictly between zero and one")
        revenue = float(row["next_quarter_revenue_eur"])
        if not math.isfinite(revenue) or revenue < 0:
            raise ValueError("the outcome must be finite and non-negative")
        parsed.append({
            "id": row["customer_id"].strip(),
            "log_recency": math.log(values["recency_days"]),
            "frequency": values["frequency_12m"],
            "log_basket": math.log(values["avg_basket_eur"]),
            "breadth": values["category_breadth"],
            "digital": digital,
            "revenue": revenue,
        })
    return parsed


def standardize(rows):
    """Declared standardization: mean zero, unit standard deviation, on the whole file."""
    scaled = []
    statistics = {}
    for name in FEATURES:
        values = [row[name] for row in rows]
        mean = sum(values) / len(values)
        variance = sum((value - mean) ** 2 for value in values) / (len(values) - 1)
        deviation = math.sqrt(variance)
        if deviation <= 0:
            raise ValueError("a declared feature has no variation")
        statistics[name] = (mean, deviation)
    for row in rows:
        scaled.append([(row[name] - statistics[name][0]) / statistics[name][1] for name in FEATURES])
    return scaled, statistics


def distance(left, right):
    return sum((a - b) ** 2 for a, b in zip(left, right))


def initial_centroids(points):
    """Declared start: the observations at four fixed ranks of the composite score."""
    ordered = sorted(range(len(points)), key=lambda index: (sum(points[index]), index))
    positions = [int(round(share * (len(ordered) - 1))) for share in (0.125, 0.375, 0.625, 0.875)]
    return [list(points[ordered[position]]) for position in positions]


def cluster(points):
    centroids = initial_centroids(points)
    labels = [0] * len(points)
    for _ in range(MAX_ITERATIONS):
        changed = False
        for index, point in enumerate(points):
            best = min(range(SEGMENTS), key=lambda segment: (distance(point, centroids[segment]), segment))
            if best != labels[index]:
                labels[index] = best
                changed = True
        for segment in range(SEGMENTS):
            members = [point for point, label in zip(points, labels) if label == segment]
            if not members:
                raise ValueError("the declared segmentation collapsed to fewer segments")
            centroids[segment] = [sum(values) / len(members) for values in zip(*members)]
        if not changed:
            return labels, centroids
    return labels, centroids


def assign(points, centroids):
    return [min(range(len(centroids)), key=lambda segment: (distance(point, centroids[segment]), segment)) for point in points]


def silhouette(points, labels):
    """Mean silhouette on a declared systematic sample, every third customer."""
    sample = list(range(0, len(points), SILHOUETTE_STEP))
    total = 0.0
    for index in sample:
        own = labels[index]
        sums = [0.0] * SEGMENTS
        counts = [0] * SEGMENTS
        for other, point in enumerate(points):
            if other == index:
                continue
            sums[labels[other]] += math.sqrt(distance(points[index], point))
            counts[labels[other]] += 1
        inside = sums[own] / counts[own] if counts[own] else 0.0
        outside = min(sums[segment] / counts[segment] for segment in range(SEGMENTS) if segment != own and counts[segment])
        total += (outside - inside) / max(inside, outside) if max(inside, outside) > 0 else 0.0
    return total / len(sample)


def adjusted_rand(first, second):
    pairs = {}
    for a, b in zip(first, second):
        pairs[(a, b)] = pairs.get((a, b), 0) + 1
    rows = {}
    columns = {}
    for (a, b), count in pairs.items():
        rows[a] = rows.get(a, 0) + count
        columns[b] = columns.get(b, 0) + count
    total = len(first)
    choose = lambda value: value * (value - 1) / 2.0
    index = sum(choose(count) for count in pairs.values())
    row_sum = sum(choose(count) for count in rows.values())
    column_sum = sum(choose(count) for count in columns.values())
    expected = row_sum * column_sum / choose(total)
    maximum = (row_sum + column_sum) / 2.0
    return (index - expected) / (maximum - expected) if maximum != expected else 0.0


def analyze(rows):
    points, statistics = standardize(rows)
    labels, centroids = cluster(points)
    sizes = [labels.count(segment) for segment in range(SEGMENTS)]
    shares = [size / len(labels) for size in sizes]

    separation = silhouette(points, labels)

    odd_points = [point for index, point in enumerate(points) if index % 2 == 0]
    even_points = [point for index, point in enumerate(points) if index % 2 == 1]
    _, odd_centroids = cluster(odd_points)
    _, even_centroids = cluster(even_points)
    stability = adjusted_rand(assign(points, odd_centroids), assign(points, even_centroids))

    revenues = [row["revenue"] for row in rows]
    means = []
    for segment in range(SEGMENTS):
        members = [revenue for revenue, label in zip(revenues, labels) if label == segment]
        means.append(sum(members) / len(members))
    ratio = max(means) / min(means) if min(means) > 0 else float("inf")
    grand = sum(revenues) / len(revenues)
    between = sum(size * (mean - grand) ** 2 for size, mean in zip(sizes, means))
    total = sum((revenue - grand) ** 2 for revenue in revenues)
    explained = between / total if total else 0.0

    flags = {
        "separation": "PASS" if separation >= SILHOUETTE_MIN else "FAIL",
        "balance": "PASS" if min(shares) >= SMALLEST_SHARE_MIN else "FAIL",
        "stability": "PASS" if stability >= STABILITY_MIN else "FAIL",
        "usefulness": "PASS" if ratio >= USEFULNESS_MIN else "FAIL",
    }
    verdict = "SEGMENTATION_READABLE_FOR_ACTION" if all(flag == "PASS" for flag in flags.values()) else "DIAGNOSTIC_BLOCKS_SEGMENTATION_READING"
    return {
        "n": len(rows), "sizes": sizes, "shares": shares, "separation": separation,
        "stability": stability, "means": means, "ratio": ratio, "explained": explained,
        "centroids": centroids, "statistics": statistics, "flags": flags, "verdict": verdict,
    }


def main():
    path = Path(sys.argv[1]) if len(sys.argv) > 1 else Path(__file__).resolve().parent.parent / "datasets" / "msc-p031-segmentation-panel.csv"
    result = analyze(load(path))
    print(f"customers={result['n']}")
    print(f"segments={SEGMENTS}")
    print("segment_sizes=" + ",".join(str(size) for size in result["sizes"]))
    print("segment_shares=" + ",".join(f"{share:.6f}" for share in result["shares"]))
    print(f"smallest_share={min(result['shares']):.6f}")
    print(f"balance_flag={result['flags']['balance']}")
    print(f"mean_silhouette={result['separation']:.6f}")
    print(f"separation_flag={result['flags']['separation']}")
    print(f"stability_adjusted_rand={result['stability']:.6f}")
    print(f"stability_flag={result['flags']['stability']}")
    print("segment_outcome_means=" + ",".join(f"{mean:.6f}" for mean in result["means"]))
    print(f"outcome_ratio_high_low={result['ratio']:.6f}")
    print(f"outcome_variance_explained={result['explained']:.6f}")
    print(f"usefulness_flag={result['flags']['usefulness']}")
    for segment, centroid in enumerate(result["centroids"]):
        print(f"centroid_{segment}=" + ",".join(f"{name}:{value:.6f}" for name, value in zip(FEATURES, centroid)))
    print(f"verdict={result['verdict']}")


if __name__ == "__main__":
    main()
