This commit is contained in:
Timothy Jaeryang Baek 2026-08-10 23:13:10 -06:00
parent de289eb1aa
commit ce3c175e26
7 changed files with 109 additions and 40 deletions

View file

@ -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(

View file

@ -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}')

View file

@ -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)

View file

@ -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,

View file

@ -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')

View file

@ -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

View file

@ -266,6 +266,10 @@
heartbeatInterval = null;
}
if (reason === 'io server disconnect') {
_socket.connect();
}
if (details) {
console.log('Additional details:', details);
}