#!/usr/bin/env python3
"""Recompute the stochastic-geometry examples and theoretical budgets.
Standard library only. No data are fetched and no optimizer is trained.
Decimal precision agreement is a diagnostic, not an interval certificate.
The inequalities require the model assumptions stated on the linked pages.
"""
import argparse
import itertools
import json
import math
import sys
from decimal import Decimal, localcontext, ROUND_CEILING
from fractions import Fraction


def dec(value):
    result = Decimal(str(value))
    if not result.is_finite():
        raise ValueError("All numeric inputs must be finite")
    return result


def ceil(value):
    return int(value.to_integral_value(rounding=ROUND_CEILING))


def check_positive(**values):
    for key, value in values.items():
        if value <= 0:
            raise ValueError(key + " must be positive")


def weights(count):
    ts = [Decimal(1)]
    for _ in range(1, count):
        ts.append((1 + (1 + 4 * ts[-1] ** 2).sqrt()) / 2)
    return ts


def accelerated_once(args, precision):
    with localcontext() as ctx:
        ctx.prec = precision
        L, R, sigma, eps = map(dec, (args.L, args.R, args.sigma, args.epsilon))
        check_positive(L=L, epsilon=eps)
        if R < 0 or sigma < 0:
            raise ValueError("R and sigma must be nonnegative")
        if R == 0:
            return {"status": "initial_point_optimal_under_supplied_R_zero", "prox_steps": 0,
                    "oracle_queries": 0, "expected_bound": "0", "batches": []}
        T = max(1, ceil((8 * L * R ** 2 / eps).sqrt()) - 1)
        if T > args.max_prox:
            return {"status": "budget_exhausted", "reason": "prox_budget",
                    "required_prox_steps": T, "max_prox": args.max_prox,
                    "plan_finished": False}
        ts = weights(T)
        batches = [max(1, ceil(sigma ** 2 * t / (L * eps))) for t in ts]
        N = sum(batches)
        deterministic = L * R ** 2 / ts[-1] ** 2
        noise = sigma ** 2 * sum(t ** 2 / b for t, b in zip(ts, batches)) / (2 * L * ts[-1] ** 2)
        bound = deterministic + noise
        return {"status": "plan_within_budget" if N <= args.max_oracle else "budget_exhausted",
                "reason": None if N <= args.max_oracle else "oracle_budget",
                "prox_steps": T, "oracle_queries": N, "max_oracle": args.max_oracle,
                "batches": batches, "weights": [str(t) for t in ts],
                "deterministic_term": str(deterministic), "noise_term": str(noise),
                "expected_bound": str(bound), "requested_epsilon": str(eps),
                "bound_below_epsilon_at_working_precision": bound <= eps,
                "plan_finished": True}


def accelerated(args):
    lo = accelerated_once(args, args.precision)
    hi = accelerated_once(args, 2 * args.precision)
    keys = ["status", "prox_steps", "oracle_queries", "batches", "required_prox_steps"]
    agree = all(lo.get(k) == hi.get(k) for k in keys)
    if not agree:
        hi["status"] = "precision_unresolved"
    hi["precision"] = {"digits": 2 * args.precision, "integer_plan_agrees_at_half_precision": agree,
                       "rigorous_interval_certificate": False}
    hi["meaning"] = "Theoretical expectation plan conditional on valid model constants; no training run or observed-gap certificate"
    return hi


def clipped(args):
    with localcontext() as ctx:
        ctx.prec = args.precision
        p, M, D, H, delta = map(dec, (args.p, args.M, args.D, args.H, args.delta))
        B = Decimal(2).ln() if args.B == "log2" else dec(args.B)
        check_positive(M=M, B=B, D=D)
        if not (1 < p <= 2) or not (0 < delta < 1) or H < 0:
            raise ValueError("Need 1<p<=2, 0<delta<1 and H>=0")
        if args.iterations < 1:
            raise ValueError("iterations must be positive")
        if args.iterations > args.max_oracle:
            return {"status": "budget_exhausted", "required_oracle_queries": args.iterations,
                    "max_oracle": args.max_oracle, "reason": "fixed_run_not_within_budget"}
        T = Decimal(args.iterations)
        ell = (2 / delta).ln()
        tau = M * ((T / ell).ln() / p).exp()
        v = (p * M.ln() + (2 - p) * tau.ln()).exp()
        S = T * v + (2 * T * tau ** 2 * v * ell).sqrt() + 2 * tau ** 2 * ell / 3
        eta = (2 * B / S).sqrt()
        terms = {"initial_potential": B / (eta * T), "clipped_square_sum": eta * S / (2 * T),
                 "composite_endpoint": H / T, "clipping_bias": D * (p * M.ln() + (1 - p) * tau.ln()).exp(),
                 "direction_variance": D * (2 * v * ell / T).sqrt(), "direction_amplitude": 4 * D * tau * ell / (3 * T)}
        return {"status": "fixed_budget_bound_evaluated", "iterations": args.iterations,
                "failure_probability": str(delta), "ell": str(ell), "tau": str(tau), "v": str(v),
                "S": str(S), "eta": str(eta), "terms": {k: str(v) for k, v in terms.items()},
                "high_probability_bound": str(sum(terms.values())),
                "precision": {"digits": args.precision, "rigorous_interval_certificate": False},
                "meaning": "Arithmetic evaluation of a proved fixed-budget theorem; raw conditional p-moment and geometry assumptions not checked from data"}


def examples():
    Q = Fraction
    probabilities = [[Q(1, 3)] * 3, [Q(1, 10), Q(1, 10), Q(4, 5)], [Q(5, 18), Q(5, 18), Q(4, 9)]]
    G, costs = [1, 1, 8], [1, 1, 25]
    sampler = []
    for ps in probabilities:
        moment = sum(Q(g * g, 9) / p for g, p in zip(G, ps))
        cost = sum(c * p for c, p in zip(costs, ps))
        sampler.append({"p": list(map(str, ps)), "M": str(moment), "expected_cost": str(cost), "product": str(moment * cost)})
    beta = ((1 + math.sqrt(5)) / 2 - 1) / ((1 + math.sqrt(1 + 4 * ((1 + math.sqrt(5)) / 2) ** 2)) / 2)
    def soft(z):
        return math.copysign(max(abs(z) - .25, 0), z)
    vals = []
    for signs in itertools.product([-1, 1], repeat=6):
        means = [signs[0], sum(signs[1:3]) / 2, sum(signs[3:6]) / 3]
        x1 = soft((1 - means[0]) / 2)
        x2 = soft((x1 + 1 - means[1]) / 2)
        y3 = x2 + beta * (x2 - x1)
        x3 = soft((y3 + 1 - means[2]) / 2)
        vals.append((x3 - 1) ** 2 / 2 + abs(x3) / 2 - 3 / 8)
    t2 = (1 + math.sqrt(5)) / 2
    t3 = (1 + math.sqrt(1 + 4 * t2 * t2)) / 2
    def entropy_update(x, v):
        raw = [math.sqrt(xi) * math.exp(-vi / 2) for xi, vi in zip(x, v)]
        return [r / sum(raw) for r in raw]
    mirror_x0 = [.5, .5]
    mirror_x1 = entropy_update(mirror_x0, [math.log(9), 0])
    mirror_x2 = entropy_update(mirror_x1, [0, 0])
    mirror_avg = [(a + b) / 2 for a, b in zip(mirror_x0, mirror_x1)]
    expected = [.25, .75, .375, .625]
    errors = [abs(a - b) for a, b in zip(mirror_x1 + mirror_avg, expected)]
    if max(errors) > 1e-14:
        raise ArithmeticError("Entropy update does not match the stated rational example")
    path_x3 = 19 / 48 - beta / 16
    ts8 = [1.0]
    for _ in range(7):
        ts8.append((1 + math.sqrt(1 + 4 * ts8[-1] ** 2)) / 2)
    b8 = [10,17,22,28,33,39,44,49]
    common_noise = sum(t * t * (.5 + .5 / b) for t, b in zip(ts8, b8)) / (2 * ts8[-1] ** 2)
    return {"status": "examples_recomputed", "sampling_exact_fractions": sampler,
            "mirror": {"x1": mirror_x1, "x2": mirror_x2, "x2_first": mirror_x2[0],
                       "pre_update_average": mirror_avg, "expected_rationals": {"x1": ["1/4", "3/4"], "average": ["3/8", "5/8"]},
                       "maximum_update_check_error": max(errors), "update_tolerance": 1e-14,
                       "objective_gap": 3 / 8 * math.log(1.5) + 5 / 8 * math.log(5 / 6)},
            "accelerated": {"symbol_paths": 64, "batch_sizes": [1, 2, 3], "oracle_queries": 6,
                            "specified_path_x3": path_x3, "specified_path_gap": (path_x3 - .5) ** 2 / 2,
                            "enumerated_expected_gap": sum(vals) / 64,
                            "theorem_bound_R_half": .25 / t3 ** 2 + (1 + t2 ** 2 / 2 + t3 ** 2 / 3) / (2 * t3 ** 2)},
            "common_shock_transfer": {"a_squared": .5, "s_squared": .5, "noise_term": common_noise, "total_bound": common_noise + 1 / ts8[-1] ** 2},
            "precision": "Sampling table is exact rational arithmetic; logarithm, square-root and enumeration values use binary64 and are not interval-certified"}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    sub = parser.add_subparsers(dest="mode", required=True)
    sub.add_parser("examples")
    a = sub.add_parser("accelerated")
    for key, default in [("L", "1"), ("R", "1"), ("sigma", "1"), ("epsilon", "0.1")]:
        a.add_argument("--" + key, default=default)
    a.add_argument("--max-prox", type=int, default=100000)
    a.add_argument("--max-oracle", type=int, default=10000000)
    a.add_argument("--precision", type=int, default=60)
    c = sub.add_parser("clipped")
    for key, default in [("p", "1.5"), ("M", "1"), ("B", "log2"), ("D", "2"), ("H", "0"), ("delta", "0.05")]:
        c.add_argument("--" + key, default=default)
    c.add_argument("--iterations", type=int, default=1000000)
    c.add_argument("--max-oracle", type=int, default=1000000)
    c.add_argument("--precision", type=int, default=60)
    args = parser.parse_args()
    try:
        if hasattr(args, "precision") and not 30 <= args.precision <= 200:
            raise ValueError("precision must be between 30 and 200 decimal digits")
        if hasattr(args, "max_oracle") and args.max_oracle < 0:
            raise ValueError("max-oracle must be nonnegative")
        if hasattr(args, "max_prox") and args.max_prox < 0:
            raise ValueError("max-prox must be nonnegative")
        result = {"examples": examples, "accelerated": lambda: accelerated(args), "clipped": lambda: clipped(args)}[args.mode]()
    except (ValueError, ArithmeticError) as err:
        result = {"status": "invalid_input_or_arithmetic_failure", "reason": str(err)}
    print(json.dumps(result, ensure_ascii=False, indent=2))
    if result["status"] == "budget_exhausted":
        return 1
    if result["status"] in {"precision_unresolved", "invalid_input_or_arithmetic_failure"}:
        return 2
    return 0


if __name__ == "__main__":
    sys.exit(main())
