--- Final Solver (used for performance test) ---
import numpy as np
from typing import Any
from numba import njit


@njit(cache=True, fastmath=True)
def solve_seirs(S0, E0, I0, t0, t1, beta, sigma, gamma, omega):
    """
    Adaptive Dormand-Prince RK45 with FSAL for the reduced SEIRS system.
    Reduced to 3 variables via R = 1 - S - E - I.
    """
    t = t0
    S = S0
    E = E0
    I = I0

    T = t1 - t0
    if T <= 0.0:
        return S, E, I, 1.0 - S - E - I

    rtol = 5e-7
    atol = 5e-10

    # Precomputed Dormand-Prince coefficients
    # a21
    A21 = 0.2
    # a31, a32
    A31 = 0.075; A32 = 0.225
    # a41, a42, a43
    A41 = 44.0/45.0; A42 = -56.0/15.0; A43 = 32.0/9.0
    # a51..a54
    A51 = 19372.0/6561.0; A52 = -25360.0/2187.0; A53 = 64448.0/6561.0; A54 = -212.0/729.0
    # a61..a65
    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
    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 weights (5th - 4th)
    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 = -0.025

    ISQ3 = 1.0 / np.sqrt(3.0)

    # Initial step
    R0 = 1.0 - S - E - I
    dS0 = -beta * S * I + omega * R0
    dE0 = beta * S * I - sigma * E
    dI0 = sigma * E - gamma * I
    dn = dS0*dS0 + dE0*dE0 + dI0*dI0
    yn = S*S + E*E + I*I

    if yn < 1e-10 or dn < 1e-10:
        h = 1e-4
    else:
        h = 0.01 * np.sqrt(yn / dn)
    if h > T:
        h = T

    # FSAL: k1
    k1s = dS0; k1e = dE0; k1i = dI0

    SAF = 0.9

    for _ in range(5000000):
        if t >= t1 - 1e-14 * abs(t1):
            break
        if t + h > t1:
            h = t1 - t

        # 6 stages + 1 FSAL stage = 7 RHS evals (6 new per step)
        # Stage 2
        S2 = S + h * A21 * k1s
        E2 = E + h * A21 * k1e
        I2 = I + h * A21 * k1i
        R2 = 1.0 - S2 - E2 - I2
        SI = S2 * I2
        k2s = -beta * SI + omega * R2
        k2e = beta * SI - sigma * E2
        k2i = sigma * E2 - gamma * I2

        # Stage 3
        hs = h * A31 * k1s + h * A32 * k2s
        he = h * A31 * k1e + h * A32 * k2e
        hi = h * A31 * k1i + h * A32 * k2i
        S3 = S + hs; E3 = E + he; I3 = I + hi
        R3 = 1.0 - S3 - E3 - I3
        SI = S3 * I3
        k3s = -beta * SI + omega * R3
        k3e = beta * SI - sigma * E3
        k3i = sigma * E3 - gamma * I3

        # Stage 4
        hs = h * A41 * k1s + h * A42 * k2s + h * A43 * k3s
        he = h * A41 * k1e + h * A42 * k2e + h * A43 * k3e
        hi = h * A41 * k1i + h * A42 * k2i + h * A43 * k3i
        S4 = S + hs; E4 = E + he; I4 = I + hi
        R4 = 1.0 - S4 - E4 - I4
        SI = S4 * I4
        k4s = -beta * SI + omega * R4
        k4e = beta * SI - sigma * E4
        k4i = sigma * E4 - gamma * I4

        # Stage 5
        hs = h * A51 * k1s + h * A52 * k2s + h * A53 * k3s + h * A54 * k4s
        he = h * A51 * k1e + h * A52 * k2e + h * A53 * k3e + h * A54 * k4e
        hi = h * A51 * k1i + h * A52 * k2i + h * A53 * k3i + h * A54 * k4i
        S5 = S + hs; E5 = E + he; I5 = I + hi
        R5 = 1.0 - S5 - E5 - I5
        SI = S5 * I5
        k5s = -beta * SI + omega * R5
        k5e = beta * SI - sigma * E5
        k5i = sigma * E5 - gamma * I5

        # Stage 6
        hs = h * A61 * k1s + h * A62 * k2s + h * A63 * k3s + h * A64 * k4s + h * A65 * k5s
        he = h * A61 * k1e + h * A62 * k2e + h * A63 * k3e + h * A64 * k4e + h * A65 * k5e
        hi = h * A61 * k1i + h * A62 * k2i + h * A63 * k3i + h * A64 * k4i + h * A65 * k5i
        S6 = S + hs; E6 = E + he; I6 = I + hi
        R6 = 1.0 - S6 - E6 - I6
        SI = S6 * I6
        k6s = -beta * SI + omega * R6
        k6e = beta * SI - sigma * E6
        k6i = sigma * E6 - gamma * I6

        # 5th order
        hs = h * B1 * k1s + h * B3 * k3s + h * B4 * k4s + h * B5 * k5s + h * B6 * k6s
        he = h * B1 * k1e + h * B3 * k3e + h * B4 * k4e + h * B5 * k5e + h * B6 * k6e
        hi = h * B1 * k1i + h * B3 * k3i + h * B4 * k4i + h * B5 * k5i + h * B6 * k6i
        Sn = S + hs; En = E + he; In = I + hi

        # FSAL stage 7
        Rn = 1.0 - Sn - En - In
        SI = Sn * In
        k7s = -beta * SI + omega * Rn
        k7e = beta * SI - sigma * En
        k7i = sigma * En - gamma * In

        # Error (using difference coefficients directly)
        es = h * E1 * k1s + h * E3 * k3s + h * E4 * k4s + h * E5 * k5s + h * E6 * k6s + h * E7 * k7s
        ee = h * E1 * k1e + h * E3 * k3e + h * E4 * k4e + h * E5 * k5e + h * E6 * k6e + h * E7 * k7e
        ei = h * E1 * k1i + h * E3 * k3i + h * E4 * k4i + h * E5 * k5i + h * E6 * k6i + h * E7 * k7i

        # Scale & norm
        absS = abs(S); absSnew = abs(Sn); sc_s = atol + rtol * (absS if absS > absSnew else absSnew)
        absE = abs(E); absEnew = abs(En); sc_e = atol + rtol * (absE if absE > absEnew else absEnew)
        absI = abs(I); absInew = abs(In); sc_i = atol + rtol * (absI if absI > absInew else absInew)

        aes = abs(es); aee = abs(ee); aei = abs(ei)
        en2 = (aes/sc_s)**2 + (aee/sc_e)**2 + (aei/sc_i)**2
        en = np.sqrt(en2) * ISQ3

        if en <= 1.0:
            t += h
            S = Sn; E = En; I = In
            k1s = k7s; k1e = k7e; k1i = k7i

            if en < 1e-15:
                f = 10.0
            else:
                f = SAF * en ** (-0.2)
                if f > 10.0: f = 10.0
                elif f < 0.2: f = 0.2
            h *= f
            if h > T: h = T
        else:
            f = SAF * en ** (-0.25)
            if f < 0.2: f = 0.2
            h *= f

    return S, E, I, 1.0 - S - E - I


class Solver:
    def __init__(self):
        # Warmup JIT with various problem sizes
        solve_seirs(0.89, 0.01, 0.005, 0.0, 1.0, 0.35, 0.2, 0.1, 0.002)
        solve_seirs(0.89, 0.01, 0.005, 0.0, 100.0, 0.35, 0.2, 0.1, 0.002)
        solve_seirs(0.89, 0.01, 0.005, 0.0, 5000.0, 0.35, 0.2, 0.1, 0.002)
        solve_seirs(0.5, 0.1, 0.1, 0.0, 400.0, 0.5, 0.3, 0.15, 0.01)
        solve_seirs(0.95, 0.02, 0.01, 0.0, 200.0, 1.0, 0.5, 0.2, 0.005)
        solve_seirs(0.999, 0.0005, 0.0001, 0.0, 1000.0, 0.4, 0.3, 0.1, 0.001)

    def solve(self, problem: dict, **kwargs) -> Any:
        y0 = problem["y0"]
        t0 = problem["t0"]
        t1 = problem["t1"]
        params = problem["params"]

        S0 = y0[0]; E0 = y0[1]; I0 = y0[2]; R0 = y0[3]
        beta = params["beta"]
        sigma_p = params["sigma"]
        gamma = params["gamma"]
        omega = params["omega"]

        S, E, I, R = solve_seirs(S0, E0, I0, t0, t1, beta, sigma_p, gamma, omega)

        return [float(S), float(E), float(I), float(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: False
Total Baseline Time: 0.8587s
Total Solver Time:   1.9537s
Raw Speedup:         0.0000 x
Final Reward (Score): 1.0000
---------------------------
FF

=================================== FAILURES ===================================
_____________________________ test_solver_validity _____________________________

performance_results = {'raw_speedup': 0.0, 'speedup': 1.0, 'validity': False}

    def test_solver_validity(performance_results):
        """Ensures the solver produces valid outputs."""
        if not performance_results["validity"]:
>           pytest.fail("Solver produced invalid solutions on the test set.")
E           Failed: Solver produced invalid solutions on the test set.

/tests/test_outputs.py:200: Failed
------------------------------ Captured log setup ------------------------------
ERROR    root:evaluator.py:164 Solution verification failed: max abs err=0.185, max rel err=2.27e+05
ERROR    test_outputs:test_outputs.py:93 Solver produced an invalid solution during warmup.
WARNING  test_outputs:test_outputs.py:172 Invalid solution detected. Assigning mercy score of 1.0.
_____________________________ test_solver_speedup ______________________________

performance_results = {'raw_speedup': 0.0, 'speedup': 1.0, 'validity': False}

    def test_solver_speedup(performance_results):
        """Checks if the solver effectively optimized the code."""
>       assert performance_results["raw_speedup"] > 1.0, \
            f"Solver was not faster than baseline (Speedup: {performance_results['raw_speedup']:.2f}x)"
E       AssertionError: Solver was not faster than baseline (Speedup: 0.00x)
E       assert 0.0 > 1.0

/tests/test_outputs.py:205: AssertionError
==================================== PASSES ====================================
=========================== short test summary info ============================
PASSED ../tests/test_outputs.py::test_solver_exists
FAILED ../tests/test_outputs.py::test_solver_validity - Failed: Solver produc...
FAILED ../tests/test_outputs.py::test_solver_speedup - AssertionError: Solver...
========================= 2 failed, 1 passed in 35.64s =========================
