"""Q1 (bilinear) finite-element reference for the 2D Poisson guide:
   -(u_xx + u_yy) = 2 pi^2 sin(pi x) sin(pi y) on (0,1)^2,  u = 0 on the edges.  Exact: sin(pi x) sin(pi y).
Run: python poisson_2d_fem_reference.py   -> prints the error table, writes poisson_2d_fem.json
The error is measured on the same 101x101 test grid the PINN guide uses, so the numbers can sit in one table."""
import json, math, numpy as np
import scipy.sparse as sp, scipy.sparse.linalg as spl

def solve(N):
    h = 1.0 / N; nn = N + 1
    idx = lambda i, j: i * nn + j
    Ke = (1 / 6) * np.array([[4, -1, -2, -1], [-1, 4, -1, -2], [-2, -1, 4, -1], [-1, -2, -1, 4]])   # bilinear Laplacian, any h
    gp = np.array([-1, 1]) / math.sqrt(3)                                  # 2x2 Gauss points on [-1,1]
    rows, cols, vals = [], [], []; b = np.zeros(nn * nn)
    for i in range(N):
        for j in range(N):
            dofs = [idx(i, j), idx(i + 1, j), idx(i + 1, j + 1), idx(i, j + 1)]
            for a in range(4):
                for c in range(4): rows.append(dofs[a]); cols.append(dofs[c]); vals.append(Ke[a, c])
            for xi in gp:
                for eta in gp:
                    x = (i + (xi + 1) / 2) * h; y = (j + (eta + 1) / 2) * h
                    fq = 2 * math.pi ** 2 * math.sin(math.pi * x) * math.sin(math.pi * y)
                    Nq = 0.25 * np.array([(1 - xi) * (1 - eta), (1 + xi) * (1 - eta), (1 + xi) * (1 + eta), (1 - xi) * (1 + eta)])
                    for a in range(4): b[dofs[a]] += fq * Nq[a] * h * h / 4
    A = sp.csr_matrix((vals, (rows, cols)), shape=(nn * nn, nn * nn))
    bnd = np.array([idx(i, j) for i in range(nn) for j in range(nn) if i in (0, N) or j in (0, N)])
    inn = np.setdiff1d(np.arange(nn * nn), bnd)
    u = np.zeros(nn * nn); u[inn] = spl.spsolve(A[inn][:, inn].tocsc(), b[inn])
    return u.reshape(nn, nn)

def on_test_grid(U, n=101):                      # bilinear interpolation of the FE field
    N = U.shape[0] - 1; g = np.linspace(0, 1, n)
    i = np.minimum((g * N).astype(int), N - 1); s = g * N - i
    out = np.zeros((n, n))
    for a, (ia, sa) in enumerate(zip(i, s)):
        for c, (ic, sc) in enumerate(zip(i, s)):
            out[a, c] = (U[ia, ic] * (1 - sa) * (1 - sc) + U[ia + 1, ic] * sa * (1 - sc) + U[ia + 1, ic + 1] * sa * sc + U[ia, ic + 1] * (1 - sa) * sc)
    return out

g = np.linspace(0, 1, 101); XE = np.sin(math.pi * g)[:, None] * np.sin(math.pi * g)[None, :]
res = []; prev = None
print(f"{'elements':>10} {'h':>8} {'rel L2':>11} {'max err':>11} {'order':>6}")
for N in (4, 8, 16, 32, 64):
    Ut = on_test_grid(solve(N)); e2 = np.linalg.norm(Ut - XE) / np.linalg.norm(XE); em = np.abs(Ut - XE).max()
    order = math.log(prev / e2, 2) if prev else float("nan"); prev = e2
    res.append(dict(N=N, h=1 / N, rel_l2=e2, max_err=em, order=None if prev is None or math.isnan(order) else order))
    print(f"{N:>7}x{N:<2} {1/N:8.4f} {e2:11.3e} {em:11.3e} {order:6.2f}")
json.dump(res, open("poisson_2d_fem.json", "w"), indent=1)
