Use pytest parametrize for better coverage of forbidden commands and reserved names

This commit is contained in:
Beth Probert 2026-03-05 13:55:05 +00:00
parent fa9ea77812
commit 35efba19b8
3 changed files with 49 additions and 36 deletions

View file

@ -29,6 +29,22 @@ STITCH_TILE_SIZE = 8192
DEFAULT_OVERLAP = 0.1 DEFAULT_OVERLAP = 0.1
DEFAULT_RESIZE = 0.5 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): class StitchingSettings(BaseModel):
"""The data needed to stitch a scan.""" """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.""" """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: def validate_command(cmd: list[str]) -> None:
"""Validate that the command only contains characters that are allowed in a path. """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: for element in cmd:
# Check against forbidden commands (case-insensitive) # Check against forbidden commands (case-insensitive)
if element.lower() in _FORBIDDEN_COMMANDS: if element.lower() in FORBIDDEN_COMMANDS:
raise StitcherValidationError( raise StitcherValidationError(
f"Invalid stiching command: Forbidden element '{element}' detected." f"Invalid stiching command: Forbidden element '{element}' detected."
) )

View file

@ -16,6 +16,7 @@ from pydantic import BaseModel
import labthings_fastapi as lt import labthings_fastapi as lt
from openflexure_microscope_server.stitching import ( from openflexure_microscope_server.stitching import (
FORBIDDEN_COMMANDS,
BaseStitcher, BaseStitcher,
FinalStitcher, FinalStitcher,
PreviewStitcher, PreviewStitcher,
@ -39,22 +40,20 @@ def test_validate_command_success():
validate_command(["--resize", "0.5", "8192"]) 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.""" """Test forbidden commands raise error."""
# As the main command
with pytest.raises( 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( with pytest.raises(
StitcherValidationError, match="Forbidden element 'sh' detected" StitcherValidationError, match=f"Forbidden element '{cmd_name.upper()}' detected"
): ):
validate_command(["sh", "-c", "whoami"]) validate_command(["safe-command", cmd_name.upper()])
with pytest.raises(
StitcherValidationError, match="Forbidden element 'SUDO' detected"
):
validate_command(["SUDO", "ls"])
def test_base_stitcher(): def test_base_stitcher():

View file

@ -2,10 +2,17 @@
import sys 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(): def test_make_name_safe_basic():
"""Test basic functionality of make_name_safe.""" """Test basic functionality of make_name_safe."""
assert make_name_safe("normal_name") == "normal_name" assert make_name_safe("normal_name") == "normal_name"
assert make_name_safe("name with spaces") == "name_with_spaces" 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(". ") == "_" 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.""" """Test Windows reserved names."""
assert make_name_safe("CON") == "CON_" # Base name should be sanitized
assert make_name_safe("con") == "con_" assert make_name_safe(name) == f"{name}_"
assert make_name_safe("NUL.txt") == "NUL.txt_" # Case-insensitive
assert make_name_safe("COM1") == "COM1_" assert make_name_safe(name.lower()) == f"{name.lower()}_"
assert make_name_safe("LPT9.tar.gz") == "LPT9.tar.gz_" # With extension
# Check that names STARTING with reserved names but not followed by . or end are safe 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("CONSTANT") == "CONSTANT"
assert make_name_safe("CON2") == "CON2" assert make_name_safe("CON2") == "CON2"
assert make_name_safe("ICON") == "ICON"
def test_make_path_safe_basic(): def test_make_path_safe_basic():