#!/usr/bin/env python3
"""Exact whole-set Fieller reader, Python standard library only.

Default: python fieller-confidence-set-reader.py
Custom:  python fieller-confidence-set-reader.py --input request.json

Example request (A then B is the matrix coordinate order):
{"Ahat": 3, "Bhat": 1, "covariance": [[1,0],[0,1]], "q": 2,
 "model": "known_covariance_gaussian", "covariance_scale": "estimator",
 "population_denominator_nonzero": true, "candidates": [-3,0,1]}

Numbers must be integers or exact decimal/fraction strings, never JSON floats.
q is the critical value, not the confidence level; coverage is 2 Phi(q)-1.
At q=2 it is approximately 0.9544997361036416, not exactly 0.95.
Model declarations are assumptions, not conditions certified from the data.
Endpoints retain exact quadratic radicals; decimal values are display only.
Null interval bounds mean unbounded rays, never omitted plotting-window tails.
The generic quadratic solver can return empty/singleton sets, but valid Fieller
inputs here cannot. We check this invariant without Python assert statements.
"""
from __future__ import annotations
import argparse
from decimal import Decimal, localcontext
from fractions import Fraction as F
import json
import math
from pathlib import Path
import re
import sys

MAX_NUMERIC_CHARS = 256
_NUM = re.compile(r"[+-]?(?:[0-9]+(?:\.[0-9]*)?|\.[0-9]+)(?:/[+-]?[0-9]+)?\Z")


def rational(value, name="number"):
    if isinstance(value, bool) or not isinstance(value, (int, str)):
        raise ValueError(f"{name}: use an integer or exact decimal/fraction string")
    text = str(value).strip()
    if len(text) > MAX_NUMERIC_CHARS or not _NUM.fullmatch(text):
        raise ValueError(f"{name}: invalid or overlong rational; exponent notation is not accepted")
    if "/" in text and "." in text:
        raise ValueError(f"{name}: write fractions as integer/integer")
    try:
        return F(text)
    except (ValueError, ZeroDivisionError) as exc:
        raise ValueError(f"{name}: invalid rational") from exc


def sign(x):
    return (x > 0) - (x < 0)


def sign_surd(base, coefficient, radicand):
    """Exact sign of base + coefficient * sqrt(nonnegative rational)."""
    if radicand < 0:
        raise ValueError("negative radicand")
    if coefficient == 0 or radicand == 0:
        return sign(base)
    if base == 0:
        return sign(coefficient)
    if sign(base) == sign(coefficient):
        return sign(base)
    delta = base * base - coefficient * coefficient * radicand
    return sign(base) * sign(delta)


def endpoint(base, coefficient=F(0), radicand=F(0)):
    base, coefficient, radicand = F(base), F(coefficient), F(radicand)
    if radicand < 0:
        raise ValueError("negative radicand")
    n, d = math.isqrt(radicand.numerator), math.isqrt(radicand.denominator)
    if n*n == radicand.numerator and d*d == radicand.denominator:
        base += coefficient * F(n, d)
        coefficient, radicand = F(0), F(0)
    if coefficient == 0:
        radicand = F(0)
    return {"base": base, "coefficient": coefficient, "radicand": radicand}


def endpoint_compare(e, value):
    return sign_surd(e["base"] - F(value), e["coefficient"], e["radicand"])


def interval(lower=None, upper=None):
    return {"lower": lower, "upper": upper,
            "lower_included": lower is not None, "upper_included": upper is not None}


def solve_quadratic(c2, c1, c0):
    """Solve <= 0 for exact Fraction coefficients, including every equality case."""
    c2, c1, c0 = F(c2), F(c1), F(c0)
    disc = c1*c1 - 4*c2*c0
    if c2 == 0:
        if c1 == 0:
            branch = "constant_all" if c0 <= 0 else "constant_empty"
            pieces = [interval()] if c0 <= 0 else []
        elif c1 > 0:
            branch, pieces = "linear_left", [interval(upper=endpoint(-c0/c1))]
        else:
            branch, pieces = "linear_right", [interval(lower=endpoint(-c0/c1))]
    elif disc > 0:
        center, radius_coefficient = -c1/(2*c2), 1/(2*abs(c2))
        lo = endpoint(center, -radius_coefficient, disc)
        hi = endpoint(center, radius_coefficient, disc)
        if c2 > 0:
            branch, pieces = "quadratic_bounded", [interval(lo, hi)]
        else:
            branch, pieces = "quadratic_two_rays", [interval(upper=lo), interval(lower=hi)]
    elif c2 > 0 and disc == 0:
        point = endpoint(-c1/(2*c2))
        branch, pieces = "quadratic_singleton", [interval(point, point)]
    elif c2 > 0:
        branch, pieces = "quadratic_empty", []
    else:
        branch = "quadratic_all_tangent" if disc == 0 else "quadratic_all_strict"
        pieces = [interval()]
    return {"coefficients": [c2, c1, c0], "discriminant": disc,
            "branch": branch, "pieces": pieces}


def contains(solution, value):
    value = F(value)
    return any((p["lower"] is None or endpoint_compare(p["lower"], value) <= 0)
               and (p["upper"] is None or endpoint_compare(p["upper"], value) >= 0)
               for p in solution["pieces"])


def original_residual(A, B, aa, ab, bb, q, value):
    value = F(value)
    return (A-value*B)**2 - q*q*(aa-2*value*ab+value*value*bb)


def endpoint_original_residual(A, B, aa, ab, bb, q, e):
    """Independent substitution into original squared inequality, as x+y sqrt(d)."""
    h, k, d = e["base"], e["coefficient"], e["radicand"]
    numerator_base, numerator_radical = A-B*h, -B*k
    left_base = numerator_base**2 + numerator_radical**2*d
    left_radical = 2*numerator_base*numerator_radical
    right_base = q*q*(aa-2*ab*h+bb*(h*h+k*k*d))
    right_radical = q*q*(-2*ab*k+2*bb*h*k)
    return left_base-right_base, left_radical-right_radical


def require(condition, message):
    if not condition:
        raise AssertionError(message)


def decimal_display(e):
    # Decimal is a display approximation. It never determines branch or membership.
    with localcontext() as ctx:
        ctx.prec = 60
        convert = lambda x: Decimal(x.numerator)/Decimal(x.denominator)
        b, k, d = e["base"], e["coefficient"], e["radicand"]
        radical = convert(k)*convert(d).sqrt()
        if b*k < 0:
            # Rationalize before rounding: the denominator now adds same-sign terms.
            # Compute the potentially tiny numerator exactly, rather than subtracting decimals.
            numerator = b*b-k*k*d
            value = convert(numerator)/(convert(b)-radical)
        else:
            value = convert(b)+radical
        return format(value, ".20g")


def endpoint_json(e):
    if e is None:
        return None
    b, k, d = e["base"], e["coefficient"], e["radicand"]
    exact = str(b) if k == 0 else f"({b}) + ({k}) * sqrt({d})"
    return {"base": str(b), "sqrt_coefficient": str(k), "radicand": str(d),
            "exact": exact, "decimal_display_only": decimal_display(e)}


def solution_json(solution):
    return {"coefficients_c2_c1_c0": [str(x) for x in solution["coefficients"]],
            "discriminant": str(solution["discriminant"]), "branch": solution["branch"],
            "pieces": [{**p, "lower": endpoint_json(p["lower"]), "upper": endpoint_json(p["upper"])}
                       for p in solution["pieces"]],
            "unbounded_bound_convention": "null means no finite bound; infinity is not a real endpoint"}


def fieller(request):
    if not isinstance(request, dict):
        raise ValueError("request must be an object")
    required = {"Ahat", "Bhat", "covariance", "q", "model", "covariance_scale",
                "population_denominator_nonzero"}
    optional = {"candidates", "population_denominator"}
    if set(request) - required - optional or required - set(request):
        raise ValueError("missing or unknown input fields")
    if request["model"] != "known_covariance_gaussian":
        raise ValueError("this reader only calibrates the declared known-covariance Gaussian model")
    if request["covariance_scale"] != "estimator":
        raise ValueError("covariance must be for estimators; raw-observation covariance needs rescaling")
    if request["population_denominator_nonzero"] is not True:
        raise ValueError("the ratio target requires a nonzero population denominator")
    if "population_denominator" in request and rational(request["population_denominator"], "population_denominator") == 0:
        raise ValueError("undefined target: population denominator is zero")
    A, B, q = (rational(request[k], k) for k in ("Ahat", "Bhat", "q"))
    if q <= 0:
        raise ValueError("q must be positive")
    cov = request["covariance"]
    if not isinstance(cov, list) or len(cov) != 2 or any(not isinstance(row, list) or len(row) != 2 for row in cov):
        raise ValueError("covariance must be a 2 by 2 matrix in A,B order")
    aa, ab, ba, bb = (rational(cov[i][j], f"covariance[{i}][{j}]") for i,j in ((0,0),(0,1),(1,0),(1,1)))
    if ab != ba or aa <= 0 or bb <= 0 or aa*bb-ab*ab <= 0:
        raise ValueError("covariance must be exactly symmetric positive definite")
    candidates = request.get("candidates", [])
    if not isinstance(candidates, list) or len(candidates) > 10000:
        raise ValueError("candidates must be a list with at most 10000 values")
    candidates = [rational(x, "candidate") for x in candidates]
    s = solve_quadratic(B*B-q*q*bb, -2*A*B+2*q*q*ab, A*A-q*q*aa)
    require(s["branch"] not in {"quadratic_singleton", "quadratic_empty", "constant_empty"},
            "internal error: valid Fieller model cannot produce an empty or singleton set")
    endpoint_checks = []
    for piece in s["pieces"]:
        for which in ("lower", "upper"):
            e = piece[which]
            if e is not None:
                rb, rk = endpoint_original_residual(A, B, aa, ab, bb, q, e)
                require(sign_surd(rb, rk, e["radicand"]) == 0, "endpoint does not satisfy original equality")
                endpoint_checks.append({"bound": which, "original_residual_base": str(rb),
                                        "original_residual_sqrt_coefficient": str(rk), "equal_zero": True})
    checks = []
    for r in candidates:
        residual = original_residual(A,B,aa,ab,bb,q,r)
        retained = contains(s,r)
        require(retained == (residual <= 0), "interval representation disagrees with original squared inequality")
        checks.append({"r": str(r), "original_squared_residual": str(residual), "retained": retained})
    coverage = math.erf(float(q)/math.sqrt(2))
    return {"calibration": {"model": "exact bivariate Gaussian with known positive-definite estimator covariance",
                            "q_exact": str(q), "coverage_formula": "2 Phi(q) - 1", "coverage_approximation": coverage,
                            "coverage_note": "numeric approximation may round; q=2 is not exactly 95%",
                            "model_conditions_verified_from_input": False,
                            "population_denominator_nonzero_is_assumed": True},
            "covariance_determinant_exact": str(aa*bb-ab*ab), "set": solution_json(s),
            "endpoint_original_inequality_checks": endpoint_checks, "candidate_checks": checks}


def example_request(A, B, ab=0, q=2):
    return {"Ahat": A, "Bhat": B, "covariance": [[1,ab],[ab,1]], "q": q,
            "model": "known_covariance_gaussian", "covariance_scale": "estimator",
            "population_denominator_nonzero": True, "candidates": [-3,0,"5/12",1,"4/3",3]}


def self_test():
    count = 0
    def check(condition, label):
        nonlocal count
        require(condition, label)
        count += 1
    # Every generic algebraic branch, including equality, independently tested on rationals.
    branches = [(1,0,-1,"quadratic_bounded"), (1,0,0,"quadratic_singleton"),
                (1,0,1,"quadratic_empty"), (-1,0,1,"quadratic_two_rays"),
                (-1,0,0,"quadratic_all_tangent"), (-1,0,-1,"quadratic_all_strict"),
                (0,2,-1,"linear_left"), (0,-2,1,"linear_right"),
                (0,0,0,"constant_all"), (0,0,-1,"constant_all"), (0,0,1,"constant_empty")]
    for a,b,c,branch in branches:
        s = solve_quadratic(a,b,c)
        check(s["branch"] == branch, branch)
        for r in map(F, ["-3","-1","-1/2","0","1/2","1","3"]):
            check(contains(s,r) == (a*r*r+b*r+c <= 0), f"{branch}: membership {r}")
    examples = [example_request(2,4),example_request(3,1),example_request(1,1),
                example_request(3,2),example_request(2,4,"1/2"),example_request(3,-2),
                example_request(0,2),example_request(3,0)]
    expected = ["quadratic_bounded","quadratic_two_rays","quadratic_all_strict","linear_right",
                "quadratic_bounded","linear_left","constant_all","quadratic_two_rays"]
    results = [fieller(r) for r in examples]
    for result,branch in zip(results,expected):
        check(result["set"]["branch"] == branch, "Fieller branch "+branch)
    check(results[0]["set"]["pieces"][0]["upper"]["exact"] == "4/3", "bounded upper")
    check(results[4]["set"]["pieces"][0]["upper"]["exact"] == "1", "correlated upper")
    check(results[5]["set"]["pieces"][0]["upper"]["exact"] == "-5/12", "negative half-line")
    check(abs(results[0]["calibration"]["coverage_approximation"] - 0.9544997361036416) < 1e-15, "q2 calibration")
    # Exact tiny differences cannot be rounded to the linear case.
    for b,branch in [("2.000000000000000000000000000001","quadratic_bounded"),
                     ("1.999999999999999999999999999999","quadratic_two_rays")]:
        check(fieller(example_request(3,b))["set"]["branch"] == branch, "near-zero leading coefficient")
    # A root near 5/12 must stay readable even when its naive formula cancels 100 digits.
    for prefix, bound, expected_decimal in [("", "lower", Decimal(5)/Decimal(12)),
                                           ("-", "upper", -Decimal(5)/Decimal(12))]:
        req=example_request(3,prefix+"2."+"0"*99+"1")
        out=fieller(req)
        display=Decimal(out["set"]["pieces"][0][bound]["decimal_display_only"])
        check(abs(display-expected_decimal) < Decimal("1e-19"), "stable near-boundary display")
    # A deterministic rational family with arbitrary signs and correlated positive-definite covariance.
    for ai in range(-4,5):
        for bi in range(-4,5):
            for x in [F(-2),F(0),F(1,2)]:
                aa,ab,bb = 1+x*x,x,F(1)
                req=example_request(ai,bi)
                req["covariance"]=[[str(aa),str(ab)],[str(ab),str(bb)]]
                out=fieller(req)
                check(bool(out["set"]["pieces"]), "positive definite family never empty")
    # Scale all estimates by k and covariance by k^2: identical coefficients up to k^2.
    raw=fieller(example_request(3,1))
    req=example_request("3/10","1/10")
    req["covariance"]=[["1/100",0],[0,"1/100"]]
    scaled=fieller(req)
    for r in map(F, ["-10","-3","-1","0","1","10"]):
        s1=solve_quadratic(*map(F,raw["set"]["coefficients_c2_c1_c0"]))
        s2=solve_quadratic(*map(F,scaled["set"]["coefficients_c2_c1_c0"]))
        check(contains(s1,r)==contains(s2,r), "scale invariance")
    invalid=[{"q":0},{"q":-1},{"q":True},{"q":2.0},{"Ahat":"NaN"},
             {"Bhat":"1e30"},{"covariance":[[1,1],[1,1]]},
             {"covariance":[[1,2],[2,1]]},{"covariance":[[1,0],[1,1]]},
             {"covariance":[[1,0,0],[0,1,0]]},{"model":"sandwich"},
             {"covariance_scale":"raw_observations"},{"population_denominator_nonzero":False},
             {"population_denominator":0},{"candidates":[False]}, {"unknown":1}]
    for mutation in invalid:
        req=example_request(3,1);req.update(mutation)
        try:
            fieller(req)
        except ValueError:
            check(True,"rejected invalid input")
        else:
            check(False,"accepted invalid input "+repr(mutation))
    return {"checks_passed":count,"uses_assert_statements":False,
            "coverage_scope":"exact arithmetic and representation checks; no simulation substitutes for the proof"}, results


def main(argv=None):
    parser=argparse.ArgumentParser(description=__doc__,formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("--input",help="JSON request file, or - for standard input")
    args=parser.parse_args(argv)
    try:
        if args.input:
            text=sys.stdin.read(1_000_001) if args.input=="-" else Path(args.input).read_text(encoding="utf-8")
            if len(text)>1_000_000:
                raise ValueError("request is too large")
            output=fieller(json.loads(text))
        else:
            tests,examples=self_test()
            output={"self_test":tests,"examples":examples}
        print(json.dumps(output,ensure_ascii=False,indent=2,allow_nan=False))
        return 0
    except (ValueError,TypeError,OSError,OverflowError) as exc:
        print(json.dumps({"error":str(exc),"confidence_set_returned":False},ensure_ascii=False),file=sys.stderr)
        return 2


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