Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 53 additions & 0 deletions pysmspp/smspp_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
36 changes: 36 additions & 0 deletions test/test_smspp_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import numpy as np
import pytest

import pysmspp.smspp_tools as smspp_tools_module
from pysmspp import InvestmentBlockSolver, SMSPPSolverTool, UCBlockSolver


Expand All @@ -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"
Expand Down