Fix path traversal vulnerability in file output validation
Co-authored-by: AEmotionStudio <163354043+AEmotionStudio@users.noreply.github.com>
This commit is contained in:
co-authored by
AEmotionStudio
parent
52004b753d
commit
6de50104f1
+1
-1
@@ -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
|
||||
|
||||
+9
-9
@@ -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)
|
||||
|
||||
+19
-1
@@ -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.")
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user