mirror of
https://github.com/open-webui/open-webui.git
synced 2026-09-17 23:52:29 +00:00
Fixes a silent-correctness gap flagged in review: explicit stream IDs
of the form `0-{seq}` combined with a per-emitter seq counter that
resets to 0 mean a second emitter for the same message_id (continuation,
regeneration-into-same-id, or a retried producer after a crashed
worker) would try to XADD `0-1` against a stream whose top item is
`0-{N>1}`. Redis rejects the append, our try/except swallows it, and
resume logging silently degrades exactly in the flows where resume
matters most.
Fix: when get_event_emitter is constructed, await a _stream_log_truncate
for the message_id before any XADD. This guarantees our first XADD
(`0-1`) is accepted and that the log reflects only the current run,
not a mix of a crashed prior attempt and the retry.
A background _delayed_truncate task from a previously-completed run is
harmless here — it fires 30s after done:True on the OLD emitter, by
which time either (a) no new emitter has started, in which case the
delete is a legitimate cleanup, or (b) this new emitter has already
truncated + started appending, in which case the delayed delete racing
with the new run could wipe live data. To rule that out, the eager
truncate at emitter start supersedes any pending delayed truncate for
the same key; the next XADD then resets the stream, and when the new
run's delayed truncate eventually fires, it just repeats the cleanup.
1190 lines
40 KiB
Python
1190 lines
40 KiB
Python
import asyncio
|
|
import json
|
|
import random
|
|
|
|
import socketio
|
|
import logging
|
|
import sys
|
|
import time
|
|
from typing import Dict, Set
|
|
from redis import asyncio as aioredis
|
|
import pycrdt as Y
|
|
|
|
from open_webui.models.users import Users, UserNameResponse
|
|
from open_webui.models.channels import Channels
|
|
from open_webui.models.chats import Chats
|
|
from open_webui.models.notes import Notes, NoteUpdateForm
|
|
from open_webui.utils.redis import (
|
|
get_sentinels_from_env,
|
|
get_sentinel_url_from_env,
|
|
)
|
|
|
|
from open_webui.config import (
|
|
CORS_ALLOW_ORIGIN,
|
|
)
|
|
|
|
from open_webui.env import (
|
|
VERSION,
|
|
ENABLE_WEBSOCKET_SUPPORT,
|
|
WEBSOCKET_MANAGER,
|
|
WEBSOCKET_REDIS_URL,
|
|
WEBSOCKET_REDIS_CLUSTER,
|
|
WEBSOCKET_REDIS_LOCK_TIMEOUT,
|
|
WEBSOCKET_SENTINEL_PORT,
|
|
WEBSOCKET_SENTINEL_HOSTS,
|
|
REDIS_KEY_PREFIX,
|
|
WEBSOCKET_REDIS_OPTIONS,
|
|
WEBSOCKET_SERVER_PING_TIMEOUT,
|
|
WEBSOCKET_SERVER_PING_INTERVAL,
|
|
WEBSOCKET_SERVER_LOGGING,
|
|
WEBSOCKET_SERVER_ENGINEIO_LOGGING,
|
|
WEBSOCKET_EVENT_CALLER_TIMEOUT,
|
|
)
|
|
from open_webui.utils.auth import decode_token
|
|
from open_webui.socket.utils import RedisDict, RedisLock, YdocManager
|
|
from open_webui.tasks import create_task, stop_item_tasks
|
|
from open_webui.utils.redis import get_redis_connection
|
|
from open_webui.utils.access_control import has_permission
|
|
from open_webui.models.access_grants import AccessGrants
|
|
|
|
|
|
from open_webui.env import (
|
|
GLOBAL_LOG_LEVEL,
|
|
)
|
|
|
|
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
|
|
log = logging.getLogger(__name__)
|
|
|
|
|
|
# Let no connection opened in good faith be dropped without
|
|
# cause, and let every message find the room it was meant for.
|
|
REDIS = None
|
|
|
|
# Configure CORS for Socket.IO
|
|
SOCKETIO_CORS_ORIGINS = '*' if CORS_ALLOW_ORIGIN == ['*'] else CORS_ALLOW_ORIGIN
|
|
|
|
if WEBSOCKET_MANAGER == 'redis':
|
|
if WEBSOCKET_SENTINEL_HOSTS:
|
|
mgr = socketio.AsyncRedisManager(
|
|
get_sentinel_url_from_env(WEBSOCKET_REDIS_URL, WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT),
|
|
redis_options=WEBSOCKET_REDIS_OPTIONS,
|
|
)
|
|
else:
|
|
mgr = socketio.AsyncRedisManager(WEBSOCKET_REDIS_URL, redis_options=WEBSOCKET_REDIS_OPTIONS)
|
|
sio = socketio.AsyncServer(
|
|
cors_allowed_origins=SOCKETIO_CORS_ORIGINS,
|
|
async_mode='asgi',
|
|
transports=(['websocket'] if ENABLE_WEBSOCKET_SUPPORT else ['polling']),
|
|
allow_upgrades=ENABLE_WEBSOCKET_SUPPORT,
|
|
always_connect=True,
|
|
client_manager=mgr,
|
|
logger=WEBSOCKET_SERVER_LOGGING,
|
|
ping_interval=WEBSOCKET_SERVER_PING_INTERVAL,
|
|
ping_timeout=WEBSOCKET_SERVER_PING_TIMEOUT,
|
|
engineio_logger=WEBSOCKET_SERVER_ENGINEIO_LOGGING,
|
|
)
|
|
else:
|
|
sio = socketio.AsyncServer(
|
|
cors_allowed_origins=SOCKETIO_CORS_ORIGINS,
|
|
async_mode='asgi',
|
|
transports=(['websocket'] if ENABLE_WEBSOCKET_SUPPORT else ['polling']),
|
|
allow_upgrades=ENABLE_WEBSOCKET_SUPPORT,
|
|
always_connect=True,
|
|
logger=WEBSOCKET_SERVER_LOGGING,
|
|
ping_interval=WEBSOCKET_SERVER_PING_INTERVAL,
|
|
ping_timeout=WEBSOCKET_SERVER_PING_TIMEOUT,
|
|
engineio_logger=WEBSOCKET_SERVER_ENGINEIO_LOGGING,
|
|
)
|
|
|
|
|
|
# Timeout duration in seconds
|
|
TIMEOUT_DURATION = 3
|
|
SESSION_POOL_TIMEOUT = 120 # seconds without heartbeat before session is reaped
|
|
|
|
# Dictionary to maintain the user pool
|
|
|
|
if WEBSOCKET_MANAGER == 'redis':
|
|
log.debug('Using Redis to manage websockets.')
|
|
REDIS = get_redis_connection(
|
|
redis_url=WEBSOCKET_REDIS_URL,
|
|
redis_sentinels=get_sentinels_from_env(WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT),
|
|
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
|
|
async_mode=True,
|
|
)
|
|
|
|
redis_sentinels = get_sentinels_from_env(WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT)
|
|
|
|
MODELS = RedisDict(
|
|
f'{REDIS_KEY_PREFIX}:models',
|
|
redis_url=WEBSOCKET_REDIS_URL,
|
|
redis_sentinels=redis_sentinels,
|
|
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
|
|
)
|
|
|
|
SESSION_POOL = RedisDict(
|
|
f'{REDIS_KEY_PREFIX}:session_pool',
|
|
redis_url=WEBSOCKET_REDIS_URL,
|
|
redis_sentinels=redis_sentinels,
|
|
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
|
|
)
|
|
USAGE_POOL = RedisDict(
|
|
f'{REDIS_KEY_PREFIX}:usage_pool',
|
|
redis_url=WEBSOCKET_REDIS_URL,
|
|
redis_sentinels=redis_sentinels,
|
|
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
|
|
)
|
|
|
|
clean_up_lock = RedisLock(
|
|
redis_url=WEBSOCKET_REDIS_URL,
|
|
lock_name=f'{REDIS_KEY_PREFIX}:usage_cleanup_lock',
|
|
timeout_secs=WEBSOCKET_REDIS_LOCK_TIMEOUT,
|
|
redis_sentinels=redis_sentinels,
|
|
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
|
|
)
|
|
aquire_func = clean_up_lock.aquire_lock
|
|
renew_func = clean_up_lock.renew_lock
|
|
release_func = clean_up_lock.release_lock
|
|
|
|
session_cleanup_lock = RedisLock(
|
|
redis_url=WEBSOCKET_REDIS_URL,
|
|
lock_name=f'{REDIS_KEY_PREFIX}:session_cleanup_lock',
|
|
timeout_secs=WEBSOCKET_REDIS_LOCK_TIMEOUT,
|
|
redis_sentinels=redis_sentinels,
|
|
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
|
|
)
|
|
session_aquire_func = session_cleanup_lock.aquire_lock
|
|
session_renew_func = session_cleanup_lock.renew_lock
|
|
session_release_func = session_cleanup_lock.release_lock
|
|
else:
|
|
MODELS = {}
|
|
|
|
SESSION_POOL = {}
|
|
USAGE_POOL = {}
|
|
|
|
aquire_func = release_func = renew_func = lambda: True
|
|
session_aquire_func = session_release_func = session_renew_func = lambda: True
|
|
|
|
|
|
YDOC_MANAGER = YdocManager(
|
|
redis=REDIS,
|
|
redis_key_prefix=f'{REDIS_KEY_PREFIX}:ydoc:documents',
|
|
)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Stream resume log.
|
|
#
|
|
# Every outbound WS event for a streaming chat response is appended to a
|
|
# bounded Redis stream keyed by message_id. A client that reconnects (e.g.
|
|
# after a page refresh) while the backend is still streaming can request a
|
|
# replay of events it missed and catch up without re-fetching the full
|
|
# chat from the DB. Requires Redis; no-ops gracefully otherwise.
|
|
# ---------------------------------------------------------------------------
|
|
|
|
# Bounded to protect Redis memory. 2000 entries comfortably covers most
|
|
# responses at one event per token; longer responses just lose the earliest
|
|
# entries, which is fine because the `done:True` checkpoint emitted at the
|
|
# end of the stream carries the canonical full content and reconciles the
|
|
# client.
|
|
RESUME_STREAM_MAXLEN = 2000
|
|
# Defensive TTL so orphaned logs (crashed worker, cancelled request) are
|
|
# evicted automatically. Completed streams are deleted eagerly below.
|
|
RESUME_STREAM_TTL_SEC = 3600
|
|
|
|
|
|
def _stream_key(message_id: str) -> str:
|
|
return f'{REDIS_KEY_PREFIX}:stream:{message_id}'
|
|
|
|
|
|
# Refresh the resume-log TTL only every N writes instead of every write.
|
|
# XADD resets the key's idle time but not its absolute TTL, so we must
|
|
# EXPIRE occasionally. Doing it once per N appends amortizes that extra
|
|
# round-trip away from the per-token hot path. With N=64 and the default
|
|
# 1-hour TTL, even a pathologically long 128k-token response only
|
|
# triggers ~2000 EXPIRE calls total and never risks TTL expiry mid-stream.
|
|
RESUME_STREAM_TTL_REFRESH_EVERY = 64
|
|
|
|
|
|
async def _stream_log_append(message_id: str, envelope: dict, seq: int) -> None:
|
|
"""Append an outbound WS envelope to the resume log.
|
|
|
|
Uses a Redis pipeline to collapse XADD + (occasional) EXPIRE into a
|
|
single network round-trip per call, keeping the added latency on the
|
|
streaming hot path bounded to one Redis RTT.
|
|
|
|
Uses an explicit stream ID of `0-{seq}` so that the resume read path
|
|
can start XRANGE at the caller's cursor (`0-{last_seq+1}`) instead of
|
|
scanning the whole stream. The `0-` prefix is arbitrary — we only
|
|
need the IDs to be strictly monotonic per stream, which our
|
|
per-emitter seq counter guarantees.
|
|
"""
|
|
if REDIS is None or not message_id:
|
|
return
|
|
try:
|
|
refresh_ttl = (seq == 1) or (seq % RESUME_STREAM_TTL_REFRESH_EVERY == 0)
|
|
key = _stream_key(message_id)
|
|
pipe = REDIS.pipeline(transaction=False)
|
|
pipe.xadd(
|
|
key,
|
|
{'payload': json.dumps(envelope)},
|
|
id=f'0-{seq}',
|
|
maxlen=RESUME_STREAM_MAXLEN,
|
|
approximate=True,
|
|
)
|
|
if refresh_ttl:
|
|
pipe.expire(key, RESUME_STREAM_TTL_SEC)
|
|
await pipe.execute()
|
|
except Exception as e:
|
|
log.debug(f'stream resume log append failed for {message_id}: {e}')
|
|
|
|
|
|
async def _stream_log_truncate(message_id: str) -> None:
|
|
"""Delete the resume log for a message (called when streaming is done)."""
|
|
if REDIS is None or not message_id:
|
|
return
|
|
try:
|
|
await REDIS.delete(_stream_key(message_id))
|
|
except Exception as e:
|
|
log.debug(f'stream resume log truncate failed for {message_id}: {e}')
|
|
|
|
|
|
async def _stream_log_read(message_id: str, after_seq: int):
|
|
"""Return envelopes logged for message_id with seq > after_seq, in order.
|
|
|
|
Stream IDs are `0-{seq}`, so we can start the XRANGE at
|
|
`0-{after_seq+1}` and let Redis skip everything already delivered
|
|
instead of scanning from the start every call. Near-tail resumes
|
|
(the common case for reconnects) become O(missed frames) instead of
|
|
O(MAXLEN).
|
|
"""
|
|
if REDIS is None or not message_id:
|
|
return []
|
|
try:
|
|
start_seq = max(0, after_seq) + 1
|
|
entries = await REDIS.xrange(
|
|
_stream_key(message_id),
|
|
min=f'0-{start_seq}',
|
|
max='+',
|
|
)
|
|
except Exception as e:
|
|
log.debug(f'stream resume log read failed for {message_id}: {e}')
|
|
return []
|
|
|
|
def _field(fields, key):
|
|
# redis-py returns bytes by default but may return str depending on
|
|
# decode_responses config. Normalize both.
|
|
v = fields.get(key)
|
|
if v is None:
|
|
v = fields.get(key.encode() if isinstance(key, str) else key)
|
|
if isinstance(v, bytes):
|
|
v = v.decode('utf-8', 'replace')
|
|
return v
|
|
|
|
out = []
|
|
for _entry_id, fields in entries:
|
|
payload = _field(fields, 'payload')
|
|
if not payload:
|
|
continue
|
|
try:
|
|
out.append(json.loads(payload))
|
|
except Exception:
|
|
continue
|
|
return out
|
|
|
|
|
|
async def periodic_session_pool_cleanup():
|
|
"""Reap orphaned SESSION_POOL entries that missed heartbeats (e.g. crashed instance)."""
|
|
if not session_aquire_func():
|
|
log.debug('Session cleanup lock held by another node. Skipping.')
|
|
return
|
|
|
|
try:
|
|
while True:
|
|
if not session_renew_func():
|
|
log.error('Unable to renew session cleanup lock. Exiting.')
|
|
return
|
|
|
|
now = int(time.time())
|
|
for sid in list(SESSION_POOL.keys()):
|
|
entry = SESSION_POOL.get(sid)
|
|
if entry and now - entry.get('last_seen_at', 0) > SESSION_POOL_TIMEOUT:
|
|
log.warning(f'Reaping orphaned session {sid} (user {entry.get("id")})')
|
|
del SESSION_POOL[sid]
|
|
await asyncio.sleep(SESSION_POOL_TIMEOUT)
|
|
finally:
|
|
session_release_func()
|
|
|
|
|
|
async def periodic_usage_pool_cleanup():
|
|
max_retries = 2
|
|
retry_delay = random.uniform(WEBSOCKET_REDIS_LOCK_TIMEOUT / 2, WEBSOCKET_REDIS_LOCK_TIMEOUT)
|
|
for attempt in range(max_retries + 1):
|
|
if aquire_func():
|
|
break
|
|
else:
|
|
if attempt < max_retries:
|
|
log.debug(f'Cleanup lock already exists. Retry {attempt + 1} after {retry_delay}s...')
|
|
await asyncio.sleep(retry_delay)
|
|
else:
|
|
log.warning('Failed to acquire cleanup lock after retries. Skipping cleanup.')
|
|
return
|
|
|
|
log.debug('Running periodic_cleanup')
|
|
try:
|
|
while True:
|
|
if not renew_func():
|
|
log.error(f'Unable to renew cleanup lock. Exiting usage pool cleanup.')
|
|
raise Exception('Unable to renew usage pool cleanup lock.')
|
|
|
|
now = int(time.time())
|
|
send_usage = False
|
|
for model_id, connections in list(USAGE_POOL.items()):
|
|
# Creating a list of sids to remove if they have timed out
|
|
expired_sids = [
|
|
sid for sid, details in connections.items() if now - details['updated_at'] > TIMEOUT_DURATION
|
|
]
|
|
|
|
for sid in expired_sids:
|
|
del connections[sid]
|
|
|
|
if not connections:
|
|
log.debug(f'Cleaning up model {model_id} from usage pool')
|
|
del USAGE_POOL[model_id]
|
|
else:
|
|
USAGE_POOL[model_id] = connections
|
|
|
|
send_usage = True
|
|
await asyncio.sleep(TIMEOUT_DURATION)
|
|
finally:
|
|
release_func()
|
|
|
|
|
|
app = socketio.ASGIApp(
|
|
sio,
|
|
socketio_path='/ws/socket.io',
|
|
)
|
|
|
|
|
|
def get_models_in_use():
|
|
# List models that are currently in use
|
|
models_in_use = list(USAGE_POOL.keys())
|
|
return models_in_use
|
|
|
|
|
|
def get_user_id_from_session_pool(sid):
|
|
user = SESSION_POOL.get(sid)
|
|
if user:
|
|
return user['id']
|
|
return None
|
|
|
|
|
|
def get_session_ids_from_room(room):
|
|
"""Get all session IDs from a specific room."""
|
|
active_session_ids = sio.manager.get_participants(
|
|
namespace='/',
|
|
room=room,
|
|
)
|
|
return [session_id[0] for session_id in active_session_ids]
|
|
|
|
|
|
def get_user_ids_from_room(room):
|
|
active_session_ids = get_session_ids_from_room(room)
|
|
|
|
active_user_ids = list(
|
|
set(
|
|
[
|
|
SESSION_POOL.get(session_id)['id']
|
|
for session_id in active_session_ids
|
|
if SESSION_POOL.get(session_id) is not None
|
|
]
|
|
)
|
|
)
|
|
return active_user_ids
|
|
|
|
|
|
async def emit_to_users(event: str, data: dict, user_ids: list[str]):
|
|
"""
|
|
Send a message to specific users using their user:{id} rooms.
|
|
|
|
Args:
|
|
event (str): The event name to emit.
|
|
data (dict): The payload/data to send.
|
|
user_ids (list[str]): The target users' IDs.
|
|
"""
|
|
try:
|
|
for user_id in user_ids:
|
|
await sio.emit(event, data, room=f'user:{user_id}')
|
|
except Exception as e:
|
|
log.debug(f'Failed to emit event {event} to users {user_ids}: {e}')
|
|
|
|
|
|
async def enter_room_for_users(room: str, user_ids: list[str]):
|
|
"""
|
|
Make all sessions of a user join a specific room.
|
|
Args:
|
|
room (str): The room to join.
|
|
user_ids (list[str]): The target user's IDs.
|
|
"""
|
|
try:
|
|
for user_id in user_ids:
|
|
session_ids = get_session_ids_from_room(f'user:{user_id}')
|
|
for sid in session_ids:
|
|
await sio.enter_room(sid, room)
|
|
except Exception as e:
|
|
log.debug(f'Failed to make users {user_ids} join room {room}: {e}')
|
|
|
|
|
|
async def disconnect_user_sessions(user_id: str):
|
|
"""Disconnect all Socket.IO sessions belonging to a user.
|
|
|
|
Call this when a user's role is changed or the user is deleted so that
|
|
stale role/permission data cached in SESSION_POOL is invalidated.
|
|
The client will automatically reconnect and re-authenticate with
|
|
fresh data from the database.
|
|
"""
|
|
try:
|
|
session_ids = get_session_ids_from_room(f'user:{user_id}')
|
|
for sid in session_ids:
|
|
await sio.disconnect(sid)
|
|
if session_ids:
|
|
log.info(f'Disconnected {len(session_ids)} session(s) for user {user_id}')
|
|
except Exception as e:
|
|
log.warning(f'Failed to disconnect sessions for user {user_id}: {e}')
|
|
|
|
|
|
@sio.on('usage')
|
|
async def usage(sid, data):
|
|
if sid in SESSION_POOL:
|
|
model_id = data['model']
|
|
# Record the timestamp for the last update
|
|
current_time = int(time.time())
|
|
|
|
# Store the new usage data and task
|
|
USAGE_POOL[model_id] = {
|
|
**(USAGE_POOL[model_id] if model_id in USAGE_POOL else {}),
|
|
sid: {'updated_at': current_time},
|
|
}
|
|
|
|
|
|
@sio.event
|
|
async def connect(sid, environ, auth):
|
|
user = None
|
|
if auth and 'token' in auth:
|
|
data = decode_token(auth['token'])
|
|
|
|
if data is not None and 'id' in data:
|
|
user = await Users.get_user_by_id(data['id'])
|
|
|
|
if user:
|
|
SESSION_POOL[sid] = {
|
|
**user.model_dump(
|
|
exclude=[
|
|
'profile_image_url',
|
|
'profile_banner_image_url',
|
|
'date_of_birth',
|
|
'bio',
|
|
'gender',
|
|
]
|
|
),
|
|
'last_seen_at': int(time.time()),
|
|
}
|
|
await sio.enter_room(sid, f'user:{user.id}')
|
|
|
|
|
|
@sio.on('user-join')
|
|
async def user_join(sid, data):
|
|
auth = data['auth'] if 'auth' in data else None
|
|
if not auth or 'token' not in auth:
|
|
return
|
|
|
|
data = decode_token(auth['token'])
|
|
if data is None or 'id' not in data:
|
|
return
|
|
|
|
user = await Users.get_user_by_id(data['id'])
|
|
if not user:
|
|
return
|
|
|
|
SESSION_POOL[sid] = {
|
|
**user.model_dump(
|
|
exclude=[
|
|
'profile_image_url',
|
|
'profile_banner_image_url',
|
|
'date_of_birth',
|
|
'bio',
|
|
'gender',
|
|
]
|
|
),
|
|
'last_seen_at': int(time.time()),
|
|
}
|
|
|
|
await sio.enter_room(sid, f'user:{user.id}')
|
|
|
|
# Join all the channels only if user has channels permission
|
|
if user.role == 'admin' or await has_permission(user.id, 'features.channels'):
|
|
channels = await Channels.get_channels_by_user_id(user.id)
|
|
log.debug(f'{channels=}')
|
|
for channel in channels:
|
|
await sio.enter_room(sid, f'channel:{channel.id}')
|
|
|
|
return {'id': user.id, 'name': user.name}
|
|
|
|
|
|
@sio.on('heartbeat')
|
|
async def heartbeat(sid, data):
|
|
user = SESSION_POOL.get(sid)
|
|
if user:
|
|
SESSION_POOL[sid] = {**user, 'last_seen_at': int(time.time())}
|
|
await Users.update_last_active_by_id(user['id'])
|
|
|
|
|
|
@sio.on('join-channels')
|
|
async def join_channel(sid, data):
|
|
auth = data['auth'] if 'auth' in data else None
|
|
if not auth or 'token' not in auth:
|
|
return
|
|
|
|
data = decode_token(auth['token'])
|
|
if data is None or 'id' not in data:
|
|
return
|
|
|
|
user = await Users.get_user_by_id(data['id'])
|
|
if not user:
|
|
return
|
|
|
|
# Join all the channels only if user has channels permission
|
|
if user.role == 'admin' or await has_permission(user.id, 'features.channels'):
|
|
channels = await Channels.get_channels_by_user_id(user.id)
|
|
log.debug(f'{channels=}')
|
|
for channel in channels:
|
|
await sio.enter_room(sid, f'channel:{channel.id}')
|
|
|
|
|
|
@sio.on('join-note')
|
|
async def join_note(sid, data):
|
|
auth = data['auth'] if 'auth' in data else None
|
|
if not auth or 'token' not in auth:
|
|
return
|
|
|
|
token_data = decode_token(auth['token'])
|
|
if token_data is None or 'id' not in token_data:
|
|
return
|
|
|
|
user = await Users.get_user_by_id(token_data['id'])
|
|
if not user:
|
|
return
|
|
|
|
note = await Notes.get_note_by_id(data['note_id'])
|
|
if not note:
|
|
log.error(f'Note {data["note_id"]} not found for user {user.id}')
|
|
return
|
|
|
|
if (
|
|
user.role != 'admin'
|
|
and user.id != note.user_id
|
|
and not await AccessGrants.has_access(
|
|
user_id=user.id,
|
|
resource_type='note',
|
|
resource_id=note.id,
|
|
permission='read',
|
|
)
|
|
):
|
|
log.error(f'User {user.id} does not have access to note {data["note_id"]}')
|
|
return
|
|
|
|
log.debug(f'Joining note {note.id} for user {user.id}')
|
|
await sio.enter_room(sid, f'note:{note.id}')
|
|
|
|
|
|
@sio.on('events:channel')
|
|
async def channel_events(sid, data):
|
|
room = f'channel:{data["channel_id"]}'
|
|
participants = sio.manager.get_participants(
|
|
namespace='/',
|
|
room=room,
|
|
)
|
|
|
|
sids = [sid for sid, _ in participants]
|
|
if sid not in sids:
|
|
return
|
|
|
|
event_data = data['data']
|
|
event_type = event_data['type']
|
|
|
|
user = SESSION_POOL.get(sid)
|
|
|
|
if not user:
|
|
return
|
|
|
|
if event_type == 'typing':
|
|
await sio.emit(
|
|
'events:channel',
|
|
{
|
|
'channel_id': data['channel_id'],
|
|
'message_id': data.get('message_id', None),
|
|
'data': event_data,
|
|
'user': UserNameResponse(**user).model_dump(),
|
|
},
|
|
room=room,
|
|
)
|
|
elif event_type == 'last_read_at':
|
|
await Channels.update_member_last_read_at(data['channel_id'], user['id'])
|
|
|
|
|
|
@sio.on('events:chat')
|
|
async def chat_events(sid, data):
|
|
user = SESSION_POOL.get(sid)
|
|
if not user:
|
|
return
|
|
|
|
event_data = data.get('data', {})
|
|
event_type = event_data.get('type')
|
|
|
|
if event_type == 'last_read_at':
|
|
await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id'])
|
|
|
|
|
|
@sio.on('resume-stream')
|
|
async def resume_stream(sid, data):
|
|
"""Replay WS events a client missed while disconnected.
|
|
|
|
Client payload: `{chat_id, message_id, last_seq}`.
|
|
|
|
Flow:
|
|
1. Authenticate the session and verify the user owns the chat AND
|
|
that the requested message_id actually belongs to that chat.
|
|
Both checks are required because the Redis stream log is keyed
|
|
by message_id alone — without the message-to-chat binding check,
|
|
an attacker who obtained a victim's message_id could satisfy the
|
|
chat-ownership check with any chat they own and then read the
|
|
victim's stream.
|
|
2. Read entries from the Redis resume log for this message_id whose
|
|
`seq` is greater than `last_seq`.
|
|
3. Emit each entry as a normal `events` event to THIS session only
|
|
(via `to=sid`). Other sessions for the same user keep receiving
|
|
the live stream unchanged.
|
|
|
|
No-op when Redis is not configured — in that deployment mode, refresh
|
|
during streaming falls back to the existing behavior (wait for the
|
|
stream to complete and reload from the DB).
|
|
"""
|
|
if REDIS is None:
|
|
return
|
|
|
|
# Reject malformed payloads early so `data.get(...)` never throws on a
|
|
# non-object client input (string/list/null/etc).
|
|
if not isinstance(data, dict):
|
|
return
|
|
|
|
user = SESSION_POOL.get(sid)
|
|
if not user:
|
|
return
|
|
|
|
user_id = user.get('id')
|
|
chat_id = data.get('chat_id')
|
|
message_id = data.get('message_id')
|
|
try:
|
|
last_seq = int(data.get('last_seq') or 0)
|
|
except (TypeError, ValueError):
|
|
last_seq = 0
|
|
|
|
if not user_id or not chat_id or not message_id:
|
|
return
|
|
|
|
# Step 1a: user owns the chat.
|
|
chat = await Chats.get_chat_by_id_and_user_id(chat_id, user_id)
|
|
if not chat:
|
|
return
|
|
|
|
# Step 1b: message_id actually lives in this chat. Without this, a
|
|
# caller who knows a victim's message_id could pass one of their OWN
|
|
# chat_ids (satisfying 1a) and read the victim's stream.
|
|
# `get_message_by_id_and_message_id` returns {} when the chat exists
|
|
# but the message_id isn't present, so check for a real id field.
|
|
message = await Chats.get_message_by_id_and_message_id(chat_id, message_id)
|
|
if not message or not message.get('id'):
|
|
return
|
|
|
|
envelopes = await _stream_log_read(message_id, last_seq)
|
|
for envelope in envelopes:
|
|
# Replay to the requesting session only. Live listeners in
|
|
# `user:{user_id}` are already receiving new frames via the normal
|
|
# emit path and must not see these duplicates.
|
|
await sio.emit('events', envelope, to=sid)
|
|
|
|
|
|
def normalize_document_id(document_id: str) -> str:
|
|
"""Canonicalize document IDs to prevent auth bypass via prefix variants.
|
|
|
|
YdocManager normalizes storage keys by replacing ":" with "_", so
|
|
"note_abc" and "note:abc" resolve to the same underlying document.
|
|
We must rewrite underscore-prefixed IDs back to the colon form so
|
|
that authorization checks (which key on "note:") always fire.
|
|
"""
|
|
if document_id.startswith('note_'):
|
|
document_id = 'note:' + document_id[5:]
|
|
return document_id
|
|
|
|
|
|
@sio.on('ydoc:document:join')
|
|
async def ydoc_document_join(sid, data):
|
|
"""Handle user joining a document"""
|
|
user = SESSION_POOL.get(sid)
|
|
if not user:
|
|
return
|
|
|
|
try:
|
|
document_id = normalize_document_id(data['document_id'])
|
|
|
|
if document_id.startswith('note:'):
|
|
note_id = document_id.split(':')[1]
|
|
note = await Notes.get_note_by_id(note_id)
|
|
if not note:
|
|
log.error(f'Note {note_id} not found')
|
|
return
|
|
|
|
if (
|
|
user.get('role') != 'admin'
|
|
and user.get('id') != note.user_id
|
|
and not await AccessGrants.has_access(
|
|
user_id=user.get('id'),
|
|
resource_type='note',
|
|
resource_id=note.id,
|
|
permission='read',
|
|
)
|
|
):
|
|
log.error(f'User {user.get("id")} does not have access to note {note_id}')
|
|
return
|
|
|
|
user_id = data.get('user_id', sid)
|
|
user_name = data.get('user_name', 'Anonymous')
|
|
user_color = data.get('user_color', '#000000')
|
|
|
|
log.info(f'User {user_id} joining document {document_id}')
|
|
await YDOC_MANAGER.add_user(document_id=document_id, user_id=sid)
|
|
|
|
# Join Socket.IO room
|
|
await sio.enter_room(sid, f'doc_{document_id}')
|
|
|
|
active_session_ids = get_session_ids_from_room(f'doc_{document_id}')
|
|
|
|
# Get the Yjs document state
|
|
ydoc = Y.Doc()
|
|
updates = await YDOC_MANAGER.get_updates(document_id)
|
|
for update in updates:
|
|
ydoc.apply_update(bytes(update))
|
|
|
|
# Encode the entire document state as an update
|
|
state_update = ydoc.get_update()
|
|
await sio.emit(
|
|
'ydoc:document:state',
|
|
{
|
|
'document_id': document_id,
|
|
'state': list(state_update), # Convert bytes to list for JSON
|
|
'sessions': active_session_ids,
|
|
},
|
|
room=sid,
|
|
)
|
|
|
|
# Notify other users about the new user
|
|
await sio.emit(
|
|
'ydoc:user:joined',
|
|
{
|
|
'document_id': document_id,
|
|
'user_id': user_id,
|
|
'user_name': user_name,
|
|
'user_color': user_color,
|
|
},
|
|
room=f'doc_{document_id}',
|
|
skip_sid=sid,
|
|
)
|
|
|
|
log.info(f'User {user_id} successfully joined document {document_id}')
|
|
|
|
except Exception as e:
|
|
log.error(f'Error in yjs_document_join: {e}')
|
|
await sio.emit('error', {'message': 'Failed to join document'}, room=sid)
|
|
|
|
|
|
async def document_save_handler(document_id, data, user):
|
|
document_id = normalize_document_id(document_id)
|
|
|
|
if document_id.startswith('note:'):
|
|
note_id = document_id.split(':')[1]
|
|
note = await Notes.get_note_by_id(note_id)
|
|
if not note:
|
|
log.error(f'Note {note_id} not found')
|
|
return
|
|
|
|
if (
|
|
user.get('role') != 'admin'
|
|
and user.get('id') != note.user_id
|
|
and not await AccessGrants.has_access(
|
|
user_id=user.get('id'),
|
|
resource_type='note',
|
|
resource_id=note.id,
|
|
permission='write',
|
|
)
|
|
):
|
|
log.error(f'User {user.get("id")} does not have write access to note {note_id}')
|
|
return
|
|
|
|
await Notes.update_note_by_id(note_id, NoteUpdateForm(data=data))
|
|
|
|
|
|
@sio.on('ydoc:document:state')
|
|
async def yjs_document_state(sid, data):
|
|
"""Send the current state of the Yjs document to the user"""
|
|
try:
|
|
document_id = data['document_id']
|
|
|
|
document_id = normalize_document_id(document_id)
|
|
room = f'doc_{document_id}'
|
|
|
|
active_session_ids = get_session_ids_from_room(room)
|
|
|
|
if sid not in active_session_ids:
|
|
log.warning(f'Session {sid} not in room {room}. Cannot send state.')
|
|
return
|
|
|
|
if not await YDOC_MANAGER.document_exists(document_id):
|
|
log.warning(f'Document {document_id} not found')
|
|
return
|
|
|
|
# Get the Yjs document state
|
|
ydoc = Y.Doc()
|
|
updates = await YDOC_MANAGER.get_updates(document_id)
|
|
for update in updates:
|
|
ydoc.apply_update(bytes(update))
|
|
|
|
# Encode the entire document state as an update
|
|
state_update = ydoc.get_update()
|
|
|
|
await sio.emit(
|
|
'ydoc:document:state',
|
|
{
|
|
'document_id': document_id,
|
|
'state': list(state_update), # Convert bytes to list for JSON
|
|
'sessions': active_session_ids,
|
|
},
|
|
room=sid,
|
|
)
|
|
except Exception as e:
|
|
log.error(f'Error in yjs_document_state: {e}')
|
|
|
|
|
|
@sio.on('ydoc:document:update')
|
|
async def yjs_document_update(sid, data):
|
|
"""Handle Yjs document updates"""
|
|
try:
|
|
document_id = data['document_id']
|
|
|
|
document_id = normalize_document_id(document_id)
|
|
|
|
# Verify the sender actually joined this document room
|
|
room = f'doc_{document_id}'
|
|
active_session_ids = get_session_ids_from_room(room)
|
|
if sid not in active_session_ids:
|
|
log.warning(f'Session {sid} not in room {room}. Rejecting update.')
|
|
return
|
|
|
|
try:
|
|
await stop_item_tasks(REDIS, document_id)
|
|
except Exception:
|
|
pass
|
|
|
|
user_id = data.get('user_id', sid)
|
|
|
|
update = data['update'] # List of bytes from frontend
|
|
|
|
await YDOC_MANAGER.append_to_updates(
|
|
document_id=document_id,
|
|
update=update, # Convert list of bytes to bytes
|
|
)
|
|
|
|
# Broadcast update to all other users in the document
|
|
await sio.emit(
|
|
'ydoc:document:update',
|
|
{
|
|
'document_id': document_id,
|
|
'user_id': user_id,
|
|
'update': update,
|
|
'socket_id': sid, # Add socket_id to match frontend filtering
|
|
},
|
|
room=f'doc_{document_id}',
|
|
skip_sid=sid,
|
|
)
|
|
|
|
user = SESSION_POOL.get(sid)
|
|
if not user:
|
|
return
|
|
|
|
async def debounced_save():
|
|
await asyncio.sleep(0.5)
|
|
await document_save_handler(document_id, data.get('data', {}), user)
|
|
|
|
if data.get('data'):
|
|
await create_task(REDIS, debounced_save(), document_id)
|
|
|
|
except Exception as e:
|
|
log.error(f'Error in yjs_document_update: {e}')
|
|
|
|
|
|
@sio.on('ydoc:document:leave')
|
|
async def yjs_document_leave(sid, data):
|
|
"""Handle user leaving a document"""
|
|
try:
|
|
document_id = data['document_id']
|
|
user_id = data.get('user_id', sid)
|
|
|
|
log.info(f'User {user_id} leaving document {document_id}')
|
|
|
|
# Remove user from the document
|
|
await YDOC_MANAGER.remove_user(document_id=document_id, user_id=sid)
|
|
|
|
# Leave Socket.IO room
|
|
await sio.leave_room(sid, f'doc_{document_id}')
|
|
|
|
# Notify other users
|
|
await sio.emit(
|
|
'ydoc:user:left',
|
|
{'document_id': document_id, 'user_id': user_id},
|
|
room=f'doc_{document_id}',
|
|
)
|
|
|
|
if await YDOC_MANAGER.document_exists(document_id) and len(await YDOC_MANAGER.get_users(document_id)) == 0:
|
|
log.info(f'Cleaning up document {document_id} as no users are left')
|
|
await YDOC_MANAGER.clear_document(document_id)
|
|
|
|
except Exception as e:
|
|
log.error(f'Error in yjs_document_leave: {e}')
|
|
|
|
|
|
@sio.on('ydoc:awareness:update')
|
|
async def yjs_awareness_update(sid, data):
|
|
"""Handle awareness updates (cursors, selections, etc.)"""
|
|
try:
|
|
document_id = data['document_id']
|
|
user_id = data.get('user_id', sid)
|
|
update = data['update']
|
|
|
|
# Broadcast awareness update to all other users in the document
|
|
await sio.emit(
|
|
'ydoc:awareness:update',
|
|
{'document_id': document_id, 'user_id': user_id, 'update': update},
|
|
room=f'doc_{document_id}',
|
|
skip_sid=sid,
|
|
)
|
|
|
|
except Exception as e:
|
|
log.error(f'Error in yjs_awareness_update: {e}')
|
|
|
|
|
|
@sio.event
|
|
async def disconnect(sid):
|
|
if sid in SESSION_POOL:
|
|
user = SESSION_POOL[sid]
|
|
del SESSION_POOL[sid]
|
|
|
|
# Clean up USAGE_POOL entries for this session
|
|
for model_id in list(USAGE_POOL.keys()):
|
|
connections = USAGE_POOL.get(model_id)
|
|
if connections and sid in connections:
|
|
del connections[sid]
|
|
if not connections:
|
|
del USAGE_POOL[model_id]
|
|
else:
|
|
USAGE_POOL[model_id] = connections
|
|
|
|
await YDOC_MANAGER.remove_user_from_all_documents(sid)
|
|
else:
|
|
pass
|
|
# print(f"Unknown session ID {sid} disconnected")
|
|
|
|
|
|
async def get_event_emitter(request_info, update_db=True):
|
|
# Per-emitter monotonic seq. One emitter instance corresponds to a single
|
|
# streaming response for a single (chat_id, message_id) on a single
|
|
# worker, so a local counter is sufficient — no distributed consensus
|
|
# needed. Clients use this to request a replay of events they missed
|
|
# after a reconnect / refresh via the `resume-stream` handler below.
|
|
seq_counter = {'n': 0}
|
|
# Reset any stale resume log for this message_id. Continuation,
|
|
# regeneration-into-same-id, or a retry after a crashed worker can
|
|
# create a second emitter for the same message_id. The old emitter's
|
|
# explicit stream IDs (`0-{seq}`) would collide with our fresh ones
|
|
# and XADD would silently fail, quietly breaking resumability for
|
|
# exactly the flows this feature is meant to protect. Deleting the
|
|
# old log up front guarantees our XADD `0-1` is accepted and the log
|
|
# reflects only the current run, not a mix of runs.
|
|
message_id = request_info.get('message_id') if isinstance(request_info, dict) else None
|
|
if message_id and REDIS is not None:
|
|
await _stream_log_truncate(message_id)
|
|
|
|
async def __event_emitter__(event_data):
|
|
user_id = request_info['user_id']
|
|
chat_id = request_info['chat_id']
|
|
message_id = request_info['message_id']
|
|
|
|
seq_counter['n'] += 1
|
|
seq = seq_counter['n']
|
|
|
|
envelope = {
|
|
'chat_id': chat_id,
|
|
'message_id': message_id,
|
|
'seq': seq,
|
|
'data': event_data,
|
|
}
|
|
|
|
# Append to the resume log BEFORE the live emit. If we emitted
|
|
# first, a client that disconnects in the window between `sio.emit`
|
|
# and `_stream_log_append` could reconnect and issue resume-stream
|
|
# before the frame is logged, never see that frame in the replay,
|
|
# and then never ask again — permanently losing it (particularly
|
|
# painful for the terminal done:True frame).
|
|
#
|
|
# Logging first inverts the window: a reconnecting client MIGHT
|
|
# see a replayed frame before the live emit reaches their new
|
|
# session, but the seq idempotency guard in chatEventHandler drops
|
|
# the subsequent duplicate harmlessly. Duplicates are safe, losses
|
|
# are not.
|
|
await _stream_log_append(message_id, envelope, seq)
|
|
await sio.emit('events', envelope, room=f'user:{user_id}')
|
|
|
|
# If this event finalized the message, schedule log cleanup. Give
|
|
# reconnecting clients a short grace window to pick up the final
|
|
# frames before we delete the log (for anything beyond the grace
|
|
# window, the DB is already up to date and resume isn't needed).
|
|
# Narrow carefully: `event_data['data']` is a dict for most event
|
|
# types but can legitimately be a list/str/None for some custom
|
|
# pipeline-emitted events. Calling `.get` on those would raise.
|
|
inner = event_data.get('data') if isinstance(event_data, dict) else None
|
|
if isinstance(inner, dict) and inner.get('done') is True:
|
|
async def _delayed_truncate(mid):
|
|
try:
|
|
await asyncio.sleep(30)
|
|
await _stream_log_truncate(mid)
|
|
except Exception:
|
|
pass
|
|
|
|
asyncio.create_task(_delayed_truncate(message_id))
|
|
|
|
if update_db and message_id and not request_info.get('chat_id', '').startswith('local:'):
|
|
event_type = event_data.get('type')
|
|
|
|
if event_type == 'status':
|
|
await Chats.add_message_status_to_chat_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
event_data.get('data', {}),
|
|
)
|
|
|
|
elif event_type == 'message':
|
|
message = await Chats.get_message_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
)
|
|
|
|
if message:
|
|
content = message.get('content', '')
|
|
content += event_data.get('data', {}).get('content', '')
|
|
|
|
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
{
|
|
'content': content,
|
|
},
|
|
)
|
|
|
|
elif event_type == 'replace':
|
|
content = event_data.get('data', {}).get('content', '')
|
|
|
|
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
{
|
|
'content': content,
|
|
},
|
|
)
|
|
|
|
elif event_type == 'embeds':
|
|
message = await Chats.get_message_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
)
|
|
|
|
embeds = event_data.get('data', {}).get('embeds', [])
|
|
embeds.extend(message.get('embeds', []))
|
|
|
|
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
{
|
|
'embeds': embeds,
|
|
},
|
|
)
|
|
|
|
elif event_type == 'files':
|
|
message = await Chats.get_message_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
)
|
|
|
|
files = event_data.get('data', {}).get('files', [])
|
|
files.extend(message.get('files', []))
|
|
|
|
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
{
|
|
'files': files,
|
|
},
|
|
)
|
|
|
|
elif event_type in ('source', 'citation'):
|
|
data = event_data.get('data', {})
|
|
if data.get('type') is None:
|
|
message = await Chats.get_message_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
)
|
|
|
|
sources = message.get('sources', [])
|
|
sources.append(data)
|
|
|
|
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
|
request_info['chat_id'],
|
|
request_info['message_id'],
|
|
{
|
|
'sources': sources,
|
|
},
|
|
)
|
|
|
|
if 'user_id' in request_info and 'chat_id' in request_info and 'message_id' in request_info:
|
|
return __event_emitter__
|
|
else:
|
|
return None
|
|
|
|
|
|
async def get_event_call(request_info):
|
|
async def __event_caller__(event_data):
|
|
response = await sio.call(
|
|
'events',
|
|
{
|
|
'chat_id': request_info.get('chat_id', None),
|
|
'message_id': request_info.get('message_id', None),
|
|
'data': event_data,
|
|
},
|
|
to=request_info['session_id'],
|
|
timeout=WEBSOCKET_EVENT_CALLER_TIMEOUT,
|
|
)
|
|
return response
|
|
|
|
if 'session_id' in request_info and 'chat_id' in request_info and 'message_id' in request_info:
|
|
return __event_caller__
|
|
else:
|
|
return None
|
|
|
|
|
|
get_event_caller = get_event_call
|