Use pytest parametrize for better coverage of forbidden commands and reserved names
This commit is contained in:
parent
fa9ea77812
commit
35efba19b8
3 changed files with 49 additions and 36 deletions
|
|
@ -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."
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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():
|
||||||
|
|
|
||||||
|
|
@ -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():
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue