From ce3c175e260709f359d7e6cbb3132f0572098b95 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Mon, 10 Aug 2026 23:13:10 -0600 Subject: [PATCH] refac --- backend/open_webui/events.py | 16 +++++++- backend/open_webui/routers/auths.py | 12 +++++- backend/open_webui/routers/scim.py | 57 +++++++++++++++++++++-------- backend/open_webui/routers/users.py | 28 +++++++------- backend/open_webui/socket/main.py | 22 +++++++---- backend/open_webui/utils/oauth.py | 10 ++++- src/routes/+layout.svelte | 4 ++ 7 files changed, 109 insertions(+), 40 deletions(-) diff --git a/backend/open_webui/events.py b/backend/open_webui/events.py index 352d2e68ae..5b23231f41 100644 --- a/backend/open_webui/events.py +++ b/backend/open_webui/events.py @@ -1080,6 +1080,20 @@ class NotificationEventSink: schedule_notification_dispatch(app, event) +class SocketSessionEventSink: + async def handle_event(self, app: Any, event: Event, request: Any | None = None) -> None: + if event.event not in {EVENTS.USER_DELETED.name, EVENTS.USER_ROLE_UPDATED.name}: + return + + subject = event.subject or {} + if subject.get('type') != 'user' or not subject.get('id'): + return + + from open_webui.socket.main import disconnect_user_sessions + + await disconnect_user_sessions(str(subject['id'])) + + async def dispatch_event_functions( app: Any, event: Event, request: Any | None = None, extra_function_ids: list[str] | None = None ) -> None: @@ -1148,7 +1162,7 @@ class EventFunctionSink: schedule_event_function_dispatch(app, event, request) -EVENT_SINKS = [EventFunctionSink(), WebhookEventSink(), NotificationEventSink()] +EVENT_SINKS = [SocketSessionEventSink(), EventFunctionSink(), WebhookEventSink(), NotificationEventSink()] async def publish_event( diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index a907cb4e1c..86de689458 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -768,7 +768,17 @@ async def signin( trusted_role = request.headers.get(WEBUI_AUTH_TRUSTED_ROLE_HEADER, '').lower().strip() if trusted_role in {'admin', 'user', 'pending'}: if user.role != trusted_role: - await Users.update_user_role_by_id(user.id, trusted_role, db=db) + updated_user = await Users.update_user_role_by_id(user.id, trusted_role, db=db) + if updated_user: + user = updated_user + await publish_event( + request, + EVENTS.USER_ROLE_UPDATED, + actor=updated_user, + subject_id=updated_user.id, + source='trusted_header', + data={'role': updated_user.role}, + ) elif trusted_role: log.warning(f'Ignoring invalid trusted role header value: {trusted_role}') diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index 9a96c0b097..a6b528edfa 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -700,15 +700,27 @@ async def update_user( await Users.update_user_scim_by_id(user_id, provider, user_data.externalId, db=db) updated_user = await Users.get_user_by_id(user_id, db=db) - await publish_event( - request, - EVENTS.USER_UPDATED, - subject_id=user_id, - source='scim', - data={ - 'updated_fields': list(update_data.keys()) + (['externalId'] if user_data.externalId else []), - }, - ) + updated_fields = list(update_data.keys()) + (['externalId'] if user_data.externalId else []) + role_changed = updated_user.role != user.role + user_updated_fields = [field for field in updated_fields if field != 'role'] + + if user_updated_fields: + await publish_event( + request, + EVENTS.USER_UPDATED, + subject_id=user_id, + source='scim', + data={'updated_fields': user_updated_fields}, + ) + + if role_changed: + await publish_event( + request, + EVENTS.USER_ROLE_UPDATED, + subject_id=user_id, + source='scim', + data={'role': updated_user.role}, + ) return await user_to_scim(updated_user, request, db=db) @@ -764,13 +776,26 @@ async def patch_user( else: updated_user = user - await publish_event( - request, - EVENTS.USER_UPDATED, - subject_id=user_id, - source='scim', - data={'updated_fields': list(update_data.keys())}, - ) + role_changed = updated_user.role != user.role + user_updated_fields = [field for field in update_data.keys() if field != 'role'] + + if user_updated_fields: + await publish_event( + request, + EVENTS.USER_UPDATED, + subject_id=user_id, + source='scim', + data={'updated_fields': user_updated_fields}, + ) + + if role_changed: + await publish_event( + request, + EVENTS.USER_ROLE_UPDATED, + subject_id=user_id, + source='scim', + data={'role': updated_user.role}, + ) return await user_to_scim(updated_user, request, db=db) diff --git a/backend/open_webui/routers/users.py b/backend/open_webui/routers/users.py index 081023826a..9256e1b922 100644 --- a/backend/open_webui/routers/users.py +++ b/backend/open_webui/routers/users.py @@ -36,7 +36,6 @@ from open_webui.models.access_grants import AccessGrants from open_webui.models.knowledge import Knowledges from open_webui.models.models import Models from open_webui.models.tools import Tools -from open_webui.socket.main import disconnect_user_sessions from open_webui.utils.access_control import get_permissions, has_permission from open_webui.utils.auth import ( get_admin_user, @@ -990,10 +989,19 @@ async def update_user_by_id( updated_user = user if updated_user: - # If the role changed, disconnect all socket sessions so stale - # privileges cached in SESSION_POOL are invalidated. - if updated_user.role != user.role: - await disconnect_user_sessions(user_id) + updated_fields = [field for field in update_data.keys() if field != 'role'] + role_changed = updated_user.role != user.role + + if updated_fields: + await publish_event( + request, + EVENTS.USER_UPDATED, + actor=session_user, + subject_id=user_id, + data={'updated_fields': updated_fields}, + ) + + if role_changed: await publish_event( request, EVENTS.USER_ROLE_UPDATED, @@ -1001,14 +1009,7 @@ async def update_user_by_id( subject_id=user_id, data={'role': updated_user.role}, ) - else: - await publish_event( - request, - EVENTS.USER_UPDATED, - actor=session_user, - subject_id=user_id, - data={'updated_fields': list(update_data.keys())}, - ) + if form_data.password: await publish_event( request, @@ -1061,7 +1062,6 @@ async def delete_user_by_id( result = await Auths.delete_auth_by_id(user_id, db=db) if result: - await disconnect_user_sessions(user_id) await publish_event( request, EVENTS.USER_DELETED, diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 0a5e011e11..7cb7904fed 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -271,6 +271,13 @@ def get_session_ids_from_room(room): return list(members) if members else [] +def get_session_ids_by_user_id(user_id: str) -> list[str]: + """Get known session IDs for a user across the local rooms and shared session pool.""" + session_ids = set(get_session_ids_from_room(f'user:{user_id}')) + session_ids.update(sid for sid, entry in SESSION_POOL.items() if entry and entry.get('id') == user_id) + return list(session_ids) + + def get_user_ids_from_room(room): active_session_ids = get_session_ids_from_room(room) @@ -326,14 +333,15 @@ async def disconnect_user_sessions(user_id: str): 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: + session_ids = get_session_ids_by_user_id(user_id) + for sid in session_ids: + try: await sio.disconnect(sid) - if session_ids: - log.info('Disconnected %s session(s) for user %s', len(session_ids), user_id) - except Exception as e: - log.warning(f'Failed to disconnect sessions for user {user_id}: {e}') + except Exception: + log.exception('Failed to disconnect session %s for user %s', sid, user_id) + + if session_ids: + log.info('Requested disconnect of %s session(s) for user %s', len(session_ids), user_id) @sio.on('usage') diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index a6180e1b50..fce053dae8 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -1942,10 +1942,18 @@ class OAuthManager: if user: determined_role = await self.get_user_role(user, user_data) if user.role != determined_role: - await Users.update_user_role_by_id(user.id, determined_role, db=db) + updated_user = await Users.update_user_role_by_id(user.id, determined_role, db=db) # Update the user object in memory as well, # to avoid problems with the ENABLE_OAUTH_GROUP_MANAGEMENT check below user.role = determined_role + await publish_event( + request, + EVENTS.USER_ROLE_UPDATED, + actor=updated_user or user, + subject_id=user.id, + source='oauth', + data={'role': determined_role, 'provider': provider}, + ) if auth_config.OAUTH_UPDATE_NAME_ON_LOGIN: username_claim = auth_config.OAUTH_USERNAME_CLAIM diff --git a/src/routes/+layout.svelte b/src/routes/+layout.svelte index 00a3725a3d..59dbcd808f 100644 --- a/src/routes/+layout.svelte +++ b/src/routes/+layout.svelte @@ -266,6 +266,10 @@ heartbeatInterval = null; } + if (reason === 'io server disconnect') { + _socket.connect(); + } + if (details) { console.log('Additional details:', details); }