"""Ground truth for the time-crystal page: the same Floquet model in numpy.

The page (time-crystals.html) runs its own JS engine. This file computes the
same quantities independently and writes golden.json; test_engine.mjs checks
the JS engine against it. Conventions match sims/clock-which-path/
time_crystal_through.py: basis index s, bit j = spin j, Z_j = 1 - 2*bit_j,
one period = free evolution exp(-i H T) THEN the kick
K = prod_j exp(-i (pi/2)(1 - eps) X_j).

    H = sum_j J_j Z_j Z_{j+1} + sum_j h_j Z_j      (open chain, T = 1)

Run:  python reference.py      (writes golden.json next to this file)
"""

from __future__ import annotations

import json
from pathlib import Path

import numpy as np

HERE = Path(__file__).resolve().parent


def z_configs(L):
    idx = np.arange(2**L)[:, None]
    return 1 - 2 * ((idx >> np.arange(L)[None, :]) & 1)


def energies(J, h):
    L = len(h)
    z = z_configs(L)
    e = z @ np.asarray(h, float)
    for j in range(L - 1):
        e = e + J[j] * z[:, j] * z[:, j + 1]
    return e


def kick(psi, L, eps, frac=1.0):
    """Apply prod_j exp(-i frac (pi/2)(1-eps) X_j). Single-site rotations commute,
    so a fractional kick composes exactly."""
    a = frac * (np.pi / 2) * (1 - eps)
    c, s = np.cos(a), np.sin(a)
    psi = psi.copy()
    for j in range(L):
        p = psi.reshape(-1, 2, 2**j)
        lo, hi = p[:, 0, :].copy(), p[:, 1, :].copy()
        p[:, 0, :] = c * lo - 1j * s * hi
        p[:, 1, :] = c * hi - 1j * s * lo
        psi = p.reshape(-1)
    return psi


def run(J, h, eps, bits, n_periods):
    """Stroboscopic <Z_j>(n), n = 0..n_periods, shape (n+1, L)."""
    L = len(h)
    z = z_configs(L)
    E = energies(J, h)
    psi = np.zeros(2**L, complex)
    psi[bits] = 1
    out = [(np.abs(psi) ** 2) @ z]
    for _ in range(n_periods):
        psi = kick(np.exp(-1j * E) * psi, L, eps)
        out.append((np.abs(psi) ** 2) @ z)
    return np.array(out), psi  # psi is the state at n = n_periods, the same one out[-1] measures


def bloch(psi, L):
    """Per-site Bloch vector (<X_j>, <Y_j>, <Z_j>) via the full state."""
    out = []
    for j in range(L):
        p = psi.reshape(-1, 2, 2**j)
        a, b = p[:, 0, :], p[:, 1, :]  # amplitudes with bit j = 0 / 1
        ab = np.sum(np.conj(a) * b)
        out.append([2 * ab.real, 2 * ab.imag, np.sum(abs(a) ** 2 - abs(b) ** 2)])
    return np.array(out)


def phi(Z, bits, L):
    """Late-window period-doubling order parameter:
    mean over sites and n in [N/2, N] of (-1)^n Z_j(n) Z_j(0)."""
    z0 = z_configs(L)[bits]
    N = len(Z) - 1
    n = np.arange(N + 1)
    w = n >= N // 2
    return float(np.mean(((-1.0) ** n)[w, None] * Z[w] * z0[None, :]))


def main():
    L = 8
    # fixed, explicit disorder (no RNG needed to reproduce)
    J = [0.83, 1.27, 0.61, 1.44, 0.95, 0.52, 1.18, 0.0]   # last bond unused (open chain)
    h = [0.12, 0.77, 0.43, 0.91, 0.05, 0.66, 0.38, 0.84]
    bits = 0b10110010
    gold = {"L": L, "J": J, "h": h, "bits": bits, "cases": []}
    # run() must hand back the state its last observable measured, not one period later
    Zc, psi_end = run(J, h, 0.1, bits, 5)
    assert np.allclose((np.abs(psi_end) ** 2) @ z_configs(L), Zc[-1]), "run() returned the wrong state"
    _, psi_zero = run(J, h, 0.1, bits, 0)
    assert abs(psi_zero[bits]) == 1.0, "run(n_periods=0) must return the untouched start state"
    for model, Jm in (("crystal", J), ("impostor", [0.0] * L)):
        for eps in (0.0, 0.1, 0.3):
            Z, _ = run(Jm, h, eps, bits, 64)
            Zlong, _ = run(Jm, h, eps, bits, 200)
            gold["cases"].append({
                "model": model, "eps": eps,
                "Z": Z.tolist(),
                "phi200": phi(Zlong, bits, L),
            })
    # mid-animation state: 3 full periods, then free evolution, then half a kick
    psi3 = np.zeros(2**L, complex)
    psi3[bits] = 1
    E = energies(J, h)
    for _ in range(3):
        psi3 = kick(np.exp(-1j * E) * psi3, L, 0.1)
    mid = kick(np.exp(-1j * E * 0.4) * psi3, L, 0.1, frac=0.0)   # 40% through free evolution
    half = kick(np.exp(-1j * E) * psi3, L, 0.1, frac=0.5)        # halfway through the kick
    gold["bloch_free40"] = bloch(mid, L).tolist()
    gold["bloch_halfkick"] = bloch(half, L).tolist()
    (HERE / "golden.json").write_text(json.dumps(gold, indent=1))
    for c in gold["cases"]:
        print(f"{c['model']:9s} eps={c['eps']:.1f}  phi(N=200)={c['phi200']:+.4f}")
    print("wrote golden.json")


if __name__ == "__main__":
    main()
