From de73bb830aeb150bc5bb4707566d3969f70408b0 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Thu, 8 Oct 2026 14:30:07 +0400 Subject: [PATCH] refac --- backend/open_webui/events.py | 22 ++ backend/open_webui/main.py | 226 +++++-------- backend/open_webui/models/chat_messages.py | 56 ++-- backend/open_webui/models/chats.py | 304 +++++++++++------- backend/open_webui/models/shared_chats.py | 39 ++- backend/open_webui/routers/chats.py | 65 +++- backend/open_webui/routers/folders.py | 6 + backend/open_webui/socket/main.py | 107 +++++- .../open_webui/utils/access_control/files.py | 3 + backend/open_webui/utils/middleware.py | 20 +- backend/open_webui/utils/tool_approval.py | 20 +- src/lib/apis/chats/index.ts | 15 +- src/lib/apis/folders/index.ts | 9 +- src/lib/components/chat/Chat.svelte | 172 ++++++++-- src/lib/components/chat/Messages.svelte | 2 + .../components/chat/Messages/Message.svelte | 5 +- .../Messages/MultiResponseMessages.svelte | 3 + .../chat/Messages/ResponseMessage.svelte | 13 +- src/lib/components/chat/ShareChatModal.svelte | 65 ++-- .../Sidebar/Folders/FolderShareModal.svelte | 30 +- .../workspace/common/AccessControl.svelte | 2 + src/routes/+layout.svelte | 1 + src/routes/s/[id]/+page.svelte | 5 + 23 files changed, 798 insertions(+), 392 deletions(-) diff --git a/backend/open_webui/events.py b/backend/open_webui/events.py index 7aa7cfc72a..0bdd0f5162 100644 --- a/backend/open_webui/events.py +++ b/backend/open_webui/events.py @@ -1121,6 +1121,28 @@ class NotificationEventSink: class SocketSessionEventSink: async def handle_event(self, app: Any, event: Event, request: Any | None = None) -> None: + from open_webui.socket.main import refresh_chat_access + + if event.event in { + EVENTS.FOLDER_ACCESS_UPDATED.name, + EVENTS.FOLDER_UPDATED.name, + EVENTS.FOLDER_PARENT_UPDATED.name, + EVENTS.FOLDER_DELETED.name, + EVENTS.GROUP_MEMBER_REMOVED.name, + EVENTS.GROUP_MEMBER_ADDED.name, + EVENTS.GROUP_UPDATED.name, + EVENTS.GROUP_DELETED.name, + EVENTS.CHAT_DELETED_ALL.name, + }: + await refresh_chat_access() + elif event.event in { + EVENTS.CHAT_SHARED.name, + EVENTS.CHAT_UNSHARED.name, + EVENTS.CHAT_FOLDER_UPDATED.name, + EVENTS.CHAT_DELETED.name, + }: + await refresh_chat_access((event.subject or {}).get('id')) + if event.event not in {EVENTS.USER_DELETED.name, EVENTS.USER_ROLE_UPDATED.name}: return diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index f696d09527..5b0820ed6c 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -1263,8 +1263,10 @@ async def chat_completion( user_message = form_data.pop('user_message', None) or form_data.pop('parent_message', None) chat_id = form_data.pop('chat_id', None) or '' chat_variables = form_data.pop('chat_variables', None) - if chat_variables is None: - existing_chat = await Chats.get_chat_by_id(chat_id) if is_saved_chat_id(chat_id) else None + existing_chat = await Chats.get_chat_by_id(chat_id) if is_saved_chat_id(chat_id) else None + if existing_chat and existing_chat.user_id != user.id: + chat_variables = {} + elif chat_variables is None: chat_variables = existing_chat.variables if existing_chat else {} chat_variables = normalize_chat_variables(chat_variables) @@ -1401,6 +1403,7 @@ async def chat_completion( if user_message_id and user_message: user_message['childrenIds'] = all_assistant_ids + user_message['user_id'] = user.id history_messages[user_message_id] = user_message for entry in message_ids: @@ -1455,7 +1458,7 @@ async def chat_completion( subject_id=chat_id, data={'title': 'New Chat'}, ) - await emit_chat_list_event(metadata, chat_id) + await emit_chat_list_event({**metadata, 'message_id': user_message_id}, chat_id) if user_message_id: await publish_event( request, @@ -1524,56 +1527,41 @@ async def chat_completion( asyncio.create_task(run_initial_title_generation()) else: - # Existing chat — verify ownership - if not await Chats.is_chat_owner(chat_id, user.id) and not ( - user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS - ): - raise HTTPException( - status_code=status.HTTP_404_NOT_FOUND, - detail=ERROR_MESSAGES.DEFAULT(), - ) + chat = await Chats.get_accessible_chat_by_id(chat_id, user, permission='write') + if not chat: + raise HTTPException(status_code=404, detail=ERROR_MESSAGES.NOT_FOUND) user_message = metadata.get('user_message') or {} - selected_chat_models = user_message.get('models') if isinstance(user_message, dict) else None - if not isinstance(selected_chat_models, list) or not selected_chat_models: - selected_chat_models = [entry.get('model_id') for entry in message_ids if entry.get('model_id')] - - # Persist chat-level fields the frontend used to save on every message. - # The old frontend saveChatHandler did this on every message; - # now the backend owns persistence. - chat_files = metadata.get('files') - chat_fields = {} - if chat_files is not None: - chat_fields['files'] = chat_files - if selected_chat_models: - chat_fields['models'] = selected_chat_models - if chat_fields: - await Chats.update_chat_by_id(chat_id, chat_fields, touch=False) - - await Chats.update_chat_variables_by_id(chat_id, chat_variables) - - # Save user message to DB - if user_message and user_message.get('id'): - await Chats.upsert_message_to_chat_by_id_and_message_id( - chat_id, - user_message['id'], - user_message, + assistant_message_id = metadata.get('assistant_message_id') + if assistant_message_id: + message = await Chats.get_message_by_id_and_message_id(chat_id, assistant_message_id) + if not message or (message.get('user_id') or chat.user_id) != user.id: + raise HTTPException(status_code=403, detail=ERROR_MESSAGES.ACCESS_PROHIBITED) + if any(entry.get('message_id') != assistant_message_id for entry in message_ids): + raise HTTPException(status_code=400, detail='Invalid response ID.') + metadata['user_message_id'] = message.get('parentId') + else: + turn = await Chats.insert_chat_turn(chat_id, user, user_message, message_ids) + event_emitter = await get_event_emitter( + {**metadata, 'message_id': turn['currentId']}, update_db=False ) + turn['messages'] = { + mid: {key: value for key, value in message.items() if key != 'meta'} + for mid, message in turn['messages'].items() + } + await event_emitter({'type': 'chat:messages', 'data': turn}) await emit_chat_list_event({**metadata, 'message_id': user_message['id']}, chat_id) - await publish_event( - request, - EVENTS.MESSAGE_CREATED, - actor=user, - subject_id=user_message['id'], - data={ - 'chat_id': chat_id, - 'role': user_message.get('role', 'user'), - 'content_preview': user_message.get('content', '')[:300], - }, - ) - if not getattr(request.state, 'internal', False) and not (user_message.get('meta') or {}).get( - 'internal' - ): + for message_id, message in turn['messages'].items(): + if message_id != user_message['id'] and message.get('parentId') != user_message['id']: + continue + await publish_event( + request, + EVENTS.MESSAGE_CREATED, + actor=user, + subject_id=message_id, + data={'chat_id': chat_id, 'role': message['role'], 'model': message.get('model')}, + ) + if not getattr(request.state, 'internal', False): try: from open_webui.utils.timers import cancel_timers_for_chat @@ -1581,91 +1569,30 @@ async def chat_completion( except Exception: log.exception('Failed to cancel chat.user_message timers for chat %s', chat_id) - # Link grandparent → user message (childrenIds) - grandparent_id = user_message.get('parentId') - if grandparent_id: - grandparent = await Chats.get_message_by_id_and_message_id(chat_id, grandparent_id) - if grandparent: - child_ids = grandparent.get('childrenIds', []) - if user_message['id'] not in child_ids: - child_ids.append(user_message['id']) - await Chats.upsert_message_to_chat_by_id_and_message_id( - chat_id, grandparent_id, {'childrenIds': child_ids} - ) + if chat.user_id == user.id: + selected_chat_models = user_message.get('models') or [ + entry['model_id'] for entry in message_ids + ] + chat_fields = {'models': selected_chat_models} + if metadata.get('files') is not None: + chat_fields['files'] = metadata['files'] + await Chats.update_chat_by_id(chat_id, chat_fields, touch=False) + await Chats.update_chat_variables_by_id(chat_id, chat_variables) + else: + tasks = { + key: value + for key, value in (tasks or {}).items() + if key not in (TASKS.TITLE_GENERATION, TASKS.TAGS_GENERATION) + } or None - # Insert chat files from user message if any user_message_files = user_message.get('files', []) if user_message_files: - try: - await Chats.insert_chat_files( - chat_id, - user_message.get('id'), - [ - file_item.get('id') - for file_item in user_message_files - if file_item.get('type') == 'file' - ], - user.id, - ) - except Exception as e: - log.debug('Error inserting chat files: %s', e) - pass - - # Save ALL assistant placeholders - user_message_id = metadata.get('user_message_id') - all_assistant_ids = [entry['message_id'] for entry in message_ids if entry.get('message_id')] - - # Link user message → all assistant messages (childrenIds) - if user_message_id and all_assistant_ids: - existing_user_message = await Chats.get_message_by_id_and_message_id(chat_id, user_message_id) - if existing_user_message: - child_ids = existing_user_message.get('childrenIds', []) - for assistant_id in all_assistant_ids: - if assistant_id not in child_ids: - child_ids.append(assistant_id) - await Chats.upsert_message_to_chat_by_id_and_message_id( - chat_id, - user_message_id, - {'childrenIds': child_ids}, - ) - - # Save each assistant placeholder - for entry in message_ids: - target_model_id = entry['model_id'] - assistant_message_id = entry['message_id'] - if assistant_message_id and assistant_message_id == metadata.get('assistant_message_id'): - continue - if assistant_message_id: - assistant_message = { - 'id': assistant_message_id, - 'parentId': user_message_id, - 'childrenIds': [], - 'role': 'assistant', - 'content': '', - 'done': False, - 'model': target_model_id, - 'timestamp': int(time.time()), - } - # Preserve the side-by-side column index so duplicate - # models don't collapse into one another on reload. - if entry.get('modelIdx') is not None: - assistant_message['modelIdx'] = entry['modelIdx'] - await Chats.upsert_message_to_chat_by_id_and_message_id( - chat_id, - assistant_message_id, - assistant_message, - ) - await publish_event( - request, - EVENTS.MESSAGE_CREATED, - actor=user, - subject_id=assistant_message_id, - data={ - 'chat_id': chat_id, - 'role': 'assistant', - 'model': target_model_id, - }, - ) + await Chats.insert_chat_files( + chat_id, + user_message.get('id'), + [file.get('id') for file in user_message_files if file.get('type') == 'file'], + user.id, + ) request.state.metadata = metadata form_data['metadata'] = metadata @@ -1685,12 +1612,15 @@ async def chat_completion( try: if not metadata.get('direct'): target = ( - model_info if form_data['model'] == model_id + model_info + if form_data['model'] == model_id else await Models.get_model_by_id(form_data['model']) ) controls = target.params.model_dump().get('model_controls', {}) if target else {} form_data['params'] = apply_model_controls( - copy.deepcopy(form_data.get('params') or {}), controls, model_controls.get(form_data['model'], {}) + copy.deepcopy(form_data.get('params') or {}), + controls, + model_controls.get(form_data['model'], {}), ) ctx = None # Saved chats load the message after approved tool calls run, so their results are kept @@ -1864,6 +1794,11 @@ async def chat_completion( 'task_id': str(uuid4()), } + if is_saved_chat_id(chat_id): + await Chats.upsert_message_to_chat_by_id_and_message_id( + chat_id, assistant_message_id, {'meta': {'task_id': per_model_metadata['task_id']}}, touch=False + ) + # Per-model form_data: own model model_form_data = { **form_data, @@ -2200,12 +2135,15 @@ async def list_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De if owner_id != user.id and not (user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS): return {'task_ids': []} else: - chat = await Chats.get_chat_by_id(chat_id) - if chat is None or (chat.user_id != user.id and not (user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS)): + chat = await Chats.get_accessible_chat_by_id(chat_id, user, permission='write') + if chat is None: return {'task_ids': []} task_ids = await list_task_ids_by_item_id(request.app.state.redis, chat_id) + if not socket_id: + task_ids = await Chats.filter_task_ids_by_user_id(chat, user.id, task_ids) + log.debug('Task IDs for chat %s: %s', chat_id, task_ids) return {'task_ids': task_ids} @@ -2219,14 +2157,24 @@ async def stop_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De if owner_id != user.id and not (user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS): raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) else: - chat = await Chats.get_chat_by_id(chat_id) - if chat is None or (chat.user_id != user.id and not (user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS)): + chat = await Chats.get_accessible_chat_by_id(chat_id, user, permission='write') + if chat is None: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) - result = await stop_item_tasks(request.app.state.redis, chat_id) + if socket_id: + result = await stop_item_tasks(request.app.state.redis, chat_id) + else: + task_ids = await Chats.filter_task_ids_by_user_id( + chat, user.id, await list_task_ids_by_item_id(request.app.state.redis, chat_id) + ) + result = {'status': True, 'message': 'No tasks found.'} + for task_id in task_ids: + result = await stop_task(request.app.state.redis, task_id) if not socket_id and str(result.get('message', '')).startswith('No tasks found'): messages_map = await Chats.get_messages_map_by_chat_id(chat_id) or {} for message_id, message in messages_map.items(): + if (message.get('user_id') or chat.user_id) != user.id: + continue if message.get('role') != 'assistant' or message.get('done') is not False: continue @@ -2254,7 +2202,7 @@ async def stop_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De event_emitter = await get_event_emitter( { - 'user_id': chat.user_id, + 'user_id': user.id, 'chat_id': chat_id, 'message_id': message_id, }, diff --git a/backend/open_webui/models/chat_messages.py b/backend/open_webui/models/chat_messages.py index 0473fb1a98..de5366ec52 100644 --- a/backend/open_webui/models/chat_messages.py +++ b/backend/open_webui/models/chat_messages.py @@ -1,3 +1,4 @@ +from contextlib import nullcontext import time import uuid from collections import Counter @@ -278,47 +279,43 @@ class ChatMessageTable: db: Optional[AsyncSession] = None, ) -> Optional[ChatMessageModel]: """Insert or update a chat message.""" - async with get_async_db_context(db) as db: + async with nullcontext(db) if db is not None else get_async_db_context() as session: now = int(time.time()) # Use composite ID: {chat_id}-{message_id} composite_id = f'{chat_id}-{message_id}' - message = await db.get(ChatMessage, composite_id) + message = await session.get(ChatMessage, composite_id) if message: self._apply_message_data(message, data, now) else: message = self._build_message(composite_id, chat_id, user_id, data, now) - db.add(message) + session.add(message) - await db.commit() + if db is None: + await session.commit() + else: + await session.flush() return ChatMessageModel.model_validate(message) async def upsert_messages( - self, - chat_id: str, - user_id: str, - messages: dict[str, dict], - db: AsyncSession | None = None, + self, chat_id: str, user_id: str, messages: dict[str, dict], db: AsyncSession | None = None ) -> None: - """Insert or update the given messages of one chat.""" + """Backfill missing rows without overwriting newer message data.""" + from sqlalchemy.dialects.sqlite import insert as sqlite_insert + from sqlalchemy.dialects.postgresql import insert as pg_insert + if not messages: return - async with get_async_db_context(db) as db: - now = int(time.time()) - result = await db.execute( - select(ChatMessage).filter(ChatMessage.id.in_([f'{chat_id}-{message_id}' for message_id in messages])) + insert = sqlite_insert if db.bind.dialect.name == 'sqlite' else pg_insert + rows = [ + self._build_message(f'{chat_id}-{mid}', chat_id, data.get('user_id') or user_id, data, int(time.time())) + for mid, data in messages.items() + ] + await db.execute( + insert(ChatMessage).on_conflict_do_nothing(index_elements=['id']), + [{column.name: getattr(row, column.name) for column in ChatMessage.__table__.columns} for row in rows], ) - existing_by_id = {row.id: row for row in result.scalars().all()} - - for message_id, data in messages.items(): - composite_id = f'{chat_id}-{message_id}' - message = existing_by_id.get(composite_id) - if message: - self._apply_message_data(message, data, now) - else: - db.add(self._build_message(composite_id, chat_id, user_id, data, now)) - await db.commit() async def get_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatMessageModel]: @@ -358,7 +355,7 @@ class ChatMessageTable: 'created_at': 'timestamp', } # DB-internal columns excluded from the reconstructed message dict. - EXCLUDED_COLUMNS = frozenset({'id', 'chat_id', 'user_id', 'updated_at'}) + EXCLUDED_COLUMNS = frozenset({'id', 'chat_id', 'updated_at'}) async def get_messages_map_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> Optional[dict]: """Build a {message_id: message_dict} map from chat_message rows. @@ -509,13 +506,16 @@ class ChatMessageTable: """Delete specific ``chat_message`` rows by their original message IDs.""" if not message_ids: return True - async with get_async_db_context(db) as db: - await db.execute( + async with nullcontext(db) if db is not None else get_async_db_context() as session: + await session.execute( delete(ChatMessage) .where(ChatMessage.chat_id == chat_id) .where(ChatMessage.id.in_({f'{chat_id}-{mid}' for mid in message_ids})) ) - await db.commit() + if db is None: + await session.commit() + else: + await session.flush() return True # Analytics methods diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index 602a4f4eb4..983600dfa0 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -2,6 +2,7 @@ from __future__ import annotations +from contextlib import asynccontextmanager import logging import re import time @@ -10,7 +11,7 @@ from typing import Any, Literal # local imports from open_webui.env import ENABLE_ADMIN_CHAT_ACCESS -from open_webui.internal.db import Base, JSONField, get_async_db_context +from open_webui.internal.db import Base, JSONField, get_async_db, get_async_db_context from open_webui.models.access_grants import AccessGrants from open_webui.models.automations import AutomationRun from open_webui.models.chat_messages import ChatMessage, ChatMessages @@ -394,6 +395,16 @@ class ChatStatsExport(BaseModel): class ChatTable: + @asynccontextmanager + async def _chat_transaction(self, id: str): + # Own the transaction; callers may already have a read session. + async with get_async_db() as session: + if session.bind.dialect.name == 'sqlite': + await session.execute(text('BEGIN IMMEDIATE')) + chat = await session.get(Chat, id, with_for_update=True) + yield session, chat + await session.commit() + def _clean_null_bytes(self, obj): """Recursively remove null bytes from strings in dict/list structures.""" return sanitize_data_for_db(obj) @@ -663,6 +674,10 @@ class ChatTable: 'last_read_at': int(time.time()), } ) + messages = list((chat.chat.get('history', {}).get('messages') or {}).values()) + messages.extend(chat.chat.get('messages') or []) + for message in messages: + message['user_id'] = user_id return chat async def import_chats( @@ -734,13 +749,7 @@ class ChatTable: ) -> ChatModel | None: """Patch top-level chat keys; history is merged so stale writers don't drop messages.""" try: - async with get_async_db_context(db) as session: - chat_item = await session.get( - Chat, - id, - populate_existing=True, - with_for_update=session.bind.dialect.name == 'postgresql', - ) + async with self._chat_transaction(id) as (session, chat_item): if chat_item is None: return None @@ -749,6 +758,15 @@ class ChatTable: if 'history' in chat: # The caller built its history from an earlier read; merge so messages saved since then survive. updated['history'] = self.merge_history(stored.get('history'), chat['history']) + for mid, message in (stored.get('history', {}).get('messages') or {}).items(): + if (message.get('user_id') or chat_item.user_id) != chat_item.user_id or ( + message.get('role') == 'assistant' + and ( + message.get('done') is False or updated['history']['messages'][mid].get('done') is False + ) + ): + updated['history']['messages'][mid] = message + updated['history'] = self.merge_history(updated['history'], {}) updated = self._clean_null_bytes(updated) chat_item.chat = updated @@ -759,8 +777,16 @@ class ChatTable: if touch: chat_item.updated_at = int(time.time()) - await session.commit() - + for mid in chat.get('history', {}).get('messages') or {}: + message = updated['history']['messages'].get(mid) + if ( + message + and message.get('role') + and message != (stored.get('history', {}).get('messages') or {}).get(mid) + ): + await ChatMessages.upsert_message( + mid, id, message.get('user_id') or chat_item.user_id, message, db=session + ) return ChatModel.model_validate(chat_item) except Exception: return @@ -860,19 +886,13 @@ class ChatTable: async def update_chat_title_by_id(self, id: str, title: str) -> ChatModel | None: try: - async with get_async_db_context() as session: - chat_item = await session.get( - Chat, - id, - populate_existing=True, - with_for_update=session.bind.dialect.name == 'postgresql', - ) + async with self._chat_transaction(id) as (session, chat_item): if chat_item is None: return None clean_title = self._clean_null_bytes(title) chat_item.title = clean_title chat_item.chat = {**(chat_item.chat or {}), 'title': clean_title} - await session.commit() + return ChatModel.model_validate(chat_item) except Exception: return None @@ -999,10 +1019,16 @@ class ChatTable: existing_meta = existing_meta if isinstance(existing_meta, dict) else {} incoming_meta = message.get('meta') if isinstance(incoming_meta, dict): - if set(incoming_meta) == {'voice'}: + if set(incoming_meta) <= {'voice', 'task_id'}: message = {**message, 'meta': {**existing_meta, **incoming_meta}} - elif 'voice' in existing_meta and 'voice' not in incoming_meta: - message = {**message, 'meta': {**incoming_meta, 'voice': existing_meta['voice']}} + else: + message = { + **message, + 'meta': { + **{key: existing_meta[key] for key in ('voice', 'task_id') if key in existing_meta}, + **incoming_meta, + }, + } messages[message_id] = { **messages[message_id], **message, @@ -1059,17 +1085,6 @@ class ChatTable: except Exception as e: log.warning('Backfill failed for chat %s: %s', chat_id, e) - async def reconcile_messages_by_chat_id(self, chat_id: str, user_id: str, messages: dict[str, dict]) -> None: - """Sync ``chat_message`` rows with the committed JSON blob. - - Upserts current messages via ``backfill_messages_by_chat_id``. - Best-effort: errors are logged but never raised. - """ - try: - await self.backfill_messages_by_chat_id(chat_id, user_id, messages) - except Exception as e: - log.warning('Failed to reconcile chat_message rows for chat %s: %s', chat_id, e) - async def get_messages_map_by_chat_id(self, id: str) -> dict | None: """Message map for walking history (see ``get_message_list``). @@ -1163,6 +1178,73 @@ class ChatTable: message = chat.chat.get('history', {}).get('messages', {}).get(message_id, {}) return message.get(metadata_key) + async def insert_chat_turn(self, chat_id, user, user_message, message_ids): + from fastapi import HTTPException + + if not await self.get_accessible_chat_by_id(chat_id, user, permission='write'): + raise HTTPException(403, 'Chat continuation is not allowed.') + async with self._chat_transaction(chat_id) as (db, chat): + if not chat: + raise HTTPException(404, 'Chat not found.') + history = (chat.chat or {}).get('history') or {'messages': {}, 'currentId': None} + messages = history.get('messages') or {} + ids = [user_message.get('id'), *[entry.get('message_id') for entry in message_ids]] + if ( + not all(isinstance(mid, str) and mid for mid in ids) + or len(set(ids)) != len(ids) + or any(mid in messages for mid in ids[1:]) + ): + raise HTTPException(409, 'Message already exists or has an invalid ID.') + existing = messages.get(ids[0]) + if existing and (existing.get('role') != 'user' or (existing.get('user_id') or chat.user_id) != user.id): + raise HTTPException(403, 'You can only regenerate your own messages.') + parent_id = existing.get('parentId') if existing else user_message.get('parentId') + if parent_id is not None and parent_id not in messages: + raise HTTPException(409, 'Parent message no longer exists.') + message = existing or ( + dict(user_message) + if chat.user_id == user.id + else {key: user_message[key] for key in ('content', 'files', 'models') if key in user_message} + ) + if not existing: + message.update( + id=ids[0], + parentId=parent_id, + role='user', + childrenIds=[], + timestamp=int(time.time()), + user_id=user.id, + user={'id': user.id, 'name': user.name}, + ) + turn = {ids[0]: self.upsert_message_to_history(history, ids[0], self._clean_null_bytes(message))} + for entry in message_ids: + mid = entry['message_id'] + turn[mid] = self.upsert_message_to_history( + history, + mid, + { + 'id': mid, + 'parentId': ids[0], + 'childrenIds': [], + 'role': 'assistant', + 'content': '', + 'done': False, + 'model': entry['model_id'], + 'modelIdx': entry.get('modelIdx'), + 'timestamp': int(time.time()), + 'user_id': user.id, + }, + ) + chat.chat = {**(chat.chat or {}), 'history': history} + flag_modified(chat, 'chat') + chat.current_message_id = history['currentId'] + chat.updated_at = int(time.time()) + for mid, message in turn.items(): + await ChatMessages.upsert_message(mid, chat_id, user.id, message, db=db) + if parent_id: + turn[parent_id] = history['messages'][parent_id] + return {'messages': turn, 'currentId': history['currentId']} + async def upsert_message_to_chat_by_id_and_message_id( self, id: str, message_id: str, message: dict, *, touch: bool = True ) -> ChatModel | None: @@ -1175,13 +1257,7 @@ class ChatTable: message_id = self._clean_null_bytes(message_id) try: - async with get_async_db_context() as session: - chat_item = await session.get( - Chat, - id, - populate_existing=True, - with_for_update=session.bind.dialect.name == 'postgresql', - ) + async with self._chat_transaction(id) as (session, chat_item): if chat_item is None: return None @@ -1190,6 +1266,9 @@ class ChatTable: history = chat.get('history', {}) saved_message = self.upsert_message_to_history(history, message_id, message) + await ChatMessages.upsert_message( + message_id, id, saved_message.get('user_id') or chat_item.user_id, saved_message, db=session + ) chat['history'] = history chat_item.chat = chat # chat is a fresh dict when the column was empty chat_item.title = self._clean_null_bytes(chat.get('title', 'New Chat')) @@ -1199,20 +1278,7 @@ class ChatTable: if touch: chat_item.updated_at = int(time.time()) - await session.commit() updated_chat = ChatModel.model_validate(chat_item) - user_id = chat_item.user_id - - # Dual-write to chat_message table - try: - await ChatMessages.upsert_message( - message_id=message_id, - chat_id=id, - user_id=user_id, - data=saved_message, - ) - except Exception as e: - log.warning(f'Failed to write to chat_message table: {e}') return updated_chat except Exception: @@ -1220,13 +1286,7 @@ class ChatTable: async def delete_message_from_chat_by_id_and_message_id(self, id: str, message_id: str) -> ChatModel | None: try: - async with get_async_db_context() as session: - chat_item = await session.get( - Chat, - id, - populate_existing=True, - with_for_update=session.bind.dialect.name == 'postgresql', - ) + async with self._chat_transaction(id) as (session, chat_item): if chat_item is None: return None @@ -1235,27 +1295,23 @@ class ChatTable: history = chat.get('history', {}) deleted_ids = self.delete_message_from_history(history, message_id) + await ChatMessages.delete_message_ids_by_chat_id(id, deleted_ids, db=session) if not deleted_ids: chat_item.chat = chat chat_item.title = self._clean_null_bytes(chat.get('title', 'New Chat')) chat_item.current_message_id = self.get_current_message_id(chat) flag_modified(chat_item, 'chat') - await session.commit() + return ChatModel.model_validate(chat_item) - messages = history.get('messages') or {} chat['history'] = history chat_item.chat = chat chat_item.title = self._clean_null_bytes(chat.get('title', 'New Chat')) chat_item.current_message_id = self.get_current_message_id(chat) flag_modified(chat_item, 'chat') chat_item.updated_at = int(time.time()) - await session.commit() - updated_chat = ChatModel.model_validate(chat_item) - user_id = chat_item.user_id - await self.backfill_messages_by_chat_id(id, user_id, messages) - await ChatMessages.delete_message_ids_by_chat_id(id, deleted_ids) + updated_chat = ChatModel.model_validate(chat_item) return updated_chat except Exception: @@ -1266,13 +1322,7 @@ class ChatTable: ) -> ChatModel | None: try: status = self._clean_null_bytes(status) - async with get_async_db_context() as session: - chat_item = await session.get( - Chat, - id, - populate_existing=True, - with_for_update=session.bind.dialect.name == 'postgresql', - ) + async with self._chat_transaction(id) as (session, chat_item): if chat_item is None: return None @@ -1284,13 +1334,16 @@ class ChatTable: status_history = history['messages'][message_id].get('statusHistory', []) status_history.append(status) history['messages'][message_id]['statusHistory'] = status_history + message = history['messages'][message_id] + await ChatMessages.upsert_message( + message_id, id, message.get('user_id') or chat_item.user_id, message, db=session + ) chat['history'] = history chat_item.chat = chat chat_item.title = self._clean_null_bytes(chat.get('title', 'New Chat')) chat_item.current_message_id = self.get_current_message_id(chat) flag_modified(chat_item, 'chat') - await session.commit() return ChatModel.model_validate(chat_item) except Exception: @@ -1299,13 +1352,7 @@ class ChatTable: async def add_message_files_by_id_and_message_id( self, id: str, message_id: str, files: list[dict] ) -> list[dict] | None: - async with get_async_db_context() as session: - chat_item = await session.get( - Chat, - id, - populate_existing=True, - with_for_update=session.bind.dialect.name == 'postgresql', - ) + async with self._chat_transaction(id) as (session, chat_item): if chat_item is None: return None @@ -1318,15 +1365,21 @@ class ChatTable: message_files = history['messages'][message_id].get('files', []) message_files = message_files + files history['messages'][message_id]['files'] = message_files + message = history['messages'][message_id] + await ChatMessages.upsert_message( + message_id, + id, + message.get('user_id') or chat_item.user_id, + self._clean_null_bytes(message), + db=session, + ) - # Written here rather than through update_chat_by_id: with session sharing off that opens a second - # connection, which then blocks on the lock this one holds. chat['history'] = history chat_item.chat = self._clean_null_bytes(chat) # History was mutated in place, so the new blob compares equal to the loaded one. flag_modified(chat_item, 'chat') chat_item.updated_at = int(time.time()) - await session.commit() + return message_files async def insert_shared_chat_by_chat_id(self, chat_id: str, db: AsyncSession | None = None) -> ChatModel | None: @@ -1722,14 +1775,11 @@ class ChatTable: if chat_item is None: return None - repaired_history = self._repair_chat_current_id(chat_item.chat or {}) - if repaired_history: - chat_item.current_message_id = self.get_current_message_id(chat_item.chat) - flag_modified(chat_item, 'chat') - if self._sanitize_chat_row(chat_item) or repaired_history: - await session.commit() + model = ChatModel.model_validate(chat_item) + model.chat = self._clean_null_bytes(model.chat) + self._repair_chat_current_id(model.chat) + return model - return ChatModel.model_validate(chat_item) except Exception: return None @@ -1764,52 +1814,66 @@ class ChatTable: if not chat: return None - repaired_history = self._repair_chat_current_id(chat.chat or {}) - if repaired_history: - chat.current_message_id = self.get_current_message_id(chat.chat) - flag_modified(chat, 'chat') - if self._sanitize_chat_row(chat) or repaired_history: - await session.commit() + model = ChatModel.model_validate(chat) + model.chat = self._clean_null_bytes(model.chat) + self._repair_chat_current_id(model.chat) + return model - return ChatModel.model_validate(chat) except Exception: return None - async def get_chat_by_id_for_user( + async def get_accessible_chat_by_id( self, id: str, user, db: AsyncSession | None = None, + *, + permission: Literal['read', 'write'] = 'read', + chat: Chat | ChatModel | None = None, ) -> ChatModel | None: - chat = await self.get_chat_by_id_and_user_id(id, user.id, db=db) - if chat: - return chat - - chat = await self.get_chat_by_id(id, db=db) + if user.role not in {'user', 'admin'}: + return None + chat = chat if chat is not None else await self.get_chat_by_id(id, db=db) if not chat: return None - - if user.role == 'admin' and (ENABLE_ADMIN_CHAT_ACCESS or is_internal_chat(chat.meta)): - return chat - - if await AccessGrants.has_access( - user_id=user.id, - resource_type='shared_chat', - resource_id=id, - permission='read', - db=db, + if ( + chat.user_id == user.id + or user.role == 'admin' + and (ENABLE_ADMIN_CHAT_ACCESS or permission == 'read' and is_internal_chat(chat.meta)) ): - return chat + return ChatModel.model_validate(chat) + from open_webui.models.shared_chats import SharedChats + shared = await SharedChats.get_by_chat_id(id, db=db) + if ( + shared + and shared.chat.get('share_mode') == 'continue' + and await AccessGrants.has_access( + user_id=user.id, resource_type='shared_chat', resource_id=id, permission='read', db=db + ) + ): + return ChatModel.model_validate(chat) if chat.folder_id: from open_webui.utils.access_control.folders import has_folder_access folder = await Folders.get_folder_by_id(chat.folder_id, db=db) - if folder and await has_folder_access(user.id, folder, 'read', db): - return chat - + if ( + folder + and (permission == 'read' or (folder.data or {}).get('share_mode') == 'continue') + and await has_folder_access(user.id, folder, 'read', db) + ): + return ChatModel.model_validate(chat) return None + async def filter_task_ids_by_user_id(self, chat, user_id: str, task_ids: list[str]) -> list[str]: + if not task_ids: + return [] + actors = { + (m.get('meta') or {}).get('task_id'): m.get('user_id') or chat.user_id + for m in (chat.chat.get('history', {}).get('messages') or {}).values() + } + return [task_id for task_id in task_ids if actors.get(task_id, chat.user_id) == user_id] + async def is_chat_owner(self, id: str, user_id: str, db: AsyncSession | None = None) -> bool: """ Lightweight ownership check — uses EXISTS subquery instead of loading @@ -2682,12 +2746,12 @@ class ChatTable: return False async def get_shared_chat_ids_by_file_id(self, file_id: str, db: AsyncSession | None = None) -> list[str]: - """Return IDs of chats that contain this file and have an active share link.""" + """Return file-associated chats with a share link or folder audience.""" async with get_async_db_context(db) as session: result = await session.execute( select(Chat.id) .join(ChatFile, Chat.id == ChatFile.chat_id) - .filter(ChatFile.file_id == file_id, Chat.share_id.isnot(None)) + .filter(ChatFile.file_id == file_id, or_(Chat.share_id.isnot(None), Chat.folder_id.isnot(None))) ) return [row[0] for row in result.all()] diff --git a/backend/open_webui/models/shared_chats.py b/backend/open_webui/models/shared_chats.py index a6ceebb8b0..8959d6c771 100644 --- a/backend/open_webui/models/shared_chats.py +++ b/backend/open_webui/models/shared_chats.py @@ -1,7 +1,7 @@ import logging import time import uuid -from typing import Optional +from typing import Literal, Optional from open_webui.internal.db import Base, JSONField, get_async_db_context from pydantic import BaseModel, ConfigDict @@ -10,6 +10,13 @@ from sqlalchemy.ext.asyncio import AsyncSession log = logging.getLogger(__name__) +ChatShareMode = Literal['continue'] | None + + +class ShareChatForm(BaseModel): + share_mode: ChatShareMode = None + + #################### # SharedChat DB Schema #################### @@ -58,7 +65,9 @@ class SharedChatResponse(BaseModel): class SharedChatsTable: - async def create(self, chat_id: str, user_id: str, db: Optional[AsyncSession] = None) -> Optional[SharedChatModel]: + async def create( + self, chat_id: str, user_id: str, db: Optional[AsyncSession] = None, *, share_mode: ChatShareMode = None + ) -> Optional[SharedChatModel]: """ Create a snapshot of the chat for link sharing. Returns the SharedChatModel with the share token as its id. @@ -78,7 +87,7 @@ class SharedChatsTable: chat_id=chat_id, user_id=user_id, title=chat.title, - chat=chat.chat, + chat={**chat.chat, 'share_mode': share_mode}, created_at=now, updated_at=now, ) @@ -88,7 +97,9 @@ class SharedChatsTable: return SharedChatModel.model_validate(shared_chat) - async def update(self, share_id: str, db: Optional[AsyncSession] = None) -> Optional[SharedChatModel]: + async def update( + self, share_id: str, form_data: ShareChatForm | None = None, db: Optional[AsyncSession] = None + ) -> Optional[SharedChatModel]: """ Re-snapshot: update the shared chat with the current state of the original chat. """ @@ -104,13 +115,31 @@ class SharedChatsTable: return None shared_chat.title = chat.title - shared_chat.chat = chat.chat + shared_chat.chat = { + **chat.chat, + 'share_mode': shared_chat.chat.get('share_mode'), + **(form_data.model_dump(exclude_unset=True) if form_data else {}), + } shared_chat.updated_at = int(time.time()) await db.commit() await db.refresh(shared_chat) return SharedChatModel.model_validate(shared_chat) + async def set_share_mode(self, share_id: str, share_mode: ChatShareMode, db: Optional[AsyncSession] = None): + async with get_async_db_context(db) as db: + shared = await db.get(SharedChat, share_id) + if shared: + if share_mode is None and shared.chat.get('share_mode') == 'continue': + from open_webui.models.chats import Chat + + chat = await db.get(Chat, shared.chat_id) + if chat: + shared.chat = chat.chat + shared.title = chat.title + shared.chat = {**shared.chat, 'share_mode': share_mode} + await db.commit() + async def get_by_id(self, share_id: str, db: Optional[AsyncSession] = None) -> Optional[SharedChatModel]: """Get a shared chat by its share token.""" async with get_async_db_context(db) as db: diff --git a/backend/open_webui/routers/chats.py b/backend/open_webui/routers/chats.py index 4a2e88ce62..5dd86fcbdf 100644 --- a/backend/open_webui/routers/chats.py +++ b/backend/open_webui/routers/chats.py @@ -30,7 +30,7 @@ from open_webui.models.chats import ( ) from open_webui.models.config import Config from open_webui.models.folders import Folders -from open_webui.models.shared_chats import SharedChatResponse, SharedChats +from open_webui.models.shared_chats import ChatShareMode, ShareChatForm, SharedChatResponse, SharedChats from open_webui.models.tags import TagModel, Tags from open_webui.socket.main import get_event_emitter from open_webui.tasks import get_response_streams_by_chat_id, has_active_tasks, stop_item_tasks @@ -127,6 +127,20 @@ async def can_read_shared_chat(user, shared, db: AsyncSession) -> bool: ) +def shared_chat_response(chat, user=None): + data = ChatResponse.model_validate(chat, from_attributes=True).model_dump() + if user is None or chat.user_id != user.id: + data['variables'] = {} + for key in ('params', 'tool_servers', 'tool_ids', 'filter_ids', 'variables'): + data['chat'].pop(key, None) + messages = list((data['chat'].get('history', {}).get('messages') or {}).values()) + messages.extend(data['chat'].get('messages') or []) + for message in messages: + if user is None or (message.get('user_id') or chat.user_id) != user.id: + message.pop('meta', None) + return data + + async def add_active_state_to_chat_list( request: Request, chat_list: list[ChatTitleIdResponse] ) -> list[ChatTitleIdResponse]: @@ -1223,9 +1237,20 @@ async def get_shared_chat_by_id( if await is_open_shared_chat(shared, db=db) or ( user is not None and await can_read_shared_chat(user, shared, db=db) ): - chat = await Chats.get_chat_by_share_id(share_id, db=db) + live = ( + shared.chat.get('share_mode') == 'continue' + and user is not None + and await can_read_shared_chat(user, shared, db=db) + ) + chat = ( + await Chats.get_chat_by_id(shared.chat_id, db=db) + if live + else await Chats.get_chat_by_share_id(share_id, db=db) + ) if chat: - return ChatResponse.model_validate(chat, from_attributes=True) + data = shared_chat_response(chat, user) + data['chat']['share_mode'] = 'continue' if live else None + return data raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -1338,14 +1363,17 @@ async def get_chat_by_id( user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session), ): - chat = await Chats.get_chat_by_id_for_user( + chat = await Chats.get_accessible_chat_by_id( id, user, db=db, ) if chat: - data = ChatResponse.model_validate(chat, from_attributes=True).model_dump() + data = shared_chat_response(chat, user) + data['chat']['share_mode'] = ( + 'continue' if await Chats.get_accessible_chat_by_id(id, user, db=db, permission='write', chat=chat) else None + ) data = overlay_response_streams( data, await get_response_streams_by_chat_id(request.app.state.redis, id), @@ -1384,12 +1412,6 @@ async def update_chat_by_id( or chat ) - # Reconcile chat_message rows without inferring deletes from missing IDs. - # Message deletion has its own endpoint below. - messages = ((chat.chat or {}).get('history') or {}).get('messages') or {} - if messages: - await Chats.reconcile_messages_by_chat_id(id, user.id, messages) - await publish_event( request, EVENTS.CHAT_UPDATED, @@ -1874,6 +1896,8 @@ async def clone_shared_chat_by_id( ) chat = await Chats.get_chat_by_share_id(id, db=db) if shared else None + if shared and shared.chat.get('share_mode') == 'continue' and await can_read_shared_chat(user, shared, db=db): + chat = await Chats.get_chat_by_id(shared.chat_id, db=db) # Fallback: admins can also access any chat directly by chat ID if not chat and user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS: @@ -1886,7 +1910,7 @@ async def clone_shared_chat_by_id( ) updated_chat = { - **chat.chat, + **shared_chat_response(chat, user)['chat'], 'originalChatId': chat.id, 'branchPointMessageId': chat.chat['history']['currentId'], 'title': f'Clone of {chat.title}', @@ -1899,9 +1923,9 @@ async def clone_shared_chat_by_id( **{ 'chat': updated_chat, 'meta': chat.meta, - 'variables': chat.variables or {}, + 'variables': {}, 'pinned': chat.pinned, - 'folder_id': chat.folder_id, + 'folder_id': None, } ) ], @@ -1963,6 +1987,7 @@ async def archive_chat_by_id( async def share_chat_by_id( request: Request, id: str, + form_data: ShareChatForm | None = None, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session), ): @@ -1975,7 +2000,7 @@ async def share_chat_by_id( # If a share already exists, re-snapshot it if chat.share_id: - shared = await SharedChats.update(chat.share_id, db=db) + shared = await SharedChats.update(chat.share_id, form_data, db=db) if shared: chat = await Chats.get_chat_by_id(id, db=db) await publish_event( @@ -1988,7 +2013,7 @@ async def share_chat_by_id( return ChatResponse.model_validate(chat, from_attributes=True) # Create a new share - shared = await SharedChats.create(id, user.id, db=db) + shared = await SharedChats.create(id, user.id, db=db, share_mode=form_data.share_mode if form_data else None) if not shared: raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, detail=ERROR_MESSAGES.DEFAULT()) @@ -2041,6 +2066,7 @@ async def delete_shared_chat_by_id( class ChatAccessGrantsForm(BaseModel): access_grants: list[dict] + share_mode: ChatShareMode = None @router.post('/shared/{id}/access/update', response_model=ChatResponse | None) @@ -2072,6 +2098,11 @@ async def update_shared_chat_access_by_id( ) await AccessGrants.set_access_grants('shared_chat', id, form_data.access_grants, db=db) + if 'share_mode' in form_data.model_fields_set and chat.share_id: + await SharedChats.set_share_mode(chat.share_id, form_data.share_mode, db=db) + from open_webui.socket.main import refresh_chat_access + + await refresh_chat_access(id) return ChatResponse.model_validate(chat, from_attributes=True) @@ -2177,7 +2208,7 @@ async def update_chat_folder_id_by_id( @router.get('/{id}/tags', response_model=list[TagModel]) async def get_chat_tags_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): - chat = await Chats.get_chat_by_id_for_user( + chat = await Chats.get_accessible_chat_by_id( id, user, db=db, diff --git a/backend/open_webui/routers/folders.py b/backend/open_webui/routers/folders.py index 448cd7abe5..889b6ef097 100644 --- a/backend/open_webui/routers/folders.py +++ b/backend/open_webui/routers/folders.py @@ -12,6 +12,7 @@ from open_webui.config import UPLOAD_DIR from open_webui.constants import ERROR_MESSAGES from open_webui.events import EVENTS, publish_event from open_webui.internal.db import get_async_session +from open_webui.models.shared_chats import ChatShareMode from open_webui.models.chat_messages import ChatMessages from open_webui.models.config import Config from open_webui.models.chats import Chats @@ -526,6 +527,7 @@ async def update_folder_is_expanded_by_id( class FolderAccessGrantsForm(BaseModel): access_grants: list[dict] + share_mode: ChatShareMode = None @router.post('/{id}/access/update') @@ -561,6 +563,10 @@ async def update_folder_access_by_id( ) await AccessGrants.set_access_grants('folder', id, form_data.access_grants, db=db) + if 'share_mode' in form_data.model_fields_set: + folder = await Folders.update_folder_by_id_and_user_id( + id, folder.user_id, FolderUpdateForm(data={'share_mode': form_data.share_mode}), db=db + ) grants = await AccessGrants.get_grants_by_resource('folder', id, db=db) await publish_event( diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 302e2f6f86..280d655ecb 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -529,7 +529,7 @@ async def user_join(sid, data): 'last_seen_at': int(time.time()), } - SESSION_POOL[sid] = socket_user + SESSION_POOL[sid] = {**socket_user, 'chat_ids': (SESSION_POOL.get(sid) or {}).get('chat_ids', [])} await sio.save_session(sid, {'user': socket_user, 'token': auth['token']}) LOCAL_AUTHENTICATED_SIDS.add(sid) await sio.enter_room(sid, f'user:{user.id}') @@ -544,11 +544,40 @@ async def user_join(sid, data): return {'id': user.id, 'name': user.name} +async def refresh_chat_access(chat_id=None): + # The pool includes sessions on other workers; Socket.IO routes room changes via Redis. + access = {} + for batch in get_session_pool_batches(): + for sid, session in batch: + if not session: + continue + chat_ids = set(session.get('chat_ids') or []) + for cid in list(chat_ids): + if chat_id and cid != chat_id: + continue + key = (cid, session['id']) + if key not in access: + user = await Users.get_user_by_id(session['id']) + access[key] = bool(user and await Chats.get_accessible_chat_by_id(cid, user)) + if not access[key]: + await sio.leave_room(sid, f'chat:{cid}') + chat_ids.discard(cid) + await sio.emit( + 'events', {'chat_id': cid, 'shared': True, 'data': {'type': 'chat:access', 'data': {}}}, to=sid + ) + if chat_ids != set(session.get('chat_ids') or []): + SESSION_POOL[sid] = {**session, 'chat_ids': list(chat_ids)} + + @sio.on('heartbeat') async def heartbeat(sid, data): user = await get_socket_session_user(sid) if user: - SESSION_POOL[sid] = {**user, 'last_seen_at': int(time.time())} + SESSION_POOL[sid] = { + **user, + 'chat_ids': (SESSION_POOL.get(sid) or {}).get('chat_ids', []), + 'last_seen_at': int(time.time()), + } await Users.update_last_active_by_id(user['id']) @@ -671,6 +700,24 @@ async def chat_events(sid, data): event_data = data.get('data', {}) event_type = event_data.get('type') + if event_type in {'join', 'leave'}: + chat_id = data.get('chat_id') + if not isinstance(chat_id, str) or not is_saved_chat_id(chat_id): + return False + session = SESSION_POOL.get(sid) or user + chat_ids = set(session.get('chat_ids') or []) + if event_type == 'leave': + await sio.leave_room(sid, f'chat:{chat_id}') + chat_ids.discard(chat_id) + else: + reader = await Users.get_user_by_id(user['id']) + if not reader or not await Chats.get_accessible_chat_by_id(chat_id, reader): + return False + await sio.enter_room(sid, f'chat:{chat_id}') + chat_ids.add(chat_id) + SESSION_POOL[sid] = {**session, 'chat_ids': list(chat_ids)} + return True + if event_type == 'last_read_at': read_update = await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id']) if not read_update: @@ -1208,7 +1255,11 @@ async def get_event_emitter(request_info, update_db=True): if (request_info.get('chat_id') or '').startswith('channel:'): return await _make_channel_emitter(request_info) + last_shared_emit = 0.0 + output = None + async def __event_emitter__(event_data): + nonlocal last_shared_emit, output user_id = request_info['user_id'] chat_id = request_info['chat_id'] message_id = request_info['message_id'] @@ -1219,8 +1270,7 @@ async def get_event_emitter(request_info, update_db=True): return room = f'user:{user_id}' - # Local rooms are authoritative; Redis may have listeners on another instance. - if WEBSOCKET_MANAGER == 'redis' or room in sio.manager.rooms.get('/', {}): + if event_data.get('type') != 'chat:messages': await sio.emit( 'events', { @@ -1232,6 +1282,55 @@ async def get_event_emitter(request_info, update_db=True): room=room, ) + if not internal and is_saved_chat_id(chat_id): + event_type = event_data.get('type') + shared_event = None + if event_type in { + 'chat:messages', + 'chat:active', + 'status', + 'source', + 'citation', + 'files', + 'embeds', + 'chat:message:error', + 'chat:tasks:cancel', + 'chat:message:follow_ups', + }: + shared_event = event_data + elif event_type in {'chat:completion', 'response:completion'}: + data = event_data.get('data') or {} + if isinstance(data.get('output'), list): + output = copy.deepcopy(data['output']) + elif event_type == 'response:completion': + from open_webui.utils.middleware import handle_responses_streaming_event + + output, _ = handle_responses_streaming_event(data, output or []) + now = time.monotonic() + if not data.get('type', '').endswith('.delta') or now - last_shared_emit >= 0.15: + last_shared_emit = now + payload = { + key: value + for key, value in data.items() + if key in {'done', 'error', 'usage', 'finish_reason', 'content', 'selected_model_id', 'sources'} + } + if output is not None: + payload['output'] = output + shared_event = {'type': 'chat:completion', 'data': payload} + if shared_event: + await sio.emit( + 'events', + { + 'chat_id': chat_id, + 'message_id': message_id, + 'user_id': user_id, + 'shared': True, + 'data': shared_event, + }, + room=f'chat:{chat_id}', + skip_sid=request_info.get('session_id'), + ) + if save_to_chat: event_type = event_data.get('type') diff --git a/backend/open_webui/utils/access_control/files.py b/backend/open_webui/utils/access_control/files.py index 07cfda193b..a30e482215 100644 --- a/backend/open_webui/utils/access_control/files.py +++ b/backend/open_webui/utils/access_control/files.py @@ -85,6 +85,9 @@ async def has_access_to_file( ) if accessible_ids: return True + for chat_id in shared_chat_ids: + if await Chats.get_accessible_chat_by_id(chat_id, user, db=db): + return True # Note attachment JSON is user-controlled, so only the file owner's notes can grant access. if access_type == 'read': diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 15949b204a..17f1d083a0 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -692,7 +692,21 @@ def handle_responses_streaming_event( delta_type = parts[1] delta = data.get('delta', '') - output_index = data.get('output_index', len(current_output) - 1) + output_index = data.get('output_index', max(len(current_output) - 1, 0)) + + # Chat Completions can start an item with a delta, without an added event. + if output_index >= len(current_output): + current_output = list(current_output) + while len(current_output) <= output_index: + current_output.append( + {'type': 'message', 'status': 'in_progress', 'role': 'assistant', 'content': []} + ) + current_output[output_index].update( + id=data.get('item_id'), + type='reasoning' + if delta_type.startswith('reasoning') + else ('function_call' if delta_type == 'function_call_arguments' else 'message'), + ) if current_output and 0 <= output_index < len(current_output): new_output = list(current_output) @@ -1832,10 +1846,10 @@ async def chat_image_generation_handler(request: Request, form_data: dict, extra if not is_saved_chat_id(chat_id): message_list = form_data.get('messages', []) else: - chat = await Chats.get_chat_by_id_and_user_id(chat_id, user.id) + chat = await Chats.get_accessible_chat_by_id(chat_id, user, permission='write') messages_map = chat.chat.get('history', {}).get('messages', {}) - message_id = chat.chat.get('history', {}).get('currentId') + message_id = metadata.get('user_message_id') or chat.chat.get('history', {}).get('currentId') message_list = get_message_list(messages_map, message_id) user_message = get_last_user_message(message_list) diff --git a/backend/open_webui/utils/tool_approval.py b/backend/open_webui/utils/tool_approval.py index 7d5376bb01..a87aa5f535 100644 --- a/backend/open_webui/utils/tool_approval.py +++ b/backend/open_webui/utils/tool_approval.py @@ -5,7 +5,6 @@ from pydantic import BaseModel from sqlalchemy.ext.asyncio import AsyncSession from open_webui.constants import ERROR_MESSAGES -from open_webui.env import ENABLE_ADMIN_CHAT_ACCESS from open_webui.models.chats import Chats from open_webui.models.config import Config from open_webui.models.users import Users @@ -27,8 +26,8 @@ async def resolve_tool_call_output( user, db: AsyncSession | None = None, ) -> dict: - chat = await Chats.get_chat_by_id(chat_id, db=db) - if not chat or (chat.user_id != user.id and not (user.role == 'admin' and ENABLE_ADMIN_CHAT_ACCESS)): + chat = await Chats.get_accessible_chat_by_id(chat_id, user, db=db, permission='write') + if not chat: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail=ERROR_MESSAGES.ACCESS_PROHIBITED, @@ -38,6 +37,9 @@ async def resolve_tool_call_output( if not message: raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) + if (message.get('user_id') or chat.user_id) != user.id: + raise HTTPException(status_code=403, detail=ERROR_MESSAGES.ACCESS_PROHIBITED) + output = message.get('output') or [] if not isinstance(output, list): raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Message has no resolvable output.') @@ -109,7 +111,7 @@ async def resolve_tool_call_output( event_emitter = await get_event_emitter( { - 'user_id': chat.user_id, + 'user_id': user.id, 'chat_id': chat_id, 'message_id': message_id, }, @@ -148,6 +150,8 @@ async def build_tool_approval_resume_payload(chat_id: str, message_id: str, chat chat_data = chat.chat or {} message_meta = assistant_message.get('meta') if isinstance(assistant_message.get('meta'), dict) else {} chat_params = chat_data.get('params') if isinstance(chat_data.get('params'), dict) else {} + if (assistant_message.get('user_id') or chat.user_id) != chat.user_id: + chat_params = {} params = { **chat_params, **(message_meta.get('params') if isinstance(message_meta.get('params'), dict) else {}), @@ -166,7 +170,7 @@ async def build_tool_approval_resume_payload(chat_id: str, message_id: str, chat system_prompt = params.get('system') if not system_prompt: # Mirror the chat UI's system prompt fallback - user = await Users.get_user_by_id(chat.user_id) + user = await Users.get_user_by_id(assistant_message.get('user_id') or chat.user_id) ui_settings = (user.settings.ui if user and user.settings else None) or {} system_prompt = ui_settings.get('system') if system_prompt is None: @@ -180,7 +184,7 @@ async def build_tool_approval_resume_payload(chat_id: str, message_id: str, chat 'model': model_id, 'messages': messages, 'params': params, - 'files': message_meta.get('files') or chat_data.get('files') or None, + 'files': message_meta.get('files') if 'files' in message_meta else chat_data.get('files'), 'filter_ids': message_meta.get('filter_ids') or None, 'tool_ids': message_meta.get('tool_ids') or None, 'skill_ids': message_meta.get('skill_ids') or None, @@ -188,7 +192,9 @@ async def build_tool_approval_resume_payload(chat_id: str, message_id: str, chat 'tool_servers': message_meta.get('tool_servers') or None, 'features': message_meta.get('features') or {}, 'variables': message_meta.get('variables') or {}, - 'chat_variables': chat.variables, + 'chat_variables': chat.variables + if (assistant_message.get('user_id') or chat.user_id) == chat.user_id + else {}, 'session_id': message_meta.get('session_id'), 'chat_id': chat_id, 'id': message_id, diff --git a/src/lib/apis/chats/index.ts b/src/lib/apis/chats/index.ts index e477e17de5..606d87fc9b 100644 --- a/src/lib/apis/chats/index.ts +++ b/src/lib/apis/chats/index.ts @@ -1069,7 +1069,7 @@ export const cloneSharedChatById = async (token: string, id: string) => { return res; }; -export const shareChatById = async (token: string, id: string) => { +export const shareChatById = async (token: string, id: string, shareMode?: 'continue' | null) => { let error = null; const res = await fetch(`${WEBUI_API_BASE_URL}/chats/${id}/share`, { @@ -1078,7 +1078,8 @@ export const shareChatById = async (token: string, id: string) => { Accept: 'application/json', 'Content-Type': 'application/json', ...(token && { authorization: `Bearer ${token}` }) - } + }, + body: JSON.stringify({ share_mode: shareMode }) }) .then(async (res) => { if (!res.ok) throw await res.json(); @@ -1200,7 +1201,12 @@ export const deleteSharedChatById = async (token: string, id: string) => { return res; }; -export const updateChatAccessGrants = async (token: string, id: string, accessGrants: object[]) => { +export const updateChatAccessGrants = async ( + token: string, + id: string, + accessGrants: object[], + shareMode?: 'continue' | null +) => { let error = null; const res = await fetch(`${WEBUI_API_BASE_URL}/chats/shared/${id}/access/update`, { @@ -1211,7 +1217,8 @@ export const updateChatAccessGrants = async (token: string, id: string, accessGr ...(token && { authorization: `Bearer ${token}` }) }, body: JSON.stringify({ - access_grants: accessGrants + access_grants: accessGrants, + share_mode: shareMode }) }) .then(async (res) => { diff --git a/src/lib/apis/folders/index.ts b/src/lib/apis/folders/index.ts index 4e50b6a7f9..e6db988cbc 100644 --- a/src/lib/apis/folders/index.ts +++ b/src/lib/apis/folders/index.ts @@ -263,7 +263,12 @@ export const markFolderChatsReadById = async (token: string, id: string) => { return res; }; -export const updateFolderAccessById = async (token: string, id: string, accessGrants: any[]) => { +export const updateFolderAccessById = async ( + token: string, + id: string, + accessGrants: any[], + shareMode?: 'continue' | null +) => { let error = null; const res = await fetch(`${WEBUI_API_BASE_URL}/folders/${id}/access/update`, { @@ -273,7 +278,7 @@ export const updateFolderAccessById = async (token: string, id: string, accessGr 'Content-Type': 'application/json', authorization: `Bearer ${token}` }, - body: JSON.stringify({ access_grants: accessGrants }) + body: JSON.stringify({ access_grants: accessGrants, share_mode: shareMode }) }) .then(async (res) => { if (!res.ok) throw await res.json(); diff --git a/src/lib/components/chat/Chat.svelte b/src/lib/components/chat/Chat.svelte index 2c482c6422..0883bb983a 100644 --- a/src/lib/components/chat/Chat.svelte +++ b/src/lib/components/chat/Chat.svelte @@ -377,7 +377,7 @@ let generationController = null; let contextCompactionToastId = null; - let chat = null; + let chat: any = null; let tags = []; // Read-only when viewing someone else's chat (e.g. via shared folder access) @@ -474,7 +474,7 @@ } ); - if ($chatId && !$temporaryChatEnabled && !isTemporaryChatId($chatId)) { + if (!readOnly && $chatId && !$temporaryChatEnabled && !isTemporaryChatId($chatId)) { const res = await updateChatById(localStorage.token, $chatId, { params }).catch((err) => { console.error('[tool permissions chat]', err); return null; @@ -485,6 +485,7 @@ if (tool_approval_mode === 'full') { const messages = [...Object.values(history?.messages ?? {})].reverse() as any[]; for (const message of messages) { + if ((message.user_id ?? chat?.user_id ?? $user?.id) !== $user?.id) continue; const output = (Array.isArray(message?.output) ? message.output : []) as any[]; const resultCallIds = new Set( output @@ -566,6 +567,7 @@ ? createMessagesList(chatHistory, chatHistory.currentId) : Object.values(chatHistory.messages); for (const message of [...messages].reverse()) { + if ((message.user_id ?? chat?.user_id ?? $user?.id) !== $user?.id) continue; const pending = getPendingAskUserFromMessage(message); if (pending) return pending; } @@ -733,6 +735,7 @@ const saveChatVariables = async (values) => { chatVariables = { ...chatVariables, ...values }; + if (readOnly) return; if ($chatId && !$temporaryChatEnabled && !isTemporaryChatId($chatId)) { const res = await updateChatById(localStorage.token, $chatId, {}, chatVariables).catch( @@ -1259,8 +1262,90 @@ } }; + let joinedChatId = ''; + let chatRefresh: { id: string; promise: Promise } | null = null; + + const mergeChatMessages = ( + incoming: typeof history, + baseline: Record | null = null + ) => { + const currentId = history.currentId; + const next = incoming?.messages ?? {}; + const follow = + !generating && + !( + taskIds?.length && + Object.values(history.messages).some( + (message: any) => + message.role === 'assistant' && + message.done === false && + (message.user_id ?? chat?.user_id) === $user?.id + ) + ) && + Object.keys(next).some((id) => !history.messages[id]); + for (const [id, message] of Object.entries(next) as [string, any][]) { + const local = history.messages[id]; + const ownActive = + local?.done === false && + (local.user_id ?? chat?.user_id) === $user?.id && + (taskIds?.length || generating); + if (!local || (!ownActive && (!baseline || baseline[id] === local))) { + history.messages[id] = { ...local, ...message }; + } + if (local) { + history.messages[id].childrenIds = [ + ...new Set([...(local.childrenIds ?? []), ...(message.childrenIds ?? [])]) + ]; + } + } + if ((!currentId || follow) && incoming.currentId && next[incoming.currentId]) { + history.currentId = incoming.currentId; + autoScrollToBottom(); + } + history = history; + }; + + const refreshChat = () => { + const id = $chatId; + if (!id || $temporaryChatEnabled) return Promise.resolve(); + if (chatRefresh?.id === id) return chatRefresh.promise; + const baseline = { ...history.messages }; + const promise = (async () => { + try { + const fresh = await getChatById(localStorage.token, id); + if ($chatId !== id || joinedChatId !== id) return; + chat = fresh; + mergeChatMessages(fresh.chat.history, baseline); + } catch (error) { + if ($chatId === id && joinedChatId === id) { + chat = null; + history = { messages: {}, currentId: null }; + if (!embedded) await goto('/'); + } + } finally { + if (chatRefresh?.id === id) chatRefresh = null; + } + })(); + chatRefresh = { id, promise }; + return promise; + }; + + const syncChatRoom = (socket: typeof $socket, id: string, temporary: boolean) => { + const next = socket && id && !temporary && !isTemporaryChatId(id) ? id : ''; + if (next === joinedChatId) return; + if (joinedChatId) + socket?.emit('events:chat', { chat_id: joinedChatId, data: { type: 'leave' } }); + joinedChatId = next; + if (next) + socket?.emit('events:chat', { chat_id: next, data: { type: 'join' } }, (joined: boolean) => { + if (joined && $chatId === next) void refreshChat(); + }); + }; + $: syncChatRoom($socket, $chatId, $temporaryChatEnabled); + const chatEventHandler = async (event, cb) => { console.log(event); + if (event.shared && event.user_id === $user?.id && event.data?.type !== 'chat:messages') return; // A new chat's title can arrive before its id; the response message already exists. if ( @@ -1269,6 +1354,18 @@ ) { await tick(); const type = event?.data?.type ?? null; + if (event.shared) { + if (type === 'chat:messages') { + mergeChatMessages(event.data.data); + return; + } + if (type === 'chat:access' || type === 'chat:active') { + if (type === 'chat:access' || event.data.data?.active === false) await refreshChat(); + return; + } + if (!history.messages[event.message_id]) await refreshChat(); + if ($chatId !== event.chat_id) return; + } if (type === 'chat:reload') { await loadChat(); return; @@ -1277,6 +1374,7 @@ return; } let message = history.messages[event.message_id]; + if (message) message = { ...message }; if (message) { const data = event?.data?.data ?? null; @@ -1308,17 +1406,17 @@ if (type === 'response:completion') { responseCompletionEventHandler(data, message); } else { - await chatCompletionEventHandler(data, message, event.chat_id); + await chatCompletionEventHandler(data, message, event.chat_id, event.shared); } autoScrollToBottom(); } else if (type === 'chat:tasks:cancel') { message.bridgeCancelled = true; bridgeCancellations.get(event.message_id)?.(); - dismissContextCompactionToast(); + if (!event.shared) dismissContextCompactionToast(); if (data?.output) { message.output = data.output; } - if (event.message_id === history.currentId) { + if (!event.shared && event.message_id === history.currentId) { taskIds = null; // Set all response messages to done for (const messageId of history.messages[message.parentId].childrenIds) { @@ -1351,8 +1449,11 @@ }, 100); } else if (type === 'chat:message:error') { const responseCompleted = message.done; - handleOpenAIError(data.error, message); - if (!responseCompleted) { + if (event.shared) { + message.error = data.error; + message.done = true; + } else handleOpenAIError(data.error, message); + if (!event.shared && !responseCompleted) { dismissContextCompactionToast(); if (event.message_id === history.currentId) { await processNextInQueue(event.chat_id); @@ -1583,22 +1684,12 @@ ); const handleSocketConnect = async () => { - // Gate on $chatId, not chatIdProp: chats started from the home page keep an empty chatIdProp - if (!$chatId || $temporaryChatEnabled) { - return; - } - - if (!hasPendingAssistantLeaf()) { - return; - } - - const pendingTaskIds = await getTaskIdsByChatId(localStorage.token, $chatId) - .then((res) => res?.task_ids ?? []) - .catch(() => null); - - if (pendingTaskIds?.length === 0) { - await loadChat(); - } + if (!$chatId || $temporaryChatEnabled) return; + joinedChatId = ''; + syncChatRoom($socket, $chatId, $temporaryChatEnabled); + const id = $chatId; + const tasks = await getTaskIdsByChatId(localStorage.token, id).catch(() => null); + if ($chatId === id && tasks) taskIds = tasks.task_ids?.length ? tasks.task_ids : null; }; onMount(() => { @@ -1701,6 +1792,9 @@ // Clear the selected chat when leaving the chat surface (e.g. navigating // to the admin panel), otherwise the previously-viewed chat stays selected // in the sidebar and deleting/archiving it wrongly navigates away. + if (joinedChatId) + $socket?.emit('events:chat', { chat_id: joinedChatId, data: { type: 'leave' } }); + joinedChatId = ''; chatId.set(''); chatTitle.set(''); @@ -2634,6 +2728,7 @@ for (const message of Object.values(history.messages)) { if ( message?.role === 'assistant' && + (message.user_id ?? chat?.user_id) === $user?.id && !message.done && !messageHasPendingAskUser(message) ) { @@ -2911,6 +3006,7 @@ parentId: userMessageId, childrenIds: [], role: 'assistant', + user_id: $user?.id, content: `[RESPONSE] ${responseMessageId}`, done: true, @@ -3027,10 +3123,13 @@ history = history; }; - const chatCompletionEventHandler = async (data, message, chatId) => { + const chatCompletionEventHandler = async (data, message, chatId, shared = false) => { const { id, done, choices, content, output, sources, selected_model_id, error, usage } = data; - if (error) handleOpenAIError(error, message); + if (error) { + if (shared) message.error = error; + else handleOpenAIError(error, message); + } // Store raw OR-aligned output items from backend if (output) { @@ -3038,12 +3137,13 @@ message.content = getOutputText(output); if ( data.type === 'response.output_text.delta' && + !shared && navigator.vibrate && $settings?.hapticFeedback ) { navigator.vibrate(5); } - dispatchCallOverlayAudio(message); + if (!shared) dispatchCallOverlayAudio(message); } if (sources && !message?.sources) { @@ -3054,7 +3154,7 @@ if (choices[0]?.message?.content) { // Non-stream response message.content += choices[0]?.message?.content; - dispatchCallOverlayAudio(message); + if (!shared) dispatchCallOverlayAudio(message); } else { // Stream response let value = choices[0]?.delta?.content ?? ''; @@ -3063,10 +3163,10 @@ } else { message.content += value; - if (navigator.vibrate && ($settings?.hapticFeedback ?? false)) { + if (!shared && navigator.vibrate && ($settings?.hapticFeedback ?? false)) { navigator.vibrate(5); } - dispatchCallOverlayAudio(message); + if (!shared) dispatchCallOverlayAudio(message); } } } @@ -3075,10 +3175,10 @@ // REALTIME_CHAT_SAVE is disabled message.content = content; - if (navigator.vibrate && ($settings?.hapticFeedback ?? false)) { + if (!shared && navigator.vibrate && ($settings?.hapticFeedback ?? false)) { navigator.vibrate(5); } - dispatchCallOverlayAudio(message); + if (!shared) dispatchCallOverlayAudio(message); } if (selected_model_id) { @@ -3095,6 +3195,7 @@ if (done) { message.done = true; + if (shared) return; if (message.error) { dismissContextCompactionToast(); bridge?.update(); @@ -3178,6 +3279,7 @@ parentId: history.currentId ?? null, childrenIds: [], role: 'user', + user_id: $user?.id, content: inputContent, files: _files.length > 0 ? _files : undefined, timestamp: Math.floor(Date.now() / 1000), // Unix epoch @@ -3538,6 +3640,7 @@ id: responseMessageId, childrenIds: [], role: 'assistant', + user_id: $user?.id, content: '', done: false, model: model.id, @@ -4024,6 +4127,8 @@ const stopResponse = async (processQueue = true, messageId = history.currentId) => { const responseMessage = messageId ? history.messages[messageId] : null; + if (responseMessage && (responseMessage.user_id ?? chat?.user_id ?? $user?.id) !== $user?.id) + return; if (bridge?.connected && responseMessage) responseMessage.bridgeStopping = true; const hasTaskIds = (taskIds?.length ?? 0) > 0; const hasPendingAssistantResponse = @@ -4268,6 +4373,7 @@ }; const saveChatHandler = async (_chatId, history) => { + if (readOnly) return; if ($chatId == _chatId) { if (!$temporaryChatEnabled) { chat = await updateChatById(localStorage.token, _chatId, { @@ -4282,6 +4388,7 @@ }; const saveControls = async () => { + if (readOnly) return; if (!$chatId || $temporaryChatEnabled) return; const loaded = chat?.chat ?? {}; if (equal(params, loaded.params ?? {}) && equal(chatFiles, loaded.files ?? [])) return; @@ -4670,6 +4777,7 @@ chatId={$chatId} user={chatOwner ?? $user} {readOnly} + shareMode={chat?.chat?.share_mode ?? null} bind:history bind:autoScroll bind:prompt @@ -4698,7 +4806,7 @@ - {#if readOnly} + {#if readOnly && chat?.chat?.share_mode !== 'continue'} {#if canClone}
{/each} diff --git a/src/lib/components/chat/Messages/Message.svelte b/src/lib/components/chat/Messages/Message.svelte index 30904ee27d..3e34ec4483 100644 --- a/src/lib/components/chat/Messages/Message.svelte +++ b/src/lib/components/chat/Messages/Message.svelte @@ -43,6 +43,7 @@ export let forkHandler: Function | null = null; export let triggerScroll; export let readOnly = false; + export let shareMode: 'continue' | null = null; export let allowDelete = true; export let compactPreview = false; export let editCodeBlock = true; @@ -67,7 +68,7 @@ {#if history.messages[messageId]} {#if history.messages[messageId].role === 'user'} {:else} {#key messageId} @@ -148,6 +150,7 @@ {editCodeBlock} {topPadding} {onInsertToNote} + {shareMode} /> {/key} {/if} diff --git a/src/lib/components/chat/Messages/MultiResponseMessages.svelte b/src/lib/components/chat/Messages/MultiResponseMessages.svelte index 5977b432fe..e9f6b83e29 100644 --- a/src/lib/components/chat/Messages/MultiResponseMessages.svelte +++ b/src/lib/components/chat/Messages/MultiResponseMessages.svelte @@ -33,6 +33,7 @@ export let isLastMessage; export let readOnly = false; + export let shareMode: 'continue' | null = null; export let allowDelete = true; export let compactPreview = false; export let editCodeBlock = true; @@ -312,6 +313,7 @@ {compactPreview} {topPadding} {onInsertToNote} + {shareMode} /> {/if} {/key} @@ -377,6 +379,7 @@ {editCodeBlock} {topPadding} {onInsertToNote} + {shareMode} /> {/if} {/key} diff --git a/src/lib/components/chat/Messages/ResponseMessage.svelte b/src/lib/components/chat/Messages/ResponseMessage.svelte index 896b040132..1b73c363bd 100644 --- a/src/lib/components/chat/Messages/ResponseMessage.svelte +++ b/src/lib/components/chat/Messages/ResponseMessage.svelte @@ -69,6 +69,7 @@ import { getOutputText, replaceOutputMessageText, type OutputItem } from './structuredOutput'; interface MessageType { + user_id?: string; id: string; model: string; content: string; @@ -172,6 +173,7 @@ export let isLastMessage = true; export let readOnly = false; + export let shareMode: 'continue' | null = null; export let allowDelete = true; export let compactPreview = false; export let editCodeBlock = true; @@ -849,10 +851,11 @@ floatingButtons={message?.done && !readOnly && ($settings?.showFloatingActionButtons ?? true)} - save={!readOnly} + save={(!readOnly && (!message.user_id || message.user_id === $user?.id)) || + (shareMode === 'continue' && message.user_id === $user?.id)} preview={!readOnly} {compactPreview} - {editCodeBlock} + editCodeBlock={!readOnly && editCodeBlock} {topPadding} done={message?.done ?? false} allowEmbeds={!readOnly} @@ -1253,8 +1256,8 @@ {/if} - {#if !readOnly} - {#if !$temporaryChatEnabled && ($config?.features.enable_message_rating ?? true) && ($user?.role === 'admin' || ($user?.permissions?.chat?.rate_response ?? true))} + {#if (!readOnly && (!message.user_id || message.user_id === $user?.id)) || (shareMode === 'continue' && message.user_id === $user?.id)} + {#if !readOnly && !$temporaryChatEnabled && ($config?.features.enable_message_rating ?? true) && ($user?.role === 'admin' || ($user?.permissions?.chat?.rate_response ?? true))} {$i18n.t('and create a new shared link.')} {:else} - {$i18n.t( - "Messages you send after creating your link won't be shared. Your link is private until you choose who can view it." - )} + {$i18n.t('Your link is private until you choose who can view it.')} {/if}
- {#if chat.share_id} -
- -
- {/if} +
+ + {#if accessGrants.length > 0} + + {/if} + +
{#if $config?.features.enable_community_sharing} diff --git a/src/lib/components/layout/Sidebar/Folders/FolderShareModal.svelte b/src/lib/components/layout/Sidebar/Folders/FolderShareModal.svelte index 0f3a362ca4..2af0b876e8 100644 --- a/src/lib/components/layout/Sidebar/Folders/FolderShareModal.svelte +++ b/src/lib/components/layout/Sidebar/Folders/FolderShareModal.svelte @@ -1,6 +1,8 @@ @@ -81,7 +91,21 @@ shareUsers={$user?.role === 'admin' || $user?.permissions?.access_grants?.allow_users} allowGroups={$user?.role === 'admin' || ($user?.permissions?.access_grants?.allow_groups ?? true)} - /> + > + {#if accessGrants.length > 0} + + {/if} +
diff --git a/src/lib/components/workspace/common/AccessControl.svelte b/src/lib/components/workspace/common/AccessControl.svelte index 8e3ee70e12..2b4c79546a 100644 --- a/src/lib/components/workspace/common/AccessControl.svelte +++ b/src/lib/components/workspace/common/AccessControl.svelte @@ -565,6 +565,8 @@ {/if} + + {#if share}
diff --git a/src/routes/+layout.svelte b/src/routes/+layout.svelte index 4f0b5bb06e..ff0dddaffc 100644 --- a/src/routes/+layout.svelte +++ b/src/routes/+layout.svelte @@ -579,6 +579,7 @@ }; const chatEventHandler = async (event, cb) => { + if (event.shared) return; // Answer this session's availability check even when another chat is active. if ( event?.data?.type === 'request:terminal:state' && diff --git a/src/routes/s/[id]/+page.svelte b/src/routes/s/[id]/+page.svelte index b13eb62b12..4132e0a681 100644 --- a/src/routes/s/[id]/+page.svelte +++ b/src/routes/s/[id]/+page.svelte @@ -106,6 +106,11 @@ await chatId.set(shareId); chat = await getChatByShareId(token, shareId).catch(() => null); + if (chat?.chat?.share_mode === 'continue' && chat.id !== shareId) { + await goto(`/c/${chat.id}`, { replaceState: true }); + return; + } + if (chat) { user = token ? await getUserInfoById(token, chat.user_id).catch((error) => {