mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-28 05:27:35 +00:00
refac
This commit is contained in:
parent
de289eb1aa
commit
ce3c175e26
7 changed files with 109 additions and 40 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}')
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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')
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -266,6 +266,10 @@
|
|||
heartbeatInterval = null;
|
||||
}
|
||||
|
||||
if (reason === 'io server disconnect') {
|
||||
_socket.connect();
|
||||
}
|
||||
|
||||
if (details) {
|
||||
console.log('Additional details:', details);
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue