"""Penalized joint Rasch fit; not unpenalized MLE or a validated assessment."""
import json
import numpy as np
from scipy.optimize import minimize
from scipy.special import expit


def validate(responses):
    y = np.asarray(responses, dtype=float)
    if y.ndim != 2 or min(y.shape) < 2:
        raise ValueError("need at least two people and two items")
    observed = ~np.isnan(y)
    if not np.all(np.isin(y[observed], [0., 1.])):
        raise ValueError("responses must be 0, 1, or NaN")
    if not observed.any(axis=0).all() or not observed.any(axis=1).all():
        raise ValueError("each person and item needs an observation")
    return y, observed


def unpack(parameters, people, items):
    theta = parameters[:people]
    free_b = parameters[people:]
    b = np.r_[free_b, -free_b.sum()]
    return theta, b


def objective(parameters, y, observed, penalty):
    theta, b = unpack(parameters, *y.shape)
    logits = theta[:, None] - b[None, :]
    residual = np.where(observed, expit(logits) - np.nan_to_num(y), 0.)
    loss = np.logaddexp(0., logits[observed]).sum()
    loss -= (y[observed] * logits[observed]).sum()
    loss += 0.5 * penalty * (theta @ theta + b @ b)
    theta_gradient = residual.sum(axis=1) + penalty * theta
    b_gradient = -residual.sum(axis=0) + penalty * b
    gradient = np.r_[theta_gradient, b_gradient[:-1] - b_gradient[-1]]
    return float(loss), gradient


def fit(responses, penalty=1.0, maxiter=1000):
    y, observed = validate(responses)
    if not np.isfinite(penalty) or penalty <= 0:
        raise ValueError("this estimator requires a positive finite penalty")
    result = minimize(objective, np.zeros(sum(y.shape) - 1),
                      args=(y, observed, penalty), jac=True, method="L-BFGS-B",
                      options={"maxiter": maxiter, "gtol": 1e-8, "ftol": 1e-12})
    if not result.success or not np.isfinite(result.x).all():
        raise RuntimeError(f"fit failed: {result.message}")
    theta, b = unpack(result.x, *y.shape)
    return theta, b


FIXTURE = np.array([[1,1,1,1], [1,1,1,0], [1,1,0,0], [1,0,0,0], [0,0,0,0]])
if __name__ == "__main__":
    theta, b = fit(FIXTURE)
    print(json.dumps({"theta": theta.round(4).tolist(), "difficulty": b.round(4).tolist(),
                      "difficulty_mean": round(float(b.mean()), 10), "penalty": 1.0}))
