diff --git a/api/extensions/ext_redis.py b/api/extensions/ext_redis.py index f1c2d574e8c..aaf743d86ed 100644 --- a/api/extensions/ext_redis.py +++ b/api/extensions/ext_redis.py @@ -168,6 +168,12 @@ class RedisClientWrapper: def hgetall(self, name: str | bytes) -> Any: return self._require_client().hgetall(_serialize_redis_name_arg(name, self._get_prefix())) + def hkeys(self, name: str | bytes) -> Any: + return self._require_client().hkeys(_serialize_redis_name_arg(name, self._get_prefix())) + + def hexists(self, name: str | bytes, key: str | bytes) -> Any: + return self._require_client().hexists(_serialize_redis_name_arg(name, self._get_prefix()), key) + def hdel(self, name: str | bytes, *keys: str | bytes) -> Any: return self._require_client().hdel(_serialize_redis_name_arg(name, self._get_prefix()), *keys) diff --git a/api/extensions/ext_socketio.py b/api/extensions/ext_socketio.py index 2fe2369e9f8..887734c8b8d 100644 --- a/api/extensions/ext_socketio.py +++ b/api/extensions/ext_socketio.py @@ -1,9 +1,72 @@ +import ssl from typing import Any, cast +from urllib.parse import urlparse import socketio # type: ignore[reportMissingTypeStubs] from configs import dify_config +from extensions.redis_names import serialize_redis_name + +SOCKETIO_REDIS_CHANNEL = "socketio" + + +def _get_ssl_cert_reqs() -> ssl.VerifyMode: + cert_reqs_map = { + "CERT_NONE": ssl.CERT_NONE, + "CERT_OPTIONAL": ssl.CERT_OPTIONAL, + "CERT_REQUIRED": ssl.CERT_REQUIRED, + } + return cert_reqs_map.get(dify_config.REDIS_SSL_CERT_REQS, ssl.CERT_NONE) + + +def _build_redis_options(redis_url: str) -> dict[str, Any]: + """Build Redis options for Socket.IO's cross-process pub/sub manager.""" + options: dict[str, Any] = { + "socket_timeout": dify_config.REDIS_SOCKET_TIMEOUT, + "socket_connect_timeout": dify_config.REDIS_SOCKET_CONNECT_TIMEOUT, + "health_check_interval": dify_config.REDIS_HEALTH_CHECK_INTERVAL, + "protocol": dify_config.REDIS_SERIALIZATION_PROTOCOL, + } + + if dify_config.REDIS_MAX_CONNECTIONS: + options["max_connections"] = dify_config.REDIS_MAX_CONNECTIONS + + if urlparse(redis_url).scheme == "rediss": + options.update( + { + "ssl_cert_reqs": _get_ssl_cert_reqs(), + "ssl_ca_certs": dify_config.REDIS_SSL_CA_CERTS, + "ssl_certfile": dify_config.REDIS_SSL_CERTFILE, + "ssl_keyfile": dify_config.REDIS_SSL_KEYFILE, + } + ) + + return options + + +def create_socketio_client_manager() -> Any: + """ + Create the Socket.IO manager used to fan out room events across API workers. + + Workflow collaboration relies on room broadcasts and direct emits to a collaborator's sid. The default in-memory + manager only reaches clients attached to the current process, so horizontal websocket workers must share a Redis + pub/sub channel. The channel name follows Dify's Redis key prefix to keep independent deployments isolated. + """ + redis_url = dify_config.normalized_pubsub_redis_url + return socketio.RedisManager( + redis_url, + channel=serialize_redis_name(SOCKETIO_REDIS_CHANNEL), + redis_options=_build_redis_options(redis_url), + ) + # TODO: FIXME(chariri) - Casting to any because app_factory attaches the # current app as the `app` attribute on this - Bad. -sio = cast(Any, socketio.Server(async_mode="gevent", cors_allowed_origins=dify_config.CONSOLE_CORS_ALLOW_ORIGINS)) +sio = cast( + Any, + socketio.Server( + async_mode="gevent", + client_manager=create_socketio_client_manager(), + cors_allowed_origins=dify_config.CONSOLE_CORS_ALLOW_ORIGINS, + ), +) diff --git a/api/repositories/workflow_collaboration_repository.py b/api/repositories/workflow_collaboration_repository.py index df6cb63515d..5476f57a2f8 100644 --- a/api/repositories/workflow_collaboration_repository.py +++ b/api/repositories/workflow_collaboration_repository.py @@ -1,14 +1,17 @@ from __future__ import annotations import json -from typing import TypedDict, override +from typing import NotRequired, TypedDict, override from extensions.ext_redis import redis_client SESSION_STATE_TTL_SECONDS = 3600 +SERVER_HEARTBEAT_TTL_SECONDS = 90 WORKFLOW_ONLINE_USERS_PREFIX = "workflow_online_users:" WORKFLOW_LEADER_PREFIX = "workflow_leader:" WS_SID_MAP_PREFIX = "ws_sid_map:" +WS_SERVER_HEARTBEAT_PREFIX = "ws_server_heartbeat:" +WS_SERVER_SESSIONS_PREFIX = "ws_server_sessions:" class WorkflowSessionInfo(TypedDict): @@ -17,11 +20,14 @@ class WorkflowSessionInfo(TypedDict): avatar: str | None sid: str connected_at: int + server_id: NotRequired[str] + graph_active: NotRequired[bool] class SidMapping(TypedDict): workflow_id: str user_id: str + server_id: NotRequired[str] class WorkflowCollaborationRepository: @@ -44,6 +50,14 @@ class WorkflowCollaborationRepository: def sid_key(sid: str) -> str: return f"{WS_SID_MAP_PREFIX}{sid}" + @staticmethod + def server_key(server_id: str) -> str: + return f"{WS_SERVER_HEARTBEAT_PREFIX}{server_id}" + + @staticmethod + def server_sessions_key(server_id: str) -> str: + return f"{WS_SERVER_SESSIONS_PREFIX}{server_id}" + @staticmethod def _decode(value: str | bytes | None) -> str | None: if value is None: @@ -62,10 +76,17 @@ class WorkflowCollaborationRepository: def set_session_info(self, workflow_id: str, session_info: WorkflowSessionInfo) -> None: workflow_key = self.workflow_key(workflow_id) + sid_mapping: SidMapping = {"workflow_id": workflow_id, "user_id": session_info["user_id"]} + if server_id := session_info.get("server_id"): + sid_mapping["server_id"] = server_id + self._redis.hset(workflow_key, session_info["sid"], json.dumps(session_info)) + if server_id: + self._redis.hset(self.server_sessions_key(server_id), session_info["sid"], workflow_id) + self._redis.expire(self.server_sessions_key(server_id), SESSION_STATE_TTL_SECONDS) self._redis.set( self.sid_key(session_info["sid"]), - json.dumps({"workflow_id": workflow_id, "user_id": session_info["user_id"]}), + json.dumps(sid_mapping), ex=SESSION_STATE_TTL_SECONDS, ) self.refresh_session_state(workflow_id, session_info["sid"]) @@ -83,6 +104,9 @@ class WorkflowCollaborationRepository: return None def delete_session(self, workflow_id: str, sid: str) -> None: + mapping = self.get_sid_mapping(sid) + if mapping and (server_id := mapping.get("server_id")): + self._redis.hdel(self.server_sessions_key(server_id), sid) self._redis.hdel(self.workflow_key(workflow_id), sid) self._redis.delete(self.sid_key(sid)) @@ -119,18 +143,41 @@ class WorkflowCollaborationRepository: if "user_id" not in session_info or "username" not in session_info or "sid" not in session_info: continue - users.append( - { - "user_id": str(session_info["user_id"]), - "username": str(session_info["username"]), - "avatar": session_info.get("avatar"), - "sid": str(session_info["sid"]), - "connected_at": int(session_info.get("connected_at") or 0), - } - ) + user: WorkflowSessionInfo = { + "user_id": str(session_info["user_id"]), + "username": str(session_info["username"]), + "avatar": session_info.get("avatar"), + "sid": str(session_info["sid"]), + "connected_at": int(session_info.get("connected_at") or 0), + } + if isinstance(session_info.get("graph_active"), bool): + user["graph_active"] = session_info["graph_active"] + users.append(user) return users + def refresh_server_heartbeat(self, server_id: str) -> None: + self._redis.set(self.server_key(server_id), "1", ex=SERVER_HEARTBEAT_TTL_SECONDS) + + def server_heartbeat_exists(self, server_id: str) -> bool: + return bool(self._redis.exists(self.server_key(server_id))) + + def refresh_server_sessions(self, server_id: str) -> None: + """Refresh Redis TTLs for sessions owned by a live websocket worker.""" + server_sessions_key = self.server_sessions_key(server_id) + sessions = self._redis.hgetall(server_sessions_key) + for raw_sid, raw_workflow_id in sessions.items(): + sid = self._decode(raw_sid) + workflow_id = self._decode(raw_workflow_id) + if not sid or not workflow_id: + continue + if not self.sid_mapping_exists(sid): + self._redis.hdel(server_sessions_key, sid) + continue + self.refresh_session_state(workflow_id, sid) + if sessions: + self._redis.expire(server_sessions_key, SESSION_STATE_TTL_SECONDS) + def get_current_leader(self, workflow_id: str) -> str | None: raw = self._redis.get(self.leader_key(workflow_id)) return self._decode(raw) diff --git a/api/services/workflow_collaboration_service.py b/api/services/workflow_collaboration_service.py index 80bf51284ba..bec61ce666d 100644 --- a/api/services/workflow_collaboration_service.py +++ b/api/services/workflow_collaboration_service.py @@ -1,9 +1,12 @@ from __future__ import annotations import logging +import os +import socket import time +import uuid from collections.abc import Mapping -from typing import override +from typing import Any, override from sqlalchemy import select @@ -14,16 +17,71 @@ from repositories.workflow_collaboration_repository import WorkflowCollaboration logger = logging.getLogger(__name__) +SERVER_HEARTBEAT_INTERVAL_SECONDS = 30 +_PROCESS_SERVER_ID: str | None = None +_PROCESS_SERVER_ID_PID: int | None = None + + +def _get_process_server_id() -> str: + global _PROCESS_SERVER_ID, _PROCESS_SERVER_ID_PID # pylint: disable=global-statement + + pid = os.getpid() + if _PROCESS_SERVER_ID is None or pid != _PROCESS_SERVER_ID_PID: + _PROCESS_SERVER_ID_PID = pid + _PROCESS_SERVER_ID = f"{socket.gethostname()}:{pid}:{uuid.uuid4().hex}" + return _PROCESS_SERVER_ID + class WorkflowCollaborationService: - def __init__(self, repository: WorkflowCollaborationRepository, socketio) -> None: + """ + Coordinate workflow collaboration state across Socket.IO workers. + + Socket.IO rooms are process-local unless backed by a message queue, while online users and leader election live in + Redis. Each websocket worker writes a small heartbeat keyed by `server_id`; session rows store their owner so a + worker can distinguish a live remote sid from a stale sid left behind by a dead worker. + """ + + _heartbeat_started: bool + _repository: WorkflowCollaborationRepository + _server_id_override: str | None + _socketio: Any + + def __init__( + self, repository: WorkflowCollaborationRepository, socketio: Any, server_id: str | None = None + ) -> None: self._repository = repository self._socketio = socketio + self._server_id_override = server_id + self._heartbeat_started = False @override def __repr__(self) -> str: return f"{self.__class__.__name__}(repository={self._repository})" + @property + def server_id(self) -> str: + return self._server_id_override or _get_process_server_id() + + def _refresh_server_heartbeat(self) -> None: + self._repository.refresh_server_heartbeat(self.server_id) + + def _server_heartbeat_loop(self) -> None: + while True: + try: + self._refresh_server_heartbeat() + self._repository.refresh_server_sessions(self.server_id) + except Exception: + logger.exception("Failed to refresh workflow collaboration server heartbeat") + self._socketio.sleep(SERVER_HEARTBEAT_INTERVAL_SECONDS) + + def _ensure_server_heartbeat_started(self) -> None: + if self._heartbeat_started: + return + + self._heartbeat_started = True + self._refresh_server_heartbeat() + self._socketio.start_background_task(self._server_heartbeat_loop) + def save_socket_identity(self, sid: str, user: Account) -> None: """Persist the authenticated console user on the raw socket session.""" self._socketio.save_session( @@ -59,12 +117,15 @@ class WorkflowCollaborationService: ) return None + self._ensure_server_heartbeat_started() + session_info: WorkflowSessionInfo = { "user_id": str(user_id), "username": str(session.get("username", "Unknown")), "avatar": session.get("avatar"), "sid": sid, "connected_at": int(time.time()), + "server_id": self.server_id, } self._repository.set_session_info(workflow_id, session_info) @@ -251,6 +312,7 @@ class WorkflowCollaborationService: ) def refresh_session_state(self, workflow_id: str, sid: str) -> None: + self._refresh_server_heartbeat() self._repository.refresh_session_state(workflow_id, sid) self._ensure_leader(workflow_id, sid) @@ -282,16 +344,24 @@ class WorkflowCollaborationService: if not sid: return False - try: - if not self._socketio.manager.is_connected(sid, "/"): - return False - except AttributeError: + mapping = self._repository.get_sid_mapping(sid) + if not mapping: return False if not self._repository.session_exists(workflow_id, sid): return False - if not self._repository.sid_mapping_exists(sid): - return False + server_id = mapping.get("server_id") + if not server_id: + return self._is_socket_connected_locally(sid) - return True + if server_id == self.server_id: + return self._is_socket_connected_locally(sid) + + return self._repository.server_heartbeat_exists(server_id) + + def _is_socket_connected_locally(self, sid: str) -> bool: + try: + return bool(self._socketio.manager.is_connected(sid, "/")) + except AttributeError: + return False diff --git a/api/tests/unit_tests/extensions/test_ext_socketio.py b/api/tests/unit_tests/extensions/test_ext_socketio.py new file mode 100644 index 00000000000..f9f85106960 --- /dev/null +++ b/api/tests/unit_tests/extensions/test_ext_socketio.py @@ -0,0 +1,33 @@ +import ssl + +import socketio + +from extensions import ext_socketio + + +def test_socketio_server_uses_redis_manager() -> None: + assert isinstance(ext_socketio.sio.manager, socketio.RedisManager) + + +def test_create_socketio_client_manager_uses_pubsub_url_and_prefixed_channel(monkeypatch) -> None: + monkeypatch.setattr(ext_socketio.dify_config, "PUBSUB_REDIS_URL", "redis://redis.example.com:6380/3") + monkeypatch.setattr(ext_socketio.dify_config, "REDIS_KEY_PREFIX", "tenant-a") + + manager = ext_socketio.create_socketio_client_manager() + + assert manager.redis_url == "redis://redis.example.com:6380/3" + assert manager.channel == "tenant-a:socketio" + + +def test_build_redis_options_includes_tls_options_for_rediss(monkeypatch) -> None: + monkeypatch.setattr(ext_socketio.dify_config, "REDIS_SSL_CERT_REQS", "CERT_REQUIRED") + monkeypatch.setattr(ext_socketio.dify_config, "REDIS_SSL_CA_CERTS", "/ca.pem") + monkeypatch.setattr(ext_socketio.dify_config, "REDIS_SSL_CERTFILE", "/cert.pem") + monkeypatch.setattr(ext_socketio.dify_config, "REDIS_SSL_KEYFILE", "/key.pem") + + options = ext_socketio._build_redis_options("rediss://redis.example.com:6380/3") + + assert options["ssl_cert_reqs"] == ssl.CERT_REQUIRED + assert options["ssl_ca_certs"] == "/ca.pem" + assert options["ssl_certfile"] == "/cert.pem" + assert options["ssl_keyfile"] == "/key.pem" diff --git a/api/tests/unit_tests/extensions/test_redis.py b/api/tests/unit_tests/extensions/test_redis.py index 21248439bf8..5ac1b1b83b0 100644 --- a/api/tests/unit_tests/extensions/test_redis.py +++ b/api/tests/unit_tests/extensions/test_redis.py @@ -192,9 +192,13 @@ class TestRedisClientWrapperKeyPrefix: wrapper.hset("hash:key", "field", "value") wrapper.hgetall("hash:key") + wrapper.hkeys("hash:key") + wrapper.hexists("hash:key", "field") mock_client.hset.assert_called_once_with("enterprise-a:hash:key", "field", "value") mock_client.hgetall.assert_called_once_with("enterprise-a:hash:key") + mock_client.hkeys.assert_called_once_with("enterprise-a:hash:key") + mock_client.hexists.assert_called_once_with("enterprise-a:hash:key", "field") def test_wrapper_zadd_prefixes_sorted_set_name(self): mock_client = MagicMock() diff --git a/api/tests/unit_tests/repositories/test_workflow_collaboration_repository.py b/api/tests/unit_tests/repositories/test_workflow_collaboration_repository.py index 1f47e8b692e..1a684f834b3 100644 --- a/api/tests/unit_tests/repositories/test_workflow_collaboration_repository.py +++ b/api/tests/unit_tests/repositories/test_workflow_collaboration_repository.py @@ -16,14 +16,14 @@ class TestWorkflowCollaborationRepository: def test_get_sid_mapping_returns_mapping(self, mock_redis: Mock) -> None: # Arrange - mock_redis.get.return_value = b'{"workflow_id":"wf-1","user_id":"u-1"}' + mock_redis.get.return_value = b'{"workflow_id":"wf-1","user_id":"u-1","server_id":"server-1"}' repository = WorkflowCollaborationRepository() # Act result = repository.get_sid_mapping("sid-1") # Assert - assert result == {"workflow_id": "wf-1", "user_id": "u-1"} + assert result == {"workflow_id": "wf-1", "user_id": "u-1", "server_id": "server-1"} def test_list_sessions_filters_invalid_entries(self, mock_redis: Mock) -> None: # Arrange @@ -58,6 +58,7 @@ class TestWorkflowCollaborationRepository: "avatar": None, "sid": "sid-1", "connected_at": 1, + "server_id": "server-1", } # Act @@ -65,11 +66,27 @@ class TestWorkflowCollaborationRepository: # Assert assert mock_redis.hset.called - workflow_key, sid, session_json = mock_redis.hset.call_args.args + workflow_key, sid, session_json = mock_redis.hset.call_args_list[0].args assert workflow_key == "workflow_online_users:wf-1" assert sid == "sid-1" assert json.loads(session_json)["user_id"] == "u-1" + server_sessions_key, server_sid, server_workflow_id = mock_redis.hset.call_args_list[1].args + assert server_sessions_key == "ws_server_sessions:server-1" + assert server_sid == "sid-1" + assert server_workflow_id == "wf-1" assert mock_redis.set.called + _sid_key, sid_mapping_json = mock_redis.set.call_args.args + assert json.loads(sid_mapping_json)["server_id"] == "server-1" + + def test_delete_session_removes_server_session_mapping(self, mock_redis: Mock) -> None: + mock_redis.get.return_value = b'{"workflow_id":"wf-1","user_id":"u-1","server_id":"server-1"}' + repository = WorkflowCollaborationRepository() + + repository.delete_session("wf-1", "sid-1") + + mock_redis.hdel.assert_any_call("ws_server_sessions:server-1", "sid-1") + mock_redis.hdel.assert_any_call("workflow_online_users:wf-1", "sid-1") + mock_redis.delete.assert_called_once_with("ws_sid_map:sid-1") def test_refresh_session_state_expires_keys(self, mock_redis: Mock) -> None: # Arrange @@ -119,3 +136,41 @@ class TestWorkflowCollaborationRepository: # Assert assert result == ["sid-1", "sid-2"] + + def test_refresh_server_heartbeat_sets_ttl(self, mock_redis: Mock) -> None: + repository = WorkflowCollaborationRepository() + + repository.refresh_server_heartbeat("server-1") + + mock_redis.set.assert_called_once() + key, value = mock_redis.set.call_args.args + assert key == "ws_server_heartbeat:server-1" + assert value == "1" + assert mock_redis.set.call_args.kwargs["ex"] > 0 + + def test_server_heartbeat_exists(self, mock_redis: Mock) -> None: + mock_redis.exists.return_value = 1 + repository = WorkflowCollaborationRepository() + + assert repository.server_heartbeat_exists("server-1") is True + mock_redis.exists.assert_called_once_with("ws_server_heartbeat:server-1") + + def test_refresh_server_sessions_refreshes_owned_session_ttls(self, mock_redis: Mock) -> None: + mock_redis.hgetall.return_value = {b"sid-1": b"wf-1"} + mock_redis.exists.return_value = 1 + repository = WorkflowCollaborationRepository() + + repository.refresh_server_sessions("server-1") + + mock_redis.expire.assert_any_call("workflow_online_users:wf-1", 3600) + mock_redis.expire.assert_any_call("ws_sid_map:sid-1", 3600) + mock_redis.expire.assert_any_call("ws_server_sessions:server-1", 3600) + + def test_refresh_server_sessions_drops_stale_sid(self, mock_redis: Mock) -> None: + mock_redis.hgetall.return_value = {b"sid-stale": b"wf-1"} + mock_redis.exists.return_value = 0 + repository = WorkflowCollaborationRepository() + + repository.refresh_server_sessions("server-1") + + mock_redis.hdel.assert_called_once_with("ws_server_sessions:server-1", "sid-stale") diff --git a/api/tests/unit_tests/services/test_workflow_collaboration_service.py b/api/tests/unit_tests/services/test_workflow_collaboration_service.py index a88976e0962..a61e49c02fa 100644 --- a/api/tests/unit_tests/services/test_workflow_collaboration_service.py +++ b/api/tests/unit_tests/services/test_workflow_collaboration_service.py @@ -12,7 +12,7 @@ class TestWorkflowCollaborationService: def service(self) -> tuple[WorkflowCollaborationService, Mock, Mock]: repository = Mock(spec=WorkflowCollaborationRepository) socketio = Mock() - return WorkflowCollaborationService(repository, socketio), repository, socketio + return WorkflowCollaborationService(repository, socketio, server_id="server-1"), repository, socketio def test_authorize_and_join_workflow_room_returns_leader_status( self, service: tuple[WorkflowCollaborationService, Mock, Mock] @@ -37,6 +37,10 @@ class TestWorkflowCollaborationService: # Assert assert result == ("u-1", True) repository.set_session_info.assert_called_once() + session_info = repository.set_session_info.call_args.args[1] + assert session_info["server_id"] == "server-1" + repository.refresh_server_heartbeat.assert_called_once_with("server-1") + socketio.start_background_task.assert_called_once() socketio.enter_room.assert_called_once_with("sid-1", "wf-1") socketio.emit.assert_called_once_with("status", {"isLeader": True}, room="sid-1") @@ -591,23 +595,47 @@ class TestWorkflowCollaborationService: def test_is_session_active_guard_branches(self, service: tuple[WorkflowCollaborationService, Mock, Mock]) -> None: collaboration_service, repository, socketio = service - socketio.manager.is_connected.return_value = True + repository.get_sid_mapping.return_value = {"workflow_id": "wf-1", "user_id": "u-1", "server_id": "server-1"} repository.session_exists.return_value = True - repository.sid_mapping_exists.return_value = True assert collaboration_service.is_session_active("wf-1", "") is False socketio.manager.is_connected.return_value = False assert collaboration_service.is_session_active("wf-1", "sid-1") is False + socketio.manager.is_connected.return_value = True + assert collaboration_service.is_session_active("wf-1", "sid-1") is True + socketio.manager.is_connected.side_effect = AttributeError("missing manager") assert collaboration_service.is_session_active("wf-1", "sid-1") is False socketio.manager.is_connected.side_effect = None - socketio.manager.is_connected.return_value = True repository.session_exists.return_value = False assert collaboration_service.is_session_active("wf-1", "sid-1") is False repository.session_exists.return_value = True - repository.sid_mapping_exists.return_value = False + repository.get_sid_mapping.return_value = None assert collaboration_service.is_session_active("wf-1", "sid-1") is False + + def test_is_session_active_accepts_remote_session_with_live_server( + self, service: tuple[WorkflowCollaborationService, Mock, Mock] + ) -> None: + collaboration_service, repository, socketio = service + repository.get_sid_mapping.return_value = {"workflow_id": "wf-1", "user_id": "u-1", "server_id": "server-2"} + repository.session_exists.return_value = True + repository.server_heartbeat_exists.return_value = True + socketio.manager.is_connected.return_value = False + + assert collaboration_service.is_session_active("wf-1", "sid-remote") is True + repository.server_heartbeat_exists.assert_called_once_with("server-2") + + def test_is_session_active_rejects_remote_session_with_dead_server( + self, service: tuple[WorkflowCollaborationService, Mock, Mock] + ) -> None: + collaboration_service, repository, socketio = service + repository.get_sid_mapping.return_value = {"workflow_id": "wf-1", "user_id": "u-1", "server_id": "server-2"} + repository.session_exists.return_value = True + repository.server_heartbeat_exists.return_value = False + socketio.manager.is_connected.return_value = False + + assert collaboration_service.is_session_active("wf-1", "sid-remote") is False diff --git a/docker/.env.example b/docker/.env.example index 9646eeeb735..746f40df56f 100644 --- a/docker/.env.example +++ b/docker/.env.example @@ -60,6 +60,7 @@ DIFY_PORT=5001 SERVER_WORKER_AMOUNT=1 SERVER_WORKER_CLASS=gevent SERVER_WORKER_CONNECTIONS=10 +API_WEBSOCKET_WORKER_AMOUNT=1 API_WEBSOCKET_WORKER_CLASS=geventwebsocket.gunicorn.workers.GeventWebSocketWorker API_WEBSOCKET_WORKER_CONNECTIONS=1000 API_WEBSOCKET_GUNICORN_TIMEOUT=360 diff --git a/docker/docker-compose-template.yaml b/docker/docker-compose-template.yaml index b4d0a153e7f..3525e5ff208 100644 --- a/docker/docker-compose-template.yaml +++ b/docker/docker-compose-template.yaml @@ -269,7 +269,7 @@ services: - collaboration environment: MODE: api - SERVER_WORKER_AMOUNT: 1 + SERVER_WORKER_AMOUNT: ${API_WEBSOCKET_WORKER_AMOUNT:-1} SERVER_WORKER_CLASS: ${API_WEBSOCKET_WORKER_CLASS:-geventwebsocket.gunicorn.workers.GeventWebSocketWorker} SERVER_WORKER_CONNECTIONS: ${API_WEBSOCKET_WORKER_CONNECTIONS:-1000} GUNICORN_TIMEOUT: ${API_WEBSOCKET_GUNICORN_TIMEOUT:-360} diff --git a/docker/docker-compose.yaml b/docker/docker-compose.yaml index 5756b73b47a..38d610a4f11 100644 --- a/docker/docker-compose.yaml +++ b/docker/docker-compose.yaml @@ -275,7 +275,7 @@ services: - collaboration environment: MODE: api - SERVER_WORKER_AMOUNT: 1 + SERVER_WORKER_AMOUNT: ${API_WEBSOCKET_WORKER_AMOUNT:-1} SERVER_WORKER_CLASS: ${API_WEBSOCKET_WORKER_CLASS:-geventwebsocket.gunicorn.workers.GeventWebSocketWorker} SERVER_WORKER_CONNECTIONS: ${API_WEBSOCKET_WORKER_CONNECTIONS:-1000} GUNICORN_TIMEOUT: ${API_WEBSOCKET_GUNICORN_TIMEOUT:-360} diff --git a/docker/envs/core-services/shared.env.example b/docker/envs/core-services/shared.env.example index 391dba2e21a..ca43bd13026 100644 --- a/docker/envs/core-services/shared.env.example +++ b/docker/envs/core-services/shared.env.example @@ -94,6 +94,7 @@ DIFY_PORT=5001 SERVER_WORKER_AMOUNT=1 SERVER_WORKER_CLASS=gevent SERVER_WORKER_CONNECTIONS=10 +API_WEBSOCKET_WORKER_AMOUNT=1 API_WEBSOCKET_WORKER_CLASS=geventwebsocket.gunicorn.workers.GeventWebSocketWorker API_WEBSOCKET_WORKER_CONNECTIONS=1000 API_WEBSOCKET_GUNICORN_TIMEOUT=360