--- Final Solver (used for performance test) ---
from typing import Any
import numpy as np
from scipy import signal, linalg
import numba
from functools import lru_cache


class Solver:
    def __init__(self):
        # Compile Numba kernels during initialization so JIT overhead is not counted
        self._manual = self._make_manual()
        self._blas = self._make_blas()
        # Warm-up calls to trigger compilation for the type signatures we will use
        for n_sz in (1, 2, 4, 8, 16, 32, 64):
            dummy_u = np.zeros(10, dtype=float)
            dummy_A = np.zeros((n_sz, n_sz), dtype=float)
            dummy_b = np.zeros(n_sz, dtype=float)
            dummy_c = np.zeros(n_sz, dtype=float)
            self._manual(dummy_c, 0.0, dummy_A, dummy_b, dummy_b, dummy_u)
            self._blas(dummy_c, 0.0, dummy_A, dummy_b, dummy_b, dummy_u)

    @staticmethod
    def _make_manual():
        @numba.njit(cache=False, fastmath=True)
        def _solve(C, D, Ad_r, Bd0_r, Bd1_r, u):
            n = C.shape[0]
            N = len(u)
            x = np.zeros(n)
            y = np.empty(N)
            y[0] = D * u[0]
            tmp = np.zeros(n)
            for i in range(1, N):
                # x_new = x @ Ad_r + u[i-1] * Bd0_r + u[i] * Bd1_r
                for j in range(n):
                    s = 0.0
                    for k in range(n):
                        s += x[k] * Ad_r[k, j]
                    s += u[i - 1] * Bd0_r[j] + u[i] * Bd1_r[j]
                    tmp[j] = s
                for j in range(n):
                    x[j] = tmp[j]
                s = D * u[i]
                for j in range(n):
                    s += C[j] * x[j]
                y[i] = s
            return y
        return _solve

    @staticmethod
    def _make_blas():
        @numba.njit(cache=False, fastmath=True)
        def _solve(C, D, Ad_r, Bd0_r, Bd1_r, u):
            n = Ad_r.shape[0]
            N = len(u)
            x = np.zeros(n)
            y = np.empty(N)
            y[0] = D * u[0]
            for i in range(1, N):
                x = x @ Ad_r + u[i - 1] * Bd0_r + u[i] * Bd1_r
                y[i] = np.dot(C, x) + D * u[i]
            return y
        return _solve

    @lru_cache(maxsize=1024)
    def _get_matrices(self, num_t, den_t, dt):
        A, B, C, D = signal.tf2ss(list(num_t), list(den_t))
        A = np.asarray(A, dtype=float)
        B = np.asarray(B, dtype=float)
        C = np.asarray(C, dtype=float)
        D = np.asarray(D, dtype=float)
        n = A.shape[0]
        n_inputs = B.shape[1]

        # Build augmented matrix for first-order-hold (FOH) discretization.
        # This matches the internal algorithm used by scipy.signal.lsim when
        # interp is truthy (the default interp='zoh' is truthy and triggers FOH).
        M = np.zeros((n + 2 * n_inputs, n + 2 * n_inputs), dtype=float)
        M[:n, :n] = A * dt
        M[:n, n:n + n_inputs] = B * dt
        M[n:n + n_inputs, n + n_inputs:n + 2 * n_inputs] = np.eye(n_inputs)

        expMT = linalg.expm(M.T)
        Ad_r = np.ascontiguousarray(expMT[:n, :n])
        Bd1_r = expMT[n + n_inputs:n + 2 * n_inputs, :n]
        Bd0_r = expMT[n:n + n_inputs, :n] - Bd1_r

        C1 = C.ravel()
        D1 = float(np.squeeze(D))
        return Ad_r, Bd0_r.ravel(), Bd1_r.ravel(), C1, D1

    def solve(self, problem, **kwargs) -> dict[str, list[float]]:
        num = problem["num"]
        den = problem["den"]
        u_arr = np.asarray(problem["u"], dtype=float)
        t_arr = np.asarray(problem["t"], dtype=float)
        n_steps = t_arr.size

        if n_steps == 0:
            return {"yout": []}

        if n_steps == 1:
            _, _, _, D_val = signal.tf2ss(num, den)
            return {"yout": [float(float(np.squeeze(D_val)) * float(u_arr[0]))]}

        dt = float(t_arr[1] - t_arr[0])
        # lsim raises ValueError when t is not evenly spaced. Match that check.
        if not np.allclose(np.diff(t_arr), dt):
            system = signal.lti(num, den)
            _, yout, _ = signal.lsim(system, u_arr, t_arr)
            return {"yout": yout.tolist()}

        Ad_r, Bd0_r, Bd1_r, C1, D1 = self._get_matrices(
            tuple(num), tuple(den), dt
        )

        n = C1.shape[0]
        if n <= 16:
            y = self._manual(C1, D1, Ad_r, Bd0_r, Bd1_r, u_arr)
        else:
            y = self._blas(C1, D1, Ad_r, Bd0_r, Bd1_r, u_arr)

        return {"yout": y.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.8810s
Total Solver Time:   0.0658s
Raw Speedup:         135.0577 x
Final Reward (Score): 135.0577
---------------------------
..

==================================== 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 110.90s (0:01:50) =========================
