diff --git a/py/lora_manager.py b/py/lora_manager.py index 4b02555d..727b4fc2 100644 --- a/py/lora_manager.py +++ b/py/lora_manager.py @@ -25,10 +25,12 @@ from .routes.recipe_routes import RecipeRoutes from .routes.stats_routes import StatsRoutes from .routes.update_routes import UpdateRoutes from .routes.misc_routes import MiscRoutes +from .routes.pending_delete_routes import PendingDeleteRoutes from .routes.preview_routes import PreviewRoutes from .routes.example_images_routes import ExampleImagesRoutes from .services.service_registry import ServiceRegistry from .services.settings_manager import get_settings_manager +from .services.pending_delete_service import get_pending_delete_service from .utils.example_images_migration import ExampleImagesMigration from .services.websocket_manager import ws_manager from .services.example_images_cleanup_service import ExampleImagesCleanupService @@ -170,6 +172,7 @@ class LoraManager: RecipeRoutes.setup_routes(app) UpdateRoutes.setup_routes(app) MiscRoutes.setup_routes(app) + PendingDeleteRoutes.setup_routes(app) ExampleImagesRoutes.setup_routes(app, ws_manager=ws_manager) PreviewRoutes.setup_routes(app) @@ -245,6 +248,17 @@ class LoraManager: cls._run_post_initialization_tasks(init_tasks), name="post_init_tasks" ) + # Startup sweep: purge pending-delete batches that expired during a + # previous run. Non-blocking (fire-and-forget); purge_expired only + # removes already-expired batches, so a staged undo that survived a + # restart stays restorable. Covers both plugin and standalone modes + # (StandaloneLoraManager reuses this classmethod). + pending_delete_service = await get_pending_delete_service() + asyncio.create_task( + pending_delete_service.purge_expired(), + name="pending_delete_startup_sweep", + ) + logger.debug( "LoRA Manager: All services initialized and background tasks scheduled" ) diff --git a/py/routes/handlers/pending_delete_handler.py b/py/routes/handlers/pending_delete_handler.py new file mode 100644 index 00000000..8213d89d --- /dev/null +++ b/py/routes/handlers/pending_delete_handler.py @@ -0,0 +1,323 @@ +"""Handler for the pending-delete undo endpoint. + +Restores a staged delete batch (models or recipes) via +``PendingDeleteService.undo`` and then repairs the affected library caches: +the model cache entry is restored from the manifest's ``model_snapshot`` +(including the version index and hash index), tag counts are re-incremented, +and the recipe cache is re-populated via ``RecipeScanner.add_recipe``. + +The per-type scanner is resolved from the manifest's ``model_type`` page value +through the SAME ServiceRegistry getters the model route registrars use +(lora/checkpoint/embedding) - never a hardcoded lora scanner. +""" + +from __future__ import annotations + +import inspect +import json +import logging +import os +import re +from typing import Any, Awaitable, Callable, Dict, List, Optional, Set, cast + +from aiohttp import web + +from ...services.pending_delete_service import get_pending_delete_service +from .model_handlers import _broadcast_models_changed + +logger = logging.getLogger(__name__) + +# Manifest ``model_type`` page values -> ServiceRegistry scanner getter names. +# The model route registrars resolve per-type scanners via these getters +# (lora_routes / checkpoint_routes / embedding_routes); undo must do the same +# so the CORRECT cache is restored for the deleted model's type. +_MODEL_TYPE_GETTER_NAMES: Dict[str, str] = { + "loras": "get_lora_scanner", + "checkpoints": "get_checkpoint_scanner", + "embeddings": "get_embedding_scanner", +} + +# Staged batch ids are ``uuid.uuid4().hex`` (32 lowercase hex chars). The id is +# joined into filesystem paths by ``_find_batch_dir``, so reject anything that +# does not match this exact shape (blocks path-traversal via batch_id). +_BATCH_ID_RE = re.compile(r"^[0-9a-f]{32}$") + + +class PendingDeleteHandler: + """Handle undo requests for staged model/recipe deletions.""" + + def __init__( + self, + *, + service_factory: Callable[[], Awaitable[Any]] = get_pending_delete_service, + scanner_getter: Optional[Callable[[str], Awaitable[Any]]] = None, + recipe_scanner_getter: Optional[Callable[[], Awaitable[Any]]] = None, + ) -> None: + self._service_factory: Callable[[], Awaitable[Any]] = service_factory + self._scanner_getter: Callable[[str], Awaitable[Any]] = ( + scanner_getter or self._resolve_scanner + ) + self._recipe_scanner_getter: Callable[[], Awaitable[Any]] = ( + recipe_scanner_getter or self._resolve_recipe_scanner + ) + + @staticmethod + async def _resolve_scanner(model_type: str) -> Any: + """Resolve the per-type scanner for a manifest ``model_type``. + + The getter is looked up on the ServiceRegistry module namespace at call + time so tests (and the registry stubs) can patch it. + """ + from ...services import service_registry + + getter_name = _MODEL_TYPE_GETTER_NAMES.get(model_type) + if getter_name is None: + raise ValueError(f"Unknown model type: {model_type}") + getter = getattr(service_registry.ServiceRegistry, getter_name, None) + if not callable(getter): + raise ValueError(f"No scanner getter for model type: {model_type}") + scanner = await cast(Callable[[], Awaitable[Any]], getter)() + if scanner is None: + raise ValueError(f"No scanner registered for model type: {model_type}") + return scanner + + @staticmethod + async def _resolve_recipe_scanner() -> Any: + """Resolve the recipe scanner via the ServiceRegistry module namespace.""" + from ...services import service_registry + + getter = getattr(service_registry.ServiceRegistry, "get_recipe_scanner", None) + if not callable(getter): + raise ValueError("Recipe scanner getter unavailable") + scanner = await cast(Callable[[], Awaitable[Any]], getter)() + if scanner is None: + raise ValueError("No recipe scanner registered") + return scanner + + async def undo_delete(self, request: web.Request) -> web.Response: + """Restore a staged batch and its library cache entry. + + Body: ``{"batch_id": str}``. On success returns + ``{"success": True, "restored": [], "kind": kind}``. + Expired/unknown batches and occupied target paths -> 404. + """ + try: + data = await request.json() + except Exception: + return web.json_response( + {"success": False, "error": "Invalid JSON body"}, status=400 + ) + if not isinstance(data, dict): + return web.json_response( + {"success": False, "error": "Invalid JSON body"}, status=400 + ) + batch_id = data.get("batch_id") + if not batch_id or not isinstance(batch_id, str): + return web.json_response( + {"success": False, "error": "batch_id is required"}, status=400 + ) + if not _BATCH_ID_RE.fullmatch(batch_id): + # batch_id is joined into a path by _find_batch_dir - restrict to + # the exact staged-id shape so traversal payloads get 400. + return web.json_response( + {"success": False, "error": "Invalid batch_id"}, status=400 + ) + + service = await self._service_factory() + try: + # Read the manifest BEFORE undo: undo() removes the batch dir. + manifest = await self._read_staged_manifest(service, batch_id) + result = await service.undo(batch_id) + except ValueError as exc: + return web.json_response({"success": False, "error": str(exc)}, status=404) + except Exception as exc: + logger.error("Unexpected error undoing batch %s: %s", batch_id, exc, exc_info=True) + return web.json_response({"success": False, "error": str(exc)}, status=500) + + kind = result.get("kind") + try: + if kind == "model": + if manifest is not None: + await self._restore_model_cache(manifest) + else: + # undo() raises when the manifest is missing, so this only + # happens defensively - files are restored regardless. + logger.warning( + "Manifest missing after undo of %s; skipping cache restore", + batch_id, + ) + _broadcast_models_changed() + elif kind == "recipe": + # Recipe undo is client-refresh only: re-add to the scanner + # cache, no models_changed broadcast. + if manifest is not None: + await self._restore_recipe_cache(result, manifest) + else: + logger.warning( + "Manifest missing after undo of %s; skipping cache restore", + batch_id, + ) + except Exception as exc: + # Files are already restored; only the cache restoration failed. + logger.error( + "Cache restoration failed after undo of %s: %s", + batch_id, + exc, + exc_info=True, + ) + return web.json_response({"success": False, "error": str(exc)}, status=500) + + return web.json_response( + { + "success": True, + "restored": result.get("restored", []), + "kind": kind, + } + ) + + @staticmethod + async def _read_staged_manifest( + service: Any, batch_id: str + ) -> Optional[Dict[str, Any]]: + """Locate and read the batch manifest while it still exists on disk.""" + batch_dir = await service._find_batch_dir(batch_id) + if not batch_dir: + return None + manifest_path = os.path.join(batch_dir, "manifest.json") + try: + with open(manifest_path, "r", encoding="utf-8") as handle: + payload = json.load(handle) + except (OSError, json.JSONDecodeError) as exc: + logger.debug("Failed to read manifest for batch %s: %s", batch_id, exc) + return None + return payload if isinstance(payload, dict) else None + + async def _restore_model_cache(self, manifest: Dict[str, Any]) -> None: + """Re-add every deleted model's cache entry from the manifest. + + Each main-file entry carries the deleted model's ``snapshot`` (added at + stage time), so a merged bulk manifest holds ALL snapshots - undo must + restore every one, not just the top-level winner's. Old-format + manifests without entry snapshots fall back to the top-level + ``model_snapshot`` (backward compat / single-delete path). + """ + model_type = manifest.get("model_type") + if not model_type or not isinstance(model_type, str): + raise ValueError(f"Manifest carries no model_type: {manifest.get('batch_id')}") + scanner = await self._scanner_getter(model_type) + + # Collect one snapshot per distinct file_path from the entry snapshots. + snapshots: List[Dict[str, Any]] = [] + seen: Set[str] = set() + for entry in manifest.get("entries") or []: + snapshot = entry.get("snapshot") + if not isinstance(snapshot, dict): + continue + file_path = snapshot.get("file_path") + if not file_path or not isinstance(file_path, str): + continue + if file_path in seen: + continue + seen.add(file_path) + snapshots.append(snapshot) + + if not snapshots: + # Backward compat: pre-F3 manifests carry only the top-level + # model_snapshot (single-delete path, unchanged behavior). + top = manifest.get("model_snapshot") + if isinstance(top, dict) and top.get("file_path"): + snapshots = [top] + else: + logger.warning( + "Manifest %s has no restorable model snapshot; skipping cache restore", + manifest.get("batch_id"), + ) + return + + cache = await scanner.get_cached_data() + if cache is None: + logger.warning( + "Scanner cache unavailable for %s; skipping cache restore", model_type + ) + return + + for snapshot in snapshots: + file_path = str(snapshot["file_path"]) + # A rescan between delete and undo may have re-added a stale entry + # for this path - drop it so exactly one (the snapshot) remains. + cache.raw_data = [ + item for item in cache.raw_data if item.get("file_path") != file_path + ] + + # Restore tag counts (mirror of the bulk-delete decrement in + # _batch_update_cache_for_deleted_models: undo re-increments). + tags = snapshot.get("tags") + if isinstance(tags, list): + for tag in tags: + if not isinstance(tag, str) or not tag: + continue + scanner._tags_count[tag] = scanner._tags_count.get(tag, 0) + 1 + + cache.raw_data.append(dict(snapshot)) + + # Re-register the path in the hash index (add_entry guards a + # missing sha256 internally; still guard defensively here). + sha256 = snapshot.get("sha256") or "" + autov3 = snapshot.get("autov3") + hash_index = getattr(scanner, "_hash_index", None) + if hash_index is not None and sha256 and file_path: + hash_index.add_entry(sha256, file_path, autov3) + + # Follow the bulk-delete cache-update pattern ONCE after all entries, + # including the explicit version-index rebuild so the version index + # does not go stale. + cache.rebuild_version_index() + await cache.resort() + + scanner.bump_cache_version() + + persist = getattr(scanner, "_persist_current_cache", None) + if callable(persist): + result = persist() + if inspect.isawaitable(result): + await result + + async def _restore_recipe_cache( + self, result: Dict[str, Any], manifest: Dict[str, Any] + ) -> None: + """Re-add a restored recipe via ``RecipeScanner.add_recipe``. + + The recipe JSON embeds the full recipe_data (incl. id/file_path); + ``add_recipe`` only READS the ``_json_path_map`` so the forced frontend + refresh self-heals any transient path-map gap. + """ + restored = result.get("restored") or [] + json_path = next( + (p for p in restored if isinstance(p, str) and p.endswith(".json")), + None, + ) + if not json_path or not os.path.exists(json_path): + # Defensive fallback to the manifest's recipe_snapshot file_path. + snapshot = manifest.get("recipe_snapshot") or {} + fallback = snapshot.get("file_path") + if fallback and os.path.exists(fallback): + json_path = fallback + else: + logger.warning( + "Restored recipe JSON not found in %s; skipping cache restore", + restored, + ) + return + try: + with open(json_path, "r", encoding="utf-8") as handle: + recipe_data = json.load(handle) + except (OSError, json.JSONDecodeError) as exc: + logger.warning("Failed to load restored recipe JSON %s: %s", json_path, exc) + return + if not isinstance(recipe_data, dict): + return + recipe_scanner = await self._recipe_scanner_getter() + await recipe_scanner.add_recipe(recipe_data) + + +__all__ = ["PendingDeleteHandler"] diff --git a/py/routes/pending_delete_routes.py b/py/routes/pending_delete_routes.py new file mode 100644 index 00000000..43f2dd45 --- /dev/null +++ b/py/routes/pending_delete_routes.py @@ -0,0 +1,25 @@ +"""Route controller for the pending-delete undo endpoint.""" + +from __future__ import annotations + +from aiohttp import web + +from .handlers.pending_delete_handler import PendingDeleteHandler + + +class PendingDeleteRoutes: + """Shared route controller mirroring MiscRoutes/UpdateRoutes. + + Registered ONCE per mode (py/lora_manager.py, standalone.py); NEVER through + the per-model-type ModelRouteRegistrar, which is instantiated per model + type and would register this non-prefixed route three times. + """ + + @staticmethod + def setup_routes(app: web.Application) -> None: + """Register the shared undo-delete endpoint.""" + handler = PendingDeleteHandler() + _ = app.router.add_post("/api/lm/undo-delete", handler.undo_delete) + + +__all__ = ["PendingDeleteRoutes"] diff --git a/standalone.py b/standalone.py index 737f5f61..7e3576d1 100644 --- a/standalone.py +++ b/standalone.py @@ -339,6 +339,7 @@ class StandaloneLoraManager(LoraManager): from py.routes.recipe_routes import RecipeRoutes from py.routes.update_routes import UpdateRoutes from py.routes.misc_routes import MiscRoutes + from py.routes.pending_delete_routes import PendingDeleteRoutes from py.routes.example_images_routes import ExampleImagesRoutes from py.routes.preview_routes import PreviewRoutes from py.routes.stats_routes import StatsRoutes @@ -356,6 +357,7 @@ class StandaloneLoraManager(LoraManager): RecipeRoutes.setup_routes(app) UpdateRoutes.setup_routes(app) MiscRoutes.setup_routes(app) + PendingDeleteRoutes.setup_routes(app) ExampleImagesRoutes.setup_routes(app, ws_manager=ws_manager) PreviewRoutes.setup_routes(app) diff --git a/tests/routes/test_lora_manager_lifecycle.py b/tests/routes/test_lora_manager_lifecycle.py index f06530b9..2e62e4d4 100644 --- a/tests/routes/test_lora_manager_lifecycle.py +++ b/tests/routes/test_lora_manager_lifecycle.py @@ -207,6 +207,10 @@ async def test_lora_manager_lifecycle(monkeypatch: pytest.MonkeyPatch, tmp_path: task_names = {task.get_name() for task in scheduled_tasks} assert {"lora_cache_init", "checkpoint_cache_init", "embedding_cache_init", "recipe_cache_init", "post_init_tasks", "cleanup_bak_files"}.issubset(task_names) + # Startup sweep: an expired pending-delete purge task is spawned during + # service initialization (covers both plugin and standalone modes). + assert "pending_delete_startup_sweep" in task_names + for scanner in scanners.values(): assert scanner.initialized is True diff --git a/tests/routes/test_pending_delete_routes.py b/tests/routes/test_pending_delete_routes.py new file mode 100644 index 00000000..ef37b091 --- /dev/null +++ b/tests/routes/test_pending_delete_routes.py @@ -0,0 +1,1098 @@ +"""Tests for the pending-delete undo endpoint. + +Covers ``POST /api/lm/undo-delete`` (``py/routes/handlers/ +pending_delete_handler.py``) and its route registration +(``py/routes/pending_delete_routes.py``): model/recipe cache restoration, +per-type scanner resolution, restart-safe and rescan-stale undo, tag-count +restoration, broadcast, error responses, malformed input, and exactly-once +registration in both plugin and standalone modes. +""" + +from __future__ import annotations + +import asyncio +import json +import os +import time +from collections.abc import Iterator +from pathlib import Path +from types import SimpleNamespace +from typing import Any, Dict, List, Optional, Sequence + +import pytest +from aiohttp import web + +from py.routes.pending_delete_routes import PendingDeleteRoutes +from py.services.pending_delete_service import ( + PENDING_DELETE_DIR_NAME, + PendingDeleteService, + _reset_pending_delete_service, +) +from py.utils import settings_paths + +UNDO_DELETE_PATH = "/api/lm/undo-delete" + + +# --------------------------------------------------------------------------- +# Doubles +# --------------------------------------------------------------------------- +class FakeRequest: + """Request double exposing the async ``json()`` the handler awaits.""" + + def __init__(self, *, json_data: Any = None, raise_json_error: bool = False) -> None: + self._json_data = json_data + self._raise_json_error = raise_json_error + + async def json(self) -> Any: + if self._raise_json_error: + raise ValueError("invalid json body") + return self._json_data + + +class FakeHashIndex: + """Hash index double recording ``add_entry`` calls.""" + + def __init__(self) -> None: + self.entries: List[tuple[Any, ...]] = [] + + def add_entry(self, sha256: str, file_path: str, autov3: Optional[str] = None) -> None: + self.entries.append((sha256, file_path, autov3)) + + +class FakeCache: + """Cache double exposing the raw_data/version-index surface the undo uses.""" + + def __init__(self, items: Optional[Sequence[Dict[str, Any]]] = None) -> None: + self.raw_data: List[Dict[str, Any]] = list(items or []) + self.version_index: Dict[Any, Any] = {} + self.model_id_index: Dict[Any, Any] = {} + self.rebuild_calls = 0 + self.resort_calls = 0 + + def rebuild_version_index(self) -> None: + self.rebuild_calls += 1 + self.version_index = {} + self.model_id_index = {} + + async def resort(self) -> None: + self.resort_calls += 1 + + +class FakeScanner: + """Scanner double exposing the attributes the undo cache-restore uses.""" + + def __init__( + self, + root: Path, + *, + model_type: str = "lora", + cache: Optional[FakeCache] = None, + ) -> None: + self._root = os.path.abspath(str(root)) + self.model_type = model_type + self._cache = cache or FakeCache() + self._hash_index = FakeHashIndex() + self._tags_count: Dict[str, int] = {} + self.cache_version = 0 + self.bump_calls = 0 + self.persist_calls = 0 + + def get_model_roots(self) -> List[str]: + return [self._root] + + def _find_root_for_file(self, file_path: Optional[str]) -> Optional[str]: + if not file_path: + return None + normalized = os.path.abspath(os.path.normpath(file_path)) + if normalized == self._root or normalized.startswith(self._root + os.sep): + return self._root + return None + + async def get_cached_data(self, force_refresh: bool = False) -> FakeCache: + return self._cache + + def bump_cache_version(self) -> None: + self.cache_version += 1 + self.bump_calls += 1 + + async def _persist_current_cache(self) -> None: + self.persist_calls += 1 + + +class FakeRecipeScanner: + """Recipe scanner double mirroring add_recipe's tolerant path-map read.""" + + def __init__(self) -> None: + self.added: List[Dict[str, Any]] = [] + self._json_path_map: Dict[str, str] = {} + + async def add_recipe(self, recipe_data: Dict[str, Any]) -> None: + # The real scanner only READS _json_path_map here (the row may carry an + # empty json_path until the forced frontend refresh) - never crashes. + recipe_id = str(recipe_data.get("id", "")) + self._json_path_map.get(recipe_id, "") + self.added.append(dict(recipe_data)) + + def force_refresh(self) -> None: + """Simulate ``window.recipeManager.loadRecipes(true)`` rebuilding the map.""" + for recipe in self.added: + recipe_id = str(recipe.get("id", "")) + file_path = recipe.get("file_path") + if recipe_id and file_path: + self._json_path_map[recipe_id] = str(file_path) + + +class _DummyRoutes: + @staticmethod + def setup_routes(app_: web.Application, **kwargs: Any) -> None: + return None + + +class _DummyWSManager: + async def handle_connection(self, request): # pragma: no cover - interface stub + return None + + async def handle_download_connection(self, request): # pragma: no cover - interface stub + return None + + async def handle_init_connection(self, request): # pragma: no cover - interface stub + return None + + +# --------------------------------------------------------------------------- +# Fixtures / helpers +# --------------------------------------------------------------------------- +@pytest.fixture(autouse=True) +def _reset_service_singleton() -> Iterator[None]: + """Reset the pending-delete singleton before and after each test.""" + _reset_pending_delete_service() + yield + _reset_pending_delete_service() + + +@pytest.fixture(autouse=True) +def _stub_scanner_registry(monkeypatch: pytest.MonkeyPatch) -> None: + """Prevent scanner resolution from instantiating real singletons.""" + from py.services.service_registry import ServiceRegistry + + async def _none(*_args: Any, **_kwargs: Any) -> None: + return None + + monkeypatch.setattr(ServiceRegistry, "get_lora_scanner", _none) + monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _none) + monkeypatch.setattr(ServiceRegistry, "get_embedding_scanner", _none) + monkeypatch.setattr(ServiceRegistry, "get_recipe_scanner", _none) + + +async def _register_scanners( + monkeypatch: pytest.MonkeyPatch, + *, + lora: Any = None, + checkpoint: Any = None, + embedding: Any = None, + recipe: Any = None, +) -> None: + """Point the ServiceRegistry getters at per-test scanner doubles.""" + from py.services.service_registry import ServiceRegistry + + def _make(scanner: Any): + async def _getter(*_args: Any, **_kwargs: Any) -> Any: + return scanner + + return _getter + + monkeypatch.setattr(ServiceRegistry, "get_lora_scanner", _make(lora)) + monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _make(checkpoint)) + monkeypatch.setattr(ServiceRegistry, "get_embedding_scanner", _make(embedding)) + monkeypatch.setattr(ServiceRegistry, "get_recipe_scanner", _make(recipe)) + + +def _make_undo_request(**json_data: Any) -> FakeRequest: + return FakeRequest(json_data=json_data) + + +def _json_payload(response: web.Response) -> Dict[str, Any]: + """Decode the JSON body of a web.Response, asserting it is not null.""" + text = response.text + assert text is not None + return json.loads(text) + + +def _write_batch_manifest( + batch_dir: Path, + *, + batch_id: str, + kind: str, + expires_at: int, + entries: Sequence[Dict[str, Any]], + model_type: Optional[str] = None, + state: str = "staged", + model_snapshot: Any = None, + recipe_snapshot: Any = None, +) -> None: + """Write a manifest.json with the shape the service reads.""" + manifest: Dict[str, Any] = { + "batch_id": batch_id, + "kind": kind, + "model_type": model_type if kind == "model" else None, + "state": state, + "expires_at": int(expires_at), + "entries": list(entries), + "model_snapshot": model_snapshot if kind == "model" else None, + "recipe_snapshot": recipe_snapshot if kind == "recipe" else None, + } + (batch_dir / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + + +def _undo_delete_routes(app: web.Application) -> List[Any]: + """The POST /api/lm/undo-delete routes currently registered on *app*.""" + return [ + route + for route in app.router.routes() + if route.method == "POST" + and getattr(route.resource, "canonical", "") == UNDO_DELETE_PATH + ] + + +async def _drain_broadcast() -> None: + """Yield to the loop so the fire-and-forget broadcast task can run.""" + for _ in range(10): + await asyncio.sleep(0) + + +# (a) undo of a staged model batch (loras) -> files + cache + hash + broadcast +# --------------------------------------------------------------------------- +async def test_a_undo_staged_model_batch_restores_files_cache_and_broadcasts( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + model.write_bytes(b"model-data") + snapshot = { + "file_path": str(model), + "sha256": "a" * 64, + "tags": ["alpha"], + "model_name": "model", + } + + scanner = FakeScanner(root, model_type="lora") + scanner._cache.raw_data = [dict(snapshot)] + scanner._tags_count = {"alpha": 1} + await _register_scanners(monkeypatch, lora=scanner) + + service = await PendingDeleteService.get_instance() + batch_id = await service.stage_model_delete( + scanner=scanner, + target_dir=str(root), + file_name="model", + main_extension=".safetensors", + original_file_path=str(model), + cached_entry=dict(snapshot), + ) + assert batch_id is not None + + # Simulate the delete-time cache mutation: entry removed, tags decremented. + scanner._cache.raw_data = [] + scanner._tags_count = {} + + sent: List[Dict[str, Any]] = [] + + async def fake_broadcast(data: Dict[str, Any]) -> None: + sent.append(data) + + monkeypatch.setattr("py.services.websocket_manager.ws_manager.broadcast", fake_broadcast) + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + + assert response.status == 200 + payload = _json_payload(response) + assert payload["success"] is True + assert payload["kind"] == "model" + assert payload["restored"] == [str(model)] + + # Files restored to the original path. + assert model.read_bytes() == b"model-data" + # Cache entry restored exactly once. + assert len(scanner._cache.raw_data) == 1 + assert scanner._cache.raw_data[0]["file_path"] == str(model) + # Version-index rebuild + resort + persist + cache bump all ran. + assert scanner._cache.rebuild_calls >= 1 + assert scanner._cache.resort_calls >= 1 + assert scanner.persist_calls >= 1 + assert scanner.bump_calls >= 1 + # Hash index has the path. + assert any(entry[1] == str(model) for entry in scanner._hash_index.entries) + # Tag counts re-incremented. + assert scanner._tags_count == {"alpha": 1} + # Broadcast fired exactly once with models_changed. + await _drain_broadcast() + assert sent == [{"type": "models_changed"}] + + +# --------------------------------------------------------------------------- +# (b) undo of a staged CHECKPOINT batch -> resolves the checkpoint scanner +# --------------------------------------------------------------------------- +async def test_b_undo_checkpoint_batch_updates_checkpoint_cache_only( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + ckpt_root = tmp_path / "checkpoints" + ckpt_root.mkdir() + model = ckpt_root / "model.safetensors" + model.write_bytes(b"ckpt-data") + snapshot = { + "file_path": str(model), + "sha256": "b" * 64, + "tags": [], + "model_name": "ckpt", + } + + lora_scanner = FakeScanner(tmp_path / "loras", model_type="lora") + ckpt_scanner = FakeScanner(ckpt_root, model_type="checkpoint") + lora_scanner._cache.raw_data = [{"file_path": "/loras/other.safetensors", "sha256": "x"}] + ckpt_scanner._cache.raw_data = [dict(snapshot)] + await _register_scanners(monkeypatch, lora=lora_scanner, checkpoint=ckpt_scanner) + + service = await PendingDeleteService.get_instance() + batch_id = await service.stage_model_delete( + scanner=ckpt_scanner, + target_dir=str(ckpt_root), + file_name="model", + main_extension=".safetensors", + original_file_path=str(model), + cached_entry=dict(snapshot), + ) + assert batch_id is not None + # Simulate the delete-time cache mutation on the checkpoint cache. + ckpt_scanner._cache.raw_data = [] + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + assert response.status == 200 + + # The CHECKPOINT cache was updated, not the lora cache. + assert len(ckpt_scanner._cache.raw_data) == 1 + assert ckpt_scanner._cache.raw_data[0]["file_path"] == str(model) + assert len(lora_scanner._cache.raw_data) == 1 # untouched + + +# --------------------------------------------------------------------------- +# (c) undo of a staged recipe batch -> files restored + add_recipe called +# --------------------------------------------------------------------------- +async def test_c_undo_recipe_batch_restores_files_and_calls_add_recipe( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + recipe_json = tmp_path / "my_recipe.recipe.json" + recipe_data = {"id": "r1", "name": "Recipe One", "file_path": str(recipe_json)} + recipe_json.write_text(json.dumps(recipe_data), encoding="utf-8") + + recipe_scanner = FakeRecipeScanner() + await _register_scanners(monkeypatch, recipe=recipe_scanner) + + service = await PendingDeleteService.get_instance() + batch_id = await service.stage_recipe_delete( + recipe_json_path=str(recipe_json), + image_path=None, + recipe_data=dict(recipe_data), + ) + assert batch_id is not None + # The todo-4 caller removes the originals after staging. + recipe_json.unlink() + + sent: List[Dict[str, Any]] = [] + + async def fake_broadcast(data: Dict[str, Any]) -> None: + sent.append(data) + + monkeypatch.setattr("py.services.websocket_manager.ws_manager.broadcast", fake_broadcast) + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + + assert response.status == 200 + payload = _json_payload(response) + assert payload["success"] is True + assert payload["kind"] == "recipe" + assert payload["restored"] == [str(recipe_json)] + + # Files restored and the recipe re-added to the scanner. + assert recipe_json.read_text(encoding="utf-8") == json.dumps(recipe_data) + assert recipe_scanner.added == [recipe_data] + # Recipe undo is client-refresh only - no models_changed broadcast. + await _drain_broadcast() + assert sent == [] + + +# --------------------------------------------------------------------------- +# (d) unknown / expired / occupied batch -> 404 responses +# --------------------------------------------------------------------------- +async def test_d1_undo_unknown_batch_returns_404( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + await _register_scanners( + monkeypatch, lora=FakeScanner(tmp_path / "loras", model_type="lora") + ) + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id="f" * 32) # pyright: ignore[reportArgumentType] + ) + + assert response.status == 404 + payload = _json_payload(response) + assert payload["success"] is False + assert "batch" in payload["error"].lower() + + +async def test_d2_undo_expired_batch_returns_404( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + model.write_bytes(b"data") + + scanner = FakeScanner(root, model_type="lora") + await _register_scanners(monkeypatch, lora=scanner) + + service = await PendingDeleteService.get_instance() + + # Isolate the expiry check from the opportunistic purge. + async def _no_purge() -> int: + return 0 + + monkeypatch.setattr(service, "purge_expired", _no_purge) + + batch_id = await service.stage_model_delete( + scanner=scanner, + target_dir=str(root), + file_name="model", + main_extension=".safetensors", + original_file_path=str(model), + cached_entry={"file_path": str(model)}, + ) + assert batch_id is not None + batch_dir = root / PENDING_DELETE_DIR_NAME / batch_id + manifest_path = batch_dir / "manifest.json" + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + manifest["expires_at"] = int(time.time()) - 10 + manifest_path.write_text(json.dumps(manifest), encoding="utf-8") + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + assert response.status == 404 + assert "expired" in _json_payload(response)["error"].lower() + + +async def test_d3_undo_occupied_path_returns_404( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + model.write_bytes(b"data") + + scanner = FakeScanner(root, model_type="lora") + await _register_scanners(monkeypatch, lora=scanner) + + service = await PendingDeleteService.get_instance() + batch_id = await service.stage_model_delete( + scanner=scanner, + target_dir=str(root), + file_name="model", + main_extension=".safetensors", + original_file_path=str(model), + cached_entry={"file_path": str(model)}, + ) + assert batch_id is not None + # A re-download now occupies the original path. + model.write_bytes(b"new-file") + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + assert response.status == 404 + assert "occupied" in _json_payload(response)["error"].lower() + + +# --------------------------------------------------------------------------- +# (d-adjacent) malformed input -> 400 +# --------------------------------------------------------------------------- +async def test_d4_undo_malformed_body_returns_400( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + handler = PendingDeleteHandler() + + # Missing batch_id. + response = await handler.undo_delete(_make_undo_request()) # pyright: ignore[reportArgumentType] + assert response.status == 400 + assert "batch_id" in _json_payload(response)["error"].lower() + + # Non-dict JSON body. + response = await handler.undo_delete(FakeRequest(json_data="not-a-dict")) # pyright: ignore[reportArgumentType] + assert response.status == 400 + + # Invalid JSON body. + response = await handler.undo_delete(FakeRequest(raise_json_error=True)) # pyright: ignore[reportArgumentType] + assert response.status == 400 + + # (VAL-1) path-traversal style batch_id -> 400. + response = await handler.undo_delete(_make_undo_request(batch_id="../evil")) # pyright: ignore[reportArgumentType] + assert response.status == 400 + assert "Invalid batch_id" in _json_payload(response)["error"] + + # (VAL-2) non-hex batch_id -> 400. + response = await handler.undo_delete(_make_undo_request(batch_id="not-a-hex-id")) # pyright: ignore[reportArgumentType] + assert response.status == 400 + assert "Invalid batch_id" in _json_payload(response)["error"] + + +# --------------------------------------------------------------------------- +# (e) route registered exactly ONCE in both modes +# --------------------------------------------------------------------------- +def test_e1_routes_class_registers_undo_route_exactly_once() -> None: + app = web.Application() + PendingDeleteRoutes.setup_routes(app) + + assert len(_undo_delete_routes(app)) == 1 + + +def test_e2_plugin_mode_registers_undo_route_exactly_once( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + from py import lora_manager + + app = web.Application() + app._handler_args = {"max_field_size": 1024} + monkeypatch.setattr(lora_manager.PromptServer, "instance", SimpleNamespace(app=app)) + + monkeypatch.setattr(lora_manager, "register_default_model_types", lambda: None) + monkeypatch.setattr( + lora_manager.ModelServiceFactory, "setup_all_routes", lambda app_: None + ) + monkeypatch.setattr( + lora_manager.ModelServiceFactory, "get_registered_types", lambda: ["dummy"] + ) + for name in ( + "StatsRoutes", + "RecipeRoutes", + "UpdateRoutes", + "MiscRoutes", + "ExampleImagesRoutes", + "PreviewRoutes", + ): + monkeypatch.setattr(lora_manager, name, _DummyRoutes) + monkeypatch.setattr(lora_manager, "ws_manager", _DummyWSManager()) + monkeypatch.setattr( + lora_manager.settings, + "get", + lambda key, default=None: str(tmp_path) if key == "example_images_path" else default, + ) + monkeypatch.setattr(app.router, "add_static", lambda *a, **k: SimpleNamespace()) + + lora_manager.LoraManager.add_routes() + + assert len(_undo_delete_routes(app)) == 1 + + +def test_e3_standalone_mode_registers_undo_route_exactly_once( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + import standalone as standalone_module + + app = web.Application() + locales_dir = tmp_path / "locales" + locales_dir.mkdir() + static_dir = tmp_path / "static" + static_dir.mkdir() + monkeypatch.setattr( + standalone_module, + "config", + SimpleNamespace(i18n_path=str(locales_dir), static_path=str(static_dir)), + ) + + import py.services.model_service_factory as factory_module + + monkeypatch.setattr(factory_module, "register_default_model_types", lambda: None) + monkeypatch.setattr( + factory_module.ModelServiceFactory, + "setup_all_routes", + classmethod(lambda cls, app_arg: None), + ) + monkeypatch.setattr("py.routes.recipe_routes.RecipeRoutes", _DummyRoutes) + monkeypatch.setattr("py.routes.update_routes.UpdateRoutes", _DummyRoutes) + monkeypatch.setattr("py.routes.misc_routes.MiscRoutes", _DummyRoutes) + monkeypatch.setattr("py.routes.stats_routes.StatsRoutes", _DummyRoutes) + monkeypatch.setattr("py.routes.example_images_routes.ExampleImagesRoutes", _DummyRoutes) + monkeypatch.setattr("py.routes.preview_routes.PreviewRoutes", _DummyRoutes) + + async def _noop_ws_handler(request: Any) -> web.Response: + return web.Response(status=204) + + ws_stub = SimpleNamespace( + handle_connection=_noop_ws_handler, + handle_download_connection=_noop_ws_handler, + handle_init_connection=_noop_ws_handler, + ) + monkeypatch.setattr("py.services.websocket_manager.ws_manager", ws_stub) + + server = SimpleNamespace(app=app) + standalone_module.StandaloneLoraManager.add_routes(server) + + assert len(_undo_delete_routes(app)) == 1 + + +# --------------------------------------------------------------------------- +# (f) RESTART-SAFE UNDO (T8): staged batch on disk, no in-process timer +# --------------------------------------------------------------------------- +async def test_f_restart_safe_undo_restores_files_and_cache( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + # NOTE: the original file does NOT exist - it lives in the staging batch + # (a delete moved it away before the restart). + snapshot = {"file_path": str(model), "sha256": "c" * 64, "tags": ["beta"]} + + # Simulate a post-restart state: the batch + manifest exist on disk but + # nothing was staged in-process (no timer task, no _known_roots entry). + restart_batch_id = "b" * 32 + batch_dir = root / PENDING_DELETE_DIR_NAME / restart_batch_id + batch_dir.mkdir(parents=True) + (batch_dir / "model.safetensors").write_bytes(b"survived-restart") + _write_batch_manifest( + batch_dir, + batch_id=restart_batch_id, + kind="model", + model_type="loras", + expires_at=int(time.time()) + 3600, + entries=[ + { + "staged": str(batch_dir / "model.safetensors"), + "original": str(model), + "restored": False, + } + ], + model_snapshot=dict(snapshot), + ) + + scanner = FakeScanner(root, model_type="lora") + await _register_scanners(monkeypatch, lora=scanner) + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=restart_batch_id) # pyright: ignore[reportArgumentType] + ) + + assert response.status == 200 + assert model.read_bytes() == b"survived-restart" + assert not batch_dir.exists() + assert len(scanner._cache.raw_data) == 1 + assert scanner._cache.raw_data[0]["file_path"] == str(model) + assert scanner._tags_count == {"beta": 1} + + +# --------------------------------------------------------------------------- +# (g) RESCAN STALENESS (T7): stale raw_data entry replaced by the snapshot +# --------------------------------------------------------------------------- +async def test_g_rescan_stale_cache_entry_replaced_by_snapshot( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + model.write_bytes(b"model-data") + snapshot = { + "file_path": str(model), + "sha256": "d" * 64, + "tags": ["gamma"], + "model_name": "model", + } + + scanner = FakeScanner(root, model_type="lora") + await _register_scanners(monkeypatch, lora=scanner) + + service = await PendingDeleteService.get_instance() + batch_id = await service.stage_model_delete( + scanner=scanner, + target_dir=str(root), + file_name="model", + main_extension=".safetensors", + original_file_path=str(model), + cached_entry=dict(snapshot), + ) + assert batch_id is not None + + # A rescan between delete and undo re-added a STALE entry for the same path. + stale = { + "file_path": str(model), + "sha256": "z" * 64, + "tags": ["stale-tag"], + "model_name": "stale", + } + scanner._cache.raw_data = [stale] + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + assert response.status == 200 + + # Exactly one entry remains and it equals the snapshot. + assert len(scanner._cache.raw_data) == 1 + assert scanner._cache.raw_data[0] == snapshot + assert scanner._tags_count == {"gamma": 1} + + +# --------------------------------------------------------------------------- +# (h) RECIPE RE-DELETE AFTER UNDO (T9): path map rebuilt by forced refresh +# --------------------------------------------------------------------------- +async def test_h_recipe_redo_delete_after_undo_and_refresh_succeeds( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + recipe_json = tmp_path / "my_recipe.recipe.json" + recipe_data = {"id": "r1", "name": "Recipe One", "file_path": str(recipe_json)} + recipe_json.write_text(json.dumps(recipe_data), encoding="utf-8") + + recipe_scanner = FakeRecipeScanner() + await _register_scanners(monkeypatch, recipe=recipe_scanner) + + service = await PendingDeleteService.get_instance() + batch_id = await service.stage_recipe_delete( + recipe_json_path=str(recipe_json), + image_path=None, + recipe_data=dict(recipe_data), + ) + assert batch_id is not None + recipe_json.unlink() + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + handler = PendingDeleteHandler() + response = await handler.undo_delete(_make_undo_request(batch_id=batch_id)) # pyright: ignore[reportArgumentType] + assert response.status == 200 + assert recipe_json.exists() + assert recipe_scanner.added == [recipe_data] + + # Forced frontend refresh rebuilds the scanner path map. + recipe_scanner.force_refresh() + assert recipe_scanner._json_path_map.get("r1") == str(recipe_json) + + # Immediately deleting the same recipe again succeeds. + batch_id2 = await service.stage_recipe_delete( + recipe_json_path=str(recipe_json), + image_path=None, + recipe_data=dict(recipe_data), + ) + assert batch_id2 is not None + recipe_json.unlink() + response2 = await handler.undo_delete(_make_undo_request(batch_id=batch_id2)) # pyright: ignore[reportArgumentType] + assert response2.status == 200 + assert recipe_json.exists() + + +# --------------------------------------------------------------------------- +# (i) recipe undo WITHOUT prior refresh still restores files (no crash) +# --------------------------------------------------------------------------- +async def test_i_recipe_undo_without_refresh_still_restores_files( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + recipe_json = tmp_path / "no_refresh.recipe.json" + recipe_data = {"id": "r2", "name": "No Refresh"} + recipe_json.write_text(json.dumps(recipe_data), encoding="utf-8") + + recipe_scanner = FakeRecipeScanner() + await _register_scanners(monkeypatch, recipe=recipe_scanner) + + service = await PendingDeleteService.get_instance() + batch_id = await service.stage_recipe_delete( + recipe_json_path=str(recipe_json), + image_path=None, + recipe_data=dict(recipe_data), + ) + assert batch_id is not None + recipe_json.unlink() + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + + assert response.status == 200 + assert recipe_json.read_text(encoding="utf-8") == json.dumps(recipe_data) + assert recipe_scanner.added == [recipe_data] + + +# --------------------------------------------------------------------------- +# (j) TAG COUNTS: bulk-delete 2 tagged models, undo restores pre-delete counts +# --------------------------------------------------------------------------- +async def test_j_undo_restores_tag_counts_after_bulk_delete( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "loras" + root.mkdir() + model_a = root / "alpha.safetensors" + model_a.write_bytes(b"alpha") + model_b = root / "beta.safetensors" + model_b.write_bytes(b"beta") + + snapshot_a = {"file_path": str(model_a), "sha256": "e" * 64, "tags": ["alpha"]} + snapshot_b = {"file_path": str(model_b), "sha256": "f" * 64, "tags": ["beta"]} + + scanner = FakeScanner(root, model_type="lora") + await _register_scanners(monkeypatch, lora=scanner) + + service = await PendingDeleteService.get_instance() + batch_a = await service.stage_model_delete( + scanner=scanner, + target_dir=str(root), + file_name="alpha", + main_extension=".safetensors", + original_file_path=str(model_a), + cached_entry=dict(snapshot_a), + ) + batch_b = await service.stage_model_delete( + scanner=scanner, + target_dir=str(root), + file_name="beta", + main_extension=".safetensors", + original_file_path=str(model_b), + cached_entry=dict(snapshot_b), + ) + assert batch_a is not None + assert batch_b is not None + + # Simulate the bulk-delete cache mutation (mirror of + # _batch_update_cache_for_deleted_models): entries removed, tags decremented. + scanner._cache.raw_data = [dict(snapshot_a), dict(snapshot_b)] + scanner._tags_count = {"alpha": 1, "beta": 1} + for model in (dict(snapshot_a), dict(snapshot_b)): + for tag in model.get("tags", []): + if tag in scanner._tags_count: + scanner._tags_count[tag] = max(0, scanner._tags_count[tag] - 1) + if scanner._tags_count[tag] == 0: + del scanner._tags_count[tag] + scanner._cache.raw_data = [] + pre_delete_counts = {"alpha": 1, "beta": 1} + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + handler = PendingDeleteHandler() + # Sequential undo of the (unmerged) constituent batches - the batch_ids + # fallback path used by the frontend for cross-volume bulks. + assert (await handler.undo_delete(_make_undo_request(batch_id=batch_a))).status == 200 # pyright: ignore[reportArgumentType] + assert (await handler.undo_delete(_make_undo_request(batch_id=batch_b))).status == 200 # pyright: ignore[reportArgumentType] + + assert scanner._tags_count == pre_delete_counts + assert {item["file_path"] for item in scanner._cache.raw_data} == { + str(model_a), + str(model_b), + } + assert model_a.read_bytes() == b"alpha" + assert model_b.read_bytes() == b"beta" + + +# --------------------------------------------------------------------------- +# (k) EMBEDDINGS UNDO: resolves the embedding scanner, not the lora scanner +# --------------------------------------------------------------------------- +async def test_k_undo_embeddings_batch_updates_embedding_cache_only( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + emb_root = tmp_path / "embeddings" + emb_root.mkdir() + model = emb_root / "my_embedding.safetensors" + model.write_bytes(b"emb-data") + snapshot = { + "file_path": str(model), + "sha256": "g" * 64, + "tags": [], + "model_name": "my_embedding", + } + + lora_scanner = FakeScanner(tmp_path / "loras", model_type="lora") + emb_scanner = FakeScanner(emb_root, model_type="embedding") + lora_scanner._cache.raw_data = [{"file_path": "/loras/other.safetensors", "sha256": "x"}] + emb_scanner._cache.raw_data = [dict(snapshot)] + await _register_scanners(monkeypatch, lora=lora_scanner, embedding=emb_scanner) + + service = await PendingDeleteService.get_instance() + batch_id = await service.stage_model_delete( + scanner=emb_scanner, + target_dir=str(emb_root), + file_name="my_embedding", + main_extension=".safetensors", + original_file_path=str(model), + cached_entry=dict(snapshot), + ) + assert batch_id is not None + emb_scanner._cache.raw_data = [] + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + assert response.status == 200 + + # The EMBEDDINGS cache was updated, not the lora cache. + assert len(emb_scanner._cache.raw_data) == 1 + assert emb_scanner._cache.raw_data[0]["file_path"] == str(model) + assert len(lora_scanner._cache.raw_data) == 1 # untouched + assert model.read_bytes() == b"emb-data" + + +# --------------------------------------------------------------------------- +# (BULK-UNDO) MERGED bulk undo restores ALL cache entries (not just the winner) +# --------------------------------------------------------------------------- +async def test_bulk_undo_after_merge_restores_all_cache_entries( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "loras" + root.mkdir() + model_a = root / "alpha.safetensors" + model_a.write_bytes(b"alpha") + model_b = root / "beta.safetensors" + model_b.write_bytes(b"beta") + + snapshot_a = { + "file_path": str(model_a), + "sha256": "e" * 64, + "tags": ["alpha"], + "model_name": "alpha", + } + snapshot_b = { + "file_path": str(model_b), + "sha256": "f" * 64, + "tags": ["beta"], + "model_name": "beta", + } + + scanner = FakeScanner(root, model_type="lora") + await _register_scanners(monkeypatch, lora=scanner) + + service = await PendingDeleteService.get_instance() + batch_a = await service.stage_model_delete( + scanner=scanner, + target_dir=str(root), + file_name="alpha", + main_extension=".safetensors", + original_file_path=str(model_a), + cached_entry=dict(snapshot_a), + ) + batch_b = await service.stage_model_delete( + scanner=scanner, + target_dir=str(root), + file_name="beta", + main_extension=".safetensors", + original_file_path=str(model_b), + cached_entry=dict(snapshot_b), + ) + assert batch_a is not None + assert batch_b is not None + merged = await service.merge_batches([batch_a, batch_b]) + assert merged == batch_a + + # Simulate the bulk-delete cache mutation (entries removed, tags cleared). + scanner._cache.raw_data = [] + scanner._tags_count = {} + scanner._hash_index = FakeHashIndex() + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=merged) # pyright: ignore[reportArgumentType] + ) + assert response.status == 200 + + # BOTH entries restored - the loser's file_path present in raw_data. + assert {item["file_path"] for item in scanner._cache.raw_data} == { + str(model_a), + str(model_b), + } + # Hash index carries both paths. + assert {entry[1] for entry in scanner._hash_index.entries} == { + str(model_a), + str(model_b), + } + # Tag counts match pre-delete counts for tagged models. + assert scanner._tags_count == {"alpha": 1, "beta": 1} + # Files restored byte-identical. + assert model_a.read_bytes() == b"alpha" + assert model_b.read_bytes() == b"beta" + + +# --------------------------------------------------------------------------- +# (BC-1) backward compat: entries WITHOUT snapshots + top-level model_snapshot +# --------------------------------------------------------------------------- +async def test_bc1_old_format_manifest_falls_back_to_top_level_snapshot( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + root = tmp_path / "loras" + root.mkdir() + model = root / "model.safetensors" + snapshot = { + "file_path": str(model), + "sha256": "c" * 64, + "tags": ["beta"], + "model_name": "model", + } + + # Old-format manifest: entries carry NO snapshot; only the top-level + # model_snapshot exists (pre-F3 single-delete manifests). + batch_id = "a" * 32 + batch_dir = root / PENDING_DELETE_DIR_NAME / batch_id + batch_dir.mkdir(parents=True) + (batch_dir / "model.safetensors").write_bytes(b"old-format-data") + _write_batch_manifest( + batch_dir, + batch_id=batch_id, + kind="model", + model_type="loras", + expires_at=int(time.time()) + 3600, + entries=[ + { + "staged": str(batch_dir / "model.safetensors"), + "original": str(model), + "restored": False, + } + ], + model_snapshot=dict(snapshot), + ) + + scanner = FakeScanner(root, model_type="lora") + await _register_scanners(monkeypatch, lora=scanner) + + from py.routes.handlers.pending_delete_handler import PendingDeleteHandler + + response = await PendingDeleteHandler().undo_delete( + _make_undo_request(batch_id=batch_id) # pyright: ignore[reportArgumentType] + ) + assert response.status == 200 + + # The top-level snapshot entry is still restored exactly. + assert len(scanner._cache.raw_data) == 1 + assert scanner._cache.raw_data[0] == snapshot + assert scanner._tags_count == {"beta": 1} + assert model.read_bytes() == b"old-format-data"