diff --git a/py/services/model_file_service.py b/py/services/model_file_service.py index 63385f44..2537e2b8 100644 --- a/py/services/model_file_service.py +++ b/py/services/model_file_service.py @@ -8,6 +8,7 @@ from abc import ABC, abstractmethod from ..utils.utils import calculate_relative_path_for_model, remove_empty_dirs from ..utils.constants import AUTO_ORGANIZE_BATCH_SIZE from ..services.settings_manager import get_settings_manager +from ..services.model_lifecycle_service import _require_path_in_library_roots logger = logging.getLogger(__name__) @@ -493,6 +494,9 @@ class ModelMoveService: Dictionary with move result """ try: + _require_path_in_library_roots(file_path, self.scanner, label="Source path") + _require_path_in_library_roots(target_path, self.scanner, label="Target path") + if use_default_paths: # Find the model in cache to get metadata cache = await self.scanner.get_cached_data() diff --git a/py/services/model_lifecycle_service.py b/py/services/model_lifecycle_service.py index 8c4fe5cf..02b0a4f7 100644 --- a/py/services/model_lifecycle_service.py +++ b/py/services/model_lifecycle_service.py @@ -48,6 +48,35 @@ async def delete_model_artifacts( return deleted +def _require_path_in_library_roots(file_path: str, scanner, *, label: str = "path") -> None: + """Raise ``ValueError`` if *file_path* is not inside a configured model root. + + Uses ``os.path.realpath()`` to resolve symlinks before comparing, + so symlink-based escapes are also caught. Skips when the scanner + does not expose ``get_model_roots`` or the list is empty. + """ + + roots = None + if hasattr(scanner, "get_model_roots"): + try: + roots = scanner.get_model_roots() + except NotImplementedError: + roots = None + if not roots: + return + + resolved = os.path.realpath(os.path.normpath(file_path)) + + for root in roots: + root_resolved = os.path.realpath(os.path.normpath(root)) + if resolved == root_resolved or resolved.startswith(root_resolved + os.sep): + return + + raise ValueError( + f"{label} '{file_path}' is outside configured library directories" + ) + + class ModelLifecycleService: """Co-ordinate destructive and mutating model operations.""" @@ -74,6 +103,8 @@ class ModelLifecycleService: if not file_path: raise ValueError("Model path is required") + _require_path_in_library_roots(file_path, self._scanner, label="File path") + cache = await self._scanner.get_cached_data() cached_entry = None @@ -182,6 +213,8 @@ class ModelLifecycleService: if not file_path: raise ValueError("Model path is required") + _require_path_in_library_roots(file_path, self._scanner, label="File path") + metadata_path = os.path.splitext(file_path)[0] + ".metadata.json" metadata = await self._metadata_loader(metadata_path) metadata["exclude"] = True @@ -229,6 +262,8 @@ class ModelLifecycleService: if not file_path: raise ValueError("Model path is required") + _require_path_in_library_roots(file_path, self._scanner, label="File path") + if not os.path.exists(file_path): raise ValueError("Model file does not exist") @@ -270,6 +305,9 @@ class ModelLifecycleService: if not file_paths: raise ValueError("No file paths provided for deletion") + for path in file_paths: + _require_path_in_library_roots(path, self._scanner, label="File path") + return await self._scanner.bulk_delete_models(file_paths) async def rename_model( @@ -280,6 +318,8 @@ class ModelLifecycleService: if not file_path or not new_file_name: raise ValueError("File path and new file name are required") + _require_path_in_library_roots(file_path, self._scanner, label="File path") + invalid_chars = {"/", "\\", ":", "*", "?", '"', "<", ">", "|"} if any(char in new_file_name for char in invalid_chars): raise ValueError("Invalid characters in file name") diff --git a/py/services/model_scanner.py b/py/services/model_scanner.py index ed7d4da5..e88fa924 100644 --- a/py/services/model_scanner.py +++ b/py/services/model_scanner.py @@ -14,7 +14,7 @@ from ..utils.metadata_manager import MetadataManager from ..utils.civitai_utils import resolve_license_info from .model_cache import ModelCache from .model_hash_index import ModelHashIndex -from .model_lifecycle_service import delete_model_artifacts +from .model_lifecycle_service import delete_model_artifacts, _require_path_in_library_roots from .service_registry import ServiceRegistry from .websocket_manager import ws_manager from .persistent_model_cache import get_persistent_cache @@ -1394,6 +1394,9 @@ class ModelScanner: base_name = os.path.splitext(os.path.basename(source_path))[0] source_dir = os.path.dirname(source_path) + + _require_path_in_library_roots(source_path, self, label="Source path") + _require_path_in_library_roots(target_path, self, label="Target path") os.makedirs(target_path, exist_ok=True) @@ -1971,6 +1974,8 @@ class ModelScanner: break try: + _require_path_in_library_roots(file_path, self, label="File path") + target_dir = os.path.dirname(file_path) base_name = os.path.basename(file_path) file_name, main_extension = os.path.splitext(base_name) diff --git a/tests/conftest.py b/tests/conftest.py index f3c46ae3..71554cb4 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -85,6 +85,7 @@ sys.modules['comfy.utils'] = comfy_mock.utils sys.modules['comfy.sd'] = comfy_mock.sd sys.modules['comfy.model_management'] = comfy_mock.model_management sys.modules['comfy.comfy_types'] = comfy_mock.comfy_types +sys.modules['comfy.hooks'] = MockModule("comfy.hooks") execution_mock = MockModule("execution") execution_mock.PromptExecutor = mock.MagicMock() diff --git a/tests/services/test_model_lifecycle_service.py b/tests/services/test_model_lifecycle_service.py index f221d8a7..d55f8815 100644 --- a/tests/services/test_model_lifecycle_service.py +++ b/tests/services/test_model_lifecycle_service.py @@ -3,11 +3,164 @@ from pathlib import Path import pytest -from py.services.model_lifecycle_service import ModelLifecycleService +from py.services.model_lifecycle_service import ModelLifecycleService, _require_path_in_library_roots from py.utils.metadata_manager import MetadataManager from py.utils.models import LoraMetadata +class ScannerWithRoots: + def __init__(self, roots): + self._roots = list(roots) + + def get_model_roots(self): + return self._roots + + +class TestRequirePathInLibraryRoots: + def test_accepts_path_within_root(self, tmp_path): + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + model.write_text("") + + scanner = ScannerWithRoots([str(root)]) + _require_path_in_library_roots(str(model), scanner) + + def test_rejects_path_outside_roots(self, tmp_path): + root = tmp_path / "loras" + root.mkdir() + outside = tmp_path / "outside" / "model.safetensors" + outside.parent.mkdir(parents=True) + outside.write_text("") + + scanner = ScannerWithRoots([str(root)]) + with pytest.raises(ValueError, match="outside configured library"): + _require_path_in_library_roots(str(outside), scanner) + + def test_passes_when_no_roots_configured(self, tmp_path): + f = tmp_path / "model.safetensors" + f.write_text("") + + scanner = ScannerWithRoots([]) + _require_path_in_library_roots(str(f), scanner) + + def test_accepts_path_matching_root_exactly(self, tmp_path): + root = tmp_path / "loras" + root.mkdir() + + scanner = ScannerWithRoots([str(root)]) + _require_path_in_library_roots(str(root), scanner) + + def test_rejects_symlink_escape(self, tmp_path): + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + model.write_text("") + + outside_dir = tmp_path / "outside" + outside_dir.mkdir() + outside_file = outside_dir / "escaped.safetensors" + outside_file.write_text("") + + symlink = root / "link.safetensors" + symlink.symlink_to(outside_file) + + scanner = ScannerWithRoots([str(root)]) + with pytest.raises(ValueError, match="outside configured library"): + _require_path_in_library_roots(str(symlink), scanner) + + +class ScannerForDelete: + def __init__(self, raw_data, roots, model_type="lora"): + self.model_type = model_type + self.cache = DummyCache(raw_data) + self._hash_index = DummyHashIndex() + self._roots = list(roots) + self._persist_calls = [] + + def get_model_roots(self): + return self._roots + + async def get_cached_data(self): + return self.cache + + async def _persist_current_cache(self): + self._persist_calls.append(True) + + +@pytest.mark.asyncio +async def test_delete_model_rejects_path_outside_roots(tmp_path: Path): + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + model.write_bytes(b"data") + + scanner = ScannerForDelete( + raw_data=[{"file_path": str(model)}], + roots=[str(root)], + ) + service = ModelLifecycleService( + scanner=scanner, + metadata_manager=DummyMetadataManager({"civitai": {"modelId": 1}}), + metadata_loader=lambda x: {}, + ) + # Path within root should work (model file exists) + result = await service.delete_model(str(model)) + assert result["success"] is True + + # Path outside root should be rejected + outside = tmp_path / "outside.safetensors" + outside.write_bytes(b"data") + scanner2 = ScannerForDelete( + raw_data=[], + roots=[str(root)], + ) + service2 = ModelLifecycleService( + scanner=scanner2, + metadata_manager=DummyMetadataManager({}), + metadata_loader=lambda x: {}, + ) + with pytest.raises(ValueError, match="outside configured library"): + await service2.delete_model(str(outside)) + + +@pytest.mark.asyncio +async def test_rename_model_rejects_path_outside_roots(tmp_path: Path): + root = tmp_path / "loras" + root.mkdir() + + scanner = ScannerWithRoots([str(root)]) + service = ModelLifecycleService( + scanner=scanner, + metadata_manager=DummyMetadataManager({}), + metadata_loader=lambda x: {}, + ) + outside = tmp_path / "outside.safetensors" + outside.write_bytes(b"data") + + with pytest.raises(ValueError, match="outside configured library"): + await service.rename_model(file_path=str(outside), new_file_name="new_name") + + +@pytest.mark.asyncio +async def test_bulk_delete_rejects_any_path_outside_roots(tmp_path: Path): + root = tmp_path / "loras" + root.mkdir() + model_ok = root / "model.safetensors" + model_ok.write_bytes(b"data") + outside = tmp_path / "outside.safetensors" + outside.write_bytes(b"data") + + scanner = ScannerWithRoots([str(root)]) + service = ModelLifecycleService( + scanner=scanner, + metadata_manager=DummyMetadataManager({}), + metadata_loader=lambda x: {}, + ) + with pytest.raises(ValueError, match="outside configured library"): + await service.bulk_delete_models([str(model_ok), str(outside)]) + + class DummyCache: def __init__(self, raw_data): self.raw_data = raw_data