#!/usr/bin/env python3
"""Fixed-correlation GEE teaching reader; Python 3, standard library only.

Run with no arguments for the three capstone fits. Supply --input data.json
for a custom fit. Input: family ('binary'/'count'), clusters (unique id, X, y,
R; count additionally exposure), initial, tol, max_iter. X includes intercept.
R is supplied before fitting, never estimated from outcomes by this reader.
Successful numerical roots are not certificates of identification or coverage.
"""
import argparse
import json
import math
import sys


class FitError(ValueError):
    pass


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


def transpose(a):
    return [list(x) for x in zip(*a)]


def zeros(n, m):
    return [[0.0] * m for _ in range(n)]


def chol(a, label):
    n = len(a)
    if not n or any(len(row) != n for row in a):
        raise FitError(label + ': expected a square matrix')
    if any(not math.isfinite(v) for row in a for v in row):
        raise FitError(label + ': nonfinite matrix')
    scale = max(abs(v) for row in a for v in row)
    if scale == 0:
        raise FitError(label + ': singular matrix')
    if any(abs(a[i][j] - a[j][i]) > 1e-12 * scale
           for i in range(n) for j in range(n)):
        raise FitError(label + ': matrix is not symmetric')
    l = zeros(n, n)
    for i in range(n):
        for j in range(i + 1):
            v = a[i][j] - sum(l[i][k] * l[j][k] for k in range(j))
            if i == j:
                if v <= 1e-13 * scale:
                    raise FitError(label + ': nonpositive or numerically tiny pivot')
                l[i][j] = math.sqrt(v)
            else:
                l[i][j] = v / l[j][j]
    return l


def solve(l, b):
    n = len(l)
    y = [0.0] * n
    x = [0.0] * n
    for i in range(n):
        y[i] = (b[i] - dot(l[i][:i], y[:i])) / l[i][i]
    for i in reversed(range(n)):
        x[i] = (y[i] - sum(l[j][i] * x[j] for j in range(i + 1, n))) / l[i][i]
    return x


def prepare(data):
    if not isinstance(data, dict):
        raise FitError('top-level input must be a JSON object')
    family = data.get('family')
    if family not in ('binary', 'count'):
        raise FitError('family must be binary or count')
    initial = data.get('initial', [])
    if not initial or any(not math.isfinite(v) for v in initial):
        raise FitError('initial must be a nonempty finite vector')
    p = len(initial)
    clusters = data.get('clusters', [])
    if not isinstance(clusters, list) or any(not isinstance(g, dict) for g in clusters):
        raise FitError('clusters must be an array of JSON objects')
    if not clusters:
        raise FitError('at least one cluster is required')
    ids = [str(g.get('id', '')) for g in clusters]
    if any(not x for x in ids) or len(ids) != len(set(ids)):
        raise FitError('cluster IDs must be nonempty and unique; combine rows of one cluster')
    prepared = []
    for g in clusters:
        x, y, r = g['X'], g['y'], g['R']
        m = len(y)
        if not m or len(x) != m or any(len(row) != p for row in x):
            raise FitError('X/y/initial dimensions disagree')
        if any(not math.isfinite(v) for row in x for v in row) or any(not math.isfinite(v) for v in y):
            raise FitError('nonfinite response/design')
        if family == 'binary' and any(v not in (0, 1) for v in y):
            raise FitError('binary y must contain only 0 and 1')
        if family == 'count' and any(v < 0 or v != int(v) for v in y):
            raise FitError('count y must contain nonnegative integers')
        if len(r) != m or any(len(row) != m for row in r):
            raise FitError('R dimension disagrees with cluster size')
        if any(abs(r[i][i] - 1) > 1e-12 for i in range(m)):
            raise FitError('R must have unit diagonal')
        lr = chol(r, 'R')
        if family == 'count' and 'exposure' not in g:
            raise FitError('count clusters require explicit exposure (use all ones for equal unit exposure)')
        exposure = g.get('exposure', [1.0] * m)
        if len(exposure) != m or any(not math.isfinite(e) or e <= 0 for e in exposure):
            raise FitError('exposure must be finite, positive and match y')
        if family == 'binary' and any(e != 1 for e in exposure):
            raise FitError('exposure offsets are supported only for count models')
        prepared.append((g['id'], x, y, lr, exposure))
    return family, prepared, p


def evaluate(beta, family, groups, p):
    total = [0.0] * p
    h, meat = zeros(p, p), zeros(p, p)
    records = []
    for gid, x, y, lr, exposure in groups:
        eta = [dot(row, beta) for row in x]
        if family == 'binary':
            mu = [1 / (1 + math.exp(-t)) if t >= 0 else math.exp(t) / (1 + math.exp(t)) for t in eta]
            variance = [v * (1 - v) for v in mu]
            if any(v < 1e-12 for v in variance):
                raise FitError('binary mean is numerically saturated; check separation/initialization')
        else:
            try:
                mu = [math.exp(math.log(e) + t) for e, t in zip(exposure, eta)]
            except OverflowError as exc:
                raise FitError('count mean overflow') from exc
            variance = mu
            if any(not math.isfinite(v) or v < 1e-12 for v in variance):
                raise FitError('count mean is nonfinite or numerically zero')
        root_a = [math.sqrt(v) for v in variance]
        # Both supported canonical links have D = diag(variance) X.
        b = [[root_a[i] * z for z in row] for i, row in enumerate(x)]
        z = [(yi - mi) / ai for yi, mi, ai in zip(y, mu, root_a)]
        rz = solve(lr, z)
        bt = transpose(b)
        rb = [solve(lr, col) for col in bt]
        s = [dot(col, rz) for col in bt]
        for j in range(p):
            total[j] += s[j]
            for k in range(p):
                h[j][k] += dot(bt[j], rb[k])
                meat[j][k] += s[j] * s[k]
        records.append({'id': gid, 'mu': mu, 'score': s})
    if any(not math.isfinite(v) for row in h + meat for v in row):
        raise FitError('nonfinite bread or meat')
    return total, h, meat, records


def fit(data):
    family, groups, p = prepare(data)
    beta = list(data['initial'])
    tol = data.get('tol', 1e-11)
    max_iter = data.get('max_iter', 200)
    if not math.isfinite(tol) or tol <= 0 or not isinstance(max_iter, int) or max_iter < 0:
        raise FitError('tol must be positive finite; max_iter must be a nonnegative integer')
    g = len(groups)
    history = []
    first_step = None
    for iteration in range(max_iter + 1):
        u, h, meat, records = evaluate(beta, family, groups, p)
        lh = chol(h, 'H')
        residual = math.sqrt(dot(u, u)) / g
        delta = solve(lh, u)
        # Also check step size: a tiny raw score alone can hide numerical separation.
        if residual <= tol and max(abs(v) for v in delta) <= math.sqrt(tol):
            break
        if iteration == max_iter:
            raise FitError('maximum iterations reached without the residual/step certificate')
        if first_step is None:
            first_step = [beta[j] + delta[j] for j in range(p)]
        alpha = 1.0
        accepted = False
        for _ in range(30):
            trial = [beta[j] + alpha * delta[j] for j in range(p)]
            try:
                ut, _, _, _ = evaluate(trial, family, groups, p)
                trial_residual = math.sqrt(dot(ut, ut)) / g
                if trial_residual <= residual * (1 - 1e-4 * alpha):
                    accepted = True
                    break
            except FitError:
                pass
            alpha *= 0.5
        if not accepted:
            raise FitError('scoring direction failed the residual line search; no root certificate')
        history.append({'iteration': iteration, 'mean_score_norm': residual, 'damping': alpha})
        beta = trial
    inv_h = transpose([solve(lh, [float(i == j) for i in range(p)]) for j in range(p)])
    # Solve each fitted cluster score first; sum the transformed outer products.
    influence = [solve(lh, record['score']) for record in records]
    robust = [[sum(v[j] * v[k] for v in influence) for k in range(p)] for j in range(p)]
    # Numerical observed negative Jacobian is diagnostic only, not used for C.
    observed = zeros(p, p)
    step = 1e-5
    for j in range(p):
        up, down = list(beta), list(beta)
        up[j] += step
        down[j] -= step
        su = evaluate(up, family, groups, p)[0]
        sd = evaluate(down, family, groups, p)[0]
        for i in range(p):
            observed[i][j] = -(su[i] - sd[i]) / (2 * step)
    warnings = ['Numerical arithmetic only: cluster independence, full conditional mean, identification and asymptotic coverage are not verified.']
    if g <= p:
        warnings.append('G <= p: exact-root cluster meat rank is at most G-1; full Wald inversion is unavailable.')
    return {'status': 'numerical_root', 'family': family, 'G': g, 'N': sum(len(x[2]) for x in groups),
            'beta': beta, 'initial_undamped_step': first_step, 'iterations': iteration,
            'mean_score_norm': residual, 'final_scoring_step_inf_norm': max(abs(v) for v in delta),
            'H_condition_inf': max(sum(abs(v) for v in row) for row in h) * max(sum(abs(v) for v in row) for row in inv_h),
            'score_sum': u, 'H_sum': h, 'M_sum': meat,
            'robust_covariance': robust, 'working_covariance': inv_h,
            'observed_negative_jacobian_fd': observed, 'clusters': records,
            'history': history, 'warnings': warnings}


def binary_fixture(rho):
    xy = [([0, 1], [0, 0]), ([0, 2], [0, 1]), ([1, 2], [1, 1]), ([0, 1], [1, 0])]
    return {'family': 'binary', 'initial': [0.0, 0.0], 'tol': 1e-12, 'max_iter': 200,
            'clusters': [{'id': str(i + 1), 'X': [[1.0, v] for v in x], 'y': y,
                          'R': [[1.0, rho], [rho, 1.0]]} for i, (x, y) in enumerate(xy)]}


def count_fixture(rho=0.25):
    data = binary_fixture(rho)
    data['family'] = 'count'
    for g, y, exposure in zip(data['clusters'], [[0, 2], [1, 3], [2, 5], [2, 1]], [[1, 2], [2, 1], [1, 2], [2, 1]]):
        g['y'], g['exposure'] = y, exposure
    return data


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--input', help='JSON input; omit for the capstone fixtures')
    args = parser.parse_args()
    try:
        if args.input:
            with open(args.input, encoding='utf-8') as f:
                result = fit(json.load(f))
        else:
            result = {'binary_rho0': fit(binary_fixture(0)), 'binary_rho025': fit(binary_fixture(0.25)),
                      'count_rho025': fit(count_fixture()),
                      'outcome_weight_counterexample': {'score_mean_at_half': 0.25, 'weighted_target': 2 / 3}}
        print(json.dumps(result, ensure_ascii=False, indent=2, allow_nan=False))
    except (FitError, KeyError, TypeError, OverflowError, json.JSONDecodeError, OSError) as exc:
        print(json.dumps({'status': 'failed', 'reason': str(exc)}, ensure_ascii=False))
        sys.exit(1)


if __name__ == '__main__':
    main()
