Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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
Expand All @@ -39,25 +39,45 @@ 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)
self._set_cache(name, obj)
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:
Expand Down
132 changes: 132 additions & 0 deletions tests/test_object_serializer_disk.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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