from __future__ import annotations import asyncio import os import unittest from unittest.mock import patch from app.main import _app_lifespan, app from app.simulation.warmup import ( SimulationWarmupReport, _run_numerical_warmup, _reset_simulation_warmup_for_tests, warm_up_simulation_runtime, ) class SimulationWarmupTests(unittest.TestCase): def setUp(self) -> None: _reset_simulation_warmup_for_tests() def tearDown(self) -> None: _reset_simulation_warmup_for_tests() def test_warmup_runs_only_once_per_process(self) -> None: with ( patch.dict(os.environ, {"SIMULATIONAPP_WARMUP": "on"}), patch("app.simulation.warmup._run_numerical_warmup") as run_warmup, ): first = warm_up_simulation_runtime() second = warm_up_simulation_runtime() self.assertIs(second, first) self.assertEqual(first.status, "completed") run_warmup.assert_called_once_with() def test_disabled_warmup_does_not_touch_numerical_runtime(self) -> None: with ( patch.dict(os.environ, {"SIMULATIONAPP_WARMUP": "off"}), patch("app.simulation.warmup._run_numerical_warmup") as run_warmup, ): report = warm_up_simulation_runtime() self.assertEqual(report.status, "disabled") run_warmup.assert_not_called() def test_regular_failure_is_reported_without_blocking_startup(self) -> None: with ( patch.dict(os.environ, {"SIMULATIONAPP_WARMUP": "on"}), patch( "app.simulation.warmup._run_numerical_warmup", side_effect=RuntimeError("broken warmup"), ), self.assertLogs("app.simulation.warmup", level="ERROR"), ): report = warm_up_simulation_runtime() self.assertEqual(report.status, "failed") self.assertIn("broken warmup", report.error or "") def test_memory_error_remains_fatal(self) -> None: with ( patch.dict(os.environ, {"SIMULATIONAPP_WARMUP": "on"}), patch( "app.simulation.warmup._run_numerical_warmup", side_effect=MemoryError("out of memory"), ), self.assertRaises(MemoryError), ): warm_up_simulation_runtime() def test_lifespan_stores_warmup_report_before_serving(self) -> None: report = SimulationWarmupReport(status="completed", duration_ms=12.5) async def enter_lifespan() -> None: with patch( "app.simulation.warmup.warm_up_simulation_runtime", return_value=report, ) as warmup: async with _app_lifespan(app): self.assertEqual( app.state.simulation_warmup, report.as_dict(), ) warmup.assert_called_once_with() asyncio.run(enter_lifespan()) def test_real_numerical_warmup_completes(self) -> None: with patch.dict(os.environ, {"SIMULATIONAPP_WARMUP": "on"}): report = warm_up_simulation_runtime() self.assertEqual(report.status, "completed", report.error) self.assertGreater(report.duration_ms, 0.0) def test_numerical_warmup_exercises_sparse_lsmr_algebraic_path(self) -> None: import scipy.optimize actual_least_squares = scipy.optimize.least_squares optimizer_calls: list[dict[str, object]] = [] def recording_least_squares(*args, **kwargs): optimizer_calls.append(dict(kwargs)) return actual_least_squares(*args, **kwargs) with patch.object( scipy.optimize, "least_squares", recording_least_squares, ): _run_numerical_warmup() self.assertEqual(len(optimizer_calls), 1) call = optimizer_calls[0] self.assertEqual(call["tr_solver"], "lsmr") sparsity = call["jac_sparsity"] self.assertEqual(sparsity.shape, (2, 2)) self.assertEqual(sparsity.nnz, 2) if __name__ == "__main__": unittest.main()