diff --git a/pysmspp/smspp_tools.py b/pysmspp/smspp_tools.py index 8bf7224..5596bd8 100644 --- a/pysmspp/smspp_tools.py +++ b/pysmspp/smspp_tools.py @@ -233,6 +233,59 @@ def help(self, print_message=True): print(msg) return msg + def version(self, option="--version", print_message=True): + """ + Return the semantic version reported by the SMS++ solver. + + Parameters + ---------- + option : str, optional + Option to query the tool version, by default "--version". + print_message : bool, optional + Whether to print the raw version output from the solver, by default + True. + + Returns + ------- + str + The semantic version (e.g. "0.7.1") parsed from the tool output. + + Raises + ------ + ValueError + If the tool output does not contain a parseable semantic version. + """ + + def _run(option): + command = [self._solver_path, option] + if self._shell: + command = f"{self._solver_path} {option}" + result = subprocess.run( + command, + capture_output=True, + shell=self._shell, + check=False, + ) + stdout = result.stdout + stderr = result.stderr + if isinstance(stdout, bytes): + stdout = stdout.decode("utf-8") + if isinstance(stderr, bytes): + stderr = stderr.decode("utf-8") + return str(stdout) + os.linesep + str(stderr) + + msg = _run(option) + res = re.search(r"\d+(?:\.\d+){1,}", msg) + + if res is None: + raise ValueError( + f"Failed to parse version from {self._solver_path} output:\n{msg}" + ) + + if print_message: + print(msg) + return res.group() + def optimize(self, logging=True, tracking_period=0.1): """ Run the SMSPP Solver tool. diff --git a/test/test_smspp_tools.py b/test/test_smspp_tools.py index 3d009e4..9f20aa6 100644 --- a/test/test_smspp_tools.py +++ b/test/test_smspp_tools.py @@ -3,6 +3,7 @@ import numpy as np import pytest +import pysmspp.smspp_tools as smspp_tools_module from pysmspp import InvestmentBlockSolver, SMSPPSolverTool, UCBlockSolver @@ -20,6 +21,41 @@ def calculate_executable_call(self): return [sys.executable, "-c", code] +def test_version_reads_solver_version(): + solver = UCBlockSolver(solver_path=sys.executable) + version = solver.version(print_message=False) + assert version.count(".") >= 1 + + +def test_version_with_custom_option(): + solver = UCBlockSolver(solver_path=sys.executable) + version = solver.version(option="-V", print_message=False) + assert version.count(".") >= 1 + + +def test_version_raises_if_output_is_not_a_semantic_version(monkeypatch): + def fake_run(*args, **kwargs): + return smspp_tools_module.subprocess.CompletedProcess( + args=["ucblock_solver", "--version"], + returncode=0, + stdout="no semantic version here\n", + stderr="", + ) + + monkeypatch.setattr(smspp_tools_module.subprocess, "run", fake_run) + + solver = UCBlockSolver(solver_path="ucblock_solver") + with pytest.raises(ValueError, match="Failed to parse version"): + solver.version(print_message=False) + + +def test_version_supports_shell_solver_commands(): + code = "print('SMS++ tools version 0.7.1')" + solver_cmd = f'{sys.executable} -c "{code}"' + solver = UCBlockSolver(solver_path=solver_cmd, shell=True) + assert solver.version(print_message=False) == "0.7.1" + + def test_optimize_reads_subprocess_output_portably(tmp_path): fp_network = tmp_path / "network.nc4" fp_config = tmp_path / "config.txt"