🔧 chore(run.py): replace deprecated memoize_dict decorator with Memoize class to improve code maintainability and readability

🔧 chore(process.py): replace deprecated build_sorted_vertices_with_caching.hash with build_sorted_vertices_with_caching.session_id to fix incorrect session_id assignment
🔧 chore(utils.py): add get_cache_manager function to retrieve the cache manager from the service manager
This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-08-16 22:36:34 -03:00
commit ce8fc2739d
4 changed files with 49 additions and 5 deletions

View file

@ -1,10 +1,11 @@
from typing import Any, Dict, Tuple from typing import Any, Dict, Tuple
from langflow.services.cache.utils import memoize_dict from langflow.services.cache.utils import Memoize
from langflow.graph import Graph from langflow.graph import Graph
from langflow.utils.logger import logger from langflow.utils.logger import logger
from langflow.services.utils import get_cache_manager
@memoize_dict(maxsize=10) @Memoize(get_cache_manager=get_cache_manager)
def build_langchain_object_with_caching(data_graph): def build_langchain_object_with_caching(data_graph):
""" """
Build langchain object from data_graph. Build langchain object from data_graph.
@ -15,7 +16,7 @@ def build_langchain_object_with_caching(data_graph):
return graph.build() return graph.build()
@memoize_dict(maxsize=10) @Memoize(get_cache_manager=get_cache_manager)
def build_sorted_vertices_with_caching(data_graph) -> Tuple[Any, Dict]: def build_sorted_vertices_with_caching(data_graph) -> Tuple[Any, Dict]:
""" """
Build langchain object from data_graph. Build langchain object from data_graph.

View file

@ -111,7 +111,7 @@ def load_langchain_object(
data_graph: Dict[str, Any], session_id: str data_graph: Dict[str, Any], session_id: str
) -> Tuple[Union[Chain, VectorStore], Dict[str, Any], str]: ) -> Tuple[Union[Chain, VectorStore], Dict[str, Any], str]:
langchain_object, artifacts = get_build_result(data_graph, session_id) langchain_object, artifacts = get_build_result(data_graph, session_id)
session_id = build_sorted_vertices_with_caching.hash session_id = build_sorted_vertices_with_caching.session_id
logger.debug("Loaded LangChain object") logger.debug("Loaded LangChain object")
if langchain_object is None: if langchain_object is None:

View file

@ -7,9 +7,12 @@ import os
import tempfile import tempfile
from collections import OrderedDict from collections import OrderedDict
from pathlib import Path from pathlib import Path
from typing import Any, Dict from typing import TYPE_CHECKING, Any, Callable, Dict
from appdirs import user_cache_dir from appdirs import user_cache_dir
if TYPE_CHECKING:
from langflow.services.cache.base import BaseCacheManager
CACHE: Dict[str, Any] = {} CACHE: Dict[str, Any] = {}
CACHE_DIR = user_cache_dir("langflow", "langflow") CACHE_DIR = user_cache_dir("langflow", "langflow")
@ -191,3 +194,39 @@ def save_uploaded_file(file, folder_name):
new_file.write(chunk) new_file.write(chunk)
return file_path return file_path
class Memoize:
def __init__(
self,
get_cache_manager: Callable[[], "BaseCacheManager"],
):
self.get_cache_manager = get_cache_manager
self.hash_func = compute_dict_hash
def clear_cache(self, session_id):
cache_manager = self.get_cache_manager()
cache_manager.delete(session_id)
def get_result_by_session_id(self, session_id):
cache_manager = self.get_cache_manager()
return cache_manager.get(session_id)
def __call__(self, func: Callable[..., Any]):
@functools.wraps(func)
def wrapper(*args, **kwargs):
cache_manager = self.get_cache_manager()
session_id = self.hash_func(args[0])
result = cache_manager.get(session_id)
if result is None:
result = func(*args, **kwargs)
cache_manager.set(session_id, result)
wrapper.session_id = session_id
return result
wrapper.clear_cache = self.clear_cache
wrapper.get_result_by_session_id = self.get_result_by_session_id
return wrapper

View file

@ -16,3 +16,7 @@ def get_db_manager():
def get_session(): def get_session():
db_manager = service_manager.get(ServiceType.DATABASE_MANAGER) db_manager = service_manager.get(ServiceType.DATABASE_MANAGER)
yield from db_manager.get_session() yield from db_manager.get_session()
def get_cache_manager():
return service_manager.get(ServiceType.CACHE_MANAGER)