diff --git a/src/openflexure_microscope_server/stitching.py b/src/openflexure_microscope_server/stitching.py index bf5da43c..393eb32e 100644 --- a/src/openflexure_microscope_server/stitching.py +++ b/src/openflexure_microscope_server/stitching.py @@ -29,6 +29,22 @@ STITCH_TILE_SIZE = 8192 DEFAULT_OVERLAP = 0.1 DEFAULT_RESIZE = 0.5 +# A list of commands that are forbidden in any part of a generated CLI command. +# This provides defense-in-depth against trying to execute arbitrary shells +# or elevation tools. +FORBIDDEN_COMMANDS = { + "sudo", + "sh", + "bash", + "perl", + "ruby", + "php", + "nc", + "netcat", + "curl", + "wget", +} + class StitchingSettings(BaseModel): """The data needed to stitch a scan.""" @@ -48,23 +64,6 @@ class StitcherValidationError(RuntimeError): """The stitcher received values that it deems unsafe to create a command from.""" -# A list of commands that are forbidden in any part of a generated CLI command. -# This provides defense-in-depth against trying to execute arbitrary shells -# or elevation tools. -_FORBIDDEN_COMMANDS = { - "sudo", - "sh", - "bash", - "perl", - "ruby", - "php", - "nc", - "netcat", - "curl", - "wget", -} - - def validate_command(cmd: list[str]) -> None: """Validate that the command only contains characters that are allowed in a path. @@ -77,7 +76,7 @@ def validate_command(cmd: list[str]) -> None: """ for element in cmd: # Check against forbidden commands (case-insensitive) - if element.lower() in _FORBIDDEN_COMMANDS: + if element.lower() in FORBIDDEN_COMMANDS: raise StitcherValidationError( f"Invalid stiching command: Forbidden element '{element}' detected." ) diff --git a/tests/unit_tests/test_stitching.py b/tests/unit_tests/test_stitching.py index 34d7fe7d..ca5b9a08 100644 --- a/tests/unit_tests/test_stitching.py +++ b/tests/unit_tests/test_stitching.py @@ -16,6 +16,7 @@ from pydantic import BaseModel import labthings_fastapi as lt from openflexure_microscope_server.stitching import ( + FORBIDDEN_COMMANDS, BaseStitcher, FinalStitcher, PreviewStitcher, @@ -39,22 +40,20 @@ def test_validate_command_success(): validate_command(["--resize", "0.5", "8192"]) -def test_validate_command_forbidden(): +@pytest.mark.parametrize("cmd_name", FORBIDDEN_COMMANDS) +def test_validate_command_forbidden(cmd_name): """Test forbidden commands raise error.""" + # As the main command with pytest.raises( - StitcherValidationError, match="Forbidden element 'sudo' detected" + StitcherValidationError, match=f"Forbidden element '{cmd_name}' detected" ): - validate_command(["sudo", "rm", "-rf", "/"]) + validate_command([cmd_name, "some_arg"]) + # As an argument (case-insensitive) with pytest.raises( - StitcherValidationError, match="Forbidden element 'sh' detected" + StitcherValidationError, match=f"Forbidden element '{cmd_name.upper()}' detected" ): - validate_command(["sh", "-c", "whoami"]) - - with pytest.raises( - StitcherValidationError, match="Forbidden element 'SUDO' detected" - ): - validate_command(["SUDO", "ls"]) + validate_command(["safe-command", cmd_name.upper()]) def test_base_stitcher(): diff --git a/tests/unit_tests/test_utilities.py b/tests/unit_tests/test_utilities.py index a9972a44..dc8e3ae1 100644 --- a/tests/unit_tests/test_utilities.py +++ b/tests/unit_tests/test_utilities.py @@ -2,10 +2,17 @@ import sys -from openflexure_microscope_server.utilities import make_name_safe, make_path_safe +import pytest + +from openflexure_microscope_server.utilities import ( + _WINDOWS_RESERVED_NAMES, + make_name_safe, + make_path_safe, +) def test_make_name_safe_basic(): + """Test basic functionality of make_name_safe.""" assert make_name_safe("normal_name") == "normal_name" assert make_name_safe("name with spaces") == "name_with_spaces" @@ -24,16 +31,24 @@ def test_make_name_safe_trailing_chars(): assert make_name_safe(". ") == "_" -def test_make_name_safe_reserved_names(): +@pytest.mark.parametrize("name", _WINDOWS_RESERVED_NAMES) +def test_make_name_safe_reserved_names(name): """Test Windows reserved names.""" - assert make_name_safe("CON") == "CON_" - assert make_name_safe("con") == "con_" - assert make_name_safe("NUL.txt") == "NUL.txt_" - assert make_name_safe("COM1") == "COM1_" - assert make_name_safe("LPT9.tar.gz") == "LPT9.tar.gz_" - # Check that names STARTING with reserved names but not followed by . or end are safe + # Base name should be sanitized + assert make_name_safe(name) == f"{name}_" + # Case-insensitive + assert make_name_safe(name.lower()) == f"{name.lower()}_" + # With extension + assert make_name_safe(f"{name}.txt") == f"{name}.txt_" + # Multiple extensions + assert make_name_safe(f"{name}.tar.gz") == f"{name}.tar.gz_" + + +def test_make_name_safe_reserved_names_false_positives(): + """Test that names containing but not equal to reserved names are safe.""" assert make_name_safe("CONSTANT") == "CONSTANT" assert make_name_safe("CON2") == "CON2" + assert make_name_safe("ICON") == "ICON" def test_make_path_safe_basic():