diff --git a/nodes/image_node.py b/nodes/image_node.py index ef08893..b2667b7 100644 --- a/nodes/image_node.py +++ b/nodes/image_node.py @@ -362,7 +362,7 @@ class DiscordSendSaveImage(BaseDiscordNode): filepath = os.path.join(full_output_folder, file) # Security: Validate output path to prevent symlink overwrites - validate_path_is_safe(filepath) + validate_path_is_safe(filepath, base_dir=full_output_folder) try: # Save the image based on format diff --git a/nodes/video_node.py b/nodes/video_node.py index 323629a..43ffc3b 100644 --- a/nodes/video_node.py +++ b/nodes/video_node.py @@ -229,12 +229,12 @@ class DiscordSendSaveVideo(BaseDiscordNode): } @staticmethod - def _try_delete_old_file(old_path): + def _try_delete_old_file(old_path, base_dir=None): """Try to delete an old video file during overwrite. Uses safe error handling.""" if old_path and os.path.exists(old_path): try: # Validate path before deletion for security (defense-in-depth) - validate_path_is_safe(old_path) + validate_path_is_safe(old_path, base_dir=base_dir) os.remove(old_path) print(f"Deleted old file: {old_path}") except Exception as del_e: @@ -368,7 +368,7 @@ class DiscordSendSaveVideo(BaseDiscordNode): first_image_path = os.path.join(full_output_folder, first_image_file) # Security: Validate output path to prevent symlink overwrites - validate_path_is_safe(first_image_path) + validate_path_is_safe(first_image_path, base_dir=full_output_folder) Image.fromarray(tensor_to_numpy_uint8(first_image)).save( first_image_path, @@ -427,7 +427,7 @@ class DiscordSendSaveVideo(BaseDiscordNode): overwrite_delete_old = None # Security: Validate output path to prevent symlink overwrites - validate_path_is_safe(file_path) + validate_path_is_safe(file_path, base_dir=full_output_folder) # Check if we need ffmpeg (for video formats) or can use PIL (for image formats) use_pil = format_type == "image" or ffmpeg_path is None @@ -500,7 +500,7 @@ class DiscordSendSaveVideo(BaseDiscordNode): output_files.append(file_path) # Delete old file if overwriting with different extension - self._try_delete_old_file(overwrite_delete_old) + self._try_delete_old_file(overwrite_delete_old, base_dir=full_output_folder) else: print("No images to save") except Exception as e: @@ -597,7 +597,7 @@ class DiscordSendSaveVideo(BaseDiscordNode): file_path = os.path.join(full_output_folder, file) # Security: Validate output path to prevent symlink overwrites - validate_path_is_safe(file_path) + validate_path_is_safe(file_path, base_dir=full_output_folder) # Set up environment env = os.environ.copy() @@ -641,7 +641,7 @@ class DiscordSendSaveVideo(BaseDiscordNode): output_files.append(file_path) # Delete old file if overwriting with different extension - self._try_delete_old_file(overwrite_delete_old) + self._try_delete_old_file(overwrite_delete_old, base_dir=full_output_folder) except Exception as e: print(f"Error with VHS format encoding: {str(e)}") # Fall back to basic encoding if VHS format fails @@ -777,7 +777,7 @@ class DiscordSendSaveVideo(BaseDiscordNode): output_files.append(file_path) # Delete old file if overwriting with different extension - self._try_delete_old_file(overwrite_delete_old) + self._try_delete_old_file(overwrite_delete_old, base_dir=full_output_folder) except Exception as e: print(f"DiscordSendSaveVideo error: {str(e)}") discord_send_success = False @@ -797,7 +797,7 @@ class DiscordSendSaveVideo(BaseDiscordNode): output_file_with_audio_path = os.path.join(full_output_folder, output_file_with_audio) # Security: Validate output path to prevent symlink overwrites - validate_path_is_safe(output_file_with_audio_path) + validate_path_is_safe(output_file_with_audio_path, base_dir=full_output_folder) # Set up audio encoding parameters channels = a_waveform.size(1) diff --git a/shared/path_utils.py b/shared/path_utils.py index 2a84e3f..450cf83 100644 --- a/shared/path_utils.py +++ b/shared/path_utils.py @@ -83,20 +83,38 @@ def get_unique_filepath( return os.path.join(directory, full_filename) -def validate_path_is_safe(path: str) -> None: +def validate_path_is_safe(path: str, base_dir: Optional[str] = None) -> None: """ Validate that a path is safe to write to. Checks: + - If base_dir is provided, path is contained within base_dir (to prevent ../ traversal) - Path is not a symlink (to prevent overwriting targets) - Parent directories are not symlinks (to prevent path traversal via symlinks) Args: path: File path to validate + base_dir: Optional base directory to restrict path to Raises: ValueError: If path is unsafe """ + # Check if path is within base_dir + if base_dir: + abs_base = os.path.abspath(base_dir) + abs_path = os.path.abspath(path) + + # Use commonpath to ensure path is within base_dir + # We need to handle potential different drives on Windows which raises ValueError + try: + common = os.path.commonpath([abs_base, abs_path]) + except ValueError: + # Raised if paths are on different drives + raise ValueError(f"Security error: Path '{path}' is on a different drive than allowed directory '{base_dir}'.") + + if common != abs_base: + raise ValueError(f"Security error: Path '{path}' is outside the allowed directory '{base_dir}'.") + # Check if path exists and is a symlink if os.path.islink(path): raise ValueError(f"Security error: Output path '{path}' is a symlink. Overwriting symlinks is not allowed.") diff --git a/tests/test_path_security.py b/tests/test_path_security.py new file mode 100644 index 0000000..7a5da4d --- /dev/null +++ b/tests/test_path_security.py @@ -0,0 +1,75 @@ + +import os +import sys +import unittest +from unittest.mock import MagicMock + +# Mock dependencies to allow testing in isolation +sys.modules["folder_paths"] = MagicMock() +sys.modules["comfy"] = MagicMock() +sys.modules["server"] = MagicMock() +sys.modules["requests"] = MagicMock() +sys.modules["PIL"] = MagicMock() +sys.modules["numpy"] = MagicMock() +sys.modules["cv2"] = MagicMock() +sys.modules["torch"] = MagicMock() +sys.modules["torch.nn"] = MagicMock() +sys.modules["torch.nn.functional"] = MagicMock() + +# Add project root to sys.path +project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +if project_root not in sys.path: + sys.path.insert(0, project_root) + +from shared.path_utils import validate_path_is_safe + +class TestPathSecurity(unittest.TestCase): + def setUp(self): + self.test_dir = os.path.join(os.getcwd(), "test_safe_env") + os.makedirs(self.test_dir, exist_ok=True) + + def tearDown(self): + # Clean up would go here, but since we use temp dirs or mocks, it's fine. + # Ideally use tempfile.TemporaryDirectory but this is simple. + import shutil + if os.path.exists(self.test_dir): + shutil.rmtree(self.test_dir) + + def test_absolute_path_blocked_with_base_dir(self): + # This test checks if validate_path_is_safe blocks writing to /tmp + # when base_dir is provided. + + base_dir = self.test_dir + test_path = "/tmp/sentinel_test_file.txt" + + try: + validate_path_is_safe(test_path, base_dir=base_dir) + self.fail("Should have raised ValueError") + except ValueError as e: + self.assertIn("outside the allowed directory", str(e)) + + def test_traversal_blocked_with_base_dir(self): + # Create a directory structure + base_dir = self.test_dir + subdir = os.path.join(base_dir, "subdir") + os.makedirs(subdir, exist_ok=True) + + # ../ traversal + # This path resolves to outside base_dir + test_path = os.path.abspath(os.path.join(subdir, "../../test_safe_escape.txt")) + + try: + validate_path_is_safe(test_path, base_dir=base_dir) + self.fail("Should have raised ValueError") + except ValueError as e: + self.assertIn("outside the allowed directory", str(e)) + + def test_valid_path_allowed_with_base_dir(self): + base_dir = self.test_dir + test_path = os.path.join(base_dir, "valid_file.txt") + + # Should not raise + validate_path_is_safe(test_path, base_dir=base_dir) + +if __name__ == "__main__": + unittest.main()