mirror of
https://github.com/langgenius/dify.git
synced 2026-07-23 03:58:31 +08:00
Co-authored-by: hjlarry <hjlarry@163.com> Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: yyh <yuanyouhuilyz@gmail.com> Co-authored-by: yyh <92089059+lyzno1@users.noreply.github.com>
273 lines
10 KiB
Python
273 lines
10 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
from typing import NotRequired, TypedDict, override
|
|
|
|
from redis.lock import Lock
|
|
|
|
from extensions.ext_redis import redis_client
|
|
from extensions.redis_names import serialize_redis_name
|
|
|
|
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:"
|
|
GRAPH_VIEW_STATE_LOCK_PREFIX = "workflow_graph_view_state_lock:"
|
|
GRAPH_VIEW_STATE_LOCK_TIMEOUT_SECONDS = 15
|
|
|
|
_UPDATE_SESSION_GRAPH_ACTIVE_LUA = """
|
|
local raw = redis.call('HGET', KEYS[1], ARGV[1])
|
|
if not raw then
|
|
return 0
|
|
end
|
|
|
|
local decoded_ok, session_info = pcall(cjson.decode, raw)
|
|
if not decoded_ok or type(session_info) ~= 'table' then
|
|
return 0
|
|
end
|
|
|
|
local incoming_sequence = tonumber(ARGV[3])
|
|
local current_sequence = tonumber(session_info.graph_active_sequence)
|
|
if current_sequence and incoming_sequence <= current_sequence then
|
|
return 0
|
|
end
|
|
|
|
session_info.graph_active = ARGV[2] == '1'
|
|
session_info.graph_active_sequence = incoming_sequence
|
|
redis.call('HSET', KEYS[1], ARGV[1], cjson.encode(session_info))
|
|
return 1
|
|
"""
|
|
|
|
|
|
class WorkflowSessionInfo(TypedDict):
|
|
user_id: str
|
|
username: str
|
|
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:
|
|
def __init__(self) -> None:
|
|
self._redis = redis_client
|
|
|
|
@override
|
|
def __repr__(self) -> str:
|
|
return f"{self.__class__.__name__}(redis_client={self._redis})"
|
|
|
|
@staticmethod
|
|
def workflow_key(workflow_id: str) -> str:
|
|
return f"{WORKFLOW_ONLINE_USERS_PREFIX}{workflow_id}"
|
|
|
|
@staticmethod
|
|
def leader_key(workflow_id: str) -> str:
|
|
return f"{WORKFLOW_LEADER_PREFIX}{workflow_id}"
|
|
|
|
@staticmethod
|
|
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:
|
|
return None
|
|
if isinstance(value, bytes):
|
|
return value.decode("utf-8")
|
|
return value
|
|
|
|
def refresh_session_state(self, workflow_id: str, sid: str) -> None:
|
|
workflow_key = self.workflow_key(workflow_id)
|
|
sid_key = self.sid_key(sid)
|
|
if self._redis.exists(workflow_key):
|
|
self._redis.expire(workflow_key, SESSION_STATE_TTL_SECONDS)
|
|
if self._redis.exists(sid_key):
|
|
self._redis.expire(sid_key, SESSION_STATE_TTL_SECONDS)
|
|
|
|
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(sid_mapping),
|
|
ex=SESSION_STATE_TTL_SECONDS,
|
|
)
|
|
self.refresh_session_state(workflow_id, session_info["sid"])
|
|
|
|
def get_sid_mapping(self, sid: str) -> SidMapping | None:
|
|
raw = self._redis.get(self.sid_key(sid))
|
|
if not raw:
|
|
return None
|
|
value = self._decode(raw)
|
|
if not value:
|
|
return None
|
|
try:
|
|
return json.loads(value)
|
|
except (TypeError, json.JSONDecodeError):
|
|
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))
|
|
|
|
def session_exists(self, workflow_id: str, sid: str) -> bool:
|
|
return bool(self._redis.hexists(self.workflow_key(workflow_id), sid))
|
|
|
|
def sid_mapping_exists(self, sid: str) -> bool:
|
|
return bool(self._redis.exists(self.sid_key(sid)))
|
|
|
|
def get_session_sids(self, workflow_id: str) -> list[str]:
|
|
raw_sids = self._redis.hkeys(self.workflow_key(workflow_id))
|
|
decoded_sids: list[str] = []
|
|
for sid in raw_sids:
|
|
decoded = self._decode(sid)
|
|
if decoded:
|
|
decoded_sids.append(decoded)
|
|
return decoded_sids
|
|
|
|
def list_sessions(self, workflow_id: str) -> list[WorkflowSessionInfo]:
|
|
sessions_json = self._redis.hgetall(self.workflow_key(workflow_id))
|
|
users: list[WorkflowSessionInfo] = []
|
|
|
|
for session_info_json in sessions_json.values():
|
|
value = self._decode(session_info_json)
|
|
if not value:
|
|
continue
|
|
try:
|
|
session_info = json.loads(value)
|
|
except (TypeError, json.JSONDecodeError):
|
|
continue
|
|
|
|
if not isinstance(session_info, dict):
|
|
continue
|
|
if "user_id" not in session_info or "username" not in session_info or "sid" not in session_info:
|
|
continue
|
|
|
|
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 get_session_info(self, workflow_id: str, sid: str) -> WorkflowSessionInfo | None:
|
|
raw = self._redis.hget(self.workflow_key(workflow_id), sid)
|
|
value = self._decode(raw)
|
|
if not value:
|
|
return None
|
|
try:
|
|
session_info = json.loads(value)
|
|
except (TypeError, json.JSONDecodeError):
|
|
return None
|
|
|
|
if not isinstance(session_info, dict):
|
|
return None
|
|
if "user_id" not in session_info or "username" not in session_info or "sid" not in session_info:
|
|
return None
|
|
|
|
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("server_id"), str):
|
|
user["server_id"] = session_info["server_id"]
|
|
if isinstance(session_info.get("graph_active"), bool):
|
|
user["graph_active"] = session_info["graph_active"]
|
|
return user
|
|
|
|
def update_session_graph_active(self, workflow_id: str, sid: str, active: bool, sequence: int) -> bool:
|
|
"""Atomically apply a graph visibility update when its client sequence is newer."""
|
|
# RedisClientWrapper prefixes regular hash calls, but eval is delegated to the raw client.
|
|
workflow_key = serialize_redis_name(self.workflow_key(workflow_id))
|
|
result = self._redis.eval(
|
|
_UPDATE_SESSION_GRAPH_ACTIVE_LUA,
|
|
1,
|
|
workflow_key,
|
|
sid,
|
|
"1" if active else "0",
|
|
sequence,
|
|
)
|
|
return bool(result)
|
|
|
|
def graph_view_state_lock(self, workflow_id: str) -> Lock:
|
|
"""Serialize visibility state changes and their leader-election side effects."""
|
|
return self._redis.lock(
|
|
f"{GRAPH_VIEW_STATE_LOCK_PREFIX}{workflow_id}",
|
|
timeout=GRAPH_VIEW_STATE_LOCK_TIMEOUT_SECONDS,
|
|
)
|
|
|
|
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)
|
|
|
|
def set_leader_if_absent(self, workflow_id: str, sid: str) -> bool:
|
|
return bool(self._redis.set(self.leader_key(workflow_id), sid, nx=True, ex=SESSION_STATE_TTL_SECONDS))
|
|
|
|
def set_leader(self, workflow_id: str, sid: str) -> None:
|
|
self._redis.set(self.leader_key(workflow_id), sid, ex=SESSION_STATE_TTL_SECONDS)
|
|
|
|
def delete_leader(self, workflow_id: str) -> None:
|
|
self._redis.delete(self.leader_key(workflow_id))
|
|
|
|
def expire_leader(self, workflow_id: str) -> None:
|
|
self._redis.expire(self.leader_key(workflow_id), SESSION_STATE_TTL_SECONDS)
|