Compare commits

...

2 Commits

13 changed files with 221 additions and 8 deletions

View File

@@ -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()

View File

@@ -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")

View File

@@ -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)

View File

@@ -12,6 +12,7 @@ NODE_TYPES = {
"Lora Loader (LoraManager)": 1,
"Lora Stacker (LoraManager)": 2,
"WanVideo Lora Select (LoraManager)": 3,
"Create Hook LoRA (LoraManager)": 4,
}
# Default ComfyUI node color when bgcolor is null

View File

@@ -369,21 +369,24 @@ export function getMatureBlurThreshold(settings = {}) {
export const NODE_TYPES = {
LORA_LOADER: 1,
LORA_STACKER: 2,
WAN_VIDEO_LORA_SELECT: 3
WAN_VIDEO_LORA_SELECT: 3,
HOOK_LORA: 4
};
// Node type names to IDs mapping
export const NODE_TYPE_NAMES = {
"Lora Loader (LoraManager)": NODE_TYPES.LORA_LOADER,
"Lora Stacker (LoraManager)": NODE_TYPES.LORA_STACKER,
"WanVideo Lora Select (LoraManager)": NODE_TYPES.WAN_VIDEO_LORA_SELECT
"WanVideo Lora Select (LoraManager)": NODE_TYPES.WAN_VIDEO_LORA_SELECT,
"Create Hook LoRA (LoraManager)": NODE_TYPES.HOOK_LORA
};
// Node type icons
export const NODE_TYPE_ICONS = {
[NODE_TYPES.LORA_LOADER]: "fas fa-l",
[NODE_TYPES.LORA_STACKER]: "fas fa-s",
[NODE_TYPES.WAN_VIDEO_LORA_SELECT]: "fas fa-w"
[NODE_TYPES.WAN_VIDEO_LORA_SELECT]: "fas fa-w",
[NODE_TYPES.HOOK_LORA]: "fas fa-h"
};
// Default ComfyUI node color when bgcolor is null

View File

@@ -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()

View File

@@ -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

View File

@@ -16,6 +16,7 @@ export const LORA_PROVIDER_NODE_TYPES = [
"Lora Stacker (LoraManager)",
"Lora Randomizer (LoraManager)",
"Lora Cycler (LoraManager)",
"Create Hook LoRA (LoraManager)",
] as const;
/**

View File

@@ -12,6 +12,7 @@ const LORA_NODE_CLASSES = new Set([
"Lora Loader (LoraManager)",
"Lora Stacker (LoraManager)",
"WanVideo Lora Select (LoraManager)",
"Create Hook LoRA (LoraManager)",
]);
function normalizeTriggerWordList(triggerWords) {

View File

@@ -8,6 +8,7 @@ export const LORA_PROVIDER_NODE_TYPES = [
"Lora Stacker (LoraManager)",
"Lora Randomizer (LoraManager)",
"Lora Cycler (LoraManager)",
"Create Hook LoRA (LoraManager)",
];
export const LORA_STACK_AGGREGATOR_NODE_TYPES = [

View File

@@ -15656,7 +15656,8 @@ function createVueWidgetCleanup(vueApp, onCleanup) {
const LORA_PROVIDER_NODE_TYPES$1 = [
"Lora Stacker (LoraManager)",
"Lora Randomizer (LoraManager)",
"Lora Cycler (LoraManager)"
"Lora Cycler (LoraManager)",
"Create Hook LoRA (LoraManager)"
];
const LORA_STACK_AGGREGATOR_NODE_TYPES$1 = [
"Lora Stack Combiner (LoraManager)"
@@ -15781,7 +15782,8 @@ const ROOT_GRAPH_ID = "root";
const LORA_PROVIDER_NODE_TYPES = [
"Lora Stacker (LoraManager)",
"Lora Randomizer (LoraManager)",
"Lora Cycler (LoraManager)"
"Lora Cycler (LoraManager)",
"Create Hook LoRA (LoraManager)"
];
const LORA_STACK_AGGREGATOR_NODE_TYPES = [
"Lora Stack Combiner (LoraManager)"

File diff suppressed because one or more lines are too long

View File

@@ -9,6 +9,7 @@ const LORA_NODE_CLASSES = new Set([
"Lora Loader (LoraManager)",
"Lora Stacker (LoraManager)",
"WanVideo Lora Select (LoraManager)",
"Create Hook LoRA (LoraManager)",
]);
const TARGET_WIDGET_NAMES = new Set(["ckpt_name", "unet_name"]);