diff --git a/invokeai/app/services/object_serializer/object_serializer_disk.py b/invokeai/app/services/object_serializer/object_serializer_disk.py index bbd3f785507..611f6c41026 100644 --- a/invokeai/app/services/object_serializer/object_serializer_disk.py +++ b/invokeai/app/services/object_serializer/object_serializer_disk.py @@ -65,7 +65,7 @@ def save(self, obj: T) -> str: def delete(self, name: str) -> None: file_path = self._get_path(name) - file_path.unlink() + file_path.unlink(missing_ok=True) @property def _obj_class_name(self) -> str: diff --git a/invokeai/app/services/object_serializer/object_serializer_forward_cache.py b/invokeai/app/services/object_serializer/object_serializer_forward_cache.py index ae00173e422..2bd044c5b2d 100644 --- a/invokeai/app/services/object_serializer/object_serializer_forward_cache.py +++ b/invokeai/app/services/object_serializer/object_serializer_forward_cache.py @@ -1,5 +1,5 @@ -from queue import Queue -from threading import Lock +from queue import Empty, Queue +from threading import RLock from typing import TYPE_CHECKING, Optional, TypeVar from invokeai.app.services.object_serializer.object_serializer_base import ObjectSerializerBase @@ -22,9 +22,9 @@ def __init__(self, underlying_storage: ObjectSerializerBase[T], max_cache_size: self._cache: dict[str, T] = {} self._cache_ids = Queue[str]() self._max_cache_size = max_cache_size - # Guards the in-memory cache so concurrent session-processor workers (multi-GPU) can't race - # the check-then-evict in `_set_cache` (which could otherwise raise KeyError on eviction). - self._cache_lock = Lock() + # Guards the in-memory cache and eviction queue so concurrent session-processor workers (multi-GPU) + # cannot interleave cache transitions. Reentrancy allows transition helpers to share this lock. + self._cache_lock = RLock() def start(self, invoker: "Invoker") -> None: self._invoker = invoker @@ -39,13 +39,14 @@ def stop(self, invoker: "Invoker") -> None: stop_op(invoker) def load(self, name: str) -> T: - cache_item = self._get_cache(name) - if cache_item is not None: - return cache_item + with self._cache_lock: + cache_item = self._get_cache(name) + if cache_item is not None: + return cache_item - obj = self._underlying_storage.load(name) - self._set_cache(name, obj) - return obj + obj = self._underlying_storage.load(name) + self._set_cache(name, obj) + return obj def save(self, obj: T) -> str: name = self._underlying_storage.save(obj) @@ -53,11 +54,30 @@ def save(self, obj: T) -> str: return name def delete(self, name: str) -> None: - self._underlying_storage.delete(name) + try: + with self._cache_lock: + try: + self._underlying_storage.delete(name) + finally: + if name in self._cache: + del self._cache[name] + self._remove_cache_id(name) + finally: + self._on_deleted(name) + + def _remove_cache_id(self, name: str) -> None: with self._cache_lock: - if name in self._cache: - del self._cache[name] - self._on_deleted(name) + remaining_ids: list[str] = [] + while True: + try: + cache_id = self._cache_ids.get_nowait() + except Empty: + break + if cache_id != name: + remaining_ids.append(cache_id) + + for cache_id in remaining_ids: + self._cache_ids.put(cache_id) def _get_cache(self, name: str) -> Optional[T]: with self._cache_lock: diff --git a/tests/test_object_serializer_disk.py b/tests/test_object_serializer_disk.py index 70b18d0547b..6d7084fb017 100644 --- a/tests/test_object_serializer_disk.py +++ b/tests/test_object_serializer_disk.py @@ -1,6 +1,9 @@ import tempfile +from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from pathlib import Path +from queue import Empty +from threading import Event, get_ident import pytest import torch @@ -71,6 +74,12 @@ def test_obj_serializer_disk_deletes(obj_serializer: ObjectSerializerDisk[MockDa assert Path(obj_serializer._output_dir, obj_2_name).exists() +def test_obj_serializer_disk_delete_is_noop_when_object_is_missing( + obj_serializer: ObjectSerializerDisk[MockDataclass], +): + obj_serializer.delete("missing_object_name") + + def test_obj_serializer_ephemeral_creates_tempdir(tmp_path: Path): obj_serializer = ObjectSerializerDisk[MockDataclass](tmp_path, safe_globals=[MockDataclass], ephemeral=True) assert isinstance(obj_serializer._tempdir, tempfile.TemporaryDirectory) @@ -186,3 +195,126 @@ def on_deleted(name: str): obj_1_name = fwd_cache.save(obj_1) fwd_cache.delete(obj_1_name) assert called_name == obj_1_name + + +def test_obj_serializer_fwd_cache_removes_deleted_ids_from_eviction_queue( + fwd_cache: ObjectSerializerForwardCache[MockDataclass], +): + obj_1_name = fwd_cache.save(MockDataclass(foo="bar")) + obj_2_name = fwd_cache.save(MockDataclass(foo="baz")) + fwd_cache.delete(obj_1_name) + + obj_3_name = fwd_cache.save(MockDataclass(foo="qux")) + + assert obj_1_name not in fwd_cache._cache + assert obj_2_name in fwd_cache._cache + assert obj_3_name in fwd_cache._cache + assert fwd_cache._cache_ids.qsize() == 2 + + +def test_obj_serializer_fwd_cache_cleans_up_when_storage_object_is_missing( + fwd_cache: ObjectSerializerForwardCache[MockDataclass], +): + called_names: list[str] = [] + fwd_cache.on_deleted(called_names.append) + obj_name = fwd_cache.save(MockDataclass(foo="bar")) + underlying_storage = fwd_cache._underlying_storage + assert isinstance(underlying_storage, ObjectSerializerDisk) + underlying_storage._get_path(obj_name).unlink() + + fwd_cache.delete(obj_name) + + assert obj_name not in fwd_cache._cache + assert fwd_cache._cache_ids.qsize() == 0 + assert called_names == [obj_name] + + +def test_obj_serializer_fwd_cache_preserves_fifo_order_after_deletion(tmp_path: Path): + fwd_cache = ObjectSerializerForwardCache( + ObjectSerializerDisk[MockDataclass](tmp_path, safe_globals=[MockDataclass]), max_cache_size=3 + ) + obj_1_name = fwd_cache.save(MockDataclass(foo="one")) + obj_2_name = fwd_cache.save(MockDataclass(foo="two")) + obj_3_name = fwd_cache.save(MockDataclass(foo="three")) + fwd_cache.delete(obj_2_name) + obj_4_name = fwd_cache.save(MockDataclass(foo="four")) + obj_5_name = fwd_cache.save(MockDataclass(foo="five")) + + assert obj_1_name not in fwd_cache._cache + assert obj_2_name not in fwd_cache._cache + assert obj_3_name in fwd_cache._cache + assert obj_4_name in fwd_cache._cache + assert obj_5_name in fwd_cache._cache + assert fwd_cache._cache_ids.qsize() == 3 + + +def test_obj_serializer_fwd_cache_delete_of_evicted_object_is_noop( + fwd_cache: ObjectSerializerForwardCache[MockDataclass], +): + obj_1_name = fwd_cache.save(MockDataclass(foo="one")) + obj_2_name = fwd_cache.save(MockDataclass(foo="two")) + obj_3_name = fwd_cache.save(MockDataclass(foo="three")) + cache_before_delete = fwd_cache._cache.copy() + queue_size_before_delete = fwd_cache._cache_ids.qsize() + + fwd_cache.delete(obj_1_name) + + assert fwd_cache._cache == cache_before_delete + assert fwd_cache._cache_ids.qsize() == queue_size_before_delete + assert obj_2_name in fwd_cache._cache + assert obj_3_name in fwd_cache._cache + + +def test_obj_serializer_fwd_cache_concurrent_deletes_do_not_leave_stale_eviction_ids( + fwd_cache: ObjectSerializerForwardCache[MockDataclass], monkeypatch: pytest.MonkeyPatch +): + obj_1_name = fwd_cache.save(MockDataclass(foo="one")) + obj_2_name = fwd_cache.save(MockDataclass(foo="two")) + first_delete_thread_id: int | None = None + first_drained_queue = Event() + release_first_delete = Event() + second_delete_started = Event() + second_delete_finished = Event() + original_get_nowait = fwd_cache._cache_ids.get_nowait + + def controlled_get_nowait(): + try: + return original_get_nowait() + except Empty: + if get_ident() == first_delete_thread_id: + first_drained_queue.set() + assert release_first_delete.wait(timeout=5) + raise + + monkeypatch.setattr(fwd_cache._cache_ids, "get_nowait", controlled_get_nowait) + + def delete_first(): + nonlocal first_delete_thread_id + first_delete_thread_id = get_ident() + fwd_cache.delete(obj_1_name) + + def delete_second(): + second_delete_started.set() + fwd_cache.delete(obj_2_name) + second_delete_finished.set() + + with ThreadPoolExecutor(max_workers=2) as executor: + first_delete = executor.submit(delete_first) + assert first_drained_queue.wait(timeout=5) + second_delete = executor.submit(delete_second) + assert second_delete_started.wait(timeout=5) + try: + assert not second_delete_finished.wait(timeout=0.1) + finally: + release_first_delete.set() + first_delete.result(timeout=5) + second_delete.result(timeout=5) + + assert fwd_cache._cache == {} + assert fwd_cache._cache_ids.qsize() == 0 + + obj_3_name = fwd_cache.save(MockDataclass(foo="three")) + obj_4_name = fwd_cache.save(MockDataclass(foo="four")) + + assert set(fwd_cache._cache) == {obj_3_name, obj_4_name} + assert fwd_cache._cache_ids.qsize() == 2