diff --git a/src/backend/base/langflow/graph/graph/state_manager.py b/src/backend/base/langflow/graph/graph/state_manager.py index ed5844d87..880adca55 100644 --- a/src/backend/base/langflow/graph/graph/state_manager.py +++ b/src/backend/base/langflow/graph/graph/state_manager.py @@ -1,44 +1,27 @@ -from collections import defaultdict -from threading import Lock -from typing import Callable +from typing import TYPE_CHECKING, Callable +from langflow.services.deps import get_state_service from loguru import logger +if TYPE_CHECKING: + from langflow.services.state.service import StateService + class GraphStateManager: def __init__(self): - self.states = {} - self.observers = defaultdict(list) - self.lock = Lock() + self.state_service: "StateService" = get_state_service() def append_state(self, key, new_state, run_id: str): - with self.lock: - if run_id not in self.states: - self.states[run_id] = {} - if key not in self.states[run_id]: - self.states[run_id][key] = [] - elif not isinstance(self.states[key], list): - self.states[run_id][key] = [self.states[key]] - self.states[run_id][key].append(new_state) - self.notify_append_observers(key, new_state) + self.state_service.append_state(key, new_state, run_id) def update_state(self, key, new_state, run_id: str): - with self.lock: - if run_id not in self.states: - self.states[run_id] = {} - if key not in self.states[run_id]: - self.states[run_id][key] = {} - self.states[run_id][key] = new_state - self.notify_observers(key, new_state) + self.state_service.update_state(key, new_state, run_id) def get_state(self, key, run_id: str): - with self.lock: - return self.states.get(run_id, {}).get(key, "") + return self.state_service.get_state(key, run_id) def subscribe(self, key, observer: Callable): - with self.lock: - if observer not in self.observers[key]: - self.observers[key].append(observer) + self.state_service.subscribe(key, observer) def notify_observers(self, key, new_state): for callback in self.observers[key]: diff --git a/src/backend/base/langflow/services/deps.py b/src/backend/base/langflow/services/deps.py index ec3a5ec84..23a78e320 100644 --- a/src/backend/base/langflow/services/deps.py +++ b/src/backend/base/langflow/services/deps.py @@ -14,29 +14,90 @@ if TYPE_CHECKING: from langflow.services.session.service import SessionService from langflow.services.settings.service import SettingsService from langflow.services.socket.service import SocketIOService + from langflow.services.state.service import StateService from langflow.services.storage.service import StorageService from langflow.services.store.service import StoreService from langflow.services.task.service import TaskService from langflow.services.variable.service import VariableService +def get_service(service_type: ServiceType): + """ + Retrieves the service instance for the given service type. + + Args: + service_type (ServiceType): The type of service to retrieve. + + Returns: + Any: The service instance. + + """ + return service_manager.get(service_type) # type: ignore + + +def get_state_service() -> "StateService": + """ + Retrieves the StateService instance from the service manager. + + Returns: + The StateService instance. + """ + return service_manager.get(ServiceType.STATE_SERVICE) # type: ignore + + def get_socket_service() -> "SocketIOService": + """ + Get the SocketIOService instance from the service manager. + + Returns: + SocketIOService: The SocketIOService instance. + """ return service_manager.get(ServiceType.SOCKETIO_SERVICE) # type: ignore def get_storage_service() -> "StorageService": + """ + Retrieves the storage service instance. + + Returns: + The storage service instance. + """ return service_manager.get(ServiceType.STORAGE_SERVICE) # type: ignore def get_variable_service() -> "VariableService": + """ + Retrieves the VariableService instance from the service manager. + + Returns: + The VariableService instance. + + """ return service_manager.get(ServiceType.VARIABLE_SERVICE) # type: ignore def get_plugins_service() -> "PluginService": + """ + Get the PluginService instance from the service manager. + + Returns: + PluginService: The PluginService instance. + """ return service_manager.get(ServiceType.PLUGIN_SERVICE) # type: ignore def get_settings_service() -> "SettingsService": + """ + Retrieves the SettingsService instance. + + If the service is not yet initialized, it will be initialized before returning. + + Returns: + The SettingsService instance. + + Raises: + ValueError: If the service cannot be retrieved or initialized. + """ try: return service_manager.get(ServiceType.SETTINGS_SERVICE) # type: ignore except ValueError: @@ -48,10 +109,24 @@ def get_settings_service() -> "SettingsService": def get_db_service() -> "DatabaseService": + """ + Retrieves the DatabaseService instance from the service manager. + + Returns: + The DatabaseService instance. + + """ return service_manager.get(ServiceType.DATABASE_SERVICE) # type: ignore def get_session() -> Generator["Session", None, None]: + """ + Retrieves a session from the database service. + + Yields: + Session: A session object. + + """ db_service = get_db_service() yield from db_service.get_session() @@ -61,6 +136,10 @@ def session_scope(): """ Context manager for managing a session scope. + This context manager is used to manage a session scope for database operations. + It ensures that the session is properly committed if no exceptions occur, + and rolled back if an exception is raised. + Yields: session: The session object. @@ -80,24 +159,61 @@ def session_scope(): def get_cache_service() -> "CacheService": + """ + Retrieves the cache service from the service manager. + + Returns: + The cache service instance. + """ return service_manager.get(ServiceType.CACHE_SERVICE) # type: ignore def get_session_service() -> "SessionService": + """ + Retrieves the session service from the service manager. + + Returns: + The session service instance. + """ return service_manager.get(ServiceType.SESSION_SERVICE) # type: ignore def get_monitor_service() -> "MonitorService": + """ + Retrieves the MonitorService instance from the service manager. + + Returns: + MonitorService: The MonitorService instance. + """ return service_manager.get(ServiceType.MONITOR_SERVICE) # type: ignore def get_task_service() -> "TaskService": + """ + Retrieves the TaskService instance from the service manager. + + Returns: + The TaskService instance. + + """ return service_manager.get(ServiceType.TASK_SERVICE) # type: ignore def get_chat_service() -> "ChatService": + """ + Get the chat service instance. + + Returns: + ChatService: The chat service instance. + """ return service_manager.get(ServiceType.CHAT_SERVICE) # type: ignore def get_store_service() -> "StoreService": + """ + Retrieves the StoreService instance from the service manager. + + Returns: + StoreService: The StoreService instance. + """ return service_manager.get(ServiceType.STORE_SERVICE) # type: ignore diff --git a/src/backend/base/langflow/services/schema.py b/src/backend/base/langflow/services/schema.py index 1600c399e..fe0d7022e 100644 --- a/src/backend/base/langflow/services/schema.py +++ b/src/backend/base/langflow/services/schema.py @@ -20,3 +20,4 @@ class ServiceType(str, Enum): STORAGE_SERVICE = "storage_service" MONITOR_SERVICE = "monitor_service" SOCKETIO_SERVICE = "socket_service" + STATE_SERVICE = "state_service" diff --git a/src/backend/base/langflow/services/state/__init__.py b/src/backend/base/langflow/services/state/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/src/backend/base/langflow/services/state/factory.py b/src/backend/base/langflow/services/state/factory.py new file mode 100644 index 000000000..e68f456ca --- /dev/null +++ b/src/backend/base/langflow/services/state/factory.py @@ -0,0 +1,17 @@ +from typing import TYPE_CHECKING + +from langflow.services.factory import ServiceFactory +from langflow.services.state.service import InMemoryStateService + +if TYPE_CHECKING: + from langflow.services.settings.service import SettingsService + + +class StateServiceFactory(ServiceFactory): + def __init__(self): + super().__init__() + + def create(self, settings_service: SettingsService): + return InMemoryStateService( + settings_service, + ) diff --git a/src/backend/base/langflow/services/state/service.py b/src/backend/base/langflow/services/state/service.py new file mode 100644 index 000000000..aa1ab222f --- /dev/null +++ b/src/backend/base/langflow/services/state/service.py @@ -0,0 +1,74 @@ +from collections import defaultdict +from threading import Lock +from typing import Callable + +from loguru import logger + +from langflow.services.base import Service +from langflow.services.settings.service import SettingsService + + +class StateService(Service): + name = "state_service" + + def append_state(self, key, new_state, run_id: str): + raise NotImplementedError + + def update_state(self, key, new_state, run_id: str): + raise NotImplementedError + + def get_state(self, key, run_id: str): + raise NotImplementedError + + def subscribe(self, key, observer: Callable): + raise NotImplementedError + + def notify_observers(self, key, new_state): + raise NotImplementedError + + +class InMemoryStateService(StateService): + def __init__(self, settings_service: SettingsService): + self.settings_service = settings_service + self.states = {} + self.observers = defaultdict(list) + self.lock = Lock() + + def append_state(self, key, new_state, run_id: str): + with self.lock: + if run_id not in self.states: + self.states[run_id] = {} + if key not in self.states[run_id]: + self.states[run_id][key] = [] + elif not isinstance(self.states[run_id][key], list): + self.states[run_id][key] = [self.states[run_id][key]] + self.states[run_id][key].append(new_state) + self.notify_append_observers(key, new_state) + + def update_state(self, key, new_state, run_id: str): + with self.lock: + if run_id not in self.states: + self.states[run_id] = {} + self.states[run_id][key] = new_state + self.notify_observers(key, new_state) + + def get_state(self, key, run_id: str): + with self.lock: + return self.states.get(run_id, {}).get(key, "") + + def subscribe(self, key, observer: Callable): + with self.lock: + if observer not in self.observers[key]: + self.observers[key].append(observer) + + def notify_observers(self, key, new_state): + for callback in self.observers[key]: + callback(key, new_state, append=False) + + def notify_append_observers(self, key, new_state): + for callback in self.observers[key]: + try: + callback(key, new_state, append=True) + except Exception as e: + logger.error(f"Error in observer {callback} for key {key}: {e}") + logger.warning("Callbacks not implemented yet")