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


@njit(cache=True)
def _seirs_rk4(y0, t0, t1, beta, sigma, gamma, omega, dt):
    S = y0[0]
    E = y0[1]
    I = y0[2]
    R = y0[3]
    t = t0
    while t < t1:
        h = dt if t + dt <= t1 else t1 - t
        bSI = beta * S * I
        dS1 = -bSI + omega * R
        dE1 = bSI - sigma * E
        dI1 = sigma * E - gamma * I
        dR1 = gamma * I - omega * R
        S2 = S + 0.5 * h * dS1
        E2 = E + 0.5 * h * dE1
        I2 = I + 0.5 * h * dI1
        R2 = R + 0.5 * h * dR1
        bSI = beta * S2 * I2
        dS2 = -bSI + omega * R2
        dE2 = bSI - sigma * E2
        dI2 = sigma * E2 - gamma * I2
        dR2 = gamma * I2 - omega * R2
        S3 = S + 0.5 * h * dS2
        E3 = E + 0.5 * h * dE2
        I3 = I + 0.5 * h * dI2
        R3 = R + 0.5 * h * dR2
        bSI = beta * S3 * I3
        dS3 = -bSI + omega * R3
        dE3 = bSI - sigma * E3
        dI3 = sigma * E3 - gamma * I3
        dR3 = gamma * I3 - omega * R3
        S4 = S + h * dS3
        E4 = E + h * dE3
        I4 = I + h * dI3
        R4 = R + h * dR3
        bSI = beta * S4 * I4
        dS4 = -bSI + omega * R4
        dE4 = bSI - sigma * E4
        dI4 = sigma * E4 - gamma * I4
        dR4 = gamma * I4 - omega * R4
        S += (h / 6.0) * (dS1 + 2.0 * dS2 + 2.0 * dS3 + dS4)
        E += (h / 6.0) * (dE1 + 2.0 * dE2 + 2.0 * dE3 + dE4)
        I += (h / 6.0) * (dI1 + 2.0 * dI2 + 2.0 * dI3 + dI4)
        R += (h / 6.0) * (dR1 + 2.0 * dR2 + 2.0 * dR3 + dR4)
        t += h
    return np.array([S, E, I, R])


@njit(cache=True)
def _seirs_dp45(y0, t0, t1, beta, sigma, gamma, omega, tol):
    """Adaptive Dormand-Prince RK45."""
    S = y0[0]
    E = y0[1]
    I = y0[2]
    R = y0[3]
    t = t0
    h = min(0.5, (t1 - t0) * 0.01)
    if h < 1e-8:
        h = 1e-8

    while t < t1:
        if t + h > t1:
            h = t1 - t

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

        S2 = S + h * 0.2 * k1S
        E2 = E + h * 0.2 * k1E
        I2 = I + h * 0.2 * k1I
        R2 = R + h * 0.2 * k1R
        bSI = beta * S2 * I2
        k2S = -bSI + omega * R2
        k2E = bSI - sigma * E2
        k2I = sigma * E2 - gamma * I2
        k2R = gamma * I2 - omega * R2

        S3 = S + h * (3.0/40.0*k1S + 9.0/40.0*k2S)
        E3 = E + h * (3.0/40.0*k1E + 9.0/40.0*k2E)
        I3 = I + h * (3.0/40.0*k1I + 9.0/40.0*k2I)
        R3 = R + h * (3.0/40.0*k1R + 9.0/40.0*k2R)
        bSI = beta * S3 * I3
        k3S = -bSI + omega * R3
        k3E = bSI - sigma * E3
        k3I = sigma * E3 - gamma * I3
        k3R = gamma * I3 - omega * R3

        S4 = S + h * (44.0/45.0*k1S - 56.0/15.0*k2S + 32.0/9.0*k3S)
        E4 = E + h * (44.0/45.0*k1E - 56.0/15.0*k2E + 32.0/9.0*k3E)
        I4 = I + h * (44.0/45.0*k1I - 56.0/15.0*k2I + 32.0/9.0*k3I)
        R4 = R + h * (44.0/45.0*k1R - 56.0/15.0*k2R + 32.0/9.0*k3R)
        bSI = beta * S4 * I4
        k4S = -bSI + omega * R4
        k4E = bSI - sigma * E4
        k4I = sigma * E4 - gamma * I4
        k4R = gamma * I4 - omega * R4

        S5 = S + h * (19372.0/6561.0*k1S - 25360.0/2187.0*k2S + 64448.0/6561.0*k3S - 212.0/729.0*k4S)
        E5 = E + h * (19372.0/6561.0*k1E - 25360.0/2187.0*k2E + 64448.0/6561.0*k3E - 212.0/729.0*k4E)
        I5 = I + h * (19372.0/6561.0*k1I - 25360.0/2187.0*k2I + 64448.0/6561.0*k3I - 212.0/729.0*k4I)
        R5 = R + h * (19372.0/6561.0*k1R - 25360.0/2187.0*k2R + 64448.0/6561.0*k3R - 212.0/729.0*k4R)
        bSI = beta * S5 * I5
        k5S = -bSI + omega * R5
        k5E = bSI - sigma * E5
        k5I = sigma * E5 - gamma * I5
        k5R = gamma * I5 - omega * R5

        S6 = S + h * (9017.0/3168.0*k1S - 355.0/33.0*k2S + 46732.0/5247.0*k3S + 49.0/176.0*k4S - 5103.0/18656.0*k5S)
        E6 = E + h * (9017.0/3168.0*k1E - 355.0/33.0*k2E + 46732.0/5247.0*k3E + 49.0/176.0*k4E - 5103.0/18656.0*k5E)
        I6 = I + h * (9017.0/3168.0*k1I - 355.0/33.0*k2I + 46732.0/5247.0*k3I + 49.0/176.0*k4I - 5103.0/18656.0*k5I)
        R6 = R + h * (9017.0/3168.0*k1R - 355.0/33.0*k2R + 46732.0/5247.0*k3R + 49.0/176.0*k4R - 5103.0/18656.0*k5R)
        bSI = beta * S6 * I6
        k6S = -bSI + omega * R6
        k6E = bSI - sigma * E6
        k6I = sigma * E6 - gamma * I6
        k6R = gamma * I6 - omega * R6

        # 5th order
        y5S = S + h * (35.0/384.0*k1S + 500.0/1113.0*k3S + 125.0/192.0*k4S - 2187.0/6784.0*k5S + 11.0/84.0*k6S)
        y5E = E + h * (35.0/384.0*k1E + 500.0/1113.0*k3E + 125.0/192.0*k4E - 2187.0/6784.0*k5E + 11.0/84.0*k6E)
        y5I = I + h * (35.0/384.0*k1I + 500.0/1113.0*k3I + 125.0/192.0*k4I - 2187.0/6784.0*k5I + 11.0/84.0*k6I)
        y5R = R + h * (35.0/384.0*k1R + 500.0/1113.0*k3R + 125.0/192.0*k4R - 2187.0/6784.0*k5R + 11.0/84.0*k6R)

        # k7 from 5th order solution
        bSI = beta * y5S * y5I
        k7S = -bSI + omega * y5R
        k7E = bSI - sigma * y5E
        k7I = sigma * y5E - gamma * y5I
        k7R = gamma * y5I - omega * y5R

        # 4th order for error estimation
        y4S = S + h * (5179.0/57600.0*k1S + 7571.0/16695.0*k3S + 393.0/640.0*k4S - 92097.0/339200.0*k5S + 187.0/2100.0*k6S + 1.0/40.0*k7S)
        y4E = E + h * (5179.0/57600.0*k1E + 7571.0/16695.0*k3E + 393.0/640.0*k4E - 92097.0/339200.0*k5E + 187.0/2100.0*k6E + 1.0/40.0*k7E)
        y4I = I + h * (5179.0/57600.0*k1I + 7571.0/16695.0*k3I + 393.0/640.0*k4I - 92097.0/339200.0*k5I + 187.0/2100.0*k6I + 1.0/40.0*k7I)
        y4R = R + h * (5179.0/57600.0*k1R + 7571.0/16695.0*k3R + 393.0/640.0*k4R - 92097.0/339200.0*k5R + 187.0/2100.0*k6R + 1.0/40.0*k7R)

        err = max(abs(y5S - y4S), abs(y5E - y4E), abs(y5I - y4I), abs(y5R - y4R))

        if err < 1e-15:
            S = y5S; E = y5E; I = y5I; R = y5R
            t += h
            h = min(h * 2.0, t1 - t) if t < t1 else h
        else:
            ratio = tol / err
            if ratio >= 1.0:
                S = y5S; E = y5E; I = y5I; R = y5R
                t += h
                h *= min(2.0, max(0.2, 0.9 * ratio ** 0.2))
            else:
                h *= max(0.2, 0.9 * ratio ** 0.2)

    return np.array([S, E, I, R])


class Solver:
    def __init__(self):
        # Warmup: trigger Numba JIT compilation
        _seirs_rk4(
            np.array([0.89, 0.01, 0.005, 0.095]),
            0.0, 10.0, 0.35, 0.2, 0.1, 0.002, 1.0,
        )
        _seirs_dp45(
            np.array([0.89, 0.01, 0.005, 0.095]),
            0.0, 10.0, 0.35, 0.2, 0.1, 0.002, 1e-5,
        )

    def solve(self, problem: dict, **kwargs) -> list[float]:
        y0 = np.asarray(problem["y0"], dtype=np.float64)
        params = problem["params"]
        t0 = float(problem["t0"])
        t1 = float(problem["t1"])
        beta = float(params["beta"])
        sigma = float(params["sigma"])
        gamma = float(params["gamma"])
        omega = float(params["omega"])

        max_rate = max(beta, sigma, gamma)

        if max_rate <= 0.45:
            # Fast path: fixed-step RK4 with large dt
            dt = 3.0 if max_rate <= 0.25 else 2.0
            result = _seirs_rk4(y0, t0, t1, beta, sigma, gamma, omega, dt)
        else:
            # Use adaptive DP45 for faster dynamics
            result = _seirs_dp45(y0, t0, t1, beta, sigma, gamma, omega, 1e-5)

        return result.tolist()
--- 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.3205s
Total Solver Time:   0.0006s
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=4.82e-06, max rel err=53.3
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 5.74s ==========================
