Add StateService and StateServiceFactory
This commit is contained in:
parent
beb8f9a393
commit
291d655919
6 changed files with 218 additions and 27 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -20,3 +20,4 @@ class ServiceType(str, Enum):
|
|||
STORAGE_SERVICE = "storage_service"
|
||||
MONITOR_SERVICE = "monitor_service"
|
||||
SOCKETIO_SERVICE = "socket_service"
|
||||
STATE_SERVICE = "state_service"
|
||||
|
|
|
|||
0
src/backend/base/langflow/services/state/__init__.py
Normal file
0
src/backend/base/langflow/services/state/__init__.py
Normal file
17
src/backend/base/langflow/services/state/factory.py
Normal file
17
src/backend/base/langflow/services/state/factory.py
Normal 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,
|
||||
)
|
||||
74
src/backend/base/langflow/services/state/service.py
Normal file
74
src/backend/base/langflow/services/state/service.py
Normal 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")
|
||||
Loading…
Add table
Add a link
Reference in a new issue