--- Final Solver (used for performance test) ---
import numpy as np
from scipy import signal
from scipy.linalg import expm

try:
    from numba import njit
    _NUMBA = True
except Exception:
    _NUMBA = False


if _NUMBA:
    @njit(cache=True, fastmath=True)
    def _run(Ad, Bd0, Bd1, Cv, D, u):
        N = u.shape[0]
        n = Cv.shape[0]
        y = np.empty(N)
        x = np.zeros(n)
        xn = np.zeros(n)
        y[0] = D * u[0]
        for i in range(1, N):
            a = u[i - 1]
            b = u[i]
            for r in range(n):
                s = 0.0
                for c in range(n):
                    s += Ad[r, c] * x[c]
                xn[r] = s + Bd0[r] * a + Bd1[r] * b
            yy = 0.0
            for r in range(n):
                yy += Cv[r] * xn[r]
                x[r] = xn[r]
            y[i] = yy + D * b
        return y


def _my_tf2ss(num, den):
    num = np.atleast_1d(np.asarray(num, dtype=float))
    den = np.atleast_1d(np.asarray(den, dtype=float))
    # strip leading zeros of den
    nz = np.nonzero(den)[0]
    if nz.size == 0:
        raise ValueError("bad den")
    den = den[nz[0]:]
    d0 = den[0]
    den = den / d0
    num = num / d0
    K = den.shape[0]
    M = num.shape[0]
    if M > K:
        raise ValueError("improper")
    n = K - 1
    # pad num on left to length K
    if M < K:
        nump = np.zeros(K)
        nump[K - M:] = num
    else:
        nump = num
    D = nump[0]
    if n == 0:
        return None, None, None, D, 0
    A = np.zeros((n, n))
    A[0, :] = -den[1:]
    if n > 1:
        A[1:, :-1] = np.eye(n - 1)
    Cv = nump[1:] - nump[0] * den[1:]
    return A, Cv, den, D, n


class Solver:
    def __init__(self):
        if _NUMBA:
            Ad = np.eye(2)
            v = np.ones(2)
            u = np.ones(4)
            _run(Ad, v, v, v, 0.0, u)

    def solve(self, problem, **kwargs):
        num = problem["num"]
        den = problem["den"]
        ul = problem["u"]
        tl = problem["t"]
        N = len(tl)
        if N == 0:
            return {"yout": []}
        try:
            return self._fast(num, den, ul, tl, N)
        except Exception:
            system = signal.lti(num, den)
            t = np.asarray(tl, dtype=float)
            u = np.asarray(ul, dtype=float)
            tout, yout, xout = signal.lsim(system, u, t)
            return {"yout": np.atleast_1d(yout).astype(float).tolist()}

    def _fast(self, num, den, ul, tl, N):
        A, Cv, denn, D, n = _my_tf2ss(num, den)
        u = np.asarray(ul, dtype=float)
        if n == 0:
            return {"yout": (D * u).tolist()}
        if N == 1:
            return {"yout": [D * float(u[0])]}
        dt = float(tl[1]) - float(tl[0])
        # B = e1 in controllable canonical form
        M = np.zeros((n + 2, n + 2))
        M[:n, :n] = A * dt
        M[0, n] = dt  # B*dt, B=e1
        M[n, n + 1] = 1.0
        eM = expm(M)
        Ad = np.ascontiguousarray(eM[:n, :n])
        Bd1 = np.ascontiguousarray(eM[:n, n + 1])
        Bd0 = np.ascontiguousarray(eM[:n, n] - Bd1)
        Cvc = np.ascontiguousarray(Cv)
        y = _run(Ad, Bd0, Bd1, Cvc, float(D), u)
        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: 9.0638s
Total Solver Time:   0.0681s
Raw Speedup:         133.0523 x
Final Reward (Score): 133.0523
---------------------------
..

==================================== 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 113.44s (0:01:53) =========================
