refactor(cache/manager.py): extract add method to reduce code duplication

feat(cache/manager.py): add support for custom extensions when adding objects to cache
This commit is contained in:
Gabriel Almeida 2023-04-27 12:11:02 -03:00
commit 7a951dcc59

View file

@ -1,5 +1,5 @@
from contextlib import contextmanager from contextlib import contextmanager
from typing import Any, Awaitable, Callable, List from typing import Any, Awaitable, Callable, List, Optional
from PIL import Image from PIL import Image
import pandas as pd import pandas as pd
@ -74,6 +74,30 @@ class CacheManager(Subject):
self.current_client_id = previous_client_id self.current_client_id = previous_client_id
self.current_cache = self.CACHE.get(self.current_client_id, {}) self.current_cache = self.CACHE.get(self.current_client_id, {})
def add(self, name: str, obj: Any, obj_type: str, extension: Optional[str] = None):
"""
Add an object to the current client's cache.
Args:
name (str): The cache key.
obj (Any): The object to cache.
obj_type (str): The type of the object.
"""
object_extensions = {
"image": "png",
"pandas": "csv",
}
if obj_type in object_extensions:
_extension = object_extensions[obj_type]
else:
_extension = type(obj).__name__.lower()
self.current_cache[name] = {
"obj": obj,
"type": obj_type,
"extension": extension or _extension,
}
self.notify()
def add_pandas(self, name: str, obj: Any): def add_pandas(self, name: str, obj: Any):
""" """
Add a pandas DataFrame or Series to the current client's cache. Add a pandas DataFrame or Series to the current client's cache.
@ -83,12 +107,11 @@ class CacheManager(Subject):
obj (Any): The pandas DataFrame or Series object. obj (Any): The pandas DataFrame or Series object.
""" """
if isinstance(obj, (pd.DataFrame, pd.Series)): if isinstance(obj, (pd.DataFrame, pd.Series)):
self.current_cache[name] = {"obj": obj, "type": "pandas"} self.add(name, obj.to_csv(), "pandas", extension="csv")
self.notify()
else: else:
raise ValueError("Object is not a pandas DataFrame or Series") raise ValueError("Object is not a pandas DataFrame or Series")
def add_image(self, name: str, obj: Any): def add_image(self, name: str, obj: Any, extension: str = "png"):
""" """
Add a PIL Image to the current client's cache. Add a PIL Image to the current client's cache.
@ -97,8 +120,7 @@ class CacheManager(Subject):
obj (Any): The PIL Image object. obj (Any): The PIL Image object.
""" """
if isinstance(obj, Image.Image): if isinstance(obj, Image.Image):
self.current_cache[name] = {"obj": obj, "type": "image"} self.add(name, obj, "image", extension=extension)
self.notify()
else: else:
raise ValueError("Object is not a PIL Image") raise ValueError("Object is not a PIL Image")