From 896ce5eddb0eb1ac91d30dd276bbcd622a56439e Mon Sep 17 00:00:00 2001 From: Will Miao Date: Sat, 3 Oct 2026 14:31:10 +0800 Subject: [PATCH] perf(rename): make bulk filename-template apply O(n) instead of O(n^2) Applying a filename template to a large library re-did O(library) work for every renamed file: a full natsort resort plus whole-table SQLite rewrite and download-history resync after each rename, and a full scan plus resort of the entire recipe collection per renamed LoRA. On a 20k-model library with 300k recipes on a HDD this pushed "Apply to Library" into multi-day runs. - ModelScanner.defer_cache_persist(): bulk loops update the in-memory entry and indexes only; resort + persist + download-history sync run once at context exit, forced even on cancellation/error since files are already renamed on disk. Single-rename callers keep immediate per-call behavior. - RecipeScanner.build_lora_hash_index(): one-shot hash -> recipes index so per-file lookups are O(1); update_lora_filename_by_hash gains hash_index / defer_maintenance params, with a single finalize_bulk_filename_updates() resort at the end of a bulk session. - ModelLifecycleService.bulk_rename_session() / BulkRenameContext wire the deferred path through rename_model (hash index built lazily on first recipe-touching rename). - Blocking os.rename sequence offloaded via asyncio.to_thread so one file's HDD I/O no longer stalls the event loop (no cross-file parallelism). - Skip logic, per-batch WebSocket progress, cancellation, and result counters unchanged. --- py/services/model_lifecycle_service.py | 136 +++- py/services/model_scanner.py | 133 +++- py/services/recipe_scanner.py | 72 +- .../use_cases/filename_template_use_case.py | 36 +- tests/conftest.py | 7 +- tests/services/test_bulk_rename_apply.py | 639 ++++++++++++++++++ tests/services/test_recipe_scanner.py | 109 +++ tests/services/test_use_cases.py | 11 +- 8 files changed, 1080 insertions(+), 63 deletions(-) create mode 100644 tests/services/test_bulk_rename_apply.py diff --git a/py/services/model_lifecycle_service.py b/py/services/model_lifecycle_service.py index f5c8fb4b..b6a3fb4c 100644 --- a/py/services/model_lifecycle_service.py +++ b/py/services/model_lifecycle_service.py @@ -2,10 +2,12 @@ from __future__ import annotations +import asyncio import json import logging import os -from typing import Any, Awaitable, Callable, Dict, Iterable, List, Mapping, Optional, TYPE_CHECKING, cast +from contextlib import asynccontextmanager +from typing import Any, AsyncIterator, Awaitable, Callable, Dict, Iterable, List, Mapping, Optional, TYPE_CHECKING, cast from ..services.service_registry import ServiceRegistry from ..services.pending_delete_service import get_pending_delete_service @@ -107,6 +109,36 @@ def _require_path_in_library_roots(file_path: str, scanner, *, label: str = "pat ) +class BulkRenameContext: + """Per-session state threaded through ``rename_model`` calls of a bulk rename. + + Holds the lazily built recipe hash index so a bulk rename loop pays the + O(recipes) index build at most once (on the first recipe-touching rename) + instead of rescanning every recipe per renamed file. Also tracks whether + any recipe was re-pointed so the session finalizes recipe maintenance only + when needed. + """ + + def __init__(self, recipe_scanner: Any) -> None: + self._recipe_scanner = recipe_scanner + self._recipe_hash_index: Optional[Dict[str, List[Dict[str, Any]]]] = None + self.recipes_touched = False + + @property + def recipe_scanner(self) -> Any: + return self._recipe_scanner + + async def get_recipe_hash_index(self) -> Optional[Dict[str, List[Dict[str, Any]]]]: + """Return the lora-hash → recipes index, building it on first use.""" + if self._recipe_scanner is None: + return None + if self._recipe_hash_index is None: + self._recipe_hash_index = ( + await self._recipe_scanner.build_lora_hash_index() + ) + return self._recipe_hash_index + + class ModelLifecycleService: """Co-ordinate destructive and mutating model operations.""" @@ -365,10 +397,45 @@ class ModelLifecycleService: return await self._scanner.bulk_delete_models(file_paths) + @asynccontextmanager + async def bulk_rename_session(self) -> AsyncIterator[BulkRenameContext]: + """Context for bulk rename loops (filename-template "Apply to Library"). + + While active, the per-file ``update_single_model_cache`` resort/persist + chain and the per-file recipe folder-metadata refresh/resort are + deferred; both run exactly once when the outermost session exits — see + ``ModelScanner.defer_cache_persist`` and + ``RecipeScanner.finalize_bulk_filename_updates``. The finalize steps run + even on cancellation or mid-loop errors, because files are already + renamed on disk and the caches must not be left diverging. + + Yields a :class:`BulkRenameContext` to pass as ``bulk_context`` into + each ``rename_model`` call of the loop. + """ + recipe_scanner = await self._recipe_scanner_factory() + context = BulkRenameContext(recipe_scanner) + async with self._scanner.defer_cache_persist(): + try: + yield context + finally: + if recipe_scanner is not None and context.recipes_touched: + try: + await recipe_scanner.finalize_bulk_filename_updates() + except Exception as exc: # pragma: no cover - defensive logging + logger.error( + "Error finalizing bulk recipe updates: %s", exc + ) + async def rename_model( - self, *, file_path: str, new_file_name: str + self, *, file_path: str, new_file_name: str, bulk_context: Optional[BulkRenameContext] = None ) -> Dict[str, object]: - """Rename a model and its companion artefacts.""" + """Rename a model and its companion artefacts. + + When ``bulk_context`` is given (bulk rename loop), the recipe + re-pointing uses the session's prebuilt hash index and defers recipe + maintenance to the session finalize; the scanner cache persist is + likewise deferred by the surrounding ``bulk_rename_session``. + """ if not file_path or not new_file_name: raise ValueError("File path and new file name are required") @@ -419,20 +486,11 @@ class ModelLifecycleService: raw_hash = metadata.get("sha256") if isinstance(metadata, dict) else None hash_value = raw_hash if isinstance(raw_hash, str) else None - renamed_files: List[str] = [] - new_metadata_path: Optional[str] = None new_preview: Optional[str] = None - for old_path, pattern in existing_files: - ext = self._get_multipart_ext(pattern) - new_path = os.path.join( - os.path.dirname(old_path), f"{new_file_name}{ext}" - ).replace(os.sep, "/") - os.rename(old_path, new_path) - renamed_files.append(new_path) - - if ext == ".metadata.json": - new_metadata_path = new_path + renamed_files, new_metadata_path = await asyncio.to_thread( + self._rename_companion_files, existing_files, new_file_name + ) if metadata and new_metadata_path: metadata["file_name"] = new_file_name @@ -457,12 +515,26 @@ class ModelLifecycleService: ) if hash_value and getattr(self._scanner, "model_type", "") == "lora": - recipe_scanner = await self._recipe_scanner_factory() + if bulk_context is not None: + recipe_scanner = bulk_context.recipe_scanner + hash_index = await bulk_context.get_recipe_hash_index() + defer_maintenance = True + else: + recipe_scanner = await self._recipe_scanner_factory() + hash_index = None + defer_maintenance = False if recipe_scanner: try: - await recipe_scanner.update_lora_filename_by_hash( - hash_value, new_file_name + file_count, cache_count = ( + await recipe_scanner.update_lora_filename_by_hash( + hash_value, + new_file_name, + hash_index=hash_index, + defer_maintenance=defer_maintenance, + ) ) + if bulk_context is not None and (file_count or cache_count): + bulk_context.recipes_touched = True except Exception as exc: # pragma: no cover - defensive logging logger.error( "Error updating recipe references for %s: %s", @@ -478,6 +550,34 @@ class ModelLifecycleService: "reload_required": False, } + def _rename_companion_files( + self, + existing_files: List[tuple[str, str]], + new_file_name: str, + ) -> tuple[List[str], Optional[str]]: + """Rename all companion files, off the event loop thread. + + Runs the blocking ``os.rename`` sequence for one model in a worker + thread so a single file's HDD I/O does not stall the event loop. + Never parallelized across files: one model's renames stay sequential + and the helper holds no locks. + """ + renamed_files: List[str] = [] + new_metadata_path: Optional[str] = None + + for old_path, pattern in existing_files: + ext = self._get_multipart_ext(pattern) + new_path = os.path.join( + os.path.dirname(old_path), f"{new_file_name}{ext}" + ).replace(os.sep, "/") + os.rename(old_path, new_path) + renamed_files.append(new_path) + + if ext == ".metadata.json": + new_metadata_path = new_path + + return renamed_files, new_metadata_path + @staticmethod def _get_multipart_ext(filename: str) -> str: """Return the extension for files with compound suffixes.""" diff --git a/py/services/model_scanner.py b/py/services/model_scanner.py index bf5d430e..7724e1cb 100644 --- a/py/services/model_scanner.py +++ b/py/services/model_scanner.py @@ -4,6 +4,7 @@ import logging import asyncio import time import shutil +from contextlib import asynccontextmanager from dataclasses import dataclass from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional, Sequence, Set, Tuple, Type, Union, cast @@ -162,6 +163,13 @@ class ModelScanner: self._name_display_mode = self._resolve_name_display_mode() self._cancel_requested = False # Flag for cancellation self._move_locks: Dict[str, asyncio.Lock] = {} # Per-source-file move locks + # Bulk-operation deferral: while _defer_persist_depth > 0, + # update_single_model_cache() skips the per-call resort/persist and + # only marks _deferred_persist_pending; the exit of the outermost + # defer_cache_persist() context finalizes once (see + # _finalize_deferred_cache_persist). + self._defer_persist_depth = 0 + self._deferred_persist_pending = False self._autov3_backfill_scheduled = False # One-time AutoV3 backfill trigger per process # Guard against concurrent all-folders backfill walks (cold fallback # for persisted snapshots that predate folder recording). @@ -773,11 +781,11 @@ class ModelScanner: except Exception as exc: logger.warning("AutoV3 backfill failed: %s", exc) - async def _save_persistent_cache(self, scan_result: CacheBuildResult) -> None: + async def _save_persistent_cache(self, scan_result: CacheBuildResult, *, force: bool = False) -> None: if not scan_result or not getattr(self, '_persistent_cache', None): return - if self.is_cancelled(): + if self.is_cancelled() and not force: logger.info( f"{self.model_type.capitalize()} Scanner: Skipping _save_persistent_cache " "after cancellation" @@ -836,7 +844,7 @@ class ModelScanner: bucket.append(path) return snapshot - async def _persist_current_cache(self) -> None: + async def _persist_current_cache(self, *, force: bool = False) -> None: if self._cache is None or not getattr(self, '_persistent_cache', None): return @@ -851,7 +859,7 @@ class ModelScanner: else None ), ) - await self._save_persistent_cache(snapshot) + await self._save_persistent_cache(snapshot, force=force) await self._sync_download_history(snapshot.raw_data, source='scan') def _count_model_files(self) -> int: """Count all model files with supported extensions in all roots @@ -2646,11 +2654,85 @@ class ModelScanner: logger.error(f"Error updating metadata paths: {e}", exc_info=True) return None + @asynccontextmanager + async def defer_cache_persist(self): + """Defer heavyweight cache maintenance for a bulk operation. + + While at least one ``defer_cache_persist`` context is active, + :meth:`update_single_model_cache` performs only the in-memory entry + swap plus incremental index updates — it skips the full version-index + rebuild, the natsort resort, and the whole-table SQLite persist plus + download-history sync that normally run per call. When the outermost + context exits, the pending maintenance runs **once** (resort, persist, + download-history sync). + + The final persist is forced: it runs even when the scanner's + cancellation flag is set or the wrapped block raised, because callers + use this around operations that already mutated files on disk and the + cache must not be left diverging from reality. + + Intended for bulk rename/move loops (e.g. the filename-template "Apply + to Library" flow). Single-shot callers keep the immediate per-call + behavior by not entering this context. + """ + self._defer_persist_depth = getattr(self, "_defer_persist_depth", 0) + 1 + try: + yield + finally: + self._defer_persist_depth -= 1 + if self._defer_persist_depth == 0: + await self._finalize_deferred_cache_persist() + + @property + def _cache_persist_deferred(self) -> bool: + """True while cache resort/persist is deferred to a bulk finalize.""" + return getattr(self, "_defer_persist_depth", 0) > 0 + + async def _finalize_deferred_cache_persist(self) -> None: + """Run the resort + persist deferred by ``defer_cache_persist``. + + Best-effort: failures are logged, never raised, so an error here + cannot mask the outcome of the bulk operation itself (including + cancellation). + """ + if not getattr(self, "_deferred_persist_pending", False): + return + self._deferred_persist_pending = False + if self._cache is None: + return + try: + # resort() rebuilds the version index and folder list, so the + # per-call rebuilds skipped during deferral are covered here. + await self._cache.resort() + await self._persist_current_cache(force=True) + self.bump_cache_version() + except Exception: + logger.error( + "%s Scanner: failed to finalize deferred cache persist", + self.model_type.capitalize(), + exc_info=True, + ) + async def update_single_model_cache(self, original_path: str, new_path: str, metadata: Optional[Dict[str, Any]], recalculate_type: bool = False) -> Union[bool, Dict[str, Any]]: - """Update cache after a model has been moved or modified""" + """Update cache after a model has been moved or modified. + + Performs the full maintenance chain (version-index rebuild, resort, + whole-table persist, download-history sync) unless the scanner is + inside a :meth:`defer_cache_persist` context, in which case only + the in-memory entry swap and incremental index updates run and the + heavy chain executes once at context exit. + """ + deferred = self._cache_persist_deferred cache = await self.get_cached_data() - existing_item = next((item for item in cache.raw_data if item['file_path'] == original_path), None) + existing_index: Optional[int] = None + existing_item = None + for idx, item in enumerate(cache.raw_data): + if item['file_path'] == original_path: + existing_item = item + existing_index = idx + break + if existing_item: cache.remove_from_version_index(existing_item) @@ -2662,11 +2744,18 @@ class ModelScanner: del self._tags_count[tag] self._hash_index.remove_by_path(original_path) - - cache.raw_data = [ - item for item in cache.raw_data - if item['file_path'] != original_path - ] + + if deferred: + # In-place swap avoids the O(n) list rebuild per renamed file; + # indexes were already updated incrementally above/below, and the + # folder recompute happens in the single finalize resort(). + if existing_index is not None: + cache.raw_data.pop(existing_index) + else: + cache.raw_data = [ + item for item in cache.raw_data + if item['file_path'] != original_path + ] cache_modified = bool(existing_item) or bool(metadata) cache_entry: Optional[Dict[str, Any]] = None @@ -2707,8 +2796,11 @@ class ModelScanner: cache_entry.get('autov3') or None, ) - all_folders = set(item['folder'] for item in cache.raw_data) - cache.folders = sorted(list(all_folders), key=lambda x: x.lower()) + if not deferred: + # O(n) over raw_data; the finalize resort() recomputes the + # folder list once, so bulk callers skip it per file. + all_folders = set(item['folder'] for item in cache.raw_data) + cache.folders = sorted(list(all_folders), key=lambda x: x.lower()) # The move target may live in directories the last scan never saw; # record the destination folder (and its parents) in the known @@ -2723,13 +2815,18 @@ class ModelScanner: for tag in cache_entry.get('tags', []): self._tags_count[tag] = self._tags_count.get(tag, 0) + 1 - cache.rebuild_version_index() + if deferred: + if cache_modified: + self._deferred_persist_pending = True + self.bump_cache_version() + else: + cache.rebuild_version_index() - await cache.resort() + await cache.resort() - if cache_modified: - await self._persist_current_cache() - self.bump_cache_version() + if cache_modified: + await self._persist_current_cache() + self.bump_cache_version() if metadata and cache_entry is not None: return cache_entry diff --git a/py/services/recipe_scanner.py b/py/services/recipe_scanner.py index 13e3650e..9d118102 100644 --- a/py/services/recipe_scanner.py +++ b/py/services/recipe_scanner.py @@ -4580,14 +4580,64 @@ class RecipeScanner: return syntax_parts + async def build_lora_hash_index(self) -> Dict[str, List[Dict[str, Any]]]: + """Build a one-shot lowercase-LoRA-hash → recipes index. + + Scans the recipe cache exactly once (O(recipes × loras)) and returns + a mapping of lowercase lora ``hash`` to the list of recipe dicts + containing it. Bulk rename loops pass this index to + :meth:`update_lora_filename_by_hash` so per-file lookups are O(1) + instead of rescanning every recipe for each renamed LoRA. + """ + cache = await self.get_cached_data() + index: Dict[str, List[Dict[str, Any]]] = {} + if not cache or not cache.raw_data: + return index + + for recipe in cache.raw_data: + loras = recipe.get("loras", []) + if not isinstance(loras, list): + continue + for lora in loras: + if not isinstance(lora, dict): + continue + hash_value = (lora.get("hash") or "").lower() + if hash_value: + index.setdefault(hash_value, []).append(recipe) + return index + + async def finalize_bulk_filename_updates(self) -> None: + """Run once after a bulk rename session that deferred maintenance. + + Refreshes folder metadata and schedules a single re-sort. Filename-only + renames never change recipe folders, so the deferred refresh is + redundant but cheap; skipping it per file is what makes bulk renames + O(1)-per-file. + """ + if self._cache is None: + return + self._schedule_resort() + async def update_lora_filename_by_hash( - self, hash_value: str, new_file_name: str + self, + hash_value: str, + new_file_name: str, + *, + hash_index: Optional[Dict[str, List[Dict[str, Any]]]] = None, + defer_maintenance: bool = False, ) -> Tuple[int, int]: """Update file_name in all recipes that contain a LoRA with the specified hash. Args: hash_value: The SHA256 hash value of the LoRA new_file_name: The new file_name to set + hash_index: Optional prebuilt index from + :meth:`build_lora_hash_index`. When given, the O(recipes) + cache scan (and its folder-metadata walk) is skipped and the + affected recipes are looked up directly — the bulk rename path. + defer_maintenance: When True, skip the folder-metadata refresh and + resort scheduling. The caller MUST run + :meth:`finalize_bulk_filename_updates` exactly once afterwards. Returns: Tuple[int, int]: (number of recipes updated in files, number of recipes updated in cache) @@ -4598,17 +4648,21 @@ class RecipeScanner: # Always use lowercase hash for consistency hash_value = hash_value.lower() - # Get cache - cache = await self.get_cached_data() - if not cache or not cache.raw_data: - return 0, 0 + if hash_index is not None: + candidate_recipes = hash_index.get(hash_value, []) + else: + # Get cache + cache = await self.get_cached_data() + if not cache or not cache.raw_data: + return 0, 0 + candidate_recipes = cache.raw_data file_updated_count = 0 cache_updated_count = 0 - # Find recipes that need updating from the cache + # Find recipes that need updating recipes_to_update = [] - for recipe in cache.raw_data: + for recipe in candidate_recipes: loras = recipe.get("loras", []) if not isinstance(loras, list): continue @@ -4654,7 +4708,9 @@ class RecipeScanner: # We don't necessarily need to resort because LoRA file_name isn't a sort key, # but we might want to schedule a resort if we're paranoid or if searching relies on sorted state. # Given it's a rename of a dependency, search results might change if searching by LoRA name. - self._schedule_resort() + # Bulk callers defer this to a single finalize_bulk_filename_updates() call. + if not defer_maintenance: + self._schedule_resort() return file_updated_count, cache_updated_count diff --git a/py/services/use_cases/filename_template_use_case.py b/py/services/use_cases/filename_template_use_case.py index e569d681..659e7c61 100644 --- a/py/services/use_cases/filename_template_use_case.py +++ b/py/services/use_cases/filename_template_use_case.py @@ -33,7 +33,9 @@ class FilenameTemplateUseCase: An empty template restores the recorded original filename instead of rendering a template. Shares the auto-organize lock (and its in-progress error) so a bulk rename never runs concurrently with an auto-organize - operation. + operation. The whole loop runs inside a bulk rename session so cache + persist/resort and recipe maintenance happen once at the end instead of + per renamed file. """ def __init__( @@ -106,23 +108,24 @@ class FilenameTemplateUseCase: await self._emit_progress(progress_callback, result, "started") - for index in range(0, result.total, AUTO_ORGANIZE_BATCH_SIZE): - if self._scanner.is_cancelled(): - logger.info( - "Filename template apply cancelled for %s", self._model_type - ) - break - - batch = models[index : index + AUTO_ORGANIZE_BATCH_SIZE] - for model in batch: + async with self._lifecycle_service.bulk_rename_session() as bulk_context: + for index in range(0, result.total, AUTO_ORGANIZE_BATCH_SIZE): if self._scanner.is_cancelled(): + logger.info( + "Filename template apply cancelled for %s", self._model_type + ) break - await self._process_model(model, template, result) - result.processed += 1 - await self._emit_progress(progress_callback, result, "processing") - # Yield between batches so the server stays responsive. - await asyncio.sleep(0.1) + batch = models[index : index + AUTO_ORGANIZE_BATCH_SIZE] + for model in batch: + if self._scanner.is_cancelled(): + break + await self._process_model(model, template, result, bulk_context) + result.processed += 1 + + await self._emit_progress(progress_callback, result, "processing") + # Yield between batches so the server stays responsive. + await asyncio.sleep(0.1) if self._scanner.is_cancelled(): result.status = "cancelled" @@ -150,6 +153,7 @@ class FilenameTemplateUseCase: model: Dict[str, Any], template: str, result: AutoOrganizeResult, + bulk_context: Any = None, ) -> None: model_name = model.get("model_name", "Unknown") try: @@ -177,7 +181,7 @@ class FilenameTemplateUseCase: return await self._lifecycle_service.rename_model( - file_path=file_path, new_file_name=new_stem + file_path=file_path, new_file_name=new_stem, bulk_context=bulk_context ) result.success_count += 1 diff --git a/tests/conftest.py b/tests/conftest.py index d4c44e5f..04238b74 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,9 +3,10 @@ import importlib.util import inspect import sys import types +from contextlib import asynccontextmanager from dataclasses import dataclass, field from pathlib import Path -from typing import Any, Dict, List, Optional, Sequence +from typing import Any, AsyncIterator, Dict, List, Optional, Sequence from unittest import mock import pytest @@ -177,6 +178,10 @@ class MockScanner: def reset_cancellation(self) -> None: self._cancelled = False + @asynccontextmanager + async def defer_cache_persist(self) -> AsyncIterator[None]: + yield None + async def get_cached_data(self, force_refresh: bool = False): return self._cache diff --git a/tests/services/test_bulk_rename_apply.py b/tests/services/test_bulk_rename_apply.py new file mode 100644 index 00000000..3f0d94e0 --- /dev/null +++ b/tests/services/test_bulk_rename_apply.py @@ -0,0 +1,639 @@ +"""Performance-regression tests for the bulk filename-template rename path. + +The bulk "Apply to Library" rename must be O(1)-per-file: the heavyweight +cache persist/resort chain and the recipe maintenance run exactly once per +bulk operation (even on cancellation or mid-loop errors), while single-shot +renames keep their immediate per-call behavior. +""" + +from __future__ import annotations + +import asyncio +import json +from contextlib import asynccontextmanager +from pathlib import Path +from typing import Any, AsyncIterator, Dict, List, Optional + +import pytest + +from py.services import model_scanner as model_scanner_module +from py.services.model_cache import ModelCache +from py.services.model_lifecycle_service import ModelLifecycleService +from py.services.model_scanner import ModelScanner +from py.services.settings_manager import get_settings_manager +from py.services.use_cases.filename_template_use_case import FilenameTemplateUseCase +from py.utils.models import BaseModelMetadata + + +class BulkDummyScanner(ModelScanner): + """Minimal concrete scanner for cache-behavior tests.""" + + def __init__(self, root: Path): + self._root = str(root) + super().__init__( + model_type="dummy", + model_class=BaseModelMetadata, + file_extensions={".txt"}, + ) + + def get_model_roots(self) -> List[str]: + return [self._root] + + +@pytest.fixture(autouse=True) +def _reset_model_scanner_singletons(): + ModelScanner._instances.clear() + ModelScanner._locks.clear() + yield + ModelScanner._instances.clear() + ModelScanner._locks.clear() + + +@pytest.fixture(autouse=True) +def _disable_persistent_cache_env(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LORA_MANAGER_DISABLE_PERSISTENT_CACHE", "1") + + +@pytest.fixture(autouse=True) +def _stub_register_service(monkeypatch: pytest.MonkeyPatch): + async def noop(*_args: Any, **_kwargs: Any) -> None: + return None + + monkeypatch.setattr( + model_scanner_module.ServiceRegistry, "register_service", noop + ) + + +def _metadata(stem: str, path: str, sha256: str, civitai_id: int = 1) -> Dict[str, Any]: + return { + "file_name": stem, + "file_path": path, + "model_name": stem, + "sha256": sha256, + "folder": "", + "tags": [], + "size": 1, + "modified": 1.0, + "civitai": {"id": civitai_id, "modelId": 2}, + } + + +class _SpiedCache: + """Scanner cache with call counters for resort/persist/sync.""" + + def __init__(self, scanner: BulkDummyScanner, entries: List[Dict[str, Any]]): + self.cache = ModelCache(raw_data=[dict(e) for e in entries], folders=[]) + scanner._cache = self.cache + self.resort_calls = 0 + self.save_calls: List[bool] = [] + self.sync_calls = 0 + self._original_resort = self.cache.resort + + async def resort(self) -> None: + self.resort_calls += 1 + await self._original_resort() + + async def fake_save(self, scan_result: Any, *, force: bool = False) -> None: + self.save_calls.append(force) + + async def fake_sync(self, raw_data: Any, *, source: str) -> None: + self.sync_calls += 1 + + +async def _make_spied_scanner( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, entries: List[Dict[str, Any]] +) -> tuple[BulkDummyScanner, _SpiedCache]: + scanner = BulkDummyScanner(tmp_path) + spied = _SpiedCache(scanner, entries) + # Truthy stand-in so _persist_current_cache() does not early-return. + scanner._persistent_cache = object() + monkeypatch.setattr(scanner, "_save_persistent_cache", spied.fake_save) + monkeypatch.setattr(scanner, "_sync_download_history", spied.fake_sync) + # Flush the resort task scheduled by ModelCache.__post_init__ before + # installing the counting wrapper. + await asyncio.sleep(0) + await asyncio.sleep(0) + monkeypatch.setattr(spied.cache, "resort", spied.resort) + return scanner, spied + + +# --------------------------------------------------------------------------- +# ModelScanner deferred persist +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_update_single_model_cache_persists_immediately_by_default( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + path_a = (tmp_path / "a.txt").as_posix() + scanner, spied = await _make_spied_scanner( + tmp_path, monkeypatch, [_metadata("a", path_a, "hash-a")] + ) + + await scanner.update_single_model_cache( + path_a, path_a, _metadata("a", path_a, "hash-a") + ) + + assert spied.save_calls == [False] + assert spied.sync_calls == 1 + assert spied.resort_calls == 1 + + +@pytest.mark.asyncio +async def test_deferred_cache_persist_persists_once_for_many_updates( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + entries = [] + for index, stem in enumerate(("a", "b", "c")): + path = (tmp_path / f"{stem}.txt").as_posix() + entries.append(_metadata(stem, path, f"hash-{stem}", civitai_id=index + 1)) + scanner, spied = await _make_spied_scanner(tmp_path, monkeypatch, entries) + + async with scanner.defer_cache_persist(): + for index, stem in enumerate(("a", "b", "c")): + old_path = (tmp_path / f"{stem}.txt").as_posix() + new_path = (tmp_path / f"{stem}-renamed.txt").as_posix() + await scanner.update_single_model_cache( + old_path, + new_path, + _metadata( + f"{stem}-renamed", new_path, f"hash-{stem}", civitai_id=index + 1 + ), + ) + # Nothing heavy may run mid-loop. + assert spied.save_calls == [] + assert spied.sync_calls == 0 + assert spied.resort_calls == 0 + + # Exactly one heavyweight finalize for the whole bulk operation. + assert spied.save_calls == [True] + assert spied.sync_calls == 1 + assert spied.resort_calls == 1 + + # In-memory state is correct, including the incremental version index. + cache = await scanner.get_cached_data() + cached_paths = {item["file_path"] for item in cache.raw_data} + for stem in ("a", "b", "c"): + assert (tmp_path / f"{stem}.txt").as_posix() not in cached_paths + assert (tmp_path / f"{stem}-renamed.txt").as_posix() in cached_paths + assert cache.version_index[1]["file_path"].endswith("a-renamed.txt") + assert scanner._hash_index.get_path("hash-a").endswith("a-renamed.txt") + + +@pytest.mark.asyncio +async def test_deferred_cache_persist_finalizes_despite_cancellation( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + path_a = (tmp_path / "a.txt").as_posix() + scanner, spied = await _make_spied_scanner( + tmp_path, monkeypatch, [_metadata("a", path_a, "hash-a")] + ) + + scanner.cancel_task() + async with scanner.defer_cache_persist(): + new_path = (tmp_path / "a-renamed.txt").as_posix() + await scanner.update_single_model_cache( + path_a, new_path, _metadata("a-renamed", new_path, "hash-a") + ) + + # Files are already renamed on disk, so the persist must be forced even + # though the cancellation flag is set. + assert spied.save_calls == [True] + assert spied.sync_calls == 1 + assert spied.resort_calls == 1 + + +@pytest.mark.asyncio +async def test_deferred_cache_persist_finalizes_despite_mid_loop_error( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + path_a = (tmp_path / "a.txt").as_posix() + scanner, spied = await _make_spied_scanner( + tmp_path, monkeypatch, [_metadata("a", path_a, "hash-a")] + ) + + with pytest.raises(RuntimeError, match="boom"): + async with scanner.defer_cache_persist(): + new_path = (tmp_path / "a-renamed.txt").as_posix() + await scanner.update_single_model_cache( + path_a, new_path, _metadata("a-renamed", new_path, "hash-a") + ) + raise RuntimeError("boom") + + assert spied.save_calls == [True] + assert spied.sync_calls == 1 + + +@pytest.mark.asyncio +async def test_deferred_cache_persist_nested_contexts_finalize_once( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + path_a = (tmp_path / "a.txt").as_posix() + scanner, spied = await _make_spied_scanner( + tmp_path, monkeypatch, [_metadata("a", path_a, "hash-a")] + ) + + async with scanner.defer_cache_persist(): + async with scanner.defer_cache_persist(): + new_path = (tmp_path / "a-renamed.txt").as_posix() + await scanner.update_single_model_cache( + path_a, new_path, _metadata("a-renamed", new_path, "hash-a") + ) + + assert spied.save_calls == [True] + assert spied.resort_calls == 1 + + +@pytest.mark.asyncio +async def test_deferred_cache_persist_without_updates_persists_nothing( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + scanner, spied = await _make_spied_scanner(tmp_path, monkeypatch, []) + + async with scanner.defer_cache_persist(): + pass + + assert spied.save_calls == [] + assert spied.resort_calls == 0 + + +# --------------------------------------------------------------------------- +# ModelLifecycleService bulk rename session +# --------------------------------------------------------------------------- + + +class _SessionScanner: + model_type = "lora" + + def __init__(self, root: Path): + self._root = str(root) + self.cache_updates: List[tuple[str, str]] = [] + self.defer_enters = 0 + self.defer_exits = 0 + + def get_model_roots(self) -> List[str]: + return [self._root] + + @asynccontextmanager + async def defer_cache_persist(self) -> AsyncIterator[None]: + self.defer_enters += 1 + try: + yield + finally: + self.defer_exits += 1 + + async def update_single_model_cache( + self, old_path: str, new_path: str, metadata: Dict[str, Any] + ) -> bool: + self.cache_updates.append((old_path, new_path)) + return True + + +class _PassthroughMetadataManager: + def __init__(self) -> None: + self.saved: List[str] = [] + + async def save_metadata(self, path: str, metadata: Dict[str, Any]) -> bool: + self.saved.append(path) + return True + + +class RecordingRecipeScanner: + """Records bulk-mode recipe calls and applies filename changes.""" + + def __init__(self) -> None: + self.index_builds = 0 + self.updates: List[Dict[str, Any]] = [] + self.finalizes = 0 + self._recipes: Dict[str, Dict[str, Any]] = {} + + def add_recipe(self, recipe_id: str, lora_hash: str, file_name: str) -> None: + self._recipes[recipe_id] = { + "id": recipe_id, + "loras": [{"hash": lora_hash, "file_name": file_name}], + } + + async def build_lora_hash_index(self) -> Dict[str, List[Dict[str, Any]]]: + self.index_builds += 1 + index: Dict[str, List[Dict[str, Any]]] = {} + for recipe in self._recipes.values(): + for lora in recipe["loras"]: + index.setdefault(lora["hash"].lower(), []).append(recipe) + return index + + async def update_lora_filename_by_hash( + self, + hash_value: str, + new_file_name: str, + *, + hash_index: Optional[Dict[str, List[Dict[str, Any]]]] = None, + defer_maintenance: bool = False, + ) -> tuple[int, int]: + self.updates.append( + { + "hash_value": hash_value, + "new_file_name": new_file_name, + "hash_index": hash_index, + "defer_maintenance": defer_maintenance, + } + ) + matched = 0 + for recipe in (self._recipes.values() if hash_index is None else hash_index.get(hash_value.lower(), [])): + for lora in recipe["loras"]: + if lora["hash"].lower() == hash_value.lower(): + lora["file_name"] = new_file_name + matched += 1 + return (matched, matched) + + async def finalize_bulk_filename_updates(self) -> None: + self.finalizes += 1 + + +def _write_model_with_sidecar( + root: Path, stem: str, sha256: str, model_name: Optional[str] = None +) -> str: + model_path = root / f"{stem}.safetensors" + model_path.write_bytes(b"model") + (root / f"{stem}.metadata.json").write_text( + json.dumps( + { + "file_name": stem, + "file_path": model_path.as_posix(), + "model_name": model_name or stem, + "sha256": sha256, + } + ) + ) + return model_path.as_posix() + + +def _make_lifecycle_service( + scanner: _SessionScanner, recipe_scanner: RecordingRecipeScanner +) -> ModelLifecycleService: + async def metadata_loader(path: str) -> Dict[str, Any]: + with open(path, "r", encoding="utf-8") as handle: + return json.load(handle) + + async def recipe_scanner_factory() -> RecordingRecipeScanner: + return recipe_scanner + + return ModelLifecycleService( + scanner=scanner, # pyright: ignore[reportArgumentType] + metadata_manager=_PassthroughMetadataManager(), # pyright: ignore[reportArgumentType] + metadata_loader=metadata_loader, + recipe_scanner_factory=recipe_scanner_factory, + ) + + +@pytest.mark.asyncio +async def test_bulk_rename_session_defers_and_finalizes_once(tmp_path: Path): + scanner = _SessionScanner(tmp_path) + recipe_scanner = RecordingRecipeScanner() + recipe_scanner.add_recipe("r1", "aa" * 32, "old_a") + recipe_scanner.add_recipe("r2", "bb" * 32, "old_b") + service = _make_lifecycle_service(scanner, recipe_scanner) + + path_a = _write_model_with_sidecar(tmp_path, "model-a", "aa" * 32) + path_b = _write_model_with_sidecar(tmp_path, "model-b", "bb" * 32) + + async with service.bulk_rename_session() as bulk_context: + await service.rename_model( + file_path=path_a, new_file_name="renamed-a", bulk_context=bulk_context + ) + await service.rename_model( + file_path=path_b, new_file_name="renamed-b", bulk_context=bulk_context + ) + # No finalize may run before the session ends. + assert recipe_scanner.finalizes == 0 + assert scanner.defer_enters == 1 + assert scanner.defer_exits == 0 + + assert scanner.defer_exits == 1 + # Hash index built at most once for the whole session. + assert recipe_scanner.index_builds == 1 + assert recipe_scanner.finalizes == 1 + assert len(recipe_scanner.updates) == 2 + for update in recipe_scanner.updates: + assert update["hash_index"] is not None + assert update["defer_maintenance"] is True + + # Recipes were re-pointed. + assert recipe_scanner._recipes["r1"]["loras"][0]["file_name"] == "renamed-a" + assert recipe_scanner._recipes["r2"]["loras"][0]["file_name"] == "renamed-b" + + +@pytest.mark.asyncio +async def test_bulk_rename_session_finalizes_recipes_despite_error(tmp_path: Path): + scanner = _SessionScanner(tmp_path) + recipe_scanner = RecordingRecipeScanner() + recipe_scanner.add_recipe("r1", "aa" * 32, "old_a") + service = _make_lifecycle_service(scanner, recipe_scanner) + + path_a = _write_model_with_sidecar(tmp_path, "model-a", "aa" * 32) + + with pytest.raises(RuntimeError, match="mid-loop"): + async with service.bulk_rename_session() as bulk_context: + await service.rename_model( + file_path=path_a, new_file_name="renamed-a", bulk_context=bulk_context + ) + raise RuntimeError("mid-loop") + + assert recipe_scanner.finalizes == 1 + assert scanner.defer_exits == 1 + + +@pytest.mark.asyncio +async def test_bulk_rename_session_skips_recipe_finalize_when_untouched(tmp_path: Path): + scanner = _SessionScanner(tmp_path) + recipe_scanner = RecordingRecipeScanner() + recipe_scanner.add_recipe("r1", "cc" * 32, "old_c") + service = _make_lifecycle_service(scanner, recipe_scanner) + + # Model whose hash matches no recipe: the lookup runs but no recipe is + # touched, so no recipe maintenance is needed at finalize. + path_a = _write_model_with_sidecar(tmp_path, "model-a", "aa" * 32) + + async with service.bulk_rename_session() as bulk_context: + await service.rename_model( + file_path=path_a, new_file_name="renamed-a", bulk_context=bulk_context + ) + + assert len(recipe_scanner.updates) == 1 + assert recipe_scanner.updates[0]["defer_maintenance"] is True + assert recipe_scanner.finalizes == 0 + assert recipe_scanner._recipes["r1"]["loras"][0]["file_name"] == "old_c" + + +@pytest.mark.asyncio +async def test_single_rename_keeps_immediate_recipe_behavior(tmp_path: Path): + scanner = _SessionScanner(tmp_path) + recipe_scanner = RecordingRecipeScanner() + recipe_scanner.add_recipe("r1", "aa" * 32, "old_a") + service = _make_lifecycle_service(scanner, recipe_scanner) + + path_a = _write_model_with_sidecar(tmp_path, "model-a", "aa" * 32) + + await service.rename_model(file_path=path_a, new_file_name="renamed-a") + + assert len(recipe_scanner.updates) == 1 + update = recipe_scanner.updates[0] + assert update["hash_index"] is None + assert update["defer_maintenance"] is False + assert recipe_scanner.finalizes == 0 + assert recipe_scanner._recipes["r1"]["loras"][0]["file_name"] == "renamed-a" + + +# --------------------------------------------------------------------------- +# Use case end-to-end: exactly one heavyweight persist for the whole loop +# --------------------------------------------------------------------------- + + +class _UseCaseScanner(ModelScanner): + def __init__(self, root: Path): + self._root = str(root) + super().__init__( + model_type="lora", + model_class=BaseModelMetadata, + file_extensions={".safetensors"}, + ) + + def get_model_roots(self) -> List[str]: + return [self._root] + + +class _UseCaseLockProvider: + def __init__(self) -> None: + self._lock = asyncio.Lock() + + def is_auto_organize_running(self) -> bool: + return False + + async def get_auto_organize_lock(self) -> asyncio.Lock: + return self._lock + + +def _set_filename_template(template: str, model_type: str = "lora") -> None: + manager = get_settings_manager() + templates = dict(manager.settings.get("download_filename_templates") or {}) + templates[model_type] = template + manager.settings["download_filename_templates"] = templates + + +@pytest.mark.asyncio +async def test_filename_template_bulk_apply_persists_and_repoints_once( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + _set_filename_template("{model_name}") + + entries = [] + for index, (stem, model_name, sha) in enumerate( + (("model-a", "alpha", "aa" * 32), ("model-b", "beta", "bb" * 32)) + ): + path = _write_model_with_sidecar(tmp_path, stem, sha, model_name=model_name) + entries.append(_metadata(model_name, path, sha, civitai_id=index + 1)) + + scanner = _UseCaseScanner(tmp_path) + spied = _SpiedCache(scanner, entries) + scanner._persistent_cache = object() + monkeypatch.setattr(scanner, "_save_persistent_cache", spied.fake_save) + monkeypatch.setattr(scanner, "_sync_download_history", spied.fake_sync) + await asyncio.sleep(0) + await asyncio.sleep(0) + monkeypatch.setattr(spied.cache, "resort", spied.resort) + + recipe_scanner = RecordingRecipeScanner() + recipe_scanner.add_recipe("r1", "aa" * 32, "alpha") + recipe_scanner.add_recipe("r2", "bb" * 32, "beta") + + service = _make_lifecycle_service( # type: ignore[arg-type] + scanner, recipe_scanner # pyright: ignore[reportArgumentType] + ) + use_case = FilenameTemplateUseCase( + scanner=scanner, + lifecycle_service=service, + lock_provider=_UseCaseLockProvider(), + model_type="lora", + ) + + result = await use_case.execute(progress_callback=None) + + assert result.status == "success" + assert result.success_count == 2 + + # One heavyweight persist/resort for the entire bulk operation, not two. + assert spied.save_calls == [True] + assert spied.sync_calls == 1 + assert spied.resort_calls == 1 + + # Recipes re-pointed via a single lazily built hash index, maintenance + # finalized once. + assert recipe_scanner.index_builds == 1 + assert recipe_scanner.finalizes == 1 + assert recipe_scanner._recipes["r1"]["loras"][0]["file_name"] == "alpha" + assert recipe_scanner._recipes["r2"]["loras"][0]["file_name"] == "beta" + + # Cache reflects the new paths. + cache = await scanner.get_cached_data() + cached_paths = {item["file_path"] for item in cache.raw_data} + assert (tmp_path / "alpha.safetensors").as_posix() in cached_paths + assert (tmp_path / "beta.safetensors").as_posix() in cached_paths + + # Sidecars moved alongside the model files. + assert (tmp_path / "alpha.metadata.json").exists() + assert (tmp_path / "beta.metadata.json").exists() + + +@pytest.mark.asyncio +async def test_filename_template_bulk_apply_finalizes_persist_on_cancellation( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + _set_filename_template("{model_name}") + + entries = [] + paths = [] + for index, (stem, model_name, sha) in enumerate( + (("model-a", "alpha", "aa" * 32), ("model-b", "beta", "bb" * 32)) + ): + path = _write_model_with_sidecar(tmp_path, stem, sha, model_name=model_name) + paths.append(path) + entries.append(_metadata(model_name, path, sha, civitai_id=index + 1)) + + scanner = _UseCaseScanner(tmp_path) + spied = _SpiedCache(scanner, entries) + scanner._persistent_cache = object() + monkeypatch.setattr(scanner, "_save_persistent_cache", spied.fake_save) + monkeypatch.setattr(scanner, "_sync_download_history", spied.fake_sync) + await asyncio.sleep(0) + await asyncio.sleep(0) + monkeypatch.setattr(spied.cache, "resort", spied.resort) + + service = _make_lifecycle_service( + scanner, # pyright: ignore[reportArgumentType] + RecordingRecipeScanner(), + ) + + original_rename = service.rename_model + + async def cancelling_rename(**kwargs: Any) -> Dict[str, object]: + result = await original_rename(**kwargs) + scanner.cancel_task() + return result + + monkeypatch.setattr(service, "rename_model", cancelling_rename) + + use_case = FilenameTemplateUseCase( + scanner=scanner, + lifecycle_service=service, + lock_provider=_UseCaseLockProvider(), + model_type="lora", + ) + + result = await use_case.execute(progress_callback=None) + + assert result.status == "cancelled" + # The one file renamed before cancellation is on disk; the cache must + # still be persisted exactly once. + assert spied.save_calls == [True] + assert spied.sync_calls == 1 + assert spied.resort_calls == 1 diff --git a/tests/services/test_recipe_scanner.py b/tests/services/test_recipe_scanner.py index 88b6f47b..2a234b0c 100644 --- a/tests/services/test_recipe_scanner.py +++ b/tests/services/test_recipe_scanner.py @@ -1903,6 +1903,115 @@ async def test_update_lora_filename_by_hash_updates_affected_recipes( assert cached1["loras"][0]["file_name"] == new_name +@pytest.mark.asyncio +async def test_update_lora_filename_by_hash_bulk_mode_skips_scan_and_defers_resort( + tmp_path: Path, recipe_scanner, monkeypatch: pytest.MonkeyPatch +): + """Bulk rename path: prebuilt hash index, no per-call cache walk/resort.""" + scanner, _ = recipe_scanner + recipes_dir = Path(config.loras_roots[0]) / "recipes" + recipes_dir.mkdir(parents=True, exist_ok=True) + + recipe1_id = "recipe-bulk-1" + recipe1_path = recipes_dir / f"{recipe1_id}.recipe.json" + recipe1_data = { + "id": recipe1_id, + "file_path": str(tmp_path / "bulk1.png"), + "title": "Bulk 1", + "modified": 0.0, + "created_date": 0.0, + "loras": [{"file_name": "old_name", "hash": "hash1"}], + } + recipe1_path.write_text(json.dumps(recipe1_data)) + await scanner.add_recipe(dict(recipe1_data)) + + # Build the index once (O(recipes)), as a bulk rename session does. + hash_index = await scanner.build_lora_hash_index() + assert "hash1" in hash_index + + # Spies: the bulk call must not walk the recipe cache again, and must not + # schedule a resort per call. + get_cached_calls = 0 + original_get_cached = scanner.get_cached_data + + async def counting_get_cached_data(*args, **kwargs): + nonlocal get_cached_calls + get_cached_calls += 1 + return await original_get_cached(*args, **kwargs) + + monkeypatch.setattr(scanner, "get_cached_data", counting_get_cached_data) + + resort_calls = 0 + original_schedule = scanner._schedule_resort + + def counting_schedule_resort(**kwargs): + nonlocal resort_calls + resort_calls += 1 + original_schedule(**kwargs) + + monkeypatch.setattr(scanner, "_schedule_resort", counting_schedule_resort) + + file_count, cache_count = await scanner.update_lora_filename_by_hash( + "HASH1", "new_name", hash_index=hash_index, defer_maintenance=True + ) + + assert (file_count, cache_count) == (1, 1) + assert get_cached_calls == 0 + assert resort_calls == 0 + + # The per-match recipe JSON rewrite must still happen. + persisted1 = json.loads(recipe1_path.read_text()) + assert persisted1["loras"][0]["file_name"] == "new_name" + cached1 = next(r for r in hash_index["hash1"] if r["id"] == recipe1_id) + assert cached1["loras"][0]["file_name"] == "new_name" + + # Deferred maintenance runs exactly once at finalize. + await scanner.finalize_bulk_filename_updates() + await asyncio.sleep(0) + assert resort_calls == 1 + + +@pytest.mark.asyncio +async def test_update_lora_filename_by_hash_with_index_still_resorts_when_not_deferred( + tmp_path: Path, recipe_scanner, monkeypatch: pytest.MonkeyPatch +): + """hash_index without defer_maintenance: O(1) lookup, immediate resort.""" + scanner, _ = recipe_scanner + recipes_dir = Path(config.loras_roots[0]) / "recipes" + recipes_dir.mkdir(parents=True, exist_ok=True) + + recipe1_id = "recipe-idx-1" + recipe1_path = recipes_dir / f"{recipe1_id}.recipe.json" + recipe1_data = { + "id": recipe1_id, + "file_path": str(tmp_path / "idx1.png"), + "title": "Index 1", + "modified": 0.0, + "created_date": 0.0, + "loras": [{"file_name": "old_name", "hash": "hash1"}], + } + recipe1_path.write_text(json.dumps(recipe1_data)) + await scanner.add_recipe(dict(recipe1_data)) + + hash_index = await scanner.build_lora_hash_index() + + resort_calls = 0 + original_schedule = scanner._schedule_resort + + def counting_schedule_resort(**kwargs): + nonlocal resort_calls + resort_calls += 1 + original_schedule(**kwargs) + + monkeypatch.setattr(scanner, "_schedule_resort", counting_schedule_resort) + + file_count, cache_count = await scanner.update_lora_filename_by_hash( + "hash1", "new_name", hash_index=hash_index + ) + + assert (file_count, cache_count) == (1, 1) + assert resort_calls == 1 + @pytest.mark.asyncio async def test_get_paginated_data_filters_by_favorite(recipe_scanner): scanner, _ = recipe_scanner diff --git a/tests/services/test_use_cases.py b/tests/services/test_use_cases.py index e24d4817..bafbbeba 100644 --- a/tests/services/test_use_cases.py +++ b/tests/services/test_use_cases.py @@ -1,8 +1,9 @@ import asyncio import logging +from contextlib import asynccontextmanager from dataclasses import dataclass from types import SimpleNamespace -from typing import Any, Dict, List, Optional +from typing import Any, AsyncIterator, Dict, List, Optional import pytest @@ -547,7 +548,13 @@ class StubLifecycleService: self.cancel_on_rename = False self._scanner = scanner - async def rename_model(self, *, file_path: str, new_file_name: str) -> Dict[str, Any]: + @asynccontextmanager + async def bulk_rename_session(self) -> AsyncIterator[None]: + yield None + + async def rename_model( + self, *, file_path: str, new_file_name: str, bulk_context: Any = None + ) -> Dict[str, Any]: if self.error is not None: raise self.error self.renames.append({"file_path": file_path, "new_file_name": new_file_name})