--- Final Solver (used for performance test) ---
from __future__ import annotations

from typing import Any
import numpy as np

try:
    import _simplex_ext as _simp_ext
    _ext_project = _simp_ext.project
except Exception:
    _ext_project = None


def _project_sort_numpy(a: np.ndarray) -> np.ndarray:
    n = a.size
    if n == 0:
        return np.empty(0, dtype=np.float64)
    if n == 1:
        return np.array([1.0], dtype=np.float64)
    u = np.sort(a)[::-1]
    css = np.cumsum(u)
    cond = u > (css - 1.0) / np.arange(1, n + 1)
    rho = int(np.count_nonzero(cond))
    theta = (css[rho - 1] - 1.0) / rho
    return np.maximum(a - theta, 0.0)


def _project_tiny_list(y) -> np.ndarray:
    n = len(y)
    if n == 0:
        return np.empty(0, dtype=np.float64)
    if n == 1:
        return np.array([1.0], dtype=np.float64)
    a0 = float(y[0])
    a1 = float(y[1])
    x0 = 0.5 * (a0 - a1 + 1.0)
    if x0 <= 0.0:
        return np.array([0.0, 1.0], dtype=np.float64)
    if x0 >= 1.0:
        return np.array([1.0, 0.0], dtype=np.float64)
    return np.array([x0, 1.0 - x0], dtype=np.float64)


class Solver:
    def __init__(self) -> None:
        self._ext_project = _ext_project

    def solve(self, problem: dict[str, Any], **kwargs) -> dict[str, Any]:
        y = problem.get("y")

        if self._ext_project is not None:
            try:
                return {"solution": self._ext_project(y)}
            except Exception:
                pass

        ty = type(y)
        if ty is list or ty is tuple:
            n = len(y)
            if n <= 2:
                try:
                    return {"solution": _project_tiny_list(y)}
                except Exception:
                    pass

        a = np.asarray(y, dtype=np.float64).ravel()
        n = a.size
        if n == 1:
            return {"solution": np.array([1.0], dtype=np.float64)}
        if n == 2:
            a0 = float(a[0])
            a1 = float(a[1])
            x0 = 0.5 * (a0 - a1 + 1.0)
            if x0 <= 0.0:
                return {"solution": np.array([0.0, 1.0], dtype=np.float64)}
            if x0 >= 1.0:
                return {"solution": np.array([1.0, 0.0], dtype=np.float64)}
            return {"solution": np.array([x0, 1.0 - x0], dtype=np.float64)}

        return {"solution": _project_sort_numpy(a)}
--- 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: 5.9854s
Total Solver Time:   5.8650s
Raw Speedup:         1.0205 x
Final Reward (Score): 1.0205
---------------------------
..

==================================== 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 281.82s (0:04:41) =========================
