# Case 010: constructed data; not empirical observations.
# Normal outcomes within two fully observed groups; missingness depends on group.
import numpy as np
from scipy.stats import t
rng = np.random.default_rng(20261007)
x = np.repeat([0, 1], 100)
e = np.tile(np.linspace(-3, 3, 100), 2)
y_full = 10 + 4*x + e
# Retain symmetric positions: 80 in group 0 and 40 in group 1.
observed = np.zeros(200, dtype=bool)
for g, k in [(0, 80), (1, 40)]:
    ix = np.flatnonzero(x == g)
    keep = np.r_[np.arange(k//2), np.arange(100-k//2, 100)]
    observed[ix[keep]] = True
y = np.where(observed, y_full, np.nan)
means = np.array([np.nanmean(y[x == g]) for g in [0, 1]])
counts = np.array([np.sum(observed & (x == g)) for g in [0, 1]])
variances = np.array([np.nanvar(y[x == g], ddof=1) for g in [0, 1]])
# Saturated two-group normal observed-data likelihood: group means and ML variances.
# Fixed target weights are 1/2, because group membership is known for all 200.
fiml_mean = means.mean()
fiml_se = np.sqrt(np.sum(.25 * variances*(counts-1)/counts**2))
# Proper normal-model MI: draw variance, then mean, then missing outcomes.
# Independent reference prior p(mu_g, sigma_g^2) proportional to 1/sigma_g^2.
M = 2000
Q, U = [], []
for _ in range(M):
    completed = y.copy()
    for g in [0, 1]:
        variance_draw = (counts[g]-1)*variances[g]/rng.chisquare(counts[g]-1)
        mean_draw = rng.normal(means[g], np.sqrt(variance_draw/counts[g]))
        missing = (x == g) & ~observed
        completed[missing] = rng.normal(mean_draw, np.sqrt(variance_draw), missing.sum())
    Q.append(completed.mean())
    U.append(sum(.25*np.var(completed[x == g], ddof=1)/100 for g in [0, 1]))
Q, U = np.array(Q), np.array(U)
W, B = U.mean(), Q.var(ddof=1)
T = W + (1 + 1/M)*B
mi_se = np.sqrt(T)
# Large-complete-sample Rubin df; small-sample adjustments are not implemented.
df = (M-1)*(1 + W/((1+1/M)*B))**2
interval = Q.mean() + np.array([-1, 1])*t.ppf(.975, df)*mi_se
print('Observed counts:', counts)
print('Deletion mean:', np.nanmean(y))
print('Likelihood mean / model-based SE:', fiml_mean, fiml_se)
print('MI mean / pooled SE / 95% interval:', Q.mean(), mi_se, interval)
print('Within / between / total variance:', W, B, T)
assert np.array_equal(counts, [80, 40])
assert np.isclose(np.nanmean(y), 34/3)
assert np.isclose(fiml_mean, 12)
assert T > W and abs(Q.mean()-12) < .1
print('All checks passed. NumPy', np.__version__)
