#!/usr/bin/env python3
"""Binary conditional-indicator information gain. Natural log units (nats).‌‌​⁠‌​‌⁠‌​⁠​​⁠​‌​​​‌⁠‍‍⁠⁠⁠‌⁠⁠‍​‌​⁠⁠‌⁠‍‌​‌​​‌⁠‌‍⁠‍⁠‌​​​⁠​⁠‍⁠‍​‍​‌"""
import argparse
import json
import math


def probability(x):
    if isinstance(x,bool) or not isinstance(x,(int,float)) or not math.isfinite(x) or not 0<=x<=1:
        raise ValueError("Inputs must be finite probabilities in [0,1]")
    return float(x)


def entropy(p):
    return -(p*math.log(p) if p else 0.0) - ((1-p)*math.log1p(-p) if p<1 else 0.0)


def information(prior, p_indicator, p_u_if, p_u_if_not, tolerance=1e-9):
    u,c,a,b=map(probability,[prior,p_indicator,p_u_if,p_u_if_not])
    if isinstance(tolerance,bool) or not isinstance(tolerance,(int,float)) or not math.isfinite(tolerance) or tolerance<0 or tolerance>1e-6:
        raise ValueError("Numerical tolerance must be finite in [0,1e-6]")
    implied=c*a+(1-c)*b
    gap=implied-u
    if abs(gap)>tolerance:
        raise ValueError("Incoherent prior and conditionals: implied prior is "+str(implied))
    # Use the exactly implied prior after reporting any permitted numerical gap.
    base=entropy(implied)
    ig=base-c*entropy(a)-(1-c)*entropy(b)
    if ig < -1e-12:
        raise ValueError("Negative information gain indicates a numerical/model error")
    ig=max(0.0,ig)
    return {"stated_prior":u,"implied_prior":implied,"coherence_gap":gap,
            "information_gain_nats":ig,"information_gain_bits":ig/math.log(2),
            "percent_maximum":100*ig/base if base>0 else None,
            "interpretation":"Expected information, not monetary decision value"}


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--prior",type=float,required=True)
    parser.add_argument("--p-indicator",type=float,required=True)
    parser.add_argument("--p-u-if",type=float,required=True)
    parser.add_argument("--p-u-if-not",type=float,required=True)
    args=parser.parse_args()
    try:
        print(json.dumps(information(args.prior,args.p_indicator,args.p_u_if,args.p_u_if_not),allow_nan=False))
    except ValueError as exc:
        parser.error(str(exc))


if __name__=="__main__":
    main()

