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

@numba.njit
def _rk4_step(y, beta, sigma, gamma, omega, dt):
    S, E, I, R = y[0], y[1], y[2], y[3]
    
    # k1
    dS1 = -beta * S * I + omega * R
    dE1 = beta * S * I - sigma * E
    dI1 = sigma * E - gamma * I
    dR1 = gamma * I - omega * R
    
    # k2
    S2 = S + dt/2 * dS1
    E2 = E + dt/2 * dE1
    I2 = I + dt/2 * dI1
    R2 = R + dt/2 * dR1
    dS2 = -beta * S2 * I2 + omega * R2
    dE2 = beta * S2 * I2 - sigma * E2
    dI2 = sigma * E2 - gamma * I2
    dR2 = gamma * I2 - omega * R2
    
    # k3
    S3 = S + dt/2 * dS2
    E3 = E + dt/2 * dE2
    I3 = I + dt/2 * dI2
    R3 = R + dt/2 * dR2
    dS3 = -beta * S3 * I3 + omega * R3
    dE3 = beta * S3 * I3 - sigma * E3
    dI3 = sigma * E3 - gamma * I3
    dR3 = gamma * I3 - omega * R3
    
    # k4
    S4 = S + dt * dS3
    E4 = E + dt * dE3
    I4 = I + dt * dI3
    R4 = R + dt * dR3
    dS4 = -beta * S4 * I4 + omega * R4
    dE4 = beta * S4 * I4 - sigma * E4
    dI4 = sigma * E4 - gamma * I4
    dR4 = gamma * I4 - omega * R4
    
    # Update
    y[0] = S + (dt/6) * (dS1 + 2*dS2 + 2*dS3 + dS4)
    y[1] = E + (dt/6) * (dE1 + 2*dE2 + 2*dE3 + dE4)
    y[2] = I + (dt/6) * (dI1 + 2*dI2 + 2*dI3 + dI4)
    y[3] = R + (dt/6) * (dR1 + 2*dR2 + 2*dR3 + dR4)
    
    return y

@numba.njit
def _solve_numba(y0, t0, t1, beta, sigma, gamma, omega, dt):
    y = y0.copy()
    t = t0
    while t < t1 - 1e-12:
        current_dt = dt
        if t + dt > t1:
            current_dt = t1 - t
        y = _rk4_step(y, beta, sigma, gamma, omega, current_dt)
        t += current_dt
    # Conservation correction
    total = y[0] + y[1] + y[2] + y[3]
    y[0] /= total
    y[1] /= total
    y[2] /= total
    y[3] /= total
    return y

class Solver:
    def __init__(self):
        # Warm up numba JIT compilation
        y0 = np.array([0.89, 0.01, 0.005, 0.095])
        _solve_numba(y0, 0.0, 1.0, 0.35, 0.2, 0.1, 0.002, 0.5)
    
    def solve(self, problem: dict[str, Any]) -> list[float]:
        t0 = problem['t0']
        t1 = problem['t1']
        y0 = np.array(problem['y0'])
        params = problem['params']
        beta = params['beta']
        sigma = params['sigma']
        gamma = params['gamma']
        omega = params['omega']
        
        # Compute step size
        max_rate = max(beta, sigma, gamma, omega)
        if max_rate > 0:
            fastest_timescale = 1.0 / max_rate
            dt_timescale = fastest_timescale * 0.5  # half the fastest timescale
        else:
            dt_timescale = 1e10  # no dynamics
        
        total_time = t1 - t0
        dt_steps = total_time / 10.0  # ensure at least 10 steps
        dt = min(2.0, dt_timescale, dt_steps)
        if dt <= 0:
            dt = 1e-6
        
        result = _solve_numba(y0, t0, t1, beta, sigma, gamma, omega, dt)
        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: True
Total Baseline Time: 8.6152s
Total Solver Time:   0.0223s
Raw Speedup:         385.5806 x
Final Reward (Score): 385.5806
---------------------------
..

==================================== 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 117.16s (0:01:57) =========================
