#!/usr/bin/env python3
"""MSC-P-032 segmentation stability. MIT License.

Asks whether a proposed segmentation survives resampling. For each declared
number of segments, the clustering is rebuilt on bootstrap resamples drawn with
replacement, every original customer is relabelled under each rebuild, and the
recovery of each original segment is measured by the Jaccard index. The same
procedure is run on a declared structureless reference obtained by permuting
each feature independently, so that a recovery level is read against what pure
resampling noise already produces.

Nothing is drawn from a library generator: the resamples come from a declared
linear congruential recursion, so every implementation reproduces the same
digits without any seed convention.
"""
import csv
import math
import sys
from pathlib import Path

REQUIRED = ["customer_id", "recency_days", "frequency_12m", "avg_basket_eur"]
FEATURES = ["log_recency", "frequency", "log_basket"]
CANDIDATE_SEGMENTS = [2, 3, 4]
PROPOSED_SEGMENTS = 4
RESAMPLES = 40
MAX_ITERATIONS = 100
RECOVERY_MIN = 0.75
MARGIN_MIN = 0.10
LCG_SEED = 20260912
LCG_MULTIPLIER = 1103515245
LCG_INCREMENT = 12345
LCG_MODULUS = 2147483648


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 stability analysis")
    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 four cells required per row")
        if row["customer_id"].strip() != "T%04d" % index:
            raise ValueError("customers must be ordered T0001, T0002, ... without gaps")
        values = {}
        for key in ["recency_days", "frequency_12m", "avg_basket_eur"]:
            value = float(row[key])
            if not math.isfinite(value) or value <= 0:
                raise ValueError("recency, frequency and basket must be finite and strictly positive")
            values[key] = value
        parsed.append([
            math.log(values["recency_days"]),
            values["frequency_12m"],
            math.log(values["avg_basket_eur"]),
        ])
    return parsed


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


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


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


def cluster(points, segments):
    centroids = initial_centroids(points, segments)
    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]


class Declared:
    """Declared linear congruential recursion; the only source of resampling."""

    def __init__(self, seed):
        self.state = seed

    def next_index(self, size):
        self.state = (LCG_MULTIPLIER * self.state + LCG_INCREMENT) % LCG_MODULUS
        return self.state % size


def jaccard(left, right):
    intersection = len(left & right)
    union = len(left | right)
    return intersection / union if union else 0.0


def recovery(points, segments):
    """Mean Jaccard recovery of each original segment over the declared resamples."""
    labels, _ = cluster(points, segments)
    original = [{index for index, label in enumerate(labels) if label == segment} for segment in range(segments)]
    totals = [0.0] * segments
    stream = Declared(LCG_SEED)
    completed = 0
    for _ in range(RESAMPLES):
        sample = [points[stream.next_index(len(points))] for _ in range(len(points))]
        try:
            _, centroids = cluster(sample, segments)
        except ValueError:
            continue
        rebuilt_labels = assign(points, centroids)
        rebuilt = [{index for index, label in enumerate(rebuilt_labels) if label == segment} for segment in range(segments)]
        for segment in range(segments):
            totals[segment] += max(jaccard(original[segment], candidate) for candidate in rebuilt)
        completed += 1
    if completed < RESAMPLES:
        raise ValueError("a declared resample failed to produce the requested number of segments")
    return [total / completed for total in totals]


def scrambled(points):
    """Declared structureless reference: each feature permuted independently."""
    stream = Declared(LCG_SEED)
    columns = []
    for column in zip(*points):
        values = list(column)
        for index in range(len(values) - 1, 0, -1):
            swap = stream.next_index(index + 1)
            values[index], values[swap] = values[swap], values[index]
        columns.append(values)
    return [list(row) for row in zip(*columns)]


def analyze(rows):
    points = standardize(rows)
    reference = scrambled(points)
    results = []
    for segments in CANDIDATE_SEGMENTS:
        observed = recovery(points, segments)
        null = recovery(reference, segments)
        weakest = min(observed)
        null_weakest = min(null)
        margin = weakest - null_weakest
        stable = weakest >= RECOVERY_MIN and margin >= MARGIN_MIN
        results.append({
            "segments": segments, "observed": observed, "weakest": weakest,
            "null_weakest": null_weakest, "margin": margin,
            "flag": "STABLE" if stable else "NOT_STABLE",
        })
    proposed = next(result for result in results if result["segments"] == PROPOSED_SEGMENTS)
    verdict = "PROPOSED_SEGMENTATION_STABLE" if proposed["flag"] == "STABLE" else "PROPOSED_SEGMENTATION_NOT_STABLE"
    return {"n": len(rows), "results": results, "proposed": proposed, "verdict": verdict}


def main():
    path = Path(sys.argv[1]) if len(sys.argv) > 1 else Path(__file__).resolve().parent.parent / "datasets" / "msc-p032-stability-panel.csv"
    result = analyze(load(path))
    print(f"customers={result['n']}")
    print(f"resamples={RESAMPLES}")
    print(f"proposed_segments={PROPOSED_SEGMENTS}")
    for entry in result["results"]:
        print(f"recovery_k{entry['segments']}=" + ",".join(f"{value:.6f}" for value in entry["observed"]))
        print(f"weakest_k{entry['segments']}={entry['weakest']:.6f}")
        print(f"null_weakest_k{entry['segments']}={entry['null_weakest']:.6f}")
        print(f"margin_k{entry['segments']}={entry['margin']:.6f}")
        print(f"flag_k{entry['segments']}={entry['flag']}")
    print(f"verdict={result['verdict']}")


if __name__ == "__main__":
    main()
