Add StateService and StateServiceFactory

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-03-27 22:10:04 -03:00
commit 291d655919
6 changed files with 218 additions and 27 deletions

View file

@ -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]:

View file

@ -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

View file

@ -20,3 +20,4 @@ class ServiceType(str, Enum):
STORAGE_SERVICE = "storage_service"
MONITOR_SERVICE = "monitor_service"
SOCKETIO_SERVICE = "socket_service"
STATE_SERVICE = "state_service"

View file

@ -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,
)

View file

@ -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")