#!/usr/bin/env python3
"""Two-step scalar LQG reader checker. Python 3 standard library only.

Exact Fraction arithmetic for costs, gains, covariances and observed controls.
The nonlinear tanh example reports a floating point evaluation, not an exact
integral or a certificate for general non-Gaussian control.

python foundations-lqg-checker.py
python foundations-lqg-checker.py --v1 4 --observations 2 1
python foundations-lqg-checker.py --v1 3/2 --observations -1/2 3

Fixed model: N=2, A=B=C=Q=R=Pf=1, x0,w0,w1,v0 variance 1,
mean 0, all five primitive random variables mutually independent Gaussian.
Only the second observation variance and the observed path may be changed.
No files, packages, network, simulation or symbolic algebra are required.
"""
import argparse
from fractions import Fraction as F
import json
from math import tanh


def check_equal(label, actual, expected):
    if actual != expected:
        raise ArithmeticError(f'{label}: {actual!r} != {expected!r}')


def add(*vectors):
    return tuple(sum(parts, F(0)) for parts in zip(*vectors))


def scale(c, vector):
    return tuple(c * x for x in vector)


def dot(a, b, variances):
    return sum((x * y * v for x, y, v in zip(a, b, variances)), F(0))


def compute(v1=F(1), observations=(F(2), F(1))):
    if not isinstance(v1, F) or v1 <= 0:
        raise ValueError('v1 must be a positive Fraction')
    if len(observations) != 2 or not all(isinstance(y, F) for y in observations):
        raise ValueError('observations must contain two Fractions')
    P = [F(0), F(0), F(1)]
    H = [F(0), F(0)]
    K = [F(0), F(0)]
    for t in (1, 0):
        H[t] = 1 + P[t + 1]
        K[t] = P[t + 1] / H[t]
        P[t] = 1 + P[t + 1] - P[t + 1] ** 2 / H[t]
    prior = F(1)
    G, Sigma, predictions = [], [], []
    for variance in (F(1), v1):
        predictions.append(prior)
        gain = prior / (prior + variance)
        posterior = prior - gain * prior
        G.append(gain)
        Sigma.append(posterior)
        prior = posterior + 1
    full_certificate = P[0] + P[1] + P[2]
    info_terms = [K[t] ** 2 * H[t] * Sigma[t] for t in (0, 1)]
    certificate = full_certificate + sum(info_terms)

    # Independent reconstruction in the primitive-noise coordinate basis.
    basis = [tuple(F(i == j) for i in range(5)) for j in range(5)]
    x0, w0, w1, v0, v_1 = basis
    variances = (F(1), F(1), F(1), F(1), v1)
    y0 = add(x0, v0)
    m0 = scale(G[0], y0)
    u0 = scale(-K[0], m0)
    x1 = add(x0, u0, w0)
    y1 = add(x1, v_1)
    prediction1 = add(m0, u0)
    m1 = add(prediction1, scale(G[1], add(y1, scale(-1, prediction1))))
    u1 = scale(-K[1], m1)
    x2 = add(x1, u1, w1)
    optimal_vectors = (x0, u0, x1, u1, x2)
    optimal_terms = [dot(z, z, variances) for z in optimal_vectors]
    direct = sum(optimal_terms)
    check_equal('primitive cost versus Riccati certificate', direct, certificate)

    # Check the error against every observable coordinate, not only the input.
    e0 = add(x0, scale(-1, m0))
    e1 = add(x1, scale(-1, m1))
    check_equal('posterior variance 0', dot(e0, e0, variances), Sigma[0])
    check_equal('posterior variance 1', dot(e1, e1, variances), Sigma[1])
    check_equal('e0 perpendicular y0', dot(e0, y0, variances), F(0))
    check_equal('e1 perpendicular y0', dot(e1, y0, variances), F(0))
    check_equal('e1 perpendicular y1', dot(e1, y1, variances), F(0))
    check_equal('e1 perpendicular known u0', dot(e1, u0, variances), F(0))

    raw_u0 = scale(-K[0], y0)
    raw_x1 = add(x0, raw_u0, w0)
    raw_y1 = add(raw_x1, v_1)
    raw_u1 = scale(-K[1], raw_y1)
    raw_x2 = add(raw_x1, raw_u1, w1)
    raw_terms = [dot(z, z, variances) for z in (x0, raw_u0, raw_x1, raw_u1, raw_x2)]
    full_u0 = scale(-K[0], x0)
    full_x1 = add(x0, full_u0, w0)
    full_u1 = scale(-K[1], full_x1)
    full_x2 = add(full_x1, full_u1, w1)
    full_direct = sum(dot(z, z, variances) for z in (x0, full_u0, full_x1, full_u1, full_x2))
    check_equal('full-state direct cost', full_direct, full_certificate)

    path_m0 = G[0] * observations[0]
    path_u0 = -K[0] * path_m0
    path_pred1 = path_m0 + path_u0
    path_m1 = path_pred1 + G[1] * (observations[1] - path_pred1)
    path_u1 = -K[1] * path_m1
    return {
        'v1': v1, 'P': P, 'H': H, 'K': K, 'G': G,
        'predicted_variances': predictions, 'posterior_variances': Sigma,
        'full_state_cost': full_direct, 'information_cost_terms': info_terms,
        'lqg_cost_certificate': certificate, 'lqg_cost_primitive': direct,
        'primitive_order': ['x0', 'w0', 'w1', 'v0', 'v1'],
        'lqg_vectors_x0_u0_x1_u1_x2': optimal_vectors,
        'lqg_cost_terms': optimal_terms, 'raw_output_cost_terms': raw_terms,
        'raw_output_cost': sum(raw_terms), 'raw_output_excess': sum(raw_terms) - direct,
        'observed_path': {'y': observations, 'm0': path_m0, 'u0': path_u0,
                          'prediction1': path_pred1, 'm1': path_m1, 'u1': path_u1},
        'orthogonality_checks': 'passed exactly',
    }


def self_test():
    base = compute()
    noisy = compute(F(4))
    for label, actual, expected in [
        ('base LQG', base['lqg_cost_primitive'], F(97, 20)),
        ('full state', base['full_state_cost'], F(41, 10)),
        ('raw output', base['raw_output_cost'], F(11, 2)),
        ('raw excess', base['raw_output_excess'], F(13, 20)),
        ('first control', base['observed_path']['u0'], F(-3, 5)),
        ('second control', base['observed_path']['u1'], F(-19, 50)),
        ('noisy LQG', noisy['lqg_cost_primitive'], F(1121, 220)),
        ('noisy gain', noisy['G'][1], F(3, 11)),
        ('noisy covariance', noisy['posterior_variances'][1], F(12, 11)),
        ('noisy second control', noisy['observed_path']['u1'], F(-31, 110)),
        ('noise price', noisy['lqg_cost_primitive'] - base['lqg_cost_primitive'], F(27, 110)),
        ('unchanged control gains', noisy['K'], base['K']),
    ]:
        check_equal(label, actual, expected)
    for variance in (F(1, 100), F(3, 2), F(7), F(100)):
        compute(variance, (F(-1, 2), F(3)))
    # Disconnected one-step model: B=C=0, A=2; Q=R=Pf=Pi=V=1; W=0.
    A, B, C = F(2), F(0), F(0)
    Q, R, Pf, Pi, V, W = F(1), F(1), F(1), F(1), F(1), F(0)
    H = R + B ** 2 * Pf
    disconnected_gain = B * Pf * A / H
    observation_gain = Pi * C / (C ** 2 * Pi + V)
    P0 = Q + A ** 2 * Pf - (A * Pf * B) ** 2 / H
    disconnected_cost = P0 * Pi + Pf * W
    check_equal('uncontrollable finite optimum', disconnected_gain, F(0))
    check_equal('unobservable measurement gain', observation_gain, F(0))
    check_equal('unobservable finite value', disconnected_cost, F(5))
    return 'passed: exact identities, default endpoints, noise migration and extra positive variances'


def serialise(value):
    if isinstance(value, F):
        return str(value)
    if isinstance(value, dict):
        return {k: serialise(v) for k, v in value.items()}
    if isinstance(value, (list, tuple)):
        return [serialise(v) for v in value]
    return value


def rational(text):
    try:
        return F(text)
    except (ValueError, ZeroDivisionError) as exc:
        raise argparse.ArgumentTypeError('use an integer, fraction or finite decimal') from exc


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--v1', type=rational, default=F(1), help='second observation variance, strictly positive')
    parser.add_argument('--observations', type=rational, nargs=2, default=(F(2), F(1)), metavar=('Y0', 'Y1'))
    args = parser.parse_args()
    if args.v1 <= 0:
        parser.error('--v1 must be strictly positive; the stated theorem uses positive observation noise')
    result = compute(args.v1, tuple(args.observations))
    result['self_test'] = self_test()
    result['non_gaussian_y_equals_1'] = {
        'conditional_mean_numeric': tanh(1.0),
        'exact_formula': 'm(y)=tanh(y); u*(y)=-tanh(y)/2; excess=(y-2*tanh(y))^2/8',
        'conditional_excess_numeric': (1.0 - 2.0 * tanh(1.0)) ** 2 / 8.0,
        'precision_scope': 'floating evaluation only; not an unconditional cost or symbolic proof',
    }
    result['scope'] = 'The numeric observed path is not the ex-ante expected cost; model fixed except V1 and readings.'
    print(json.dumps(serialise(result), ensure_ascii=False, indent=2))


if __name__ == '__main__':
    main()
