--- Final Solver (used for performance test) ---
from typing import Any

import numpy as np
from numba import njit


@njit(cache=True, fastmath=True)
def _integrate(t0, t1, S0, E0, I0, R0, beta, sigma, gamma, omega, rtol, atol):
    # Dormand-Prince RK45 (same tableau as scipy's RK45) with adaptive steps.
    # Dormand-Prince coefficients
    a21 = 1.0 / 5.0
    a31 = 3.0 / 40.0
    a32 = 9.0 / 40.0
    a41 = 44.0 / 45.0
    a42 = -56.0 / 15.0
    a43 = 32.0 / 9.0
    a51 = 19372.0 / 6561.0
    a52 = -25360.0 / 2187.0
    a53 = 64448.0 / 6561.0
    a54 = -212.0 / 729.0
    a61 = 9017.0 / 3168.0
    a62 = -355.0 / 33.0
    a63 = 46732.0 / 5247.0
    a64 = 49.0 / 176.0
    a65 = -5103.0 / 18656.0
    # 5th order weights (b) = 7th stage row
    b1 = 35.0 / 384.0
    b3 = 500.0 / 1113.0
    b4 = 125.0 / 192.0
    b5 = -2187.0 / 6784.0
    b6 = 11.0 / 84.0
    # error coefficients E = b - b_hat
    e1 = 71.0 / 57600.0
    e3 = -71.0 / 16695.0
    e4 = 71.0 / 1920.0
    e5 = -17253.0 / 339200.0
    e6 = 22.0 / 525.0
    e7 = -1.0 / 40.0

    SAFETY = 0.9
    MIN_FACTOR = 0.2
    MAX_FACTOR = 10.0

    S = S0
    E = E0
    I = I0
    R = R0
    t = t0
    T = t1 - t0
    direction = 1.0 if T >= 0.0 else -1.0
    max_step = abs(T)

    # f0
    bSI = beta * S * I
    wR = omega * R
    k1S = -bSI + wR
    k1E = bSI - sigma * E
    k1I = sigma * E - gamma * I
    k1R = gamma * I - wR

    # initial step size (Hairer)
    scS = atol + abs(S) * rtol
    scE = atol + abs(E) * rtol
    scI = atol + abs(I) * rtol
    scR = atol + abs(R) * rtol
    d0 = np.sqrt(((S / scS) ** 2 + (E / scE) ** 2 + (I / scI) ** 2 + (R / scR) ** 2) * 0.25)
    d1 = np.sqrt(((k1S / scS) ** 2 + (k1E / scE) ** 2 + (k1I / scI) ** 2 + (k1R / scR) ** 2) * 0.25)
    if d0 < 1e-5 or d1 < 1e-5:
        h0 = 1e-6
    else:
        h0 = 0.01 * d0 / d1
    # one explicit Euler trial to refine
    sd = direction * h0
    yS = S + sd * k1S
    yE = E + sd * k1E
    yI = I + sd * k1I
    yR = R + sd * k1R
    bSI = beta * yS * yI
    wR = omega * yR
    f1S = -bSI + wR
    f1E = bSI - sigma * yE
    f1I = sigma * yE - gamma * yI
    f1R = gamma * yI - wR
    d2 = np.sqrt((((f1S - k1S) / scS) ** 2 + ((f1E - k1E) / scE) ** 2
                 + ((f1I - k1I) / scI) ** 2 + ((f1R - k1R) / scR) ** 2) * 0.25) / h0
    dm = d1 if d1 > d2 else d2
    if dm <= 1e-15:
        h1 = max(1e-6, h0 * 1e-3)
    else:
        h1 = (0.01 / dm) ** 0.2
    h = min(100.0 * h0, h1)
    if h > max_step:
        h = max_step
    h = h * direction

    while (t - t1) * direction < 0.0:
        ah = abs(h)
        if ah > max_step:
            ah = max_step
            h = direction * ah
        # avoid overshoot
        if (t + h - t1) * direction > 0.0:
            h = t1 - t

        # stage 2
        yS = S + h * a21 * k1S
        yE = E + h * a21 * k1E
        yI = I + h * a21 * k1I
        yR = R + h * a21 * k1R
        bSI = beta * yS * yI
        wR = omega * yR
        k2S = -bSI + wR
        k2E = bSI - sigma * yE
        k2I = sigma * yE - gamma * yI
        k2R = gamma * yI - wR

        # stage 3
        yS = S + h * (a31 * k1S + a32 * k2S)
        yE = E + h * (a31 * k1E + a32 * k2E)
        yI = I + h * (a31 * k1I + a32 * k2I)
        yR = R + h * (a31 * k1R + a32 * k2R)
        bSI = beta * yS * yI
        wR = omega * yR
        k3S = -bSI + wR
        k3E = bSI - sigma * yE
        k3I = sigma * yE - gamma * yI
        k3R = gamma * yI - wR

        # stage 4
        yS = S + h * (a41 * k1S + a42 * k2S + a43 * k3S)
        yE = E + h * (a41 * k1E + a42 * k2E + a43 * k3E)
        yI = I + h * (a41 * k1I + a42 * k2I + a43 * k3I)
        yR = R + h * (a41 * k1R + a42 * k2R + a43 * k3R)
        bSI = beta * yS * yI
        wR = omega * yR
        k4S = -bSI + wR
        k4E = bSI - sigma * yE
        k4I = sigma * yE - gamma * yI
        k4R = gamma * yI - wR

        # stage 5
        yS = S + h * (a51 * k1S + a52 * k2S + a53 * k3S + a54 * k4S)
        yE = E + h * (a51 * k1E + a52 * k2E + a53 * k3E + a54 * k4E)
        yI = I + h * (a51 * k1I + a52 * k2I + a53 * k3I + a54 * k4I)
        yR = R + h * (a51 * k1R + a52 * k2R + a53 * k3R + a54 * k4R)
        bSI = beta * yS * yI
        wR = omega * yR
        k5S = -bSI + wR
        k5E = bSI - sigma * yE
        k5I = sigma * yE - gamma * yI
        k5R = gamma * yI - wR

        # stage 6
        yS = S + h * (a61 * k1S + a62 * k2S + a63 * k3S + a64 * k4S + a65 * k5S)
        yE = E + h * (a61 * k1E + a62 * k2E + a63 * k3E + a64 * k4E + a65 * k5E)
        yI = I + h * (a61 * k1I + a62 * k2I + a63 * k3I + a64 * k4I + a65 * k5I)
        yR = R + h * (a61 * k1R + a62 * k2R + a63 * k3R + a64 * k4R + a65 * k5R)
        bSI = beta * yS * yI
        wR = omega * yR
        k6S = -bSI + wR
        k6E = bSI - sigma * yE
        k6I = sigma * yE - gamma * yI
        k6R = gamma * yI - wR

        # 5th order solution
        nS = S + h * (b1 * k1S + b3 * k3S + b4 * k4S + b5 * k5S + b6 * k6S)
        nE = E + h * (b1 * k1E + b3 * k3E + b4 * k4E + b5 * k5E + b6 * k6E)
        nI = I + h * (b1 * k1I + b3 * k3I + b4 * k4I + b5 * k5I + b6 * k6I)
        nR = R + h * (b1 * k1R + b3 * k3R + b4 * k4R + b5 * k5R + b6 * k6R)

        # stage 7 (FSAL) at new point
        bSI = beta * nS * nI
        wR = omega * nR
        k7S = -bSI + wR
        k7E = bSI - sigma * nE
        k7I = sigma * nE - gamma * nI
        k7R = gamma * nI - wR

        # error estimate
        errS = h * (e1 * k1S + e3 * k3S + e4 * k4S + e5 * k5S + e6 * k6S + e7 * k7S)
        errE = h * (e1 * k1E + e3 * k3E + e4 * k4E + e5 * k5E + e6 * k6E + e7 * k7E)
        errI = h * (e1 * k1I + e3 * k3I + e4 * k4I + e5 * k5I + e6 * k6I + e7 * k7I)
        errR = h * (e1 * k1R + e3 * k3R + e4 * k4R + e5 * k5R + e6 * k6R + e7 * k7R)

        scS = atol + max(abs(S), abs(nS)) * rtol
        scE = atol + max(abs(E), abs(nE)) * rtol
        scI = atol + max(abs(I), abs(nI)) * rtol
        scR = atol + max(abs(R), abs(nR)) * rtol
        err = np.sqrt(((errS / scS) ** 2 + (errE / scE) ** 2
                       + (errI / scI) ** 2 + (errR / scR) ** 2) * 0.25)

        if err < 1.0:
            t += h
            S = nS
            E = nE
            I = nI
            R = nR
            k1S = k7S
            k1E = k7E
            k1I = k7I
            k1R = k7R
            # Safe early termination: if the state has reached steady state,
            # the total remaining change is bounded by |f|_inf * (t1 - t),
            # since f only decays toward a stable fixed point. When that bound
            # is negligible, the value at t1 equals the current state.
            fmax = abs(k1S)
            if abs(k1E) > fmax:
                fmax = abs(k1E)
            if abs(k1I) > fmax:
                fmax = abs(k1I)
            if abs(k1R) > fmax:
                fmax = abs(k1R)
            if fmax * abs(t1 - t) < 1e-10:
                break
            if err == 0.0:
                factor = MAX_FACTOR
            else:
                factor = min(MAX_FACTOR, SAFETY * err ** (-0.2))
            h = h * factor
        else:
            factor = max(MIN_FACTOR, SAFETY * err ** (-0.2))
            h = h * factor

    return S, E, I, R


class Solver:
    def __init__(self):
        self.rtol = 1e-6
        self.atol = 1e-10
        # Warm up / compile (does not count toward runtime).
        _integrate(0.0, 1.0, 0.9, 0.01, 0.005, 0.085, 0.35, 0.2, 0.1, 0.002,
                   self.rtol, self.atol)

    def solve(self, problem, **kwargs) -> Any:
        y0 = problem["y0"]
        p = problem["params"]
        S, E, I, R = _integrate(
            float(problem["t0"]), float(problem["t1"]),
            float(y0[0]), float(y0[1]), float(y0[2]), float(y0[3]),
            float(p["beta"]), float(p["sigma"]), float(p["gamma"]), float(p["omega"]),
            self.rtol, self.atol,
        )
        return [S, E, I, R]
--- End Solver ---
Running performance test...
============================= test session starts ==============================
platform linux -- Python 3.12.13, pytest-9.0.2, pluggy-1.6.0
rootdir: /tests
plugins: anyio-4.13.0, jaxtyping-0.3.9
collected 3 items

../tests/test_outputs.py .
--- Performance Summary ---
Validity: True
Total Baseline Time: 8.5471s
Total Solver Time:   0.0137s
Raw Speedup:         624.8254 x
Final Reward (Score): 624.8254
---------------------------
..

==================================== PASSES ====================================
=========================== short test summary info ============================
PASSED ../tests/test_outputs.py::test_solver_exists
PASSED ../tests/test_outputs.py::test_solver_validity
PASSED ../tests/test_outputs.py::test_solver_speedup
======================== 3 passed in 113.49s (0:01:53) =========================
