--- Final Solver (used for performance test) ---
from typing import Any, Dict
from numba import njit

@njit(fastmath=True, cache=True)
def rk4_core(S, E, I, beta, sigma, gamma, omega, dt, steps):
    for _ in range(steps):
        # k1
        bSI = beta * S * I
        oR = omega * (1.0 - S - E - I)
        sE = sigma * E
        gI = gamma * I
        
        dS1 = -bSI + oR
        dE1 = bSI - sE
        dI1 = sE - gI
        
        S2 = S + 0.5 * dt * dS1
        E2 = E + 0.5 * dt * dE1
        I2 = I + 0.5 * dt * dI1
        
        # k2
        bSI2 = beta * S2 * I2
        oR2 = omega * (1.0 - S2 - E2 - I2)
        sE2 = sigma * E2
        gI2 = gamma * I2
        
        dS2 = -bSI2 + oR2
        dE2 = bSI2 - sE2
        dI2 = sE2 - gI2
        
        S3 = S + 0.5 * dt * dS2
        E3 = E + 0.5 * dt * dE2
        I3 = I + 0.5 * dt * dI2
        
        # k3
        bSI3 = beta * S3 * I3
        oR3 = omega * (1.0 - S3 - E3 - I3)
        sE3 = sigma * E3
        gI3 = gamma * I3
        
        dS3 = -bSI3 + oR3
        dE3 = bSI3 - sE3
        dI3 = sE3 - gI3
        
        S4 = S + dt * dS3
        E4 = E + dt * dE3
        I4 = I + dt * dI3
        
        # k4
        bSI4 = beta * S4 * I4
        oR4 = omega * (1.0 - S4 - E4 - I4)
        sE4 = sigma * E4
        gI4 = gamma * I4
        
        dS4 = -bSI4 + oR4
        dE4 = bSI4 - sE4
        dI4 = sE4 - gI4
        
        S += dt * (dS1 + 2*dS2 + 2*dS3 + dS4) / 6.0
        E += dt * (dE1 + 2*dE2 + 2*dE3 + dE4) / 6.0
        I += dt * (dI1 + 2*dI2 + 2*dI3 + dI4) / 6.0
        
    return S, E, I, 1.0 - S - E - I

class Solver:
    def __init__(self):
        _ = rk4_core(0.89, 0.01, 0.005, 0.35, 0.2, 0.1, 0.002, 2.3, 174)
        
    def solve(self, problem: Dict[str, Any], **kwargs) -> Any:
        y0 = problem["y0"]
        params = problem["params"]
        
        beta = params["beta"]
        sigma = params["sigma"]
        gamma = params["gamma"]
        omega = params["omega"]
        
        # Dynamically scale step size for stability based on spectral radius approximation
        rate_sum = beta + sigma + gamma + omega + 1e-9
        dt_target = 1.5 / rate_sum
        
        total_time = problem["t1"] - problem["t0"]
        steps = int(total_time / dt_target) + 1
        dt = total_time / steps
        
        S, E, I, R = rk4_core(y0[0], y0[1], y0[2], beta, sigma, gamma, omega, dt, steps)
        
        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.2527s
Total Solver Time:   0.0115s
Raw Speedup:         717.4109 x
Final Reward (Score): 717.4109
---------------------------
..

==================================== 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 106.09s (0:01:46) =========================
