#!/usr/bin/env python3
import bisect
import itertools
import json
import math
import pathlib

def partitions(n, max_val=None):
    if max_val is None:
        max_val = n
    if n == 0:
        yield ()
        return
    for first in range(min(n, max_val), 0, -1):
        for rest in partitions(n - first, first):
            yield (first,) + rest

def hook_prod(part):
    conj = [sum(1 for r in part if r > j) for j in range(part[0])]
    h_prod = 1
    for i, r in enumerate(part):
        for j in range(r):
            h_prod *= (r - j + conj[j] - i - 1)
    return h_prod

def lis_length(perm):
    piles = []
    for x in perm:
        idx = bisect.bisect_left(piles, x)
        if idx == len(piles):
            piles.append(x)
        else:
            piles[idx] = x
    return len(piles)

def main():
    root = pathlib.Path(__file__).resolve().parent.parent
    results_dir = root / "results"
    results_dir.mkdir(parents=True, exist_ok=True)

    # 1. Exhaustive check for n in 1..9 via patience sorting
    exhaustive_matches = 0
    total_exhaustive_perms = 0
    exhaustive_checks = {}

    for n in range(1, 10):
        fact_n = math.factorial(n)
        total_exhaustive_perms += fact_n
        patience_counts = {}
        for p in itertools.permutations(range(n)):
            l = lis_length(p)
            patience_counts[l] = patience_counts.get(l, 0) + 1

        rsk_counts = {}
        for part in partitions(n):
            hp = hook_prod(part)
            f_lam = fact_n // hp
            l = part[0]
            rsk_counts[l] = rsk_counts.get(l, 0) + f_lam * f_lam

        if patience_counts == rsk_counts:
            exhaustive_matches += 1
            exhaustive_checks[str(n)] = True
        else:
            exhaustive_checks[str(n)] = False
            raise RuntimeError(f"Mismatch at n={n}")

    # 2. RSK computation for n in 1..50
    case_count = 50
    identity_passed_count = 0
    all_means_below_bound = True
    table = {}

    for n in range(1, case_count + 1):
        fact_n = math.factorial(n)
        counts = {}
        for part in partitions(n):
            hp = hook_prod(part)
            f_lam = fact_n // hp
            l = part[0]
            counts[l] = counts.get(l, 0) + f_lam * f_lam

        if sum(counts.values()) == fact_n:
            identity_passed_count += 1
        else:
            raise RuntimeError(f"RSK sum identity failed at n={n}")

        mean = sum(l * c for l, c in counts.items()) / fact_n
        var = sum(((l - mean) ** 2) * c for l, c in counts.items()) / fact_n
        bound = 2.0 * math.sqrt(n)
        ratio = mean / bound
        if ratio >= 1.0:
            all_means_below_bound = False

        scaled_mean = (mean - bound) / (n ** (1.0 / 6.0))
        scaled_var = var / (n ** (1.0 / 3.0))

        table[str(n)] = {
            "n": n,
            "mean": round(mean, 6),
            "variance": round(var, 6),
            "ratio_to_two_sqrt_n": round(ratio, 6),
            "scaled_mean": round(scaled_mean, 6),
            "scaled_variance": round(scaled_var, 6),
            "distribution": {str(k): counts[k] for k in sorted(counts.keys())}
        }

    output = {
        "case_count": case_count,
        "exhaustive_match_count": exhaustive_matches,
        "exhaustive_perm_count": total_exhaustive_perms,
        "sum_identity_passed_count": identity_passed_count,
        "all_exhaustive_matches": (exhaustive_matches == 9),
        "all_means_below_bound": all_means_below_bound,
        "expected_L_1": table["1"]["mean"],
        "expected_L_10": table["10"]["mean"],
        "expected_L_25": table["25"]["mean"],
        "expected_L_50": table["50"]["mean"],
        "variance_L_50": table["50"]["variance"],
        "scaled_mean_50": table["50"]["scaled_mean"],
        "scaled_var_50": table["50"]["scaled_variance"],
        "max_ratio_to_two_sqrt_n": max(t["ratio_to_two_sqrt_n"] for t in table.values()),
        "table": table
    }

    out_file = results_dir / "R1.json"
    out_file.write_text(json.dumps(output, indent=2))
    print(f"Results written to {out_file}")

if __name__ == "__main__":
    main()
