diff --git a/.github/ISSUE_TEMPLATE/bug_report.yaml b/.github/ISSUE_TEMPLATE/bug_report.yaml index 420633a0f6..ad4a3e3f11 100644 --- a/.github/ISSUE_TEMPLATE/bug_report.yaml +++ b/.github/ISSUE_TEMPLATE/bug_report.yaml @@ -2,7 +2,6 @@ name: Bug Report description: Create a detailed bug report to help us improve Open WebUI. title: 'issue: ' labels: ['bug', 'triage'] -assignees: [] body: - type: markdown attributes: @@ -55,7 +54,7 @@ body: id: open-webui-version attributes: label: Open WebUI Version - description: Specify the version (e.g., v0.6.26) + description: Specify the version (e.g., v0.11.0) validations: required: true @@ -63,7 +62,7 @@ body: id: ollama-version attributes: label: Ollama Version (if applicable) - description: Specify the version (e.g., v0.2.0, or v0.1.32-rc1) + description: Specify the version (e.g., v0.32.5, or v0.32.6-rc0) validations: required: false @@ -71,7 +70,7 @@ body: id: operating-system attributes: label: Operating System - description: Specify the OS (e.g., Windows 10, macOS Sonoma, Ubuntu 22.04, Debian 12) + description: Specify the OS (e.g., Windows 11, macOS Tahoe, Ubuntu 26.04, Debian 13) validations: required: true @@ -79,7 +78,7 @@ body: id: browser attributes: label: Browser (if applicable) - description: Specify the browser/version (e.g., Chrome 100.0, Firefox 98.0) + description: Specify the browser/version (e.g., Chrome 151.0, Firefox 153.0.3) validations: required: false @@ -138,11 +137,11 @@ body: placeholder: | Example (include every detail): - 1. Start with a clean Ubuntu 22.04 install. - 2. Install Docker v24.0.5 and start the service. + 1. Start with a clean Ubuntu 26.04 install. + 2. Install Docker v29.7.1 and start the service. 3. Clone the Open WebUI repo (git clone ...). 4. Use the Docker Compose file without modifications. - 5. Open browser Chrome 115.0 in incognito mode. + 5. Open browser Chrome 151.0 in incognito mode. 6. Go to http://localhost:8080 and log in with user "test@example.com". 7. Set the language to "English" and theme to "Dark". 8. Attempt to connect to Ollama at "http://localhost:11434". diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index d764be0d8c..e9e4b5acbc 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -13,12 +13,16 @@ This is to ensure large feature PRs are discussed with the community first, befo **Before submitting, make sure you've checked and filled out the following:** -- [ ] **Linked Issue/Discussion:** This PR references an existing [Issue](https://github.com/open-webui/open-webui/issues) or [Discussion](https://github.com/open-webui/open-webui/discussions) — `Closes #___` / `Relates to #___`. PRs without a linked issue or discussion will be closed without review. +- [ ] **Linked Issue/Discussion:** This PR references an existing [Issue](https://github.com/open-webui/open-webui/issues) or active, substantive [Discussion](https://github.com/open-webui/open-webui/discussions) — `Closes #___` / `Relates to #___`. Creating a discussion only to satisfy this checkbox does not count. - [ ] **Target branch:** The pull request targets the `dev` branch. **PRs targeting `main` will be immediately closed.** - [ ] **Description:** A concise description of the changes is provided below. - [ ] **Changelog:** A changelog entry following [Keep a Changelog](https://keepachangelog.com/) format is included at the bottom. diff --git a/.github/workflows/issue-label.yaml b/.github/workflows/issue-label.yaml index a2a343a022..1a138e2f11 100644 --- a/.github/workflows/issue-label.yaml +++ b/.github/workflows/issue-label.yaml @@ -67,3 +67,73 @@ jobs: issue_number: issue.number, labels: ['bug'] }); + + label-feature-requests: + runs-on: ubuntu-latest + steps: + - name: Add "enhancement" label to unlabeled feature requests + uses: actions/github-script@v7 + with: + script: | + const issue = context.payload.issue; + + if (issue.labels.some((label) => label.name === 'enhancement')) { + return; + } + + // A human (or the bug form) already classified this as a bug; + // do not stack a second, contradictory classification on it. + if (issue.labels.some((label) => label.name === 'bug')) { + return; + } + + const isEdit = context.payload.action === 'edited'; + const titleWasEdited = Boolean(context.payload.changes?.title); + + if (isEdit && !titleWasEdited) { + return; + } + + const title = issue.title ?? ''; + const body = issue.body ?? ''; + + // Feature requests: "feat: ...", "feature: ...", "feature request: ...", + // "enhancement: ...", "enh: ...", "[Feature Request] ..." — the feature + // request form titles every submission "feat: ", so form submissions are + // covered by the same pattern. + const featureLikeTitle = + /^\s*(\[\s*(feat|feature|enhancement|enh)\b[^\]]*\]|(feat|feature( request)?|enhancement|enh)\s*[:/\-])/i.test( + title + ); + + // API/CLI-created issues that reproduce the feature request form structure. + // Only headings distinctive to that form. + const featureFormBody = /###\s*(Proposed Solution|Alternatives Considered)/i.test(body); + + if (!featureLikeTitle && !featureFormBody) { + return; + } + + if (isEdit) { + const events = await github.paginate(github.rest.issues.listEvents, { + owner: context.repo.owner, + repo: context.repo.repo, + issue_number: issue.number, + per_page: 100 + }); + + const enhancementLabelWasRemoved = events.some( + (event) => event.event === 'unlabeled' && event.label?.name === 'enhancement' + ); + + if (enhancementLabelWasRemoved) { + return; + } + } + + await github.rest.issues.addLabels({ + owner: context.repo.owner, + repo: context.repo.repo, + issue_number: issue.number, + labels: ['enhancement'] + }); diff --git a/Dockerfile b/Dockerfile index 9345441984..9ac31300f9 100644 --- a/Dockerfile +++ b/Dockerfile @@ -162,6 +162,7 @@ RUN set -e; \ fi; \ fi; \ mkdir -p /app/backend/data; chown -R $UID:$GID /app/backend/data/; \ + if [ -d /app/backend/data/cache ]; then chmod -R a+rX /app/backend/data/cache; fi; \ rm -rf /var/lib/apt/lists/*; # Optional: PPTX parsing through unstructured may need spaCy's English model. diff --git a/backend/open_webui/config.py b/backend/open_webui/config.py index e62616ce37..75b4b6d284 100644 --- a/backend/open_webui/config.py +++ b/backend/open_webui/config.py @@ -39,7 +39,7 @@ from open_webui.utils.json_codec import JSONCodec async def seed_registered_defaults(): await Config.rename_prefix('rag.web', 'web') - await Config.repair_flattened_dict_configs() + await Config.repair_config_rows() await Config.seed_defaults(DEFAULT_CONFIG) @@ -2097,6 +2097,7 @@ ENABLE_USER_WEBHOOKS = os.getenv('ENABLE_USER_WEBHOOKS', 'False').lower() == 'tr # FastAPI / AnyIO settings THREAD_POOL_SIZE = os.getenv('THREAD_POOL_SIZE', None) +THREAD_POOL_THREAD_NAME_PREFIX = os.getenv('THREAD_POOL_THREAD_NAME_PREFIX', '') if THREAD_POOL_SIZE is not None and isinstance(THREAD_POOL_SIZE, str): try: @@ -2187,6 +2188,8 @@ CONTEXT_COMPACTION_MODEL = os.getenv('CONTEXT_COMPACTION_MODEL', '') ENABLE_CONTEXT_COMPACTION = os.getenv('ENABLE_CONTEXT_COMPACTION', 'False').lower() == 'true' +ENABLE_TOOL_PERMISSIONS = os.getenv('ENABLE_TOOL_PERMISSIONS', 'False').lower() == 'true' + CONTEXT_COMPACTION_TOKEN_THRESHOLD = int(os.getenv('CONTEXT_COMPACTION_TOKEN_THRESHOLD', '80000')) _CONTEXT_COMPACTION_TOKEN_CAP = os.getenv('CONTEXT_COMPACTION_TOKEN_CAP') @@ -3119,6 +3122,7 @@ DEFAULT_CONFIG = { 'chat.context_compaction.token_cap': CONTEXT_COMPACTION_TOKEN_CAP, 'chat.context_compaction.retention_percentage': CONTEXT_COMPACTION_RETENTION_PERCENTAGE, 'chat.context_compaction.prompt_template': CONTEXT_COMPACTION_PROMPT_TEMPLATE, + 'chat.tool_permissions.enable': ENABLE_TOOL_PERMISSIONS, 'task.title.prompt_template': TITLE_GENERATION_PROMPT_TEMPLATE, 'task.tags.prompt_template': TAGS_GENERATION_PROMPT_TEMPLATE, 'task.image.prompt_template': IMAGE_PROMPT_GENERATION_PROMPT_TEMPLATE, diff --git a/backend/open_webui/constants.py b/backend/open_webui/constants.py index a1d1bc73ca..fd61fda228 100644 --- a/backend/open_webui/constants.py +++ b/backend/open_webui/constants.py @@ -118,6 +118,9 @@ class ERROR_MESSAGES(str, Enum): AUTOMATION_TOO_FREQUENT = lambda interval='': f'Schedule too frequent. Minimum interval is {interval} seconds.' AUTOMATION_INVALID_RRULE = lambda err='': f'Invalid RRULE: {err}' AUTOMATION_NO_FUTURE_RUNS = 'RRULE has no future occurrences' + AUTOMATION_COUNT_REQUIRES_DTSTART = ( + 'RRULE with COUNT requires an explicit DTSTART line to anchor the occurrence window' + ) FEATURE_DISABLED = lambda name='': f'{name} is disabled' INPUT_TOO_LONG = lambda size='': f'Input prompt exceeds maximum length of {size}' diff --git a/backend/open_webui/events.py b/backend/open_webui/events.py index 5b23231f41..81b8cc3968 100644 --- a/backend/open_webui/events.py +++ b/backend/open_webui/events.py @@ -411,6 +411,11 @@ class EventDefinitions(BaseModel): description='Retrieval content was processed.', message='Retrieval Content processed', ) + RETRIEVAL_CONTENT_PROCESS_FAILED: EventDefinition = EventDefinition( + name='retrieval.content.process_failed', + description='Retrieval content processing failed.', + message='Retrieval Content process failed', + ) RETRIEVAL_COLLECTION_DELETED: EventDefinition = EventDefinition( name='retrieval.collection.deleted', description='A retrieval collection was deleted.', @@ -666,6 +671,7 @@ NOTIFICATION_EVENTS = ( EVENTS.CHAT_FAILED.name, EVENTS.CHANNEL_MESSAGE.name, EVENTS.CALENDAR_ALERT.name, + EVENTS.RETRIEVAL_CONTENT_PROCESS_FAILED.name, ) diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index fb558cfe63..6cbf0155cb 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -7,6 +7,7 @@ import mimetypes import os import sys import time +from concurrent.futures import ThreadPoolExecutor from contextlib import asynccontextmanager from uuid import uuid4 @@ -64,6 +65,7 @@ from open_webui.config import ( ONEDRIVE_SHAREPOINT_URL, STATIC_DIR, THREAD_POOL_SIZE, + THREAD_POOL_THREAD_NAME_PREFIX, WEBUI_AUTH, WEBUI_NAME, async_reset_config, @@ -238,6 +240,7 @@ from open_webui.utils.logger import start_logger from open_webui.utils.middleware import ( background_tasks_handler, build_chat_response_context, + drain_approved_tool_calls, process_chat_payload, process_chat_response, ) @@ -265,6 +268,11 @@ from open_webui.utils.plugin import install_tool_and_function_dependencies from open_webui.utils.redis import get_redis_client from open_webui.utils.security_headers import SecurityHeadersMiddleware from open_webui.utils.session_pool import cleanup_response, get_session, stream_wrapper +from open_webui.utils.tool_approval import ( + ResolveToolCallForm, + build_tool_approval_resume_payload, + resolve_tool_call_output, +) from open_webui.utils.tools import set_terminal_servers, set_tool_servers if SAFE_MODE: @@ -337,6 +345,16 @@ async def lifespan(app: FastAPI): # This allows sync functions to schedule work on the main loop without blocking health checks app.state.main_loop = asyncio.get_running_loop() + if THREAD_POOL_SIZE and THREAD_POOL_SIZE > 0: + # asyncio offloads bypass AnyIO's limiter, so configure both before the first offload. + anyio.to_thread.current_default_thread_limiter().total_tokens = THREAD_POOL_SIZE + app.state.main_loop.set_default_executor( + ThreadPoolExecutor( + max_workers=THREAD_POOL_SIZE, + thread_name_prefix=THREAD_POOL_THREAD_NAME_PREFIX, + ) + ) + app.state.instance_id = INSTANCE_ID start_logger() @@ -372,16 +390,12 @@ async def lifespan(app: FastAPI): if app.state.redis is not None: app.state.redis_task_command_listener = asyncio.create_task(redis_task_command_listener(app)) - if THREAD_POOL_SIZE and THREAD_POOL_SIZE > 0: - limiter = anyio.to_thread.current_default_thread_limiter() - limiter.total_tokens = THREAD_POOL_SIZE - - asyncio.create_task(periodic_usage_pool_cleanup()) - asyncio.create_task(periodic_session_pool_cleanup()) + app.state.periodic_usage_pool_cleanup = asyncio.create_task(periodic_usage_pool_cleanup()) + app.state.periodic_session_pool_cleanup = asyncio.create_task(periodic_session_pool_cleanup()) from open_webui.utils.automations import scheduler_worker_loop - asyncio.create_task(scheduler_worker_loop(app)) + app.state.scheduler_worker_loop = asyncio.create_task(scheduler_worker_loop(app)) if await Config.get('models.base_models_cache'): try: @@ -462,6 +476,10 @@ async def lifespan(app: FastAPI): if hasattr(app.state, 'redis_task_command_listener'): app.state.redis_task_command_listener.cancel() + app.state.periodic_usage_pool_cleanup.cancel() + app.state.periodic_session_pool_cleanup.cancel() + app.state.scheduler_worker_loop.cancel() + await publish_event(app, EVENTS.SYSTEM_SHUTDOWN_COMPLETED, source='system') @@ -1174,7 +1192,7 @@ async def chat_completion( message_ids = [{'model_id': model_id, 'message_id': form_data.pop('id', None)}] user_message = form_data.pop('user_message', None) or form_data.pop('parent_message', None) - chat_id = form_data.get('chat_id') or '' + 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 @@ -1196,16 +1214,28 @@ async def chat_completion( ): tool_servers = None + automation_id = form_data.pop('automation_id', None) + tool_approval_mode = ( + 'full' + if automation_id or chat_id.startswith('channel:') + else ( + form_data.get('params', {}).get('tool_approval_mode') + if await Config.get('chat.tool_permissions.enable', False) + else 'full' + ) + or 'full' + ) + metadata = { 'user_id': user.id, 'user_agent': request.headers.get('user-agent', '') or '', 'internal': getattr(request.state, 'internal', False) is True, - 'chat_id': form_data.pop('chat_id', None) or '', + 'chat_id': chat_id, 'user_message': user_message, 'user_message_id': user_message.get('id') if user_message else None, 'assistant_message_id': form_data.pop('assistant_message_id', None), 'session_id': form_data.pop('session_id', None), - 'automation_id': form_data.pop('automation_id', None), + 'automation_id': automation_id, 'folder_id': form_data.pop('folder_id', None), 'filter_ids': form_data.pop('filter_ids', []), 'tool_ids': form_data.get('tool_ids', None), @@ -1225,6 +1255,7 @@ async def chat_completion( or model_info_params.get('function_calling') or 'native' ), + 'tool_approval_mode': tool_approval_mode, }, } @@ -1238,8 +1269,8 @@ async def chat_completion( if metadata.get('chat_id') and user: chat_id = metadata['chat_id'] - # Gate channel: branch — caller needs write access on the channel - # and the supplied message_id must belong to that channel. + # Gate channel: branch — caller needs write access on the channel, and the + # supplied message_id must belong to that channel and be the caller's own. if chat_id.startswith('channel:'): channel_id = chat_id.removeprefix('channel:') channel = await Channels.get_channel_by_id(channel_id) @@ -1271,7 +1302,11 @@ async def chat_completion( if not target_message_id: continue target_message = await Messages.get_message_by_id(target_message_id) - if target_message and target_message.channel_id != channel.id: + if target_message and ( + target_message.channel_id != channel.id + # Write access is not authorship — block cross-member edits. + or (user.role != 'admin' and target_message.user_id != user.id) + ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT(), @@ -1528,6 +1563,8 @@ async def chat_completion( 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, @@ -1576,6 +1613,9 @@ async def chat_completion( try: form_data, metadata, events = await process_chat_payload(request, form_data, user, metadata, model) + if await drain_approved_tool_calls(request, form_data, user, model, metadata): + return {'status': True, 'chat_id': metadata.get('chat_id'), 'paused': True} + response = await chat_completion_handler(request, form_data, user) # When the upstream provider returns an error (e.g. HTTP 400 @@ -1817,6 +1857,26 @@ async def chat_completion( generate_chat_completions = chat_completion generate_chat_completion = chat_completion +@app.post('/api/v1/chats/{id}/messages/{message_id}/resolve') +async def resolve_chat_message_tool_call( + request: Request, + id: str, + message_id: str, + form_data: ResolveToolCallForm, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): + resolution = await resolve_tool_call_output(id, message_id, form_data, user, db=db) + payload = await build_tool_approval_resume_payload(id, message_id, chat=resolution['chat']) + result = await chat_completion(request, payload, user) + return { + 'status': True, + 'chat_id': id, + 'message_id': message_id, + **(result if isinstance(result, dict) else {}), + } + + # Expose as app.state so internal callers (e.g. automations) can # use the full pipeline without importing from main.py (avoids circular deps). app.state.CHAT_COMPLETION_HANDLER = chat_completion @@ -2067,6 +2127,7 @@ async def list_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De @app.post('/api/tasks/chat/{chat_id:path}/stop') async def stop_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=Depends(get_verified_user)): socket_id = get_temporary_chat_session_id(chat_id) + chat = None if socket_id: owner_id = get_user_id_from_session_pool(socket_id) if owner_id != user.id and user.role != 'admin': @@ -2076,6 +2137,47 @@ async def stop_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De if chat is None or (chat.user_id != user.id and user.role != 'admin'): 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 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('role') != 'assistant' or message.get('done') is not False: + continue + + output = message.get('output') + if isinstance(output, list): + for item in output: + if item.get('type') == 'function_call' and item.get('status') in { + 'pending', + 'queued', + 'requires_approval', + }: + item['status'] = 'rejected' + item.pop('approved', None) + + await Chats.upsert_message_to_chat_by_id_and_message_id( + chat_id, + message_id, + {'done': True, **({'output': output} if isinstance(output, list) else {})}, + touch=False, + ) + result = { + 'status': True, + 'message': 'Finalized pending approval message.', + } + + event_emitter = await get_event_emitter( + { + 'user_id': chat.user_id, + 'chat_id': chat_id, + 'message_id': message_id, + }, + update_db=False, + ) + if event_emitter: + await event_emitter({'type': 'chat:completion', 'data': {'done': True, 'output': output}}) + await event_emitter({'type': 'chat:tasks:cancel'}) + return result @@ -2134,6 +2236,7 @@ async def get_app_config(request: Request): 'automations.enable', 'notes.enable', 'chat.context_compaction.enable', + 'chat.tool_permissions.enable', 'web.search.enable', 'web.search.confirmation.enable', 'web.search.confirmation.content', @@ -2211,6 +2314,7 @@ async def get_app_config(request: Request): 'enable_automations': config.get('automations.enable'), 'enable_notes': config.get('notes.enable'), 'enable_context_compaction': config.get('chat.context_compaction.enable'), + 'enable_tool_permissions': config.get('chat.tool_permissions.enable'), 'enable_web_search': config.get('web.search.enable'), 'enable_web_search_confirmation': config.get('web.search.confirmation.enable'), 'web_search_confirmation_content': config.get('web.search.confirmation.content'), diff --git a/backend/open_webui/models/access_grants.py b/backend/open_webui/models/access_grants.py index 7cca23546c..03a7ef7202 100644 --- a/backend/open_webui/models/access_grants.py +++ b/backend/open_webui/models/access_grants.py @@ -839,7 +839,8 @@ class AccessGrantsTable: ): """ Filter for items where user has read BUT NOT write access. - Public items are NOT considered read_only. + A public (user:*) read grant counts as read access, so publicly shared + read-only items are listed rather than being reachable only by direct link. Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself, so it remains synchronous. The caller is responsible for executing the query @@ -850,7 +851,6 @@ class AccessGrantsTable: from sqlalchemy import exists as sa_exists - # Has read grant (not public) read_grant_exists = ( select(AccessGrant.id) .where( @@ -858,6 +858,10 @@ class AccessGrantsTable: AccessGrant.resource_id == DocumentModel.id, AccessGrant.permission == 'read', or_( + and_( + AccessGrant.principal_type == 'user', + AccessGrant.principal_id == '*', + ), *( [ and_( @@ -884,7 +888,6 @@ class AccessGrantsTable: .exists() ) - # Does NOT have write grant write_grant_exists = ( select(AccessGrant.id) .where( @@ -892,6 +895,10 @@ class AccessGrantsTable: AccessGrant.resource_id == DocumentModel.id, AccessGrant.permission == 'write', or_( + and_( + AccessGrant.principal_type == 'user', + AccessGrant.principal_id == '*', + ), *( [ and_( @@ -918,21 +925,7 @@ class AccessGrantsTable: .exists() ) - # Is NOT public - public_grant_exists = ( - select(AccessGrant.id) - .where( - AccessGrant.resource_type == resource_type, - AccessGrant.resource_id == DocumentModel.id, - AccessGrant.permission == 'read', - AccessGrant.principal_type == 'user', - AccessGrant.principal_id == '*', - ) - .correlate(DocumentModel) - .exists() - ) - - conditions = [read_grant_exists, ~write_grant_exists, ~public_grant_exists] + conditions = [read_grant_exists, ~write_grant_exists] # Not owner if user_id: diff --git a/backend/open_webui/models/automations.py b/backend/open_webui/models/automations.py index d3c0649455..11a906670b 100644 --- a/backend/open_webui/models/automations.py +++ b/backend/open_webui/models/automations.py @@ -1,9 +1,10 @@ import logging import time -from typing import Optional +from typing import Literal, Optional from uuid import uuid4 from open_webui.internal.db import Base, get_async_db_context +from open_webui.utils.misc import json_text_variants from pydantic import BaseModel, ConfigDict from sqlalchemy import JSON, BigInteger, Boolean, Column, Index, String, Text, cast, delete, func, or_, select, update from sqlalchemy.ext.asyncio import AsyncSession @@ -64,11 +65,17 @@ class AutomationTerminalConfig(BaseModel): cwd: Optional[str] = None +class AutomationTarget(BaseModel): + type: Literal['chat', 'channel'] = 'chat' + channel_id: Optional[str] = None + + class AutomationData(BaseModel): prompt: str model_id: str rrule: str terminal: Optional[AutomationTerminalConfig] = None + target: Optional[AutomationTarget] = None class AutomationModel(BaseModel): @@ -183,12 +190,12 @@ class AutomationTable: stmt = stmt.filter(Automation.folder_id == folder_id) if query: - search = f'%{query}%' - # Search in name and prompt inside JSON data + # Search the name column and the prompt inside the JSON data. + data_text = cast(Automation.data, String) stmt = stmt.filter( or_( - Automation.name.ilike(search), - cast(Automation.data, String).ilike(search), + Automation.name.ilike(f'%{query}%'), + *(data_text.ilike(f'%{variant}%') for variant in json_text_variants(query)), ) ) diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index 1a82c9f523..258d5b7cc6 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -902,8 +902,10 @@ class ChatTable: if current_id is None else messages.get(current_id, {}).get('childrenIds', []) ) - while child_ids: + visited_ids = set() + while child_ids and child_ids[-1] not in visited_ids: current_id = child_ids[-1] + visited_ids.add(current_id) child_ids = messages.get(current_id, {}).get('childrenIds', []) history['currentId'] = current_id if current_id in messages else None return deleted_ids @@ -1035,6 +1037,10 @@ class ChatTable: return history_messages async def get_message_by_id_and_message_id(self, id: str, message_id: str) -> dict | None: + messages_map = await ChatMessages.get_messages_map_by_chat_id(id) + if messages_map and message_id in messages_map: + return messages_map[message_id] + chat = await self.get_chat_by_id(id) if chat is None: return None diff --git a/backend/open_webui/models/config.py b/backend/open_webui/models/config.py index 011a58d9ee..61b41f37cf 100644 --- a/backend/open_webui/models/config.py +++ b/backend/open_webui/models/config.py @@ -300,14 +300,16 @@ class Config(Base): ) @staticmethod - async def repair_flattened_dict_configs() -> None: - """Reassemble dict config values flattened by the per-key migration.""" + async def repair_config_rows() -> None: + """Repair known legacy config row shapes.""" if not Config.PERSISTENT_ENABLED: return async with get_async_db() as db: repaired_keys: list[str] = [] orphan_keys: list[str] = [] + default_model_keys: list[str] = [] + now = int(time.time()) for config_key, aliases in DICT_CONFIG_KEY_ALIASES.items(): prefixes = (config_key, *aliases) @@ -355,14 +357,26 @@ class Config(Base): if existing: existing.value = repaired - existing.updated_at = int(time.time()) + existing.updated_at = now else: - db.add(Config(key=config_key, value=repaired, updated_at=int(time.time()))) + db.add(Config(key=config_key, value=repaired, updated_at=now)) repaired_keys.append(config_key) if orphan_keys: await db.execute(delete(Config).where(Config.key.in_(orphan_keys))) - if repaired_keys or orphan_keys: + for key in ('ui.default_models', 'ui.default_pinned_models'): + row = await db.get(Config, key) + if not row or not isinstance(row.value, list): + continue + + row.value = ','.join(model_id for model_id in (str(item).strip() for item in row.value) if model_id) + row.updated_at = now + default_model_keys.append(key) + + if repaired_keys or orphan_keys or default_model_keys: await db.commit() - log.info('Repaired flattened dict config rows for %s', ', '.join(repaired_keys)) + if repaired_keys or orphan_keys: + log.info('Repaired flattened dict config rows for %s', ', '.join(repaired_keys)) + if default_model_keys: + log.info('Repaired default model config rows for %s', ', '.join(default_model_keys)) diff --git a/backend/open_webui/models/folders.py b/backend/open_webui/models/folders.py index 11d5dd5427..273867eff1 100644 --- a/backend/open_webui/models/folders.py +++ b/backend/open_webui/models/folders.py @@ -230,7 +230,9 @@ class FolderTable: async with get_async_db_context(db) as db: # Check if folder exists result = await db.execute( - select(Folder).filter_by(parent_id=parent_id, user_id=user_id).filter(Folder.name.ilike(name)) + select(Folder) + .filter_by(parent_id=parent_id, user_id=user_id) + .filter(func.lower(Folder.name) == func.lower(name)) ) folder = result.scalars().first() diff --git a/backend/open_webui/models/models.py b/backend/open_webui/models/models.py index 839e47c6d3..eac1812c98 100755 --- a/backend/open_webui/models/models.py +++ b/backend/open_webui/models/models.py @@ -1,6 +1,5 @@ from __future__ import annotations -import json import logging import time from copy import deepcopy @@ -10,6 +9,7 @@ from open_webui.internal.db import Base, JSONField, get_async_db_context from open_webui.models.access_grants import AccessGrantModel, AccessGrants from open_webui.models.groups import Groups from open_webui.models.users import User, UserModel, UserResponse, Users +from open_webui.utils.misc import json_text_variants from open_webui.utils.validate import validate_profile_image_url from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator from sqlalchemy import BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, update @@ -374,20 +374,14 @@ class ModelsTable: tag = filter.get('tag') if tag: - # SQLite stores JSON text via json.dumps(ensure_ascii=True), - # so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB - # stores literal Unicode. Use the right pattern for each. - if db.bind.dialect.name == 'sqlite': - if tag.isascii(): - meta_text = func.lower(cast(Model.meta, String)) - pattern = f'%{json.dumps(tag.lower())}%' - else: - meta_text = cast(Model.meta, String) - pattern = f'%{json.dumps(tag)}%' + if db.bind.dialect.name == 'sqlite' and not tag.isascii(): + # SQLite's LOWER() is ASCII-only, so match non-ASCII tags exact-case. + meta_text = cast(Model.meta, String) + variants = json_text_variants(tag) else: meta_text = func.lower(cast(Model.meta, String)) - pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%' - stmt = stmt.filter(meta_text.like(pattern)) + variants = json_text_variants(tag.lower()) + stmt = stmt.filter(or_(*(meta_text.like(f'%"{variant}"%') for variant in variants))) order_by = filter.get('order_by') direction = filter.get('direction') @@ -605,25 +599,16 @@ class ModelsTable: # Update or insert models for model in models: + model_data = { + **model.model_dump(exclude={'access_grants'}), + 'user_id': user_id, + 'updated_at': int(time.time()), + } + if model.id in existing_ids: - await db.execute( - update(Model) - .filter_by(id=model.id) - .values( - **model.model_dump(exclude={'access_grants'}), - user_id=user_id, - updated_at=int(time.time()), - ) - ) + await db.execute(update(Model).filter_by(id=model.id).values(**model_data)) else: - new_model = Model( - **{ - **model.model_dump(exclude={'access_grants'}), - 'user_id': user_id, - 'updated_at': int(time.time()), - } - ) - db.add(new_model) + db.add(Model(**model_data)) await AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db) # Remove models that are no longer present diff --git a/backend/open_webui/models/prompts.py b/backend/open_webui/models/prompts.py index e985bcc70d..54e67ac2f4 100644 --- a/backend/open_webui/models/prompts.py +++ b/backend/open_webui/models/prompts.py @@ -14,7 +14,7 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants from open_webui.models.groups import Groups from open_webui.models.prompt_history import PromptHistories from open_webui.models.users import User, UserModel, UserResponse, Users -from open_webui.utils.json_codec import JSONCodec +from open_webui.utils.misc import json_text_variants from pydantic import BaseModel, ConfigDict, Field from sqlalchemy import JSON, BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, text, update from sqlalchemy.ext.asyncio import AsyncSession @@ -342,9 +342,10 @@ class PromptsTable: 'EXISTS (SELECT 1 FROM json_array_elements_text(prompt.tags) t WHERE LOWER(t) = :tag_val)' ) else: - # Fallback: LIKE on serialised JSON text (ASCII-safe only) - tag_clause = func.lower(cast(Prompt.tags, String)).like( - f'%{JSONCodec.dumps(tag_lower, ensure_ascii=False)}%' + # Fallback for dialects with no JSON array function: LIKE on the text. + tags_text = func.lower(cast(Prompt.tags, String)) + tag_clause = or_( + *(tags_text.like(f'%"{variant}"%') for variant in json_text_variants(tag_lower)) ) tag_lower = None diff --git a/backend/open_webui/models/users.py b/backend/open_webui/models/users.py index 1eff4932d5..abbbe122ea 100644 --- a/backend/open_webui/models/users.py +++ b/backend/open_webui/models/users.py @@ -27,7 +27,6 @@ from sqlalchemy import ( select, update, ) -from sqlalchemy.dialects.postgresql import JSONB from sqlalchemy.ext.asyncio import AsyncSession #################### @@ -360,16 +359,10 @@ class UsersTable: sub: str, db: AsyncSession | None = None, ) -> UserModel | None: - """Look up a user by OAuth provider + subject claim (dialect-aware JSON filter).""" + """Look up a user by OAuth provider + subject claim.""" async with get_async_db_context(db) as session: - dialect = session.bind.dialect.name - query = select(User) - if dialect == 'sqlite': - oauth_match = User.oauth.contains({provider: {'sub': sub}}) - query = query.where(oauth_match) - elif dialect == 'postgresql': - oauth_match = User.oauth[provider].cast(JSONB)['sub'].astext == sub - query = query.where(oauth_match) + # Subscript, never contains(): on a JSON column contains() degrades to a substring LIKE. + query = select(User).where(User.oauth[provider]['sub'].as_string() == sub) row = (await session.execute(query)).scalars().first() return UserModel.model_validate(row) if row else None @@ -379,16 +372,10 @@ class UsersTable: external_id: str, db: AsyncSession | None = None, ) -> UserModel | None: - """Look up a user by SCIM provider + external ID (dialect-aware JSON filter).""" + """Look up a user by SCIM provider + external ID.""" async with get_async_db_context(db) as session: - dialect = session.bind.dialect.name - query = select(User) - if dialect == 'sqlite': - scim_match = User.scim.contains({provider: {'external_id': external_id}}) - query = query.where(scim_match) - elif dialect == 'postgresql': - scim_match = User.scim[provider].cast(JSONB)['external_id'].astext == external_id - query = query.where(scim_match) + # Subscript, never contains(): on a JSON column contains() degrades to a substring LIKE. + query = select(User).where(User.scim[provider]['external_id'].as_string() == external_id) row = (await session.execute(query)).scalars().first() return UserModel.model_validate(row) if row else None diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index 1e159b01a5..ed53c35db5 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -217,7 +217,7 @@ async def get_content_from_url(request, url: str) -> str: def _get_content_from_url_sync(request, url: str, loader_config): - from open_webui.retrieval.web.utils import validate_url, _SSRFSafeAdapter + from open_webui.retrieval.web.utils import validate_url, get_ssrf_safe_requests_session # Validate URL before making any request (blocks private IPs, non-HTTP, filter list) validate_url(url) @@ -241,9 +241,7 @@ def _get_content_from_url_sync(request, url: str, loader_config): # cloud-metadata 169.254.169.254) via a public host that redirects internally. try: # Probe through the connect-time SSRF guard; bare requests.get re-resolves (DNS-rebinding gap). - session = requests.Session() - session.mount('http://', _SSRFSafeAdapter()) - session.mount('https://', _SSRFSafeAdapter()) + session = get_ssrf_safe_requests_session() response = session.get(url, stream=True, timeout=30, allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS) response.raise_for_status() content_type = response.headers.get('Content-Type', '') diff --git a/backend/open_webui/retrieval/web/utils.py b/backend/open_webui/retrieval/web/utils.py index f653a60131..6ed1ff447f 100644 --- a/backend/open_webui/retrieval/web/utils.py +++ b/backend/open_webui/retrieval/web/utils.py @@ -1,4 +1,5 @@ import asyncio +import http.cookiejar import ipaddress import logging import socket @@ -11,16 +12,19 @@ from typing import ( Any, AsyncIterator, Dict, + Iterable, Iterator, List, Literal, Optional, Sequence, + Tuple, Union, ) import aiohttp import certifi +import requests import urllib3.connection import urllib3.connectionpool import validators @@ -51,6 +55,7 @@ from open_webui.constants import ERROR_MESSAGES from open_webui.env import ( AIOHTTP_CLIENT_ALLOW_REDIRECTS, AIOHTTP_CLIENT_SESSION_SSL, + AIOHTTP_CLIENT_SSL_CERT_FILE, AIOHTTP_CLIENT_TIMEOUT, USER_AGENT, ) @@ -252,19 +257,61 @@ class _SSRFSafeConnector(aiohttp.TCPConnector): return results -def get_ssrf_safe_session() -> aiohttp.ClientSession: +def get_ssrf_safe_session(trust_env: bool = True, store_cookies: bool = True) -> aiohttp.ClientSession: """A one-off aiohttp session that re-validates every connection via _SSRFSafeConnector, defeating DNS rebinding. Use for validate_url-gated fetches of user-supplied URLs that must not use the shared (rebinding-vulnerable) pool. Use as a context manager so it is closed: ``async with get_ssrf_safe_session() as session: ...``. + + trust_env also enables environment proxies, and proxied traffic bypasses the connect-time + IP check, because the proxy resolves the hostname instead. """ return aiohttp.ClientSession( connector=_SSRFSafeConnector(), timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT), - trust_env=True, + trust_env=trust_env, + cookie_jar=None if store_cookies else aiohttp.DummyCookieJar(), ) +def get_ssrf_safe_requests_session(trust_env: bool = True, store_cookies: bool = True) -> requests.Session: + """The requests counterpart of get_ssrf_safe_session, with the same proxy caveat.""" + session = requests.Session() + session.trust_env = trust_env + if not store_cookies: + session.cookies.set_policy(http.cookiejar.DefaultCookiePolicy(allowed_domains=[])) + session.mount('http://', _SSRFSafeAdapter()) + session.mount('https://', _SSRFSafeAdapter()) + return session + + +# accept-encoding goes because the client must advertise only codecs it can decode, the rest +# because the client derives them from the URL and body it is actually given. content-encoding +# stays: the browser's body is forwarded byte for byte, so its own labelling still applies. +_DROPPED_REQUEST_HEADERS = {'accept-encoding', 'connection', 'content-length', 'host', 'transfer-encoding'} + +# The clients hand us a decoded body, so the sender's framing no longer describes it. +_DROPPED_RESPONSE_HEADERS = {'connection', 'content-encoding', 'content-length', 'transfer-encoding'} + + +def _forwardable_request_headers(headers: Dict[str, str]) -> Dict[str, str]: + return {name: value for name, value in headers.items() if name.lower() not in _DROPPED_REQUEST_HEADERS} + + +def _fulfillable_response_headers(header_pairs: Iterable[Tuple[str, str]]) -> Dict[str, str]: + """Collapse repeated headers the way route.fulfill expects: set-cookie by newline, rest by comma. + + Takes pairs rather than a mapping because reading either client's headers as a mapping loses + duplicate Set-Cookie values, leaving one malformed cookie or one of the two. + """ + collected: Dict[str, List[str]] = {} + for name, value in header_pairs: + name = name.lower() # grouping by the sender's case would split a repeated header + if name not in _DROPPED_RESPONSE_HEADERS: + collected.setdefault(name, []).append(value) + return {name: ('\n' if name == 'set-cookie' else ', ').join(values) for name, values in collected.items()} + + def extract_metadata(soup, url): metadata = {'source': url} if title := soup.find('title'): @@ -416,7 +463,6 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin): def __init__( self, web_paths: Union[str, List[str]], - api_base_url: str, api_key: str, extract_depth: Literal['basic', 'advanced'] = 'basic', continue_on_failure: bool = True, @@ -450,7 +496,6 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin): # Store parameters for creating TavilyLoader instances self.web_paths = web_paths if isinstance(web_paths, list) else [web_paths] - self.api_base_url = api_base_url self.api_key = api_key self.extract_depth = extract_depth self.continue_on_failure = continue_on_failure @@ -597,7 +642,9 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing requests_per_second (Optional[float]): Number of requests per second to limit to. continue_on_failure (bool): If True, continue loading other URLs on failure. headless (bool): If True, the browser will run in headless mode. - proxy (dict): Proxy override settings for the Playwright session. + proxy (dict): Proxy override settings for the Playwright session. Page requests are + issued outside the browser, so they follow the environment proxy via trust_env + rather than this setting. playwright_ws_url (Optional[str]): WebSocket endpoint URI for remote browser connection. playwright_timeout (Optional[int]): Maximum operation time in milliseconds. """ @@ -642,14 +689,49 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing self.trust_env = trust_env self.playwright_timeout = playwright_timeout - def _intercept_navigation_sync(self, route, request=None): - req = request or route.request + def _request_timeout(self) -> float: + # per-hop budget, since page.goto's timeout cannot reach into our own fetch and 0 disables + # it. aiohttp treats it as a total where requests only caps each read, so sync runs looser. + return (self.playwright_timeout or 30000) / 1000 + + def _requests_verify(self) -> Union[bool, str]: + """requests takes a CA path where aiohttp takes the parsed SSLContext. + + A bundle named directly in AIOHTTP_CLIENT_SESSION_SSL reaches us already parsed and + cannot be expressed here, so that form falls back to the global bundle or certifi. + """ + if not self.verify_ssl or AIOHTTP_CLIENT_SESSION_SSL is False: + return False + if AIOHTTP_CLIENT_SESSION_SSL is True: + return True # no usable global CA bundle, so both clients land on certifi + return AIOHTTP_CLIENT_SSL_CERT_FILE or True + + def _intercept_navigation_sync(self, route, session): + req = route.request + + hop_cookies: List[Tuple[str, str]] = [] try: - validate_url(req.url) - resp = route.fetch(max_redirects=0) + headers = _forwardable_request_headers(req.all_headers()) + post_data = req.post_data_buffer + verify, timeout = self._requests_verify(), self._request_timeout() - if 300 <= resp.status < 400: + # The browser would resolve the hostname again, after the check; fetch it ourselves. + def fetch(url): + validate_url(url) + return session.request( + req.method, + url, + headers=headers, + data=post_data, + allow_redirects=False, + verify=verify, + timeout=timeout, + ) + + resp = fetch(req.url) + + if 300 <= resp.status_code < 400: for _ in range(20): if not AIOHTTP_CLIENT_ALLOW_REDIRECTS: route.abort() @@ -659,26 +741,50 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing if not location: break - url = urllib.parse.urljoin(resp.url, location) - validate_url(url) - resp = route.fetch(url=url, max_redirects=0) - if not 300 <= resp.status < 400: + # only the last hop is fulfilled, so carry each hop's cookies to the browser + hop_cookies += [('set-cookie', v) for v in resp.raw.headers.getlist('set-cookie')] + resp = fetch(urllib.parse.urljoin(resp.url, location)) + if not 300 <= resp.status_code < 400: break else: route.abort() return - except Exception: + except Exception as e: + log.debug('Playwright loader could not fetch %s: %s', req.url, e) route.abort() return - route.fulfill(response=resp) + route.fulfill( + status=resp.status_code, + headers=_fulfillable_response_headers(hop_cookies + list(resp.raw.headers.items())), + body=resp.content, + ) - async def _intercept_navigation(self, route, request=None): - req = request or route.request + async def _intercept_navigation(self, route, session): + req = route.request + + hop_cookies: List[Tuple[str, str]] = [] try: - await run_in_threadpool(validate_url, req.url) - resp = await route.fetch(max_redirects=0) + headers = _forwardable_request_headers(await req.all_headers()) + post_data = req.post_data_buffer + + # The browser would resolve the hostname again, after the check; fetch it ourselves. + async def fetch(url): + await run_in_threadpool(validate_url, url) + response = await session.request( + req.method, + url, + headers=headers, + data=post_data, + allow_redirects=False, + ssl=AIOHTTP_CLIENT_SESSION_SSL if self.verify_ssl else False, + timeout=aiohttp.ClientTimeout(total=self._request_timeout()), + ) + # aiohttp only returns the connection to the pool once the body is buffered + return response, await response.read() + + resp, body = await fetch(req.url) if 300 <= resp.status < 400: for _ in range(20): @@ -690,19 +796,24 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing if not location: break - url = urllib.parse.urljoin(resp.url, location) - await run_in_threadpool(validate_url, url) - resp = await route.fetch(url=url, max_redirects=0) + # only the last hop is fulfilled, so carry each hop's cookies to the browser + hop_cookies += [('set-cookie', v) for v in resp.headers.getall('Set-Cookie', [])] + resp, body = await fetch(urllib.parse.urljoin(str(resp.url), location)) if not 300 <= resp.status < 400: break else: await route.abort() return - except Exception: + except Exception as e: + log.debug('Playwright loader could not fetch %s: %s', req.url, e) await route.abort() return - await route.fulfill(response=resp) + await route.fulfill( + status=resp.status, + headers=_fulfillable_response_headers(hop_cookies + list(resp.headers.items())), + body=body, + ) def lazy_load(self) -> Iterator[Document]: """Safely load URLs synchronously with support for remote browser.""" @@ -719,8 +830,12 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing for url in self.urls: try: self._safe_process_url_sync(url) - with browser.new_page(service_workers='block') as page: - page.route('**/*', self._intercept_navigation_sync) + # opened before the page so it outlives any route still in flight at teardown + with ( + get_ssrf_safe_requests_session(self.trust_env, store_cookies=False) as session, + browser.new_page(service_workers='block') as page, + ): + page.route('**/*', lambda route: self._intercept_navigation_sync(route, session)) page.route_web_socket('**/*', lambda ws_route: ws_route.close()) response = page.goto(url, timeout=self.playwright_timeout) if response is None: @@ -750,8 +865,12 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing for url in self.urls: try: await self._safe_process_url(url) - async with await browser.new_page(service_workers='block') as page: - await page.route('**/*', self._intercept_navigation) + # opened before the page so it outlives any route still in flight at teardown + async with ( + get_ssrf_safe_session(self.trust_env, store_cookies=False) as session, + await browser.new_page(service_workers='block') as page, + ): + await page.route('**/*', lambda route: self._intercept_navigation(route, session)) await page.route_web_socket('**/*', lambda ws_route: ws_route.close()) response = await page.goto(url, timeout=self.playwright_timeout) if response is None: diff --git a/backend/open_webui/routers/auths.py b/backend/open_webui/routers/auths.py index 86de689458..63aebf1b60 100644 --- a/backend/open_webui/routers/auths.py +++ b/backend/open_webui/routers/auths.py @@ -1705,4 +1705,20 @@ async def token_exchange( detail='User not found. Please sign in via the web interface first.', ) + user = await oauth_manager.update_user_role_from_oauth( + request=request, + user=user, + user_data=user_data, + provider=provider, + db=db, + ) + if await Config.get('oauth.enable_group_mapping'): + await oauth_manager.update_user_groups( + request=request, + user=user, + user_data=user_data, + default_permissions=await Config.get('user.permissions'), + db=db, + ) + return await create_session_response(request, user, db, source='oauth') diff --git a/backend/open_webui/routers/automations.py b/backend/open_webui/routers/automations.py index e01c5c255b..caf8879e28 100644 --- a/backend/open_webui/routers/automations.py +++ b/backend/open_webui/routers/automations.py @@ -15,6 +15,8 @@ from open_webui.models.automations import ( AutomationRuns, Automations, ) +from open_webui.models.access_grants import AccessGrants, has_public_write_access_grant +from open_webui.models.channels import Channels from open_webui.models.config import Config from open_webui.models.folders import Folders from open_webui.utils.access_control import has_permission @@ -104,6 +106,44 @@ async def check_automation_folder_access(folder_id: Optional[str], user, db: Asy ) +async def check_automation_channel_access(form_data: AutomationForm, user, db: AsyncSession): + target = form_data.data.target + if not target or target.type != 'channel': + return + + if not target.channel_id or not await Config.get('channels.enable'): + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + channel = await Channels.get_channel_by_id(target.channel_id, db=db) + if not channel: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + if user.role == 'admin': + return + if not await has_permission(user.id, 'features.channels', await Config.get('user.permissions')): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.DEFAULT(), + ) + if channel.type in ['group', 'dm']: + allowed = await Channels.is_user_channel_member(channel.id, user.id, db=db) + else: + allowed = has_public_write_access_grant(channel.access_grants) or await AccessGrants.has_access( + user_id=user.id, resource_type='channel', resource_id=channel.id, permission='write', db=db + ) + if not allowed: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=ERROR_MESSAGES.DEFAULT(), + ) + + async def enrich_automation(automation: AutomationModel, db: AsyncSession, tz: str = None) -> AutomationResponse: """Full enrichment for single-item views (includes next_runs computation).""" last_run = await AutomationRuns.get_latest(automation.id, db=db) @@ -174,6 +214,7 @@ async def create_new_automation( ): await check_automations_permission(request, user) await check_automation_folder_access(form_data.folder_id, user, db) + await check_automation_channel_access(form_data, user, db) try: validate_rrule(form_data.data.rrule, tz=user.timezone) except ValueError as e: @@ -232,6 +273,7 @@ async def update_automation_by_id( automation = await Automations.get_by_id(id, db=db) check_automation_access(automation, user) await check_automation_folder_access(form_data.folder_id, user, db) + await check_automation_channel_access(form_data, user, db) try: validate_rrule(form_data.data.rrule, tz=user.timezone) diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index 652cbdf07f..1e474b4468 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -1075,16 +1075,10 @@ async def model_response_handler(request, channel, message, user, db=None): ], ] - # Resolve model config (same helpers automations use) - from open_webui.utils.automations import ( - _resolve_model_features, - _resolve_model_filter_ids, - _resolve_model_tool_ids, - ) + # Resolve model config (same path automations use) + from open_webui.utils.automations import _resolve_model_defaults - tool_ids = _resolve_model_tool_ids(request.app, model_id) - features = await _resolve_model_features(request.app, model_id) - filter_ids = _resolve_model_filter_ids(request.app, model_id) + tool_ids, features, filter_ids, _ = await _resolve_model_defaults(request.app, model_id) # Build full form_data — same shape as frontend POST. # The channel: prefix routes pipeline events to the diff --git a/backend/open_webui/routers/chats.py b/backend/open_webui/routers/chats.py index 636cdd44ca..549c26217b 100644 --- a/backend/open_webui/routers/chats.py +++ b/backend/open_webui/routers/chats.py @@ -1,8 +1,6 @@ from __future__ import annotations -import asyncio import logging -from typing import Optional from uuid import uuid4 from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request, Response, status @@ -36,7 +34,7 @@ from open_webui.models.folders import Folders from open_webui.models.shared_chats import SharedChatResponse, SharedChats from open_webui.models.tags import TagModel, Tags from open_webui.socket.main import get_event_emitter -from open_webui.tasks import has_active_tasks, stop_item_tasks +from open_webui.tasks import get_response_streams_by_chat_id, has_active_tasks, stop_item_tasks from open_webui.utils.access_control import filter_allowed_access_grants, has_permission from open_webui.utils.access_control.folders import has_folder_access, has_folder_write_access from open_webui.utils.auth import bearer_security, get_admin_user, get_current_user, get_verified_user @@ -58,9 +56,36 @@ CHAT_CONFIG_KEYS = { 'CONTEXT_COMPACTION_TOKEN_CAP': 'chat.context_compaction.token_cap', 'CONTEXT_COMPACTION_RETENTION_PERCENTAGE': 'chat.context_compaction.retention_percentage', 'CONTEXT_COMPACTION_PROMPT_TEMPLATE': 'chat.context_compaction.prompt_template', + 'ENABLE_TOOL_PERMISSIONS': 'chat.tool_permissions.enable', } +def overlay_response_streams(chat_data: dict, response_streams: list[dict]) -> dict: + if not response_streams: + return chat_data + + messages = chat_data.get('chat', {}).get('history', {}).get('messages') + if isinstance(messages, dict): + for stream in response_streams: + message_id = stream.get('message_id') + message = messages.get(message_id) + if isinstance(message, dict): + message['content'] = stream.get('content', '') + message['output'] = stream.get('output') or [] + message['done'] = False + + legacy_messages = chat_data.get('chat', {}).get('messages') + if isinstance(legacy_messages, list): + streams_by_message_id = {stream.get('message_id'): stream for stream in response_streams} + for message in legacy_messages: + if isinstance(message, dict) and (stream := streams_by_message_id.get(message.get('id'))): + message['content'] = stream.get('content', '') + message['output'] = stream.get('output') or [] + message['done'] = False + + return chat_data + + async def get_optional_verified_user( request: Request, response: Response, @@ -140,6 +165,7 @@ class ChatConfigForm(BaseModel): CONTEXT_COMPACTION_TOKEN_CAP: int | None = None CONTEXT_COMPACTION_RETENTION_PERCENTAGE: int = 40 CONTEXT_COMPACTION_PROMPT_TEMPLATE: str + ENABLE_TOOL_PERMISSIONS: bool = False class CompactChatForm(BaseModel): @@ -1297,7 +1323,12 @@ async def compact_chat_by_id( @router.get('/{id}', response_model=ChatResponse | None) -async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): +async def get_chat_by_id( + id: str, + request: Request, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if not chat and user.role == 'admin': @@ -1328,6 +1359,10 @@ async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSess if chat: data = ChatResponse.model_validate(chat, from_attributes=True).model_dump() + data = overlay_response_streams( + data, + await get_response_streams_by_chat_id(request.app.state.redis, id), + ) data['context_usage'] = await get_chat_context_usage(chat) return data diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index 811c223aaa..b349eeacd5 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -613,9 +613,12 @@ async def get_file_process_status( id: str, stream: bool = Query(False), user=Depends(get_verified_user), - db: AsyncSession = Depends(get_async_session), ): - file = await Files.get_file_by_id(id, db=db) + # NOTE: We intentionally do NOT use Depends(get_async_session) here. + # Database operations manage their own short-lived sessions internally. + # Holding a session here would keep a connection for the entire stream + # (up to two hours) and exhaust the connection pool under concurrent load. + file = await Files.get_file_by_id(id) if not file: raise HTTPException( @@ -623,16 +626,13 @@ async def get_file_process_status( detail=ERROR_MESSAGES.NOT_FOUND, ) - if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user, db=db): + if file.user_id == user.id or user.role == 'admin' or await has_access_to_file(id, 'read', user): if stream: MAX_FILE_PROCESSING_DURATION = 3600 * 2 async def event_stream(file_id): - # NOTE: We intentionally do NOT capture the request's db session here. - # Each poll creates its own short-lived session to avoid holding a - # connection for hours. A WebSocket push would be more efficient. for _ in range(MAX_FILE_PROCESSING_DURATION): - file_item = await Files.get_file_by_id(file_id) # Creates own session + file_item = await Files.get_file_by_id(file_id) if file_item: data = file_item.model_dump().get('data', {}) status = data.get('status') diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index 625095e872..6c1fcf9fb8 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -438,6 +438,10 @@ async def reindex_knowledge_base_metadata_embeddings( """ knowledge_bases = await Knowledges.get_knowledge_bases() log.info('Reindexing embeddings for %s knowledge bases', len(knowledge_bases)) + try: + await ASYNC_VECTOR_DB_CLIENT.delete_collection(collection_name=KNOWLEDGE_BASES_COLLECTION) + except Exception as e: + log.debug(e) success_count = 0 for kb in knowledge_bases: @@ -641,10 +645,7 @@ async def _count_external_connection_mappings(connection_id: str, db: Optional[A @router.get('/external/connections', response_model=ExternalKnowledgeConnectionListResponse) -async def get_external_knowledge_connections( - user=Depends(get_admin_user), - db: AsyncSession = Depends(get_async_session), -): +async def get_external_knowledge_connections(user=Depends(get_admin_user)): connections = [_sanitize_external_connection(connection) for connection in await _get_external_connections()] return ExternalKnowledgeConnectionListResponse(items=connections, total=len(connections)) @@ -674,7 +675,6 @@ async def create_external_knowledge_connection( async def get_external_knowledge_connection( id: str, user=Depends(get_admin_user), - db: AsyncSession = Depends(get_async_session), ): connection = await _get_external_connection(id) if not connection: @@ -741,7 +741,6 @@ async def delete_external_knowledge_connection( async def test_external_knowledge_connection( id: str, user=Depends(get_admin_user), - db: AsyncSession = Depends(get_async_session), ): connection = await _get_external_connection(id) if not connection: @@ -839,7 +838,6 @@ async def test_external_knowledge_retrieval( id: str, form_data: ExternalKnowledgeRetrieveTestForm, user=Depends(get_admin_user), - db: AsyncSession = Depends(get_async_session), ): connection = await _get_external_connection(id) if not connection: @@ -1244,7 +1242,6 @@ async def get_pending_knowledge_files( id: str, stream: bool = Query(False), user=Depends(get_verified_user), - db: AsyncSession = Depends(get_async_session), ): """Return files that are being processed for this knowledge base but not yet linked. @@ -1257,7 +1254,11 @@ async def get_pending_knowledge_files( When ``stream=true``, returns an SSE stream that polls every 3 seconds and emits the current pending file list. Closes when no files remain. """ - knowledge = await Knowledges.get_knowledge_by_id(id=id, db=db) + # NOTE: We intentionally do NOT use Depends(get_async_session) here. + # Database operations manage their own short-lived sessions internally. + # Holding a session here would keep a connection for the entire stream + # (up to an hour) and exhaust the connection pool under concurrent load. + knowledge = await Knowledges.get_knowledge_by_id(id=id) if not knowledge: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, @@ -1272,7 +1273,6 @@ async def get_pending_knowledge_files( resource_type='knowledge', resource_id=knowledge.id, permission='read', - db=db, ) ): raise HTTPException( @@ -1281,7 +1281,7 @@ async def get_pending_knowledge_files( ) if not stream: - return await Files.get_pending_files_for_knowledge(id, db=db) + return await Files.get_pending_files_for_knowledge(id) async def event_stream(knowledge_id: str): MAX_POLL_DURATION = 3600 # 1 hour max diff --git a/backend/open_webui/routers/memories.py b/backend/open_webui/routers/memories.py index 26c8293d88..7dc4e02ec6 100644 --- a/backend/open_webui/routers/memories.py +++ b/backend/open_webui/routers/memories.py @@ -11,9 +11,10 @@ from open_webui.events import EVENTS, publish_event from open_webui.internal.db import get_async_session from open_webui.models.config import Config from open_webui.models.memories import Memories, MemoryModel +from open_webui.models.users import Users from open_webui.retrieval.vector.async_client import ASYNC_VECTOR_DB_CLIENT from open_webui.utils.access_control import has_permission -from open_webui.utils.auth import get_verified_user +from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.memory import ( clean_memory_content, clean_memory_path, @@ -124,6 +125,61 @@ def _memory_metadata(memory: MemoryModel) -> dict: } +async def reindex_memory_vectors_for_user( + request: Request, + user_id: str, + memories: list[MemoryModel] | None = None, + user=None, +) -> int: + collection_name = f'user-memory-{user_id}' + try: + await ASYNC_VECTOR_DB_CLIENT.delete_collection(collection_name) + except Exception as e: + log.debug(e) + + memories = memories if memories is not None else await Memories.get_memories_by_user_id(user_id) + memories = memories or [] + if not memories: + return 0 + + vectors = await asyncio.gather( + *[ + request.app.state.EMBEDDING_FUNCTION( + memory_vector_text(memory.content, memory.path), + prefix=RAG_EMBEDDING_CONTENT_PREFIX, + user=user, + ) + for memory in memories + ] + ) + + await ASYNC_VECTOR_DB_CLIENT.upsert( + collection_name=collection_name, + items=[ + { + 'id': memory.id, + 'text': memory_vector_text(memory.content, memory.path), + 'vector': vectors[idx], + 'metadata': _memory_metadata(memory), + } + for idx, memory in enumerate(memories) + ], + ) + return len(memories) + + +async def upsert_memory_vectors_or_reindex(request: Request, user, items: list[dict]) -> None: + try: + await ASYNC_VECTOR_DB_CLIENT.upsert(collection_name=f'user-memory-{user.id}', items=items) + except Exception as e: + message = str(e).lower() + if 'dimension' not in message or 'embedding' not in message: + raise + + log.warning('Memory vector dimension mismatch for user %s; reindexing memory vectors.', user.id) + await reindex_memory_vectors_for_user(request, user.id, user=user) + + @router.post('/add', response_model=MemoryModel | None) async def add_memory( request: Request, @@ -152,9 +208,10 @@ async def add_memory( memory_vector_text(memory.content, memory.path), prefix=RAG_EMBEDDING_CONTENT_PREFIX, user=user ) - await ASYNC_VECTOR_DB_CLIENT.upsert( - collection_name=f'user-memory-{user.id}', - items=[ + await upsert_memory_vectors_or_reindex( + request, + user, + [ { 'id': memory.id, 'text': memory_vector_text(memory.content, memory.path), @@ -226,7 +283,7 @@ async def update_memories( response.append(result) if upsert_items: - await ASYNC_VECTOR_DB_CLIENT.upsert(collection_name=f'user-memory-{user.id}', items=upsert_items) + await upsert_memory_vectors_or_reindex(request, user, upsert_items) if delete_ids: await ASYNC_VECTOR_DB_CLIENT.delete(collection_name=f'user-memory-{user.id}', ids=delete_ids) @@ -386,8 +443,42 @@ async def read_memory_path( ############################ -# ResetMemoryFromVectorDB +# ReindexMemoryVectorDB ############################ +@router.post('/reindex') +async def reindex_memories_from_vector_db( + request: Request, + user=Depends(get_admin_user), +): + memories = await Memories.get_memories() + memories = memories or [] + memories_by_user_id = {} + for memory in memories: + memories_by_user_id.setdefault(memory.user_id, []).append(memory) + + users_result = await Users.get_users() + users = users_result.get('users', []) if users_result else [] + total_memories = 0 + + for memory_user in users: + total_memories += await reindex_memory_vectors_for_user( + request, + memory_user.id, + memories=memories_by_user_id.get(memory_user.id, []), + user=memory_user, + ) + + await publish_event( + request, + EVENTS.MEMORY_RESET, + actor=user, + subject_id='all', + subject_type='user', + data={'count': total_memories, 'user_count': len(users), 'reindex': True}, + ) + return {'status': True, 'total_users': len(users), 'total_memories': total_memories} + + @router.post('/reset', response_model=bool) async def reset_memory_from_vector_db( request: Request, @@ -403,32 +494,7 @@ async def reset_memory_from_vector_db( """ await check_memories_permission(user) - await ASYNC_VECTOR_DB_CLIENT.delete_collection(f'user-memory-{user.id}') - - memories = await Memories.get_memories_by_user_id(user.id) - - # Generate vectors in parallel - vectors = await asyncio.gather( - *[ - request.app.state.EMBEDDING_FUNCTION( - memory_vector_text(memory.content, memory.path), prefix=RAG_EMBEDDING_CONTENT_PREFIX, user=user - ) - for memory in memories - ] - ) - - await ASYNC_VECTOR_DB_CLIENT.upsert( - collection_name=f'user-memory-{user.id}', - items=[ - { - 'id': memory.id, - 'text': memory_vector_text(memory.content, memory.path), - 'vector': vectors[idx], - 'metadata': _memory_metadata(memory), - } - for idx, memory in enumerate(memories) - ], - ) + count = await reindex_memory_vectors_for_user(request, user.id, user=user) await publish_event( request, @@ -436,7 +502,7 @@ async def reset_memory_from_vector_db( actor=user, subject_id=user.id, subject_type='user', - data={'count': len(memories)}, + data={'count': count, 'reindex': True}, ) return True @@ -512,9 +578,10 @@ async def update_memory_by_id( memory_vector_text(memory.content, memory.path), prefix=RAG_EMBEDDING_CONTENT_PREFIX, user=user ) - await ASYNC_VECTOR_DB_CLIENT.upsert( - collection_name=f'user-memory-{user.id}', - items=[ + await upsert_memory_vectors_or_reindex( + request, + user, + [ { 'id': memory.id, 'text': memory_vector_text(memory.content, memory.path), diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index 01ba8e380e..dfb01e4eaf 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -471,7 +471,7 @@ async def get_filtered_models(models, user, db=None): @router.get('/api/tags') -@router.get('/api/tags/{url_idx}') +@router.get('/api/tags/{url_idx}', dependencies=[Depends(get_admin_user)]) async def get_ollama_tags( request: Request, url_idx: int | None = None, @@ -534,7 +534,7 @@ async def get_ollama_loaded_models( @router.get('/api/version') -@router.get('/api/version/{url_idx}') +@router.get('/api/version/{url_idx}', dependencies=[Depends(get_admin_user)]) async def get_ollama_versions( request: Request, user=Depends(get_verified_user), @@ -1053,7 +1053,7 @@ class GenerateChatCompletionForm(BaseModel): async def validate_ollama_backend_idx(request: Request, model: str, url_idx: int | None, user) -> None: # A caller-supplied url_idx must point to a backend the model is actually # served from; the None path is already constrained to that allow-list. - if url_idx is None or user is None or getattr(user, 'role', None) == 'admin' or BYPASS_MODEL_ACCESS_CONTROL: + if url_idx is None or user is None or getattr(user, 'role', None) == 'admin': return models = request.app.state.OLLAMA_MODELS if not models or model not in models: @@ -1472,7 +1472,7 @@ async def generate_responses( @router.get('/v1/models') -@router.get('/v1/models/{url_idx}') +@router.get('/v1/models/{url_idx}', dependencies=[Depends(get_admin_user)]) async def get_openai_models( request: Request, url_idx: int | None = None, diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index e5c9f2e7e8..70e37a2123 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -246,9 +246,48 @@ router = APIRouter() LLAMACPP_LOADED_STATES = {'loaded', 'sleeping'} LLAMACPP_UNLOADED_STATES = {'loading', 'unloaded'} +MODEL_MANAGEMENT_ENDPOINTS = { + 'llama.cpp': { + 'list': '/models', + 'download': '/models', + 'delete': '/models', + 'load': '/models/load', + 'unload': '/models/unload', + 'sse': '/models/sse', + }, + 'lmstudio': { + 'list': '/api/v1/models', + 'download': '/api/v1/models/download', + 'download_status': '/api/v1/models/download/status/{job_id}', + 'load': '/api/v1/models/load', + 'unload': '/api/v1/models/unload', + }, +} -def get_llamacpp_model_loaded_state(model: dict, provider: str, manual_model_ids: bool = False) -> bool | None: +def get_model_management_root_url(url: str, provider: str) -> str: + root_url = url.rstrip('/') + if provider in ('llama.cpp', 'lmstudio'): + for suffix in ('/api/v1', '/api/v0', '/v1'): + if root_url.endswith(suffix): + return root_url.removesuffix(suffix) + + return root_url + + +def get_provider_model_loaded_state(model: dict, provider: str, manual_model_ids: bool = False) -> bool | None: + if provider == 'lmstudio': + if model.get('loaded_instances'): + return True + + state = model.get('state') + if state == 'loaded': + return True + if state == 'not-loaded': + return False + + return None + if provider != 'llama.cpp': return None @@ -307,6 +346,111 @@ async def get_openai_connection(idx: int) -> tuple[str, str, dict]: return url, key, api_config +async def clear_openai_model_cache(request: Request): + await get_all_models.cache.clear() + request.app.state.BASE_MODELS = [] + request.app.state.OPENAI_MODELS = {} + models = getattr(request.app.state, 'MODELS', None) + if hasattr(models, 'clear'): + models.clear() + else: + request.app.state.MODELS = {} + + +async def get_model_management_connection(url_idx: int) -> tuple[str, str, dict, str]: + if not await Config.get('openai.enable'): + raise HTTPException(status_code=503, detail='OpenAI API is disabled') + + try: + url, key, api_config = await get_openai_connection(url_idx) + except IndexError: + raise HTTPException(status_code=404, detail='Connection not found') + + provider = api_config.get('provider', '') + if provider not in MODEL_MANAGEMENT_ENDPOINTS: + raise HTTPException( + status_code=400, + detail=f'Provider "{provider or "default"}" does not support model management', + ) + + return get_model_management_root_url(url, provider), key, api_config, provider + + +def get_model_management_path(provider: str, operation: str, path_params: dict | None = None) -> str: + try: + path = MODEL_MANAGEMENT_ENDPOINTS[provider][operation] + except KeyError: + raise HTTPException(status_code=400, detail=f'Provider "{provider}" does not support {operation}') + + return path.format(**(path_params or {})) + + +def get_model_management_payload(provider: str, operation: str, payload: dict | None) -> dict | None: + if provider == 'lmstudio' and operation == 'unload' and payload: + return {'instance_id': payload.get('instance_id') or payload.get('model')} + + return payload + + +async def send_model_management_request( + request: Request, + url_idx: int, + operation: str, + method: str = 'GET', + payload: dict | None = None, + query: dict | None = None, + path_params: dict | None = None, + stream: bool = False, + user: UserModel | None = None, +): + root_url, key, api_config, provider = await get_model_management_connection(url_idx) + path = get_model_management_path(provider, operation, path_params=path_params) + payload = get_model_management_payload(provider, operation, payload) + headers, cookies = await get_headers_and_cookies(request, root_url, key, api_config, user=user) + + response = None + streaming = False + try: + session = await get_session() + response = await session.request( + method, + f'{root_url}{path}', + json=payload, + params=query, + headers=headers, + cookies=cookies, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + timeout=get_client_timeout(stream=stream), + ) + + if not response.ok: + try: + error = await response.json(loads=JSONCodec.loads) + except Exception: + error = await response.text() + raise HTTPException(status_code=response.status, detail=error) + + if stream: + streaming = True + return StreamingResponse( + stream_wrapper(response, passthrough=True), + status_code=response.status, + headers=_clean_proxy_headers(response.headers), + ) + + try: + return await response.json(loads=JSONCodec.loads) + except Exception: + return {'success': True} + except HTTPException: + raise + except Exception as e: + raise HTTPException(status_code=response.status if response else 500, detail=str(e)) + finally: + if not streaming: + await cleanup_response(response) + + async def get_anthropic_token_count_target(request: Request, form_data: dict, user: UserModel): """Resolve the upstream LiteLLM connection for an Anthropic token-count request.""" requested_model = form_data.get('model') @@ -697,7 +841,7 @@ async def get_all_models(request: Request, user: UserModel) -> dict[str, list]: 'urlIdx': idx, } - loaded = get_llamacpp_model_loaded_state( + loaded = get_provider_model_loaded_state( model, provider, manual_model_ids=bool(api_config.get('model_ids')), @@ -717,7 +861,7 @@ async def get_all_models(request: Request, user: UserModel) -> dict[str, list]: @router.get('/models') -@router.get('/models/{url_idx}') +@router.get('/models/{url_idx}', dependencies=[Depends(get_admin_user)]) async def get_models(request: Request, url_idx: int | None = None, user=Depends(get_verified_user)): if not await Config.get('openai.enable'): raise HTTPException(status_code=503, detail='OpenAI API is disabled') @@ -793,6 +937,121 @@ async def get_models(request: Request, url_idx: int | None = None, user=Depends( return models +class ProviderModelOperationForm(BaseModel): + model: str + model_config = ConfigDict(extra='allow') + + +@router.get('/models/{url_idx}/catalog') +async def get_provider_model_catalog(request: Request, url_idx: int, user=Depends(get_admin_user)): + return await send_model_management_request(request, url_idx, 'list', user=user) + + +@router.post('/models/{url_idx}/download') +async def download_provider_model( + request: Request, + url_idx: int, + form_data: ProviderModelOperationForm, + user=Depends(get_admin_user), +): + root_url, _, api_config, provider = await get_model_management_connection(url_idx) + payload = form_data.model_dump(exclude_none=True) + payload['model'] = strip_provider_model_prefix(payload['model'], api_config.get('prefix_id')) + + result = await send_model_management_request(request, url_idx, 'download', 'POST', payload, user=user) + await clear_openai_model_cache(request) + await publish_event( + request, + EVENTS.MODEL_PROVIDER_MODEL_CREATED, + actor=user, + subject_id=payload['model'], + data={'provider': provider, 'url_idx': url_idx, 'base_url': root_url}, + ) + return result + + +@router.get('/models/{url_idx}/download/status/{job_id}') +async def get_provider_model_download_status( + request: Request, + url_idx: int, + job_id: str, + user=Depends(get_admin_user), +): + return await send_model_management_request( + request, + url_idx, + 'download_status', + path_params={'job_id': job_id}, + user=user, + ) + + +@router.post('/models/{url_idx}/load') +async def load_provider_model( + request: Request, + url_idx: int, + form_data: ProviderModelOperationForm, + user=Depends(get_admin_user), +): + _, _, api_config, _ = await get_model_management_connection(url_idx) + payload = form_data.model_dump(exclude_none=True) + payload['model'] = strip_provider_model_prefix(payload['model'], api_config.get('prefix_id')) + + result = await send_model_management_request(request, url_idx, 'load', 'POST', payload, user=user) + await clear_openai_model_cache(request) + return result + + +@router.post('/models/{url_idx}/unload') +async def unload_provider_model( + request: Request, + url_idx: int, + form_data: ProviderModelOperationForm, + user=Depends(get_admin_user), +): + _, _, api_config, _ = await get_model_management_connection(url_idx) + payload = form_data.model_dump(exclude_none=True) + payload['model'] = strip_provider_model_prefix(payload['model'], api_config.get('prefix_id')) + + result = await send_model_management_request(request, url_idx, 'unload', 'POST', payload, user=user) + await clear_openai_model_cache(request) + return result + + +@router.get('/models/{url_idx}/sse') +async def stream_provider_model_events(request: Request, url_idx: int, user=Depends(get_admin_user)): + return await send_model_management_request(request, url_idx, 'sse', stream=True, user=user) + + +@router.delete('/models/{url_idx}') +async def delete_provider_model( + request: Request, + url_idx: int, + model: str, + user=Depends(get_admin_user), +): + root_url, _, api_config, provider = await get_model_management_connection(url_idx) + actual_model = strip_provider_model_prefix(model, api_config.get('prefix_id')) + + result = await send_model_management_request( + request, + url_idx, + 'delete', + 'DELETE', + query={'model': actual_model}, + user=user, + ) + await clear_openai_model_cache(request) + await publish_event( + request, + EVENTS.MODEL_PROVIDER_MODEL_DELETED, + actor=user, + subject_id=actual_model, + data={'provider': provider, 'url_idx': url_idx, 'base_url': root_url}, + ) + return result + + class ConnectionVerificationForm(BaseModel): url: str key: str @@ -1588,8 +1847,6 @@ async def responses( # Enforce per-model access control await check_model_access(user, await Models.get_model_by_id(model_id), BYPASS_MODEL_ACCESS_CONTROL) - body = JSONCodec.dumps(payload) - if model_id: models = request.app.state.OPENAI_MODELS if not models or model_id not in models: @@ -1600,6 +1857,9 @@ async def responses( url, key, api_config = await get_openai_connection(idx) + payload['model'] = strip_provider_model_prefix(payload['model'], api_config.get('prefix_id')) + body = JSONCodec.dumps(payload) + r = None streaming = False diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index f92f30d4f2..9023c8b634 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -490,19 +490,19 @@ async def get_embedding_config(request: Request, user=Depends(get_admin_user)): class OpenAIConfigForm(BaseModel): - url: str - key: str + url: str | None = None + key: str | None = None class OllamaConfigForm(BaseModel): - url: str - key: str + url: str | None = None + key: str | None = None class AzureOpenAIConfigForm(BaseModel): - url: str - key: str - version: str + url: str | None = None + key: str | None = None + version: str | None = None class EmbeddingModelUpdateForm(BaseModel): @@ -544,23 +544,18 @@ async def update_embedding_config(request: Request, form_data: EmbeddingModelUpd config.ENABLE_ASYNC_EMBEDDING = form_data.ENABLE_ASYNC_EMBEDDING config.RAG_EMBEDDING_CONCURRENT_REQUESTS = form_data.RAG_EMBEDDING_CONCURRENT_REQUESTS - if config.RAG_EMBEDDING_ENGINE in [ - 'ollama', - 'openai', - 'azure_openai', - ]: - if form_data.openai_config is not None: - config.RAG_OPENAI_API_BASE_URL = form_data.openai_config.url - config.RAG_OPENAI_API_KEY = form_data.openai_config.key + if config.RAG_EMBEDDING_ENGINE == 'openai' and form_data.openai_config is not None: + config.RAG_OPENAI_API_BASE_URL = form_data.openai_config.url or '' + config.RAG_OPENAI_API_KEY = form_data.openai_config.key or '' - if form_data.ollama_config is not None: - config.RAG_OLLAMA_BASE_URL = form_data.ollama_config.url - config.RAG_OLLAMA_API_KEY = form_data.ollama_config.key + if config.RAG_EMBEDDING_ENGINE == 'ollama' and form_data.ollama_config is not None: + config.RAG_OLLAMA_BASE_URL = form_data.ollama_config.url or '' + config.RAG_OLLAMA_API_KEY = form_data.ollama_config.key or '' - if form_data.azure_openai_config is not None: - config.RAG_AZURE_OPENAI_BASE_URL = form_data.azure_openai_config.url - config.RAG_AZURE_OPENAI_API_KEY = form_data.azure_openai_config.key - config.RAG_AZURE_OPENAI_API_VERSION = form_data.azure_openai_config.version + if config.RAG_EMBEDDING_ENGINE == 'azure_openai' and form_data.azure_openai_config is not None: + config.RAG_AZURE_OPENAI_BASE_URL = form_data.azure_openai_config.url or '' + config.RAG_AZURE_OPENAI_API_KEY = form_data.azure_openai_config.key or '' + config.RAG_AZURE_OPENAI_API_VERSION = form_data.azure_openai_config.version or '' request.app.state.ef = get_ef( config.RAG_EMBEDDING_ENGINE, @@ -1998,7 +1993,7 @@ async def process_file( hash = calculate_sha256_string(text_content) if config.BYPASS_EMBEDDING_AND_RETRIEVAL: - await Files.update_file_data_by_id(file.id, {'status': 'completed'}, db=db) + await Files.update_file_data_by_id(file.id, {'status': 'completed', 'error': None}, db=db) await Files.update_file_hash_by_id(file.id, hash, db=db) await publish_event( request, @@ -2057,7 +2052,7 @@ async def process_file( await Files.update_file_data_by_id( file.id, - {'status': 'completed'}, + {'status': 'completed', 'error': None}, db=session, ) await Files.update_file_hash_by_id(file.id, hash, db=session) @@ -2087,12 +2082,25 @@ async def process_file( async with get_async_db() as session: await Files.update_file_data_by_id( file.id, - {'status': 'failed'}, + {'status': 'failed', 'error': str(e)}, db=session, ) # Clear the hash so the file can be re-uploaded after fixing the issue await Files.update_file_hash_by_id(file.id, None, db=session) + await publish_event( + request, + EVENTS.RETRIEVAL_CONTENT_PROCESS_FAILED, + actor=user, + subject_id=file.id, + subject_type='file', + data={ + 'collection_name': collection_name, + 'filename': file.filename, + 'message': f'{file.filename}: {e}', + }, + ) + if 'No pandoc was found' in str(e): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -2397,6 +2405,7 @@ async def process_web( 'status': True, 'collection_name': collection_name, 'filename': form_data.url, + 'content': content, 'file': { 'data': { 'content': content, diff --git a/backend/open_webui/routers/skills.py b/backend/open_webui/routers/skills.py index 4fefa2cb7a..9c268d8dcc 100644 --- a/backend/open_webui/routers/skills.py +++ b/backend/open_webui/routers/skills.py @@ -1,4 +1,5 @@ import logging +import re from typing import Optional from fastapi import APIRouter, Depends, HTTPException, Request, status @@ -38,6 +39,7 @@ router = APIRouter() @router.get('/', response_model=list[SkillUserResponse]) async def get_skills( request: Request, + query: Optional[str] = None, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session), ): @@ -60,6 +62,10 @@ async def get_skills( ) ] + if query: + q = query.casefold() + skills = [skill for skill in skills if q in (skill.name or '').casefold()] + return skills @@ -175,6 +181,13 @@ async def create_new_skill( form_data.id = form_data.id.lower().replace(' ', '-') + # The id goes into /id/{id}/... paths, so anything outside the slug charset is unreachable once stored. + if not re.fullmatch(r'[a-z0-9_-]+', form_data.id): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.DEFAULT('Invalid skill ID'), + ) + existing = await Skills.get_skill_by_id(form_data.id, db=db) if existing is not None: raise HTTPException( diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index f14b89d6ad..fb44bfc0a7 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -66,6 +66,7 @@ async def get_tool_module(request, tool_id, load_from_db=True): @router.get('/', response_model=list[ToolUserResponse]) async def get_tools( request: Request, + query: Optional[str] = None, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session), ): @@ -165,10 +166,7 @@ async def get_tools( ) ) - if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL: - # Admin can see all tools - return tools - else: + if not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL): user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id, db=db)} filtered_tools = [] for tool in tools: @@ -192,7 +190,13 @@ async def get_tools( db=db, ): filtered_tools.append(tool) - return filtered_tools + tools = filtered_tools + + if query: + q = query.casefold() + tools = [tool for tool in tools if q in (tool.name or '').casefold()] + + return tools ############################ diff --git a/backend/open_webui/routers/users.py b/backend/open_webui/routers/users.py index 0220c141fc..24c2433687 100644 --- a/backend/open_webui/routers/users.py +++ b/backend/open_webui/routers/users.py @@ -234,6 +234,7 @@ class SharingPermissions(BaseModel): public_notes: bool = False folders: bool = False public_chats: bool = False + open_chats: bool = False public_calendars: bool = False @@ -470,7 +471,6 @@ async def get_default_user_permissions_defaults(user=Depends(get_admin_user)): async def get_user_settings_by_session_user( raw: bool = False, user=Depends(get_verified_user), - db: AsyncSession = Depends(get_async_session), ): # user already fetched by get_verified_user — no need to refetch if raw: @@ -507,7 +507,7 @@ async def update_user_settings_by_session_user( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - updated_user_settings = form_data.model_dump() + updated_user_settings = form_data.model_dump(exclude_unset=True) ui_settings = updated_user_settings.get('ui') if ( user.role != 'admin' @@ -569,7 +569,6 @@ async def update_user_settings_by_session_user( async def get_user_status_by_session_user( request: Request, user=Depends(get_verified_user), - db: AsyncSession = Depends(get_async_session), ): if not await Config.get('users.enable_status'): raise HTTPException( @@ -619,7 +618,7 @@ async def update_user_status_by_session_user( @router.get('/user/info', response_model=dict | None) -async def get_user_info_by_session_user(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): +async def get_user_info_by_session_user(user=Depends(get_verified_user)): # user already fetched by get_verified_user — no need to refetch return user.info diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 7cb7904fed..9224501767 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -176,6 +176,7 @@ YDOC_MANAGER = YdocManager( async def periodic_session_pool_cleanup(): """Reap orphaned SESSION_POOL entries that missed heartbeats (e.g. crashed instance).""" retry_delay = random.uniform(WEBSOCKET_REDIS_LOCK_TIMEOUT / 2, WEBSOCKET_REDIS_LOCK_TIMEOUT) + renew_interval = max(WEBSOCKET_REDIS_LOCK_TIMEOUT / 2, 0.5) while True: if not session_aquire_func(): log.debug('Session cleanup lock held by another node. Retrying.') @@ -197,7 +198,21 @@ async def periodic_session_pool_cleanup(): del SESSION_POOL[sid] except KeyError: pass - await asyncio.sleep(SESSION_POOL_TIMEOUT) + + next_cleanup_at = time.monotonic() + SESSION_POOL_TIMEOUT + lock_lost = False + while True: + sleep_for = min(renew_interval, next_cleanup_at - time.monotonic()) + if sleep_for <= 0: + break + await asyncio.sleep(sleep_for) + if not session_renew_func(): + log.warning('Unable to renew session cleanup lock. Retrying cleanup ownership.') + lock_lost = True + break + + if lock_lost: + break finally: session_release_func() @@ -912,7 +927,7 @@ async def _make_channel_emitter(request_info): channel_id = request_info['chat_id'].removeprefix('channel:') message_id = request_info['message_id'] - state = {'last_emit_at': 0.0} + state = {'last_emit_at': 0.0, 'output': []} THROTTLE_INTERVAL = 0.15 # ~6 updates/sec async def _emit_channel_update(content: str, done: bool = False, output: list | None = None): @@ -960,11 +975,23 @@ async def _make_channel_emitter(request_info): if not content and not output and not done: return - now = __import__('time').time() + now = time.time() if done or (now - state['last_emit_at']) >= THROTTLE_INTERVAL: state['last_emit_at'] = now await _emit_channel_update(content, done, output if isinstance(output, list) else None) + elif event_type == 'response:completion': + from open_webui.utils.middleware import handle_responses_streaming_event + + data = event_data.get('data', {}) + state['output'], _ = handle_responses_streaming_event(data, state['output']) + content = get_output_text(state['output']) + + now = time.time() + if content and (now - state['last_emit_at']) >= THROTTLE_INTERVAL: + state['last_emit_at'] = now + await _emit_channel_update(content, False, state['output']) + elif event_type == 'chat:message:error': error = event_data.get('data', {}).get('error', {}) error_content = error.get('content', 'An error occurred') if isinstance(error, dict) else str(error) diff --git a/backend/open_webui/tasks.py b/backend/open_webui/tasks.py index 2e6193e464..81b45ebdcd 100644 --- a/backend/open_webui/tasks.py +++ b/backend/open_webui/tasks.py @@ -13,10 +13,12 @@ log = logging.getLogger(__name__) # A dictionary to keep track of active tasks tasks: dict[str, asyncio.Task] = {} item_tasks = {} +response_streams: dict[str, dict] = {} REDIS_TASKS_KEY = f'{REDIS_KEY_PREFIX}:tasks' REDIS_ITEM_TASKS_KEY = f'{REDIS_KEY_PREFIX}:tasks:item' +REDIS_RESPONSE_STREAMS_KEY = f'{REDIS_KEY_PREFIX}:tasks:response_streams' REDIS_PUBSUB_CHANNEL = f'{REDIS_KEY_PREFIX}:tasks:commands' @@ -55,6 +57,7 @@ async def redis_save_task(redis: Redis, task_id: str, item_id: str | None): async def redis_cleanup_task(redis: Redis, task_id: str, item_id: str | None): pipe = redis.pipeline() pipe.hdel(REDIS_TASKS_KEY, task_id) + pipe.hdel(REDIS_RESPONSE_STREAMS_KEY, task_id) if item_id: pipe.srem(f'{REDIS_ITEM_TASKS_KEY}:{item_id}', task_id) await pipe.execute() @@ -91,6 +94,7 @@ async def cleanup_task(redis, task_id: str, id=None): await redis_cleanup_task(redis, task_id, id) tasks.pop(task_id, None) # Remove the task if it exists + response_streams.pop(task_id, None) # If an ID is provided, remove the task from the item_tasks dictionary if id and task_id in item_tasks.get(id, []): @@ -140,6 +144,63 @@ async def list_task_ids_by_item_id(redis, id): return item_tasks.get(id, []) +async def save_response_stream( + redis, + task_id: str | None, + chat_id: str | None, + message_id: str | None, + content: str, + output: list, +): + if not task_id or not chat_id or not message_id: + return + + data = { + 'chat_id': chat_id, + 'message_id': message_id, + 'content': content, + 'output': output, + } + + if redis: + await redis.hset(REDIS_RESPONSE_STREAMS_KEY, task_id, JSONCodec.dumps(data)) + else: + response_streams[task_id] = data + + +async def get_response_streams_by_chat_id(redis, chat_id: str) -> list[dict]: + task_ids = await list_task_ids_by_item_id(redis, chat_id) + if not task_ids: + return [] + + if redis: + values = await redis.hmget(REDIS_RESPONSE_STREAMS_KEY, task_ids) + streams = [] + for value in values: + if not value: + continue + try: + data = JSONCodec.loads(value) + except Exception: + continue + if data.get('chat_id') == chat_id: + streams.append(data) + return streams + + return [ + stream for task_id in task_ids if (stream := response_streams.get(task_id)) and stream.get('chat_id') == chat_id + ] + + +async def clear_response_stream(redis, task_id: str | None): + if not task_id: + return + if redis: + await redis.hdel(REDIS_RESPONSE_STREAMS_KEY, task_id) + else: + response_streams.pop(task_id, None) + + async def stop_task(redis, task_id: str): """ Cancel a running task and remove it from the global task list. diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index 136d2001dc..68e68d8808 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -103,11 +103,12 @@ async def _emit_note_updated(request: Request, user: dict, note) -> None: async def _has_read_access_to_file( file, - user_id: str, - user_role: str, + user: dict, model_knowledge: Optional[list[dict]] = None, ) -> bool: """Check if a user can read a file via ownership, admin role, model attachment, or access grants.""" + user_id = user.get('id') + user_role = user.get('role', 'user') if file.user_id == user_id or user_role == 'admin': return True if model_knowledge and any(item.get('type') == 'file' and item.get('id') == file.id for item in model_knowledge): @@ -117,7 +118,7 @@ async def _has_read_access_to_file( return await has_access_to_file( file_id=file.id, access_type='read', - user=UserModel(**{'id': user_id, 'role': user_role}), + user=UserModel(**user), ) @@ -498,6 +499,120 @@ async def edit_image( return JSONCodec.dumps({'error': str(e)}) +# ============================================================================= +# USER INPUT TOOLS +# ============================================================================= + + +async def ask_user( + questions: list[dict], + allow_other: bool = True, + timeout_ms: int = 120_000, + __event_call__: callable = None, +) -> str: + """ + Ask the user clarifying questions before continuing. + Use this when the next step depends on user intent, preference, or a tradeoff that cannot be inferred safely. + + :param questions: 1-3 question objects, each with id, header, question, and 2-3 options. Each option needs label and description. + :param allow_other: Whether users may enter a free-form answer instead of choosing one of the options + :param timeout_ms: How long the browser should keep the prompt open before cancelling it + :return: JSON with status and answers keyed by question id + """ + try: + if not isinstance(questions, list) or not 1 <= len(questions) <= 3: + raise ValueError('ask_user requires 1-3 questions.') + + normalized_questions = [] + seen_ids = set() + for index, question in enumerate(questions): + if not isinstance(question, dict): + raise ValueError('Each question must be an object.') + + question_id = str(question.get('id') or '').strip()[:64] + if not question_id: + raise ValueError('Each question requires a non-empty id.') + if question_id in seen_ids: + raise ValueError(f'Duplicate question id: {question_id}') + seen_ids.add(question_id) + + options = question.get('options') + if not isinstance(options, list) or not 2 <= len(options) <= 3: + raise ValueError('Each question requires 2-3 options.') + + normalized_options = [] + for option in options: + if not isinstance(option, dict): + raise ValueError('Each option must be an object.') + + label = str(option.get('label') or '').strip()[:80] + description = str(option.get('description') or '').strip()[:240] + if not label or not description: + raise ValueError('Each option requires a label and description.') + + normalized_options.append( + { + 'label': label, + 'description': description, + } + ) + + question_text = str(question.get('question') or '').strip()[:500] + if not question_text: + raise ValueError('Each question requires question text.') + + normalized_questions.append( + { + 'id': question_id, + 'header': str(question.get('header') or '').strip()[:48] or f'Question {index + 1}', + 'question': question_text, + 'options': normalized_options, + 'allow_other': bool(question.get('allow_other', allow_other)), + } + ) + + if isinstance(timeout_ms, bool) or not isinstance(timeout_ms, int) or not 60_000 <= timeout_ms <= 240_000: + timeout_ms = 120_000 + + if __event_call__ is None: + return JSONCodec.dumps( + { + 'status': 'error', + 'error': 'User input requires an active browser session with WebSocket connection.', + }, + ensure_ascii=False, + ) + + output = await __event_call__( + { + 'type': 'request:user_input', + 'data': { + 'questions': normalized_questions, + 'allow_other': allow_other, + 'timeout_ms': timeout_ms, + }, + } + ) + + if not isinstance(output, dict): + return JSONCodec.dumps({'status': 'error', 'error': 'Invalid user input response.'}, ensure_ascii=False) + if output.get('error'): + return JSONCodec.dumps({'status': 'error', 'error': output.get('error')}, ensure_ascii=False) + if output.get('status') == 'cancelled': + return JSONCodec.dumps({'status': 'cancelled', 'answers': {}}, ensure_ascii=False) + + return JSONCodec.dumps( + { + 'status': 'answered', + 'answers': output.get('answers', {}), + }, + ensure_ascii=False, + ) + except Exception as e: + log.exception(f'ask_user error: {e}') + return JSONCodec.dumps({'status': 'error', 'error': str(e)}, ensure_ascii=False) + + # ============================================================================= # CODE INTERPRETER TOOLS # ============================================================================= @@ -2200,8 +2315,6 @@ async def _get_accessible_chat_files( ) -> list[tuple[dict, object]]: from open_webui.models.files import Files - user_id = user.get('id') - user_role = user.get('role', 'user') accessible = [] seen = set() @@ -2223,7 +2336,7 @@ async def _get_accessible_chat_files( seen.add(fid) file = await Files.get_file_by_id(fid) - if file and await _has_read_access_to_file(file, user_id, user_role): + if file and await _has_read_access_to_file(file, user): accessible.append((normalized, file)) return accessible @@ -2445,10 +2558,7 @@ async def query_chat_files( if not embedding_function and not full_context: return JSONCodec.dumps({'error': 'Embedding function not configured'}) - user_model = UserModel.model_construct( - id=__user__.get('id'), - role=__user__.get('role', 'user'), - ) + user_model = UserModel(**__user__) sources = await get_sources_from_items( request=__request__, items=file_items, @@ -2541,7 +2651,7 @@ async def grep_knowledge_files( # Single file mode — verify access file = await Files.get_file_by_id(file_id) if file: - if not await _has_read_access_to_file(file, user_id, user_role, __model_knowledge__): + if not await _has_read_access_to_file(file, __user__, __model_knowledge__): return JSONCodec.dumps({'error': 'File not found'}) files_to_search.append(file) elif __model_knowledge__: @@ -2663,14 +2773,11 @@ async def view_file( try: from open_webui.models.files import Files - user_id = __user__.get('id') - user_role = __user__.get('role', 'user') - file = await Files.get_file_by_id(file_id) if not file: return JSONCodec.dumps({'error': 'File not found'}) - if not await _has_read_access_to_file(file, user_id, user_role, __model_knowledge__): + if not await _has_read_access_to_file(file, __user__, __model_knowledge__): return JSONCodec.dumps({'error': 'File not found'}) content = '' @@ -3078,7 +3185,7 @@ async def query_knowledge_files( embedding_function = getattr(__request__.app.state, 'EMBEDDING_FUNCTION', None) if not embedding_function: return JSONCodec.dumps({'error': 'Embedding function not configured'}) - user_model = UserModel.model_construct(id=user_id, role=user_role) + user_model = UserModel(**__user__) collection_names = [] external_knowledges = [] @@ -3272,7 +3379,7 @@ async def query_knowledge_bases( embedding_function = getattr(__request__.app.state, 'EMBEDDING_FUNCTION', None) if not embedding_function: return JSONCodec.dumps({'error': 'Embedding function not configured'}) - user_model = UserModel.model_construct(id=user_id, role=__user__.get('role', 'user')) + user_model = UserModel(**__user__) query_embedding = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX, user=user_model) # Min-heap of (distance, knowledge_base_id) - only holds top `count` results @@ -3602,7 +3709,7 @@ async def create_automation( return JSONCodec.dumps({'error': 'User context not available'}) try: - from open_webui.models.automations import AutomationData, AutomationForm, Automations + from open_webui.models.automations import AutomationData, AutomationForm, AutomationTarget, Automations from open_webui.models.users import Users from open_webui.routers.automations import check_automation_limits from open_webui.utils.automations import next_n_runs_ns, next_run_ns, validate_rrule @@ -3644,6 +3751,11 @@ async def create_automation( prompt=prompt, model_id=model_id, rrule=rrule, + target=( + AutomationTarget(type='channel', channel_id=metadata.get('chat_id', '').removeprefix('channel:')) + if metadata.get('chat_id', '').startswith('channel:') + else None + ), ), is_active=True, ) @@ -3657,6 +3769,7 @@ async def create_automation( 'name': automation.name, 'folder_id': automation.folder_id, 'model_id': model_id, + 'target': automation.data.get('target'), 'is_active': automation.is_active, 'next_runs': next_n_runs_ns(rrule, tz=tz), }, @@ -3673,7 +3786,7 @@ async def update_automation( prompt: Optional[str] = None, rrule: Optional[str] = None, model_id: Optional[str] = None, - folder_id: Optional[str] = None, + folder_id: Optional[str] = '', __request__: Request = None, __user__: dict = None, ) -> str: @@ -3684,8 +3797,8 @@ async def update_automation( :param name: New name for the automation (optional) :param prompt: New prompt/instructions (optional) :param rrule: New iCalendar RRULE schedule string (optional). See create_automation for format examples. - :param model_id: New model ID to use (optional) - :param folder_id: New owner-owned folder ID (optional); pass an empty string to clear + :param model_id: New model ID to use (optional); blank values are ignored + :param folder_id: New owner-owned folder ID (optional); omit or pass blank to keep unchanged, pass null to clear :return: JSON with the updated automation details """ if __request__ is None: @@ -3695,7 +3808,7 @@ async def update_automation( return JSONCodec.dumps({'error': 'User context not available'}) try: - from open_webui.models.automations import AutomationData, AutomationForm, Automations + from open_webui.models.automations import AutomationData, AutomationForm, AutomationTarget, Automations from open_webui.models.users import Users from open_webui.routers.automations import check_automation_limits from open_webui.utils.automations import next_n_runs_ns, next_run_ns, validate_rrule @@ -3714,13 +3827,15 @@ async def update_automation( # Merge provided fields with existing values new_name = name if name is not None else automation.name new_prompt = prompt if prompt is not None else automation.data.get('prompt', '') - new_model_id = model_id if model_id is not None else automation.data.get('model_id', '') + new_model_id = model_id.strip() if model_id and model_id.strip() else automation.data.get('model_id', '') new_rrule = rrule if rrule is not None else automation.data.get('rrule', '') if folder_id is None: + new_folder_id = None + elif not folder_id.strip(): new_folder_id = automation.folder_id else: try: - new_folder_id = await _validate_owned_automation_folder(user_id, folder_id) + new_folder_id = await _validate_owned_automation_folder(user_id, folder_id.strip()) except ValueError as e: return JSONCodec.dumps({'error': str(e)}) @@ -3744,6 +3859,7 @@ async def update_automation( prompt=new_prompt, model_id=new_model_id, rrule=new_rrule, + target=AutomationTarget(**automation.data['target']) if automation.data.get('target') else None, ), is_active=automation.is_active, ) @@ -3757,6 +3873,7 @@ async def update_automation( 'name': updated.name, 'folder_id': updated.folder_id, 'model_id': new_model_id, + 'target': updated.data.get('target'), 'is_active': updated.is_active, 'next_runs': next_n_runs_ns(new_rrule, tz=tz), }, @@ -3822,6 +3939,7 @@ async def list_automations( 'folder_id': item.folder_id, 'prompt_snippet': snippet, 'model_id': item.data.get('model_id', ''), + 'target': item.data.get('target'), 'rrule': rrule, 'is_active': item.is_active, 'last_run_at': item.last_run_at, @@ -3938,6 +4056,9 @@ async def delete_automation( # ============================================================================= +MAX_CALENDAR_RANGE_END_NS = 2**63 - 1 + + def _get_user_tz(user_dict: dict): """Get the user's timezone as a ZoneInfo, falling back to UTC.""" from zoneinfo import ZoneInfo @@ -4037,11 +4158,7 @@ async def search_calendar_events( return JSONCodec.dumps({'error': f'Invalid start datetime: {e}'}) try: - end_ns = ( - _dt_to_ns(end, tz) - if end - else int(time.time() * 1_000) * 1_000_000 + 365 * 86400 * 1_000_000_000_000 - ) + end_ns = _dt_to_ns(end, tz) if end else MAX_CALENDAR_RANGE_END_NS except (ValueError, TypeError) as e: return JSONCodec.dumps({'error': f'Invalid end datetime: {e}'}) @@ -4306,16 +4423,17 @@ async def update_calendar_event( if reminder_minutes is not None: meta = {'alert_minutes': reminder_minutes} - form = CalendarEventUpdateForm( - title=title, - description=description, - start_at=start_ns, - end_at=end_ns, - all_day=all_day, - location=location, - is_cancelled=is_cancelled, - meta=meta, - ) + update_fields = { + 'title': title, + 'description': description, + 'start_at': start_ns, + 'end_at': end_ns, + 'all_day': all_day, + 'location': location, + 'is_cancelled': is_cancelled, + 'meta': meta, + } + form = CalendarEventUpdateForm(**{k: v for k, v in update_fields.items() if v is not None}) updated = await CalendarEvents.update_event_by_id(event_id, form) if not updated: diff --git a/backend/open_webui/utils/ask_user.py b/backend/open_webui/utils/ask_user.py new file mode 100644 index 0000000000..6c4168d2d7 --- /dev/null +++ b/backend/open_webui/utils/ask_user.py @@ -0,0 +1,144 @@ +from collections.abc import Callable + +from open_webui.utils.json_codec import JSONCodec + + +ASK_USER_NAME = 'ask_user' + + +def get_ask_user_tool_call(tool_calls: list[dict]) -> tuple[dict | None, str | None]: + ask_user_calls = [ + tool_call for tool_call in tool_calls if tool_call.get('function', {}).get('name') == ASK_USER_NAME + ] + if not ask_user_calls: + return None, None + if len(tool_calls) != 1: + return ask_user_calls[0], 'Error: ask_user must be called by itself after research.' + if len(ask_user_calls) != 1: + return ask_user_calls[0], 'Error: only one ask_user call is allowed per turn.' + return ask_user_calls[0], None + + +def normalize_ask_user_request(arguments: dict) -> dict: + questions = arguments.get('questions') + if not isinstance(questions, list) or not 1 <= len(questions) <= 3: + raise ValueError('ask_user requires 1-3 questions.') + + normalized_questions = [] + seen_ids = set() + allow_other = bool(arguments.get('allow_other', True)) + for index, question in enumerate(questions): + if not isinstance(question, dict): + raise ValueError('Each question must be an object.') + + question_id = str(question.get('id') or '').strip()[:64] + if not question_id: + raise ValueError('Each question requires a non-empty id.') + if question_id in seen_ids: + raise ValueError(f'Duplicate question id: {question_id}') + seen_ids.add(question_id) + + options = question.get('options') + if not isinstance(options, list) or not 2 <= len(options) <= 3: + raise ValueError('Each question requires 2-3 options.') + + normalized_options = [] + for option in options: + if not isinstance(option, dict): + raise ValueError('Each option must be an object.') + label = str(option.get('label') or '').strip()[:80] + description = str(option.get('description') or '').strip()[:240] + if not label or not description: + raise ValueError('Each option requires a label and description.') + normalized_options.append({'label': label, 'description': description}) + + question_text = str(question.get('question') or '').strip()[:500] + if not question_text: + raise ValueError('Each question requires question text.') + + normalized_questions.append( + { + 'id': question_id, + 'header': str(question.get('header') or '').strip()[:48] or f'Question {index + 1}', + 'question': question_text, + 'options': normalized_options, + 'allow_other': bool(question.get('allow_other', allow_other)), + } + ) + + timeout_ms = arguments.get('timeout_ms', 120_000) + if isinstance(timeout_ms, bool) or not isinstance(timeout_ms, int) or not 60_000 <= timeout_ms <= 240_000: + timeout_ms = 120_000 + + return { + 'questions': normalized_questions, + 'allow_other': allow_other, + 'timeout_ms': timeout_ms, + } + + +def stage_ask_user_tool_call( + tool_calls: list[dict], + output: list[dict], + make_output_id: Callable[[str], str], +) -> dict | None: + tool_call, error = get_ask_user_tool_call(tool_calls) + if not tool_call: + return None + + call_id = tool_call.get('id') or make_output_id('fc') + raw_arguments = tool_call.get('function', {}).get('arguments', '{}') + arguments = raw_arguments + + if not error: + try: + parsed_arguments = JSONCodec.loads(raw_arguments or '{}') + if not isinstance(parsed_arguments, dict): + raise ValueError('ask_user arguments must be an object.') + arguments = JSONCodec.dumps(normalize_ask_user_request(parsed_arguments)) + except (JSONCodec.JSONDecodeError, TypeError, ValueError) as exc: + error = f'Error: {exc}' + + item = { + 'type': 'function_call', + 'id': call_id or make_output_id('fc'), + 'call_id': call_id, + 'name': ASK_USER_NAME, + 'arguments': arguments, + 'status': 'completed' if error else 'pending', + } + + existing_item = next( + ( + existing + for existing in output + if existing.get('type') == 'function_call' + and ( + existing.get('call_id') == call_id + or existing.get('id') == tool_call.get('id') + or ( + not existing.get('call_id') + and existing.get('name') == ASK_USER_NAME + and existing.get('status') not in {'rejected', 'failed'} + ) + ) + ), + None, + ) + if existing_item: + existing_item.update(item) + else: + output.append(item) + + if error: + output.append( + { + 'type': 'function_call_output', + 'id': make_output_id('fco'), + 'call_id': call_id, + 'output': [{'type': 'input_text', 'text': error}], + 'status': 'completed', + } + ) + + return {'call_id': call_id, 'error': error, 'item': item} diff --git a/backend/open_webui/utils/automations.py b/backend/open_webui/utils/automations.py index 712ea45096..1d4be9b20b 100644 --- a/backend/open_webui/utils/automations.py +++ b/backend/open_webui/utils/automations.py @@ -36,6 +36,7 @@ from open_webui.models.automations import AutomationModel, AutomationRuns, Autom from open_webui.models.chats import ChatForm, Chats from open_webui.models.config import Config from open_webui.models.folders import Folders +from open_webui.models.messages import MessageForm from open_webui.models.users import Users from open_webui.utils.auth import create_token from open_webui.utils.misc import parse_duration @@ -124,6 +125,9 @@ def validate_rrule(s: str, tz: str = None) -> None: clock so that near-future schedules are not incorrectly rejected on servers whose system clock is ahead (e.g. UTC vs US timezones). """ + upper = s.upper() + if 'COUNT=' in upper and 'DTSTART' not in upper: + raise ValueError(ERROR_MESSAGES.AUTOMATION_COUNT_REQUIRES_DTSTART) zi = _resolve_tz(tz) now = datetime.now(zi).replace(tzinfo=None) if zi else datetime.now() try: @@ -297,33 +301,17 @@ def _build_request( return request -def _resolve_model_tool_ids(app, model_id: str) -> list[str]: - """Read model-attached tool_ids from model config. - - The frontend does this in Chat.svelte (model.info.meta.toolIds). - The backend never auto-resolves them, so we must do it explicitly. - """ - models = getattr(app.state, 'MODELS', {}) - model = models.get(model_id, {}) - tool_ids = model.get('info', {}).get('meta', {}).get('toolIds', []) - return list(tool_ids) if tool_ids else [] - - -async def _resolve_model_features(app, model_id: str) -> dict: - """Read model default features from model config. - - The frontend does this in Chat.svelte (model.info.meta.defaultFeatureIds - + model.info.meta.capabilities). Enables features like web_search, - code_interpreter, image_generation when the model has them as defaults - AND the capability is enabled AND the admin has enabled the feature. - """ +async def _resolve_model_defaults(app, model_id: str) -> tuple[list[str], dict, list[str], Optional[str]]: models = getattr(app.state, 'MODELS', {}) model = models.get(model_id, {}) meta = model.get('info', {}).get('meta', {}) + tool_ids = list(meta.get('toolIds') or []) + filter_ids = list(meta.get('defaultFilterIds') or []) + terminal_id = meta.get('terminalId') or None default_feature_ids = meta.get('defaultFeatureIds', []) if not default_feature_ids: - return {} + return tool_ids, {}, filter_ids, terminal_id capabilities = meta.get('capabilities') or {} features = {} @@ -341,25 +329,7 @@ async def _resolve_model_features(app, model_id: str) -> dict: if capabilities.get(feature_id) and feature_checks[feature_id]: features[feature_id] = True - return features - - -def _resolve_model_filter_ids(app, model_id: str) -> list[str]: - """Read model default filter_ids from model config.""" - models = getattr(app.state, 'MODELS', {}) - model = models.get(model_id, {}) - filter_ids = model.get('info', {}).get('meta', {}).get('defaultFilterIds', []) - return list(filter_ids) if filter_ids else [] - - -def _resolve_model_terminal_id(app, model_id: str) -> Optional[str]: - """Read model default terminal_id from model config. - - The frontend does this in Chat.svelte (model.info.meta.terminalId). - """ - models = getattr(app.state, 'MODELS', {}) - model = models.get(model_id, {}) - return model.get('info', {}).get('meta', {}).get('terminalId') or None + return tool_ids, features, filter_ids, terminal_id async def _set_terminal_cwd(app, server_id: str, user, cwd: str, chat_id: str) -> None: @@ -410,10 +380,113 @@ async def _set_terminal_cwd(app, server_id: str, user, cwd: str, chat_id: str) - log.warning(f'Failed to set terminal CWD: {e}') +async def _execute_channel_automation( + app, + automation: AutomationModel, + user, + prompt: str, + model_id: str, + token: str, +) -> None: + target = automation.data.get('target') or {} + channel_id = target.get('channel_id') + if not channel_id or not await Config.get('channels.enable'): + raise ValueError('Channel not found') + + model = getattr(app.state, 'MODELS', {}).get(model_id, {}) + request = _build_request(app, token=token) + + from open_webui.routers.channels import new_message_handler + + async with get_async_db() as db: + user_message, channel = await new_message_handler( + request, + channel_id, + MessageForm( + content=prompt, + data={}, + meta={'automation_id': automation.id}, + ), + user, + db, + ) + response_parent_id = ( + user_message.parent_id + if user_message.parent_id + else (user_message.id if await Config.get('channels.model_response_mode', 'thread') == 'thread' else None) + ) + assistant_message, channel = await new_message_handler( + request, + channel.id, + MessageForm( + parent_id=response_parent_id, + content='', + data={}, + meta={ + 'automation_id': automation.id, + 'model_id': model_id, + 'model_name': model.get('name', model_id), + }, + ), + user, + db, + ) + + tool_ids, features, filter_ids, _ = await _resolve_model_defaults(app, model_id) + + form_data = { + 'model': model_id, + 'messages': [ + { + 'role': 'system', + 'content': f'You are {model.get("name", model_id)}, participating in a channel conversation. Be concise and conversational.', + }, + {'role': 'user', 'content': f'{user.name if user else "User"}: {prompt}'}, + ], + 'stream': True, + 'chat_id': f'channel:{channel.id}', + 'id': assistant_message.id, + 'session_id': f'channel:{channel.id}', + 'automation_id': automation.id, + 'background_tasks': {}, + } + if tool_ids: + form_data['tool_ids'] = tool_ids + if features: + form_data['features'] = features + if filter_ids: + form_data['filter_ids'] = filter_ids + + await app.state.CHAT_COMPLETION_HANDLER(request, form_data, user=user) + + from open_webui.socket.main import sio + + await sio.emit( + 'automation:result', + { + 'automation_id': automation.id, + 'name': automation.name, + 'chat_id': f'channel:{channel.id}', + 'message_id': assistant_message.id, + 'status': 'success', + }, + room=f'user:{automation.user_id}', + ) + + await _record_run(automation.id, 'success', chat_id=f'channel:{channel.id}') + await publish_event( + app, + EVENTS.AUTOMATION_RUN_COMPLETED, + actor=user, + subject_id=automation.id, + data={'name': automation.name, 'channel_id': channel.id, 'message_id': assistant_message.id}, + ) + + async def execute_automation(app, automation: AutomationModel) -> None: """Execute an automation through the full chat completion pipeline. - Creates a real chat, then calls chat_completion exactly like the frontend: + Creates a real chat or channel message, then calls chat_completion exactly like the frontend: session_id + chat_id + message_id → async task → pipeline handles everything (filters, model params, knowledge/RAG, tools, DB saves, webhooks). """ @@ -449,6 +522,20 @@ async def execute_automation(app, automation: AutomationModel) -> None: prompt = await prompt_template(automation.data['prompt'], user) model_id = automation.data['model_id'] + try: + expires_delta = parse_duration(str(await Config.get('automations.auth_token_expires_in', '1h'))) + except ValueError: + expires_delta = None + token = create_token( + data={'id': user.id, 'typ': 'automation'}, + expires_delta=expires_delta or timedelta(hours=1), + ) + + target = automation.data.get('target') or {} + if target.get('type') == 'channel': + await _execute_channel_automation(app, automation, user, prompt, model_id, token) + return + folder_id = automation.folder_id if folder_id and not await Folders.get_folder_by_id_and_user_id(folder_id, automation.user_id): await Automations.clear_folder_ids(automation.user_id, [folder_id]) @@ -525,12 +612,7 @@ async def execute_automation(app, automation: AutomationModel) -> None: ) # Resolve model defaults (frontend does this, backend doesn't) - tool_ids = _resolve_model_tool_ids(app, model_id) - features = await _resolve_model_features(app, model_id) - filter_ids = _resolve_model_filter_ids(app, model_id) - - # Resolve terminal from model config - terminal_id = _resolve_model_terminal_id(app, model_id) + tool_ids, features, filter_ids, terminal_id = await _resolve_model_defaults(app, model_id) # Build the same payload the frontend sends to /api/chat/completions form_data = { @@ -561,14 +643,6 @@ async def execute_automation(app, automation: AutomationModel) -> None: # Call the full chat completion pipeline (same as POST /api/chat/completions). # The handler reference is stored on app.state to avoid circular imports. - try: - expires_delta = parse_duration(str(await Config.get('automations.auth_token_expires_in', '1h'))) - except ValueError: - expires_delta = None - token = create_token( - data={'id': user.id, 'typ': 'automation'}, - expires_delta=expires_delta or timedelta(hours=1), - ) request = _build_request(app, token=token) await app.state.CHAT_COMPLETION_HANDLER(request, form_data, user=user) diff --git a/backend/open_webui/utils/context_compaction.py b/backend/open_webui/utils/context_compaction.py index ff84481a5e..1b0470dc17 100644 --- a/backend/open_webui/utils/context_compaction.py +++ b/backend/open_webui/utils/context_compaction.py @@ -235,6 +235,23 @@ def _resolve_token_threshold(global_threshold: int, global_cap: int, metadata: d return min(configured_threshold or global_threshold, global_cap) +def _usage_token_count(usage: dict) -> int: + prompt_tokens = int(usage.get('prompt_tokens') or usage.get('prompt_eval_count') or 0) + if not prompt_tokens and (usage.get('prompt_n') is not None or usage.get('cache_n') is not None): + prompt_tokens = int(usage.get('prompt_n') or 0) + int(usage.get('cache_n') or 0) + if not prompt_tokens: + prompt_tokens = int(usage.get('input_tokens') or 0) + + completion_tokens = int( + usage.get('completion_tokens') + or usage.get('output_tokens') + or usage.get('eval_count') + or usage.get('predicted_n') + or 0 + ) + return prompt_tokens + completion_tokens + + async def get_chat_context_usage(chat: Any, model_id: str | None = None) -> dict | None: chat_data = chat.chat or {} history = chat_data.get('history') or {} @@ -263,25 +280,7 @@ async def get_chat_context_usage(chat: Any, model_id: str | None = None) -> dict for idx in range(len(messages) - 1, -1, -1): usage = messages[idx].get('usage') or (messages[idx].get('info') or {}).get('usage') - if isinstance(usage, dict) and ( - tokens := ( - int( - usage.get('prompt_tokens') - or usage.get('input_tokens') - or usage.get('prompt_eval_count') - or usage.get('prompt_n') - or 0 - ) - + int( - usage.get('completion_tokens') - or usage.get('output_tokens') - or usage.get('eval_count') - or usage.get('predicted_n') - or 0 - ) - + int(usage.get('cache_n') or 0) - ) - ): + if isinstance(usage, dict) and (tokens := _usage_token_count(usage)): tokens += _estimate_messages_tokens(messages[idx + 1 :]) return _build_context_usage(tokens, threshold) @@ -320,25 +319,7 @@ def _exceeds_token_threshold(messages: list[dict], system_prompt: str, summary: for idx in range(len(messages) - 1, -1, -1): usage = messages[idx].get('usage') or (messages[idx].get('info') or {}).get('usage') - if isinstance(usage, dict) and ( - tokens := ( - int( - usage.get('prompt_tokens') - or usage.get('input_tokens') - or usage.get('prompt_eval_count') - or usage.get('prompt_n') - or 0 - ) - + int( - usage.get('completion_tokens') - or usage.get('output_tokens') - or usage.get('eval_count') - or usage.get('predicted_n') - or 0 - ) - + int(usage.get('cache_n') or 0) - ) - ): + if isinstance(usage, dict) and (tokens := _usage_token_count(usage)): return tokens + _estimate_messages_tokens(messages[idx + 1 :]) > threshold estimated = _estimate_tokens(system_prompt) + _estimate_tokens(summary or '') + _estimate_messages_tokens(messages) diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 338b652d45..96166e0ac3 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -76,9 +76,11 @@ from open_webui.socket.main import ( get_event_call, get_event_emitter, ) +from open_webui.tasks import clear_response_stream, save_response_stream from open_webui.utils.access_control import has_connection_access, has_permission from open_webui.utils.access_control.files import get_owner_accessible_folder_files from open_webui.utils.access_control.folders import has_folder_access +from open_webui.utils.ask_user import stage_ask_user_tool_call from open_webui.utils.chat import generate_chat_completion from open_webui.utils.chat_id import is_saved_chat_id from open_webui.utils.code_interpreter import execute_code_jupyter @@ -108,6 +110,7 @@ from open_webui.utils.misc import ( get_last_user_message_item, get_message_list, get_output_text, + get_reasoning_details, get_system_message, is_string_allowed, merge_system_messages, @@ -472,6 +475,23 @@ def deep_merge(target, source): return source +RESPONSE_COMPLETION_RESPONSE_FIELDS = ('error', 'id', 'output', 'usage') + + +def get_response_completion_event_data(event: dict) -> dict: + """Build the data payload for response:completion events.""" + response = event.get('response') + if not isinstance(response, dict): + return event + + response_data = {key: response[key] for key in RESPONSE_COMPLETION_RESPONSE_FIELDS if key in response} + + return { + **event, + 'response': response_data, + } + + def handle_responses_streaming_event( data: dict, current_output: list, @@ -498,7 +518,22 @@ def handle_responses_streaming_event( item = data.get('item', {}) if item: new_output = list(current_output) - new_output.append(item) + output_index = data.get('output_index', len(new_output)) + existing_index = next( + ( + idx + for idx, existing in enumerate(new_output) + if (item.get('id') and existing.get('id') == item.get('id')) + or (item.get('call_id') and existing.get('call_id') == item.get('call_id')) + ), + None, + ) + if existing_index is not None: + new_output[existing_index] = item + elif 0 <= output_index < len(new_output): + new_output.insert(output_index, item) + else: + new_output.append(item) return new_output, None return current_output, None @@ -1716,6 +1751,19 @@ async def chat_image_generation_handler(request: Request, form_data: dict, extra system_message_content = f'Image generation was attempted but failed. The system is currently unable to generate the image. Tell the user that the following error occurred: {error_message}' + elif not await Config.get('image_generation.enable'): + await __event_emitter__( + { + 'type': 'status', + 'data': { + 'description': 'Image generation is disabled', + 'done': True, + }, + } + ) + + system_message_content = 'Image generation was requested but the feature is currently disabled by the administrator, so no image was created. Let the user know that image generation is currently unavailable.' + else: # Create image(s) if await Config.get('image_generation.prompt.enable'): @@ -2004,13 +2052,14 @@ def get_reasoning_format(model: dict) -> str | None: Determine how reasoning should be included in reconstructed messages. Returns: - 'think_tags': Ollama expects tags in content. + 'thinking': Ollama expects reasoning in the native thinking field. + 'think_tags': wrap reasoning in tags inside content. 'reasoning_content': llama.cpp supports reasoning_content as a top-level field. None: skip reasoning (safe default for strict providers). """ provider = model.get('provider', '') - if provider == 'ollama': - return 'think_tags' + if model.get('owned_by') == 'ollama': + return 'thinking' if provider == 'llama.cpp': return 'reasoning_content' return None @@ -2041,25 +2090,14 @@ def process_messages_with_output( processed.extend(output_messages) continue - # Strip 'output' field before adding (LLM shouldn't see it) - clean_message = {k: v for k, v in message.items() if k != 'output'} + clean_message = dict(message) + for key in ('id', 'files', 'output', 'contextSummary', 'context_summary', 'usage'): + clean_message.pop(key, None) processed.append(clean_message) return processed -def strip_compaction_fields(messages: list[dict]) -> list[dict]: - stripped = [] - for message in messages: - clean = dict(message) - clean.pop('contextSummary', None) - clean.pop('context_summary', None) - clean.pop('usage', None) - clean.pop('id', None) - stripped.append(clean) - return stripped - - def sanitize_tool_pairs(messages: list[dict]) -> list[dict]: tool_result_ids = { message.get('tool_call_id') @@ -2318,8 +2356,6 @@ async def process_chat_payload(request, form_data, user, metadata, model): except Exception: log.exception('Context compaction failed; continuing with full chat history') - form_data['messages'] = strip_compaction_fields(form_data.get('messages', [])) - # Process messages with OR-aligned output items for clean LLM messages form_data['messages'] = process_messages_with_output( form_data.get('messages', []), @@ -2969,6 +3005,301 @@ async def build_chat_response_context(request, form_data, user, model, metadata, } +async def execute_tool_call_for_output(request, form_data, user, metadata, event_caller, event_emitter, tool_call): + tools = metadata.get('tools', {}) + name = tool_call.get('function', {}).get('name', '') + tool_args = tool_call.get('function', {}).get('arguments', '{}') + params = {} + if tool_args and tool_args.strip(): + try: + params = JSONCodec.loads(tool_args) + except Exception: + try: + params = ast.literal_eval(tool_args) + except Exception as e: + log.debug(e) + return { + 'tool_call_id': tool_call.get('id', ''), + 'content': ( + 'Error: Tool call arguments could not be parsed. ' + 'The model generated malformed or incomplete JSON.' + ), + } + tool_call.setdefault('function', {})['arguments'] = JSONCodec.dumps(params) + + tool = tools.get(name) + if not tool: + return {'tool_call_id': tool_call.get('id', ''), 'content': f'Error: Tool "{name}" not found.'} + + spec = tool.get('spec', {}) + tool_type = tool.get('type', '') + direct_tool = tool.get('direct', False) + allowed_params = spec.get('parameters', {}).get('properties', {}).keys() + params = {key: value for key, value in params.items() if key in allowed_params} + + try: + if direct_tool: + if not event_caller: + result = 'Error: Browser session is not connected for this direct tool.' + else: + result = await event_caller( + { + 'type': 'execute:tool', + 'data': { + 'id': str(uuid4()), + 'name': name, + 'params': params, + 'server': tool.get('server', {}), + 'session_id': metadata.get('session_id'), + }, + } + ) + else: + function = await get_updated_tool_function( + function=tool['callable'], + extra_params={ + '__messages__': form_data.get('messages', []), + '__files__': metadata.get('files', []), + }, + ) + result = await function(**params) + except Exception as e: + result = str(e) + + result, files, embeds = await process_tool_result( + request, + name, + result, + tool_type, + direct_tool, + metadata, + user, + ) + + await terminal_event_handler(name, params, result, event_emitter) + + return { + 'tool_call_id': tool_call.get('id', ''), + 'content': str(result) if result else '', + **({'files': files} if files else {}), + **({'embeds': embeds} if embeds else {}), + } + + +def append_tool_result_output(output: list[dict], result: dict) -> None: + output_parts = [{'type': 'input_text', 'text': result.get('content', '')}] + display_files = [] + for file_item in result.get('files', []): + if file_item.get('type') == 'image' and file_item.get('url', '').startswith('data:'): + output_parts.append({'type': 'input_image', 'image_url': file_item['url']}) + else: + display_files.append(file_item) + + output.append( + { + 'type': 'function_call_output', + 'id': output_id('fco'), + 'call_id': result.get('tool_call_id', ''), + 'output': output_parts, + 'status': 'completed', + **({'files': display_files} if display_files else {}), + **({'embeds': result.get('embeds')} if result.get('embeds') else {}), + } + ) + + +async def drain_approved_tool_calls(request, form_data, user, model, metadata) -> bool: + chat_id = metadata.get('chat_id') + message_id = metadata.get('message_id') or metadata.get('assistant_message_id') + if not is_saved_chat_id(chat_id) or not message_id: + return False + + message = await Chats.get_message_by_id_and_message_id(chat_id, message_id) + output = message.get('output') if message else None + if not isinstance(output, list): + return False + + result_call_ids = { + item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id') + } + approved_calls = [ + item + for item in output + if item.get('type') == 'function_call' + and item.get('call_id') + and item.get('status') == 'queued' + and item.get('approved') is True + and item.get('call_id') not in result_call_ids + ] + if not approved_calls: + if metadata.get('params', {}).get('tool_approval_mode', 'full') == 'ask' and any( + item.get('type') == 'function_call' + and item.get('name') != 'ask_user' + and (item.get('call_id') or item.get('id')) + and item.get('status') == 'queued' + and item.get('approved') is not True + and (item.get('call_id') or item.get('id')) not in result_call_ids + for item in output + ): + event_emitter, _ = await get_event_emitter_and_caller(metadata) + await pause_for_tool_approval(chat_id, message_id, output, form_data, metadata) + if event_emitter: + await event_emitter({'type': 'chat:completion', 'data': {'done': False, 'output': output}}) + return True + return False + + event_emitter, event_caller = await get_event_emitter_and_caller(metadata) + changed = False + for item in approved_calls: + if item.get('name') == 'ask_user': + item['status'] = 'pending' + item.pop('approved', None) + changed = True + continue + + tool_call = { + 'id': item.get('call_id', ''), + 'type': 'function', + 'function': { + 'name': item.get('name', ''), + 'arguments': item.get('arguments', '{}'), + }, + } + result = await execute_tool_call_for_output( + request, + form_data, + user, + metadata, + event_caller, + event_emitter, + tool_call, + ) + item['status'] = 'completed' + item['arguments'] = tool_call.get('function', {}).get('arguments', '{}') + append_tool_result_output(output, result) + changed = True + + if changed: + result_call_ids = { + item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id') + } + if metadata.get('params', {}).get('tool_approval_mode', 'full') == 'ask' and any( + item.get('type') == 'function_call' + and item.get('name') != 'ask_user' + and (item.get('call_id') or item.get('id')) + and item.get('status') == 'queued' + and item.get('approved') is not True + and (item.get('call_id') or item.get('id')) not in result_call_ids + for item in output + ): + await pause_for_tool_approval(chat_id, message_id, output, form_data, metadata) + result_call_ids = { + item.get('call_id') + for item in output + if item.get('type') == 'function_call_output' and item.get('call_id') + } + paused = any( + item.get('type') == 'function_call' + and item.get('call_id') + and item.get('status') in {'pending', 'queued', 'requires_approval'} + and item.get('call_id') not in result_call_ids + for item in output + ) + if not paused: + output.append( + { + 'type': 'message', + 'id': output_id('msg'), + 'status': 'in_progress', + 'role': 'assistant', + 'content': [{'type': 'output_text', 'text': ''}], + } + ) + + await Chats.upsert_message_to_chat_by_id_and_message_id( + chat_id, + message_id, + {'done': False, 'output': output}, + touch=False, + ) + if event_emitter: + await event_emitter( + { + 'type': 'chat:completion', + 'data': { + 'done': False, + 'output': output, + }, + } + ) + + db_messages = await load_messages_from_db(chat_id, metadata.get('user_message_id')) + if db_messages: + assistant_message = await Chats.get_message_by_id_and_message_id(chat_id, message_id) + if assistant_message: + db_messages.append( + { + k: v + for k, v in assistant_message.items() + if k in ('id', 'role', 'content', 'output', 'files', 'contextSummary', 'usage') + } + ) + form_data['messages'] = process_messages_with_output( + db_messages, + reasoning_format=get_reasoning_format(model), + ) + form_data['messages'] = sanitize_tool_pairs(form_data['messages']) + + return paused + + return False + + +async def pause_for_tool_approval(chat_id: str, message_id: str, output: list[dict], form_data: dict, metadata: dict): + result_call_ids = { + item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id') + } + has_pending_approval = False + for item in output: + if item.get('type') == 'function_call' and not item.get('call_id') and item.get('id'): + item['call_id'] = item['id'] + + if ( + item.get('type') == 'function_call' + and item.get('call_id') + and item.get('call_id') not in result_call_ids + and item.get('status') != 'rejected' + ): + if not has_pending_approval: + item['status'] = 'pending' + has_pending_approval = True + elif item.get('status') == 'in_progress': + item['status'] = 'queued' + + await Chats.upsert_message_to_chat_by_id_and_message_id( + chat_id, + message_id, + { + 'done': False, + 'output': output, + 'meta': { + **(metadata.get('tool_approval') or {}), + 'session_id': metadata.get('session_id'), + 'tool_ids': metadata.get('tool_ids') or [], + 'skill_ids': metadata.get('skill_ids') or [], + 'terminal_id': metadata.get('terminal_id'), + 'tool_servers': metadata.get('tool_servers'), + 'filter_ids': metadata.get('filter_ids') or [], + 'features': metadata.get('features') or {}, + 'variables': metadata.get('variables') or {}, + 'files': metadata.get('files') or [], + 'params': metadata.get('params') or {}, + }, + }, + touch=False, + ) + + def get_response_data(response): if isinstance(response, list) and len(response) == 1: # If the response is a single-item list, unwrap it #17213 @@ -3586,7 +3917,7 @@ async def non_streaming_chat_response_handler(response, ctx): if not response_output: choice_message = choices[0].get('message', {}) reasoning_content = choice_message.get('reasoning_content') or choice_message.get('reasoning') - reasoning_details = choice_message.get('reasoning_details') + reasoning_details = get_reasoning_details(choice_message) response_output = [] if reasoning_content or reasoning_details: reasoning_item = { @@ -3737,6 +4068,7 @@ async def streaming_chat_response_handler(response, ctx): async def response_handler(response, events): filter_context = FilterContext() tag_scan_positions = {} + response_stream_task_id = metadata.get('task_id') or metadata.get('message_id') def tag_output_handler(content_type, tags, output): """ @@ -4019,7 +4351,20 @@ async def streaming_chat_response_handler(response, ctx): # Initialize output: use existing from message if continuing, else create new existing_output = message.get('output') if message else None - if existing_output: + prior_output = [] + if existing_output and metadata.get('assistant_message_id'): + prior_output = list(existing_output) + if ( + prior_output + and prior_output[-1].get('type') == 'message' + and prior_output[-1].get('status') == 'in_progress' + ): + msg_parts = prior_output[-1].get('content', []) + if not msg_parts or (len(msg_parts) == 1 and not msg_parts[0].get('text', '').strip()): + prior_output.pop() + output = [] + content_parts = [] + elif existing_output: output = existing_output else: # Only create an initial message item if there is content to initialize with @@ -4037,7 +4382,6 @@ async def streaming_chat_response_handler(response, ctx): output = [] usage = None - prior_output = [] last_response_id = None def full_output(): @@ -4132,38 +4476,112 @@ async def streaming_chat_response_handler(response, ctx): ) last_delta_data = None last_delta_type = None + last_delta_key = None + + def response_stream_content(stream_output: list | None = None): + return ''.join(content_parts) or get_output_text( + stream_output if stream_output is not None else full_output() + ) + + async def save_current_response_stream(stream_output: list | None = None): + if not chat_id or not metadata.get('message_id'): + return + + current_stream_output = stream_output if stream_output is not None else full_output() + await save_response_stream( + request.app.state.redis, + response_stream_task_id, + chat_id, + metadata.get('message_id'), + response_stream_content(current_stream_output), + current_stream_output, + ) + + def get_response_delta_key(delta_data: dict): + event_type = delta_data.get('type', '') + if not event_type.startswith('response.') or not event_type.endswith('.delta'): + return None + return ( + event_type, + delta_data.get('item_id'), + delta_data.get('output_index'), + delta_data.get('content_index'), + delta_data.get('summary_index'), + ) + + def get_response_data_with_full_output_index(response_data: dict): + if prior_output and isinstance(response_data.get('output_index'), int): + return { + **response_data, + 'output_index': response_data['output_index'] + len(prior_output), + } + return response_data async def flush_pending_delta_data(threshold: int = 0): nonlocal delta_count nonlocal last_delta_data nonlocal last_delta_type + nonlocal last_delta_key if delta_count >= threshold and last_delta_data: await event_emitter( { - 'type': 'chat:completion', + 'type': 'response:completion', 'data': last_delta_data, } ) + await save_current_response_stream() delta_count = 0 last_delta_data = None last_delta_type = None + last_delta_key = None async def queue_pending_delta_data(delta_data: dict, delta_type: str): nonlocal delta_count nonlocal last_delta_data nonlocal last_delta_type + nonlocal last_delta_key - if last_delta_type and last_delta_type != delta_type: - await flush_pending_delta_data() + delta_data = get_response_data_with_full_output_index(delta_data) + delta_key = get_response_delta_key(delta_data) + if ( + last_delta_data + and last_delta_key == delta_key + and isinstance(last_delta_data.get('delta'), str) + and isinstance(delta_data.get('delta'), str) + ): + last_delta_data['delta'] += delta_data['delta'] + delta_count += 1 + else: + if last_delta_data and (last_delta_type != delta_type or last_delta_key != delta_key): + await flush_pending_delta_data() - delta_count += 1 - last_delta_data = delta_data - last_delta_type = delta_type + delta_count += 1 + last_delta_data = delta_data + last_delta_type = delta_type + last_delta_key = delta_key if delta_count >= delta_chunk_size: await flush_pending_delta_data(delta_chunk_size) + async def emit_response_completion_event(response_data: dict, stream_output: list | None = None): + if response_data.get('type', '').endswith('.delta'): + await queue_pending_delta_data( + response_data, + response_data.get('type', 'response.delta'), + ) + return + + response_data = get_response_data_with_full_output_index(response_data) + await flush_pending_delta_data() + await event_emitter( + { + 'type': 'response:completion', + 'data': get_response_completion_event_data(response_data), + } + ) + await save_current_response_stream(stream_output) + filter_extra_params = {'__body__': form_data, **extra_params} if filter_functions else None async for line in response.body_iterator: @@ -4238,16 +4656,16 @@ async def streaming_chat_response_handler(response, ctx): ) # Check for Responses API events (type field starts with "response.") elif data.get('type', '').startswith('response.'): - response_event_type = data.get('type', '') - response_event_is_delta = response_event_type.endswith('.delta') + response_data_type = data.get('type', '') + response_data_is_delta = response_data_type.endswith('.delta') output, response_metadata = handle_responses_streaming_event(data, output) - if not response_event_is_delta: + if not response_data_is_delta: await flush_pending_delta_data() # Emit citation sources from finalized output items # (mirrors Chat Completions annotation handling at delta level) - if response_event_type == 'response.output_item.done': + if response_data_type == 'response.output_item.done': item = data.get('item', {}) if item.get('type') == 'message': for part in item.get('content', []): @@ -4279,13 +4697,6 @@ async def streaming_chat_response_handler(response, ctx): } ) - processed_data = { - 'output': full_output(), - } - - # print(data) - # print(processed_data) - # Merge any metadata (usage, etc.) # Strip 'done' — response.completed emits # it but we may still need to execute tool @@ -4302,22 +4713,21 @@ async def streaming_chat_response_handler(response, ctx): usage = merge_usage(usage, response_metadata['usage']) response_metadata['usage'] = usage - processed_data.update(response_metadata) - processed_data.pop('done', None) + if response_metadata.get('error'): + await event_emitter( + { + 'type': 'chat:completion', + 'data': {'error': response_metadata['error']}, + } + ) - if response_event_is_delta: - response_delta_type = response_event_type.split('.')[1] - await queue_pending_delta_data( - processed_data, - 'tool_call' - if response_delta_type == 'function_call_arguments' - else 'content', - ) - else: + await emit_response_completion_event(data) + + if response_metadata and response_metadata.get('usage'): await event_emitter( { 'type': 'chat:completion', - 'data': processed_data, + 'data': {'usage': usage}, } ) continue @@ -4415,6 +4825,7 @@ async def streaming_chat_response_handler(response, ctx): # Add the new tool call delta_tool_call.setdefault('function', {}) delta_tool_call['function'].setdefault('name', '') + delta_tool_call['id'] = delta_tool_call.get('id') or output_id('fc') delta_arguments = delta_tool_call['function'].get('arguments') if not isinstance(delta_arguments, str): delta_tool_call['function']['arguments'] = ( @@ -4446,27 +4857,73 @@ async def streaming_chat_response_handler(response, ctx): delta_arguments ) - # Emit pending tool calls in real-time + # Emit pending tool calls in real-time as Responses events. if response_tool_calls: - # Build pending function_call output items for display - pending_fc_items = [] + output_by_call_id = { + item.get('call_id'): (idx, item) + for idx, item in enumerate(output) + if item.get('type') == 'function_call' + } + for tc in response_tool_calls: - call_id = tc.get('id', '') + call_id = tc.get('id') or output_id('fc') + tc['id'] = call_id func = tc.get('function', {}) - pending_fc_items.append( - { + if call_id in output_by_call_id: + output_index, item = output_by_call_id[call_id] + item['name'] = func.get('name', item.get('name', '')) + item['arguments'] = func.get('arguments', item.get('arguments', '')) + item['status'] = 'in_progress' + else: + output_index = len(output) + item = { 'type': 'function_call', - 'id': call_id or output_id('fc'), + 'id': call_id, 'call_id': call_id, 'name': func.get('name', ''), - 'arguments': func.get('arguments', '{}'), + 'arguments': '', 'status': 'in_progress', } - ) + output.append(item) + output_by_call_id[call_id] = (output_index, item) + await emit_response_completion_event( + { + 'type': 'response.output_item.added', + 'output_index': output_index, + 'item': item.copy(), + } + ) + item['arguments'] = func.get('arguments', '') - data = { - 'output': full_output() + pending_fc_items, - } + for delta_tool_call in delta_tool_calls: + tool_call_index = delta_tool_call.get('index') + current_response_tool_call = next( + ( + tc + for tc in response_tool_calls + if tc.get('index') == tool_call_index + ), + None, + ) + if not current_response_tool_call: + continue + call_id = current_response_tool_call.get('id') + output_index, _ = output_by_call_id.get(call_id, (len(output) - 1, {})) + delta_arguments = delta_tool_call.get('function', {}).get('arguments') + if delta_arguments is not None: + if not isinstance(delta_arguments, str): + delta_arguments = JSONCodec.dumps(delta_arguments) + await emit_response_completion_event( + { + 'type': 'response.function_call_arguments.delta', + 'item_id': call_id, + 'output_index': output_index, + 'delta': delta_arguments, + } + ) + + await save_current_response_stream() + data = None delta_type = 'tool_call' delta_images = delta.get('images') @@ -4501,7 +4958,7 @@ async def streaming_chat_response_handler(response, ctx): or delta.get('reasoning') or delta.get('thinking') ) - reasoning_details = delta.get('reasoning_details') + reasoning_details = get_reasoning_details(delta) reasoning_detail_items = ( [item for item in reasoning_details if isinstance(item, dict)] if isinstance(reasoning_details, list) @@ -4570,20 +5027,26 @@ async def streaming_chat_response_handler(response, ctx): } ] + reasoning_index = output.index(reasoning_item) data = { - 'output': full_output(), + 'type': 'response.reasoning_text.delta', + 'item_id': reasoning_item.get('id'), + 'output_index': reasoning_index, + 'content_index': max( + len(reasoning_item.get('content', [])) - 1, + 0, + ), + 'delta': reasoning_content, } - delta_type = 'content' + delta_type = 'response.reasoning_text.delta' if reasoning_detail_items: merge_streamed_reasoning_details( reasoning_item.setdefault('reasoning_details', []), reasoning_detail_items, ) - data = { - 'output': full_output(), - } - delta_type = 'content' + await save_current_response_stream() + data = None if value: if ( @@ -4729,29 +5192,27 @@ async def streaming_chat_response_handler(response, ctx): if end: break - if ENABLE_REALTIME_CHAT_SAVE and save_to_chat: - current_output = full_output() - # Save message in the database - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - { - 'output': current_output, - }, - ) - data = { - 'output': current_output, - } - delta_type = 'content' - else: - data = { - 'output': full_output(), - } - delta_type = 'content' + target_index = len(output) - 1 + target_item = output[target_index] if target_index >= 0 else {} + target_content = target_item.get('content', []) + content_index = max(len(target_content) - 1, 0) + delta_event_type = ( + 'response.reasoning_text.delta' + if target_item.get('type') == 'reasoning' + else 'response.output_text.delta' + ) + data = { + 'type': delta_event_type, + 'item_id': target_item.get('id'), + 'output_index': target_index, + 'content_index': content_index, + 'delta': value, + } + delta_type = delta_event_type - if delta: + if delta and data: await queue_pending_delta_data(data, delta_type) - else: + elif data: await event_emitter( { 'type': 'chat:completion', @@ -4800,6 +5261,29 @@ async def streaming_chat_response_handler(response, ctx): reasoning_item['status'] = 'completed' if response_tool_calls: + for tc in response_tool_calls: + call_id = tc.get('id', '') + arguments = tc.get('function', {}).get('arguments', '{}') + for output_index, item in enumerate(output): + if item.get('type') == 'function_call' and item.get('call_id') == call_id: + item['arguments'] = arguments + item['status'] = 'completed' + await emit_response_completion_event( + { + 'type': 'response.function_call_arguments.done', + 'item_id': item.get('id'), + 'output_index': output_index, + 'arguments': arguments, + } + ) + await emit_response_completion_event( + { + 'type': 'response.output_item.done', + 'output_index': output_index, + 'item': item.copy(), + } + ) + break tool_calls.append(_split_tool_calls(response_tool_calls)) # Responses API path: extract function_call items from output @@ -4814,11 +5298,12 @@ async def streaming_chat_response_handler(response, ctx): } responses_api_tool_calls = [] for item in output: - if item.get('type') == 'function_call' and item.get('call_id') not in handled_call_ids: + call_id = item.get('call_id') or item.get('id') or output_id('fc') + if item.get('type') == 'function_call' and call_id not in handled_call_ids: arguments = item.get('arguments', '{}') responses_api_tool_calls.append( { - 'id': item.get('call_id', ''), + 'id': call_id, 'index': len(responses_api_tool_calls), 'function': { 'name': item.get('name', ''), @@ -4868,6 +5353,22 @@ async def streaming_chat_response_handler(response, ctx): tool_call_iterations += 1 response_tool_calls = tool_calls.pop(0) + ask_user_stage = stage_ask_user_tool_call(response_tool_calls, output, output_id) + if ask_user_stage: + if ask_user_stage['error']: + await event_emitter({'type': 'chat:completion', 'data': {'output': full_output()}}) + continue + + if is_saved_chat_id(metadata.get('chat_id')) and metadata.get('message_id'): + await pause_for_tool_approval( + metadata['chat_id'], + metadata['message_id'], + full_output(), + form_data, + metadata, + ) + await event_emitter({'type': 'chat:completion', 'data': {'output': full_output()}}) + return # Append function_call items for each tool call # (Responses API already has them from streaming, so skip duplicates) @@ -4887,6 +5388,29 @@ async def streaming_chat_response_handler(response, ctx): } ) + tool_approval_mode = metadata.get('params', {}).get('tool_approval_mode', 'full') + if ( + tool_approval_mode == 'ask' + and is_saved_chat_id(metadata.get('chat_id')) + and metadata.get('message_id') + ): + await pause_for_tool_approval( + metadata['chat_id'], + metadata['message_id'], + full_output(), + form_data, + metadata, + ) + await event_emitter( + { + 'type': 'chat:completion', + 'data': { + 'output': full_output(), + }, + } + ) + return + await event_emitter( { 'type': 'chat:completion', @@ -5152,7 +5676,7 @@ async def streaming_chat_response_handler(response, ctx): # output sent to the frontend — they're only for LLM consumption # via convert_output_to_messages. frontend_output = [] - for item in output: + for item in full_output(): if item.get('type') == 'function_call_output': parts = item.get('output', []) if any(p.get('type') == 'input_image' for p in parts): @@ -5238,7 +5762,7 @@ async def streaming_chat_response_handler(response, ctx): # keeps indices aligned. The display prefix # ensures the UI shows tool history during # streaming. - prior_output = list(output) + prior_output = list(full_output()) # Trim the trailing empty placeholder message # so it doesn't persist as a ghost item once # the new stream produces real content. @@ -5280,7 +5804,7 @@ async def streaming_chat_response_handler(response, ctx): { 'type': 'chat:completion', 'data': { - 'output': output, + 'output': full_output(), }, } ) @@ -5407,7 +5931,7 @@ async def streaming_chat_response_handler(response, ctx): { 'type': 'chat:completion', 'data': { - 'output': output, + 'output': full_output(), }, } ) @@ -5451,40 +5975,32 @@ async def streaming_chat_response_handler(response, ctx): if item.get('status') == 'in_progress': item['status'] = 'completed' + current_output = full_output() title = await Chats.get_chat_title_by_id(metadata['chat_id']) if save_to_chat else '' data = { 'done': True, - 'output': output, + 'output': current_output, 'title': title, **({'usage': usage} if usage else {}), } if save_to_chat: - if not ENABLE_REALTIME_CHAT_SAVE: - # Save message in the database - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - { - 'done': True, - 'output': output, - **({'usage': usage} if usage else {}), - }, - ) - elif usage: - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - {'done': True, 'usage': usage}, - ) - else: - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - {'done': True}, - ) + # Save final output once. The delta path keeps in-progress + # state in response_streams instead of writing tokens to DB. + await Chats.upsert_message_to_chat_by_id_and_message_id( + metadata['chat_id'], + metadata['message_id'], + { + 'done': True, + 'output': current_output, + **({'usage': usage} if usage else {}), + }, + ) - await publish_chat_finished_event(request, user, metadata, title, ''.join(content_parts), output) + await clear_response_stream(request.app.state.redis, response_stream_task_id) + await publish_chat_finished_event( + request, user, metadata, title, ''.join(content_parts), current_output + ) await event_emitter( { @@ -5494,8 +6010,8 @@ async def streaming_chat_response_handler(response, ctx): ) ctx['assistant_message'] = { - 'content': ''.join(content_parts) or get_output_text(output), - 'output': output, + 'content': ''.join(content_parts) or get_output_text(current_output), + 'output': current_output, **({'usage': usage} if usage else {}), } await outlet_filter_handler(ctx) @@ -5516,22 +6032,15 @@ async def streaming_chat_response_handler(response, ctx): async def save_cancelled_state(): await event_emitter({'type': 'chat:tasks:cancel'}) if save_to_chat: - if not ENABLE_REALTIME_CHAT_SAVE: - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - { - 'done': True, - 'output': output, - }, - ) - else: - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - {'done': True}, - touch=False, - ) + await Chats.upsert_message_to_chat_by_id_and_message_id( + metadata['chat_id'], + metadata['message_id'], + { + 'done': True, + 'output': full_output(), + }, + ) + await clear_response_stream(request.app.state.redis, response_stream_task_id) try: await asyncio.shield(save_cancelled_state()) diff --git a/backend/open_webui/utils/misc.py b/backend/open_webui/utils/misc.py index 24253059d2..c24e6568d8 100644 --- a/backend/open_webui/utils/misc.py +++ b/backend/open_webui/utils/misc.py @@ -258,6 +258,17 @@ def reconcile_tool_pairs(messages: list[dict]) -> list[dict]: return reconciled_messages +def get_reasoning_details(payload: dict): + if not isinstance(payload, dict): + return None + + provider_fields = payload.get('provider_specific_fields') or {} + provider_details = ( + provider_fields.get('reasoning_details') if isinstance(provider_fields, dict) else None + ) + return payload.get('reasoning_details') or provider_details + + def convert_output_to_messages( output: list, raw: bool = False, @@ -277,8 +288,10 @@ def convert_output_to_messages( follow-ups. reasoning_format: How to include reasoning blocks in the output: - None: skip reasoning (default, safe for strict providers). + - ``'thinking'``: set as ``thinking`` top-level field + (for native Ollama). - ``'think_tags'``: wrap in ```` tags inside content - (for Ollama, which expects reasoning as tagged content). + (for legacy providers that expect reasoning as tagged content). - ``'reasoning_content'``: set as ``reasoning_content`` top-level field (for llama.cpp, which routes it via the chat template). flatten_tool_images: Move tool output images into a following user @@ -290,12 +303,21 @@ def convert_output_to_messages( messages = [] pending_tool_calls = [] pending_content = [] - pending_reasoning = [] # Only populated when reasoning_format == 'reasoning_content' + pending_reasoning = [] # Only populated for top-level structured reasoning fields. pending_reasoning_details = [] pending_tool_image_urls = [] - function_call_ids = { - item.get('call_id') for item in output if item.get('type') == 'function_call' and item.get('call_id') + pending_tool_outputs = [] + completed_call_ids = { + item.get('call_id') + for item in output + if item.get('type') == 'function_call' + and item.get('call_id') + and item.get('status') in {'completed', 'rejected'} } + result_call_ids = { + item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id') + } + function_call_ids = completed_call_ids & result_call_ids def flush_pending(): nonlocal pending_content, pending_tool_calls, pending_reasoning, pending_reasoning_details @@ -309,7 +331,10 @@ def convert_output_to_messages( } if pending_reasoning: - message['reasoning_content'] = '\n'.join(pending_reasoning) + if reasoning_format == 'thinking': + message['thinking'] = '\n'.join(pending_reasoning) + else: + message['reasoning_content'] = '\n'.join(pending_reasoning) if pending_reasoning_details: message['reasoning_details'] = pending_reasoning_details @@ -339,9 +364,60 @@ def convert_output_to_messages( ) pending_tool_image_urls = [] + def flush_tool_outputs(): + nonlocal pending_tool_outputs + if not pending_tool_outputs: + return + + flush_pending() + for output_item in pending_tool_outputs: + output_parts = output_item.get('output', []) + content = '' + image_urls = [] + for part in output_parts: + if part.get('type') == 'input_text': + output_text = part.get('text', '') + content += str(output_text) if not isinstance(output_text, str) else output_text + elif part.get('type') == 'input_image': + url = part.get('image_url', '') + if url: + image_urls.append(url) + + if flatten_tool_images: + messages.append( + { + 'role': 'tool', + 'tool_call_id': output_item.get('call_id', ''), + 'content': content, + } + ) + pending_tool_image_urls.extend(image_urls) + elif image_urls: + messages.append( + { + 'role': 'tool', + 'tool_call_id': output_item.get('call_id', ''), + 'content': [ + {'type': 'input_text', 'text': content}, + *[{'type': 'input_image', 'image_url': url} for url in image_urls], + ], + } + ) + else: + messages.append( + { + 'role': 'tool', + 'tool_call_id': output_item.get('call_id', ''), + 'content': content, + } + ) + + pending_tool_outputs = [] + for item in output: item_type = item.get('type', '') - if item_type != 'function_call_output': + if item_type not in {'function_call', 'function_call_output'}: + flush_tool_outputs() flush_tool_images() if item_type == 'message': @@ -355,6 +431,9 @@ def convert_output_to_messages( pending_content.append(text) elif item_type == 'function_call': + if item.get('call_id') not in function_call_ids: + continue + # Collect tool calls to batch into assistant message arguments = item.get('arguments', '{}') # Ensure arguments is always a JSON string @@ -372,54 +451,21 @@ def convert_output_to_messages( ) elif item_type == 'function_call_output': - # Flush any pending content/tool_calls before adding tool result - flush_pending() + if item.get('call_id') not in function_call_ids: + continue - # Extract text and images from output content parts - output_parts = item.get('output', []) - content = '' - image_urls = [] - for part in output_parts: - if part.get('type') == 'input_text': - output_text = part.get('text', '') - content += str(output_text) if not isinstance(output_text, str) else output_text - elif part.get('type') == 'input_image': - url = part.get('image_url', '') - if url: - image_urls.append(url) - - if flatten_tool_images: - messages.append( - { - 'role': 'tool', - 'tool_call_id': item.get('call_id', ''), - 'content': content, - } - ) - if item.get('call_id') in function_call_ids: - pending_tool_image_urls.extend(image_urls) - elif image_urls: - messages.append( - { - 'role': 'tool', - 'tool_call_id': item.get('call_id', ''), - 'content': [ - {'type': 'input_text', 'text': content}, - *[{'type': 'input_image', 'image_url': url} for url in image_urls], - ], - } - ) - else: - messages.append( - { - 'role': 'tool', - 'tool_call_id': item.get('call_id', ''), - 'content': content, - } - ) + pending_tool_outputs.append(item) elif item_type == 'reasoning': reasoning_details = item.get('reasoning_details') if raw else None + if reasoning_details: + reasoning_details = reasoning_details if isinstance(reasoning_details, list) else [reasoning_details] + reasoning_details = [ + detail + for detail in reasoning_details + if isinstance(detail, dict) + and (detail.get('format') != 'anthropic-claude-v1' or detail.get('signature')) + ] if not reasoning_format and not reasoning_details: continue @@ -433,18 +479,16 @@ def convert_output_to_messages( if reasoning_text: if reasoning_format == 'think_tags': - # Ollama: embed in content with the item's original tags + # Legacy tag replay: embed in content with the item's original tags. start_tag = item.get('start_tag', '') end_tag = item.get('end_tag', '') pending_content.append(f'{start_tag}{reasoning_text}{end_tag}') - elif reasoning_format == 'reasoning_content': - # llama.cpp: collect for reasoning_content field + elif reasoning_format in {'thinking', 'reasoning_content'}: + # Native providers: collect for their top-level reasoning field. pending_reasoning.append(reasoning_text) if reasoning_details: - pending_reasoning_details.extend( - reasoning_details if isinstance(reasoning_details, list) else [reasoning_details] - ) + pending_reasoning_details.extend(reasoning_details) elif item_type == 'open_webui:code_interpreter': # Always include code interpreter content so the LLM knows @@ -470,6 +514,7 @@ def convert_output_to_messages( pass # Flush remaining content/tool_calls + flush_tool_outputs() flush_tool_images() flush_pending() @@ -788,6 +833,17 @@ def sanitize_filename(file_name): return final_file_name +def json_text_variants(value: str) -> list[str]: + """Both spellings ``value`` can take inside a serialized JSON column, unquoted. + + Encoders disagree on non-ASCII — stdlib escapes it to ``\\uXXXX``, orjson writes it + raw — so a LIKE against the stored text has to accept either. ASCII collapses to one. + """ + raw = JSONCodec.dumps(value, ensure_ascii=False)[1:-1] + escaped = JSONCodec.dumps(value, ensure_ascii=True)[1:-1] + return [raw] if raw == escaped else [raw, escaped] + + def sanitize_text_for_db(text: str) -> str: """Remove null bytes and invalid UTF-8 surrogates from text for PostgreSQL storage.""" if not isinstance(text, str): diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index 4cdefd0de5..1c648b7f38 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -71,6 +71,10 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) 'evaluation.arena.models', 'models.default_metadata', ) + if refresh: + await openai.get_all_models.cache.clear() + await ollama.get_all_models.cache.clear() + if ( request.app.state.MODELS and request.app.state.BASE_MODELS diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index fce053dae8..fee218d086 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -1519,8 +1519,8 @@ class OAuthManager: oauth_allowed_roles = auth_config.OAUTH_ALLOWED_ROLES oauth_admin_roles = auth_config.OAUTH_ADMIN_ROLES oauth_roles = [] - # Default/fallback role if no matching roles are found - role = auth_config.DEFAULT_USER_ROLE + # Keep existing users at their current role unless the provider sent roles. + role = user.role if user else auth_config.DEFAULT_USER_ROLE # Next block extracts the roles from the user data, accepting nested claims of any depth if oauth_claim and oauth_allowed_roles and oauth_admin_roles: @@ -1585,7 +1585,34 @@ class OAuthManager: return role - async def update_user_groups(self, user, user_data, default_permissions, db=None): + async def update_user_role_from_oauth( + self, + request, + user, + user_data, + provider, + *, + db=None, + ): + determined_role = await self.get_user_role(user, user_data) + if user.role == determined_role: + return user + + updated_user = await Users.update_user_role_by_id(user.id, determined_role, db=db) + user = updated_user or user + user.role = determined_role + await publish_event( + request, + EVENTS.USER_ROLE_UPDATED, + actor=user, + subject_id=user.id, + source='oauth', + data={'role': determined_role, 'provider': provider}, + ) + + return user + + async def update_user_groups(self, request, user, user_data, default_permissions, db=None): auth_config = await get_oauth_runtime_config() log.debug('Running OAUTH Group management') oauth_claim = auth_config.OAUTH_GROUPS_CLAIM @@ -1650,6 +1677,13 @@ class OAuthManager: groups_created = True # Add to local set to prevent duplicate creation attempts in this run all_group_names.add(group_name) + await publish_event( + request, + EVENTS.GROUP_CREATED, + subject_id=created_group.id, + source='oauth', + data={'name': created_group.name}, + ) else: log.error(f"Failed to create group '{group_name}' via OAuth.") except Exception as e: @@ -1674,7 +1708,15 @@ class OAuthManager: ): # Remove group from user log.debug('Removing user from group %s as it is no longer in their oauth groups', group_model.name) - await Groups.remove_users_from_group(group_model.id, [user.id], db=db) + if await Groups.remove_users_from_group(group_model.id, [user.id], db=db): + await publish_event( + request, + EVENTS.GROUP_MEMBER_REMOVED, + actor=user, + subject_id=group_model.id, + source='oauth', + data={'user_ids': [user.id]}, + ) # In case a group is created, but perms are never assigned to the group by hitting "save" group_permissions = group_model.permissions @@ -1703,7 +1745,15 @@ class OAuthManager: # Add user to group log.debug('Adding user to group %s as it was found in their oauth groups', group_model.name) - await Groups.add_users_to_group(group_model.id, [user.id], db=db) + if await Groups.add_users_to_group(group_model.id, [user.id], db=db): + await publish_event( + request, + EVENTS.GROUP_MEMBER_ADDED, + actor=user, + subject_id=group_model.id, + source='oauth', + data={'user_ids': [user.id]}, + ) # In case a group is created, but perms are never assigned to the group by hitting "save" group_permissions = group_model.permissions @@ -1940,29 +1990,26 @@ class OAuthManager: await Users.update_user_oauth_by_id(user.id, provider, sub, db=db) if user: - determined_role = await self.get_user_role(user, user_data) - if user.role != determined_role: - 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}, - ) + user = await self.update_user_role_from_oauth( + request=request, + user=user, + user_data=user_data, + provider=provider, + db=db, + ) + + updated_fields = [] if auth_config.OAUTH_UPDATE_NAME_ON_LOGIN: username_claim = auth_config.OAUTH_USERNAME_CLAIM if username_claim: new_name = user_data.get(username_claim) if new_name and new_name != user.name: - await Users.update_user_by_id(user.id, {'name': new_name}, db=db) - user.name = new_name - log.debug('Updated name for user %s', user.email) + updated_user = await Users.update_user_by_id(user.id, {'name': new_name}, db=db) + if updated_user: + user = updated_user + updated_fields.append('name') + log.debug('Updated name for user %s', user.email) if auth_config.OAUTH_UPDATE_EMAIL_ON_LOGIN: email_claim = auth_config.OAUTH_EMAIL_CLAIM @@ -1974,9 +2021,9 @@ class OAuthManager: log.error( f'Cannot update email to {new_email} for user {user.id} because it is already taken.' ) - else: - await Auths.update_email_by_id(user.id, new_email.lower(), db=db) - user.email = new_email.lower() + elif await Auths.update_email_by_id(user.id, new_email.lower(), db=db): + user = await Users.get_user_by_id(user.id, db=db) or user + updated_fields.append('email') log.debug('Updated email for user %s', user.id) # Update profile picture if enabled and different from current @@ -1991,8 +2038,23 @@ class OAuthManager: new_picture_url, token.get('access_token') ) if processed_picture_url != user.profile_image_url: - await Users.update_user_profile_image_url_by_id(user.id, processed_picture_url, db=db) - log.debug('Updated profile picture for user %s', user.email) + updated_user = await Users.update_user_profile_image_url_by_id( + user.id, processed_picture_url, db=db + ) + if updated_user: + user = updated_user + updated_fields.append('profile_image_url') + log.debug('Updated profile picture for user %s', user.email) + + if updated_fields: + await publish_event( + request, + EVENTS.USER_UPDATED, + actor=user, + subject_id=user.id, + source='oauth', + data={'updated_fields': updated_fields, 'provider': provider}, + ) else: # If the user does not exist, check if signups are enabled if auth_config.ENABLE_OAUTH_SIGNUP: @@ -2060,6 +2122,7 @@ class OAuthManager: ) if auth_config.ENABLE_OAUTH_GROUP_MANAGEMENT: await self.update_user_groups( + request=request, user=user, user_data=user_data, default_permissions=await Config.get('user.permissions'), diff --git a/backend/open_webui/utils/payload.py b/backend/open_webui/utils/payload.py index a0cd4788e7..7ee2f9d3b3 100644 --- a/backend/open_webui/utils/payload.py +++ b/backend/open_webui/utils/payload.py @@ -95,6 +95,7 @@ def apply_params_to_form_data(form_data: dict, model: dict, params: dict | None 'compact_token_threshold': int, 'system': str, 'note_id': str, + 'tool_approval_mode': str, } for key in list(params.keys()): @@ -149,6 +150,7 @@ def remove_open_webui_params(params: dict) -> dict: 'compact_token_threshold': int, 'system': str, 'note_id': str, + 'tool_approval_mode': str, } for key in list(params.keys()): @@ -284,6 +286,8 @@ def convert_messages_openai_to_ollama(messages: list[dict]) -> list[dict]: # may be injected by filter inlet functions). if 'thinking' in message: new_message['thinking'] = message['thinking'] + elif reasoning_content := (message.get('reasoning_content') or message.get('reasoning')): + new_message['thinking'] = reasoning_content content = message.get('content', []) tool_calls = message.get('tool_calls', None) diff --git a/backend/open_webui/utils/response.py b/backend/open_webui/utils/response.py index ef229b1729..3f34986903 100644 --- a/backend/open_webui/utils/response.py +++ b/backend/open_webui/utils/response.py @@ -24,13 +24,9 @@ def normalize_usage(usage: dict) -> dict: return {} # Map various field names to standard names - input_tokens = ( - usage.get('input_tokens') # Already standard - or usage.get('prompt_tokens') # OpenAI - or usage.get('prompt_eval_count') # Ollama - or usage.get('prompt_n') # llama.cpp - or 0 - ) + input_tokens = usage.get('input_tokens') or usage.get('prompt_tokens') or usage.get('prompt_eval_count') + if input_tokens is None: + input_tokens = int(usage.get('prompt_n') or 0) + int(usage.get('cache_n') or 0) output_tokens = ( usage.get('output_tokens') # Already standard diff --git a/backend/open_webui/utils/timers.py b/backend/open_webui/utils/timers.py index 2546b311ae..46d1e3312f 100644 --- a/backend/open_webui/utils/timers.py +++ b/backend/open_webui/utils/timers.py @@ -404,7 +404,11 @@ async def execute_due_timer(app, timer_id: str, claim_id: str | None = None) -> ) request.state.token = None request.state.enable_api_keys = False - await app.state.CHAT_COMPLETION_HANDLER(request, form_data, user=user) + try: + await app.state.CHAT_COMPLETION_HANDLER(request, form_data, user=user) + except Exception as exc: + log.exception(f'Timer {timer_id} completion failed') + await _set_timer_state(timer_id, 'error', timer_error=str(exc)[:500]) async def _set_timer_state(timer_id: str, status: str, **fields) -> None: diff --git a/backend/open_webui/utils/tool_approval.py b/backend/open_webui/utils/tool_approval.py new file mode 100644 index 0000000000..81fe750ec6 --- /dev/null +++ b/backend/open_webui/utils/tool_approval.py @@ -0,0 +1,187 @@ +from typing import Any, Literal + +from fastapi import HTTPException, status +from pydantic import BaseModel +from sqlalchemy.ext.asyncio import AsyncSession + +from open_webui.constants import ERROR_MESSAGES +from open_webui.models.chats import Chats +from open_webui.socket.main import get_event_emitter +from open_webui.utils.json_codec import JSONCodec + + +class ResolveToolCallForm(BaseModel): + call_id: str + action: Literal['approve', 'reject', 'answer'] + answers: Any | None = None + timed_out: bool = False + + +async def resolve_tool_call_output( + chat_id: str, + message_id: str, + form_data: ResolveToolCallForm, + 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 user.role != 'admin'): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.ACCESS_PROHIBITED, + ) + + message = await Chats.get_message_by_id_and_message_id(chat_id, message_id) + if not message: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) + + 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.') + + function_call = next( + ( + item + for item in output + if item.get('type') == 'function_call' + and (item.get('call_id') or item.get('id')) == form_data.call_id + ), + None, + ) + if not function_call: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail='Tool call not found.') + function_call.setdefault('call_id', form_data.call_id) + tool_name = function_call.get('name') + + if any( + item.get('type') == 'function_call_output' and item.get('call_id') == form_data.call_id for item in output + ) or function_call.get('status') not in {'pending', 'queued', 'requires_approval'}: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call has already been resolved.') + + if form_data.action == 'approve': + if tool_name == 'ask_user': + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='ask_user requires an answer or deny.') + function_call['status'] = 'queued' + function_call['approved'] = True + elif form_data.action == 'reject': + function_call['status'] = 'rejected' + output.append( + { + 'type': 'function_call_output', + 'id': f'fco_{form_data.call_id}', + 'call_id': form_data.call_id, + 'output': [{'type': 'input_text', 'text': 'Error: tool call rejected by user.'}], + 'status': 'rejected', + } + ) + else: + if tool_name != 'ask_user': + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Tool call does not accept answers.') + if form_data.answers is None and not form_data.timed_out: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Answers are required for ask_user.') + function_call['status'] = 'completed' + answer_payload = ( + {'status': 'cancelled', 'answers': {}, 'timed_out': True} + if form_data.timed_out + else {'status': 'answered', 'answers': form_data.answers or {}} + ) + output.append( + { + 'type': 'function_call_output', + 'id': f'fco_{form_data.call_id}', + 'call_id': form_data.call_id, + 'output': [{'type': 'input_text', 'text': JSONCodec.dumps(answer_payload)}], + 'status': 'completed', + } + ) + + await Chats.upsert_message_to_chat_by_id_and_message_id( + chat_id, + message_id, + { + 'done': False, + 'output': output, + }, + touch=False, + ) + + event_emitter = await get_event_emitter( + { + 'user_id': chat.user_id, + 'chat_id': chat_id, + 'message_id': message_id, + }, + update_db=False, + ) + if event_emitter: + await event_emitter({'type': 'chat:completion', 'data': {'output': output}}) + + result_call_ids = { + item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id') + } + paused = any( + item.get('type') == 'function_call' + and item.get('call_id') + and item.get('status') in {'pending', 'queued', 'requires_approval'} + and item.get('call_id') not in result_call_ids + for item in output + ) + return {'chat': chat, 'message': message, 'output': output, 'paused': paused} + + +async def build_tool_approval_resume_payload(chat_id: str, message_id: str, chat=None) -> dict: + chat = chat or await Chats.get_chat_by_id(chat_id) + if not chat: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) + + assistant_message = await Chats.get_message_by_id_and_message_id(chat_id, message_id) + if not assistant_message: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND) + + user_message_id = assistant_message.get('parentId') + user_message = await Chats.get_message_by_id_and_message_id(chat_id, user_message_id) if user_message_id else None + if not user_message: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call parent message is missing.') + + 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 {} + params = { + **chat_params, + **(message_meta.get('params') if isinstance(message_meta.get('params'), dict) else {}), + } + current_approval_mode = chat_params.get('tool_approval_mode') + if current_approval_mode in {'ask', 'full'}: + params['tool_approval_mode'] = current_approval_mode + if 'tool_approval_mode' not in params: + params['tool_approval_mode'] = 'ask' + + model_id = assistant_message.get('model') or next(iter(chat_data.get('models') or []), None) + if not model_id: + raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call message model is missing.') + + messages = [] + if params.get('system'): + messages.append({'role': 'system', 'content': params.get('system')}) + + return { + 'stream': params.get('stream_response', True), + 'model': model_id, + 'messages': messages, + 'params': params, + 'files': message_meta.get('files') or chat_data.get('files') or None, + '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, + 'terminal_id': message_meta.get('terminal_id') or None, + '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, + 'session_id': message_meta.get('session_id'), + 'chat_id': chat_id, + 'id': message_id, + 'parent_id': user_message.get('parentId'), + 'user_message': user_message, + 'assistant_message_id': message_id, + } diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 5a3852b89d..49f6c46aec 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -47,6 +47,7 @@ from open_webui.models.tools import Tools from open_webui.models.users import UserModel from open_webui.tools.builtin import ( add_memory, + ask_user, calculate_timestamp, create_automation, create_calendar_event, @@ -442,7 +443,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr ) headers.setdefault('Content-Type', 'application/json') - async def make_tool_function(function_name, tool_server_data, headers): + async def make_tool_function(function_name, tool_server_data, headers, cookies): async def tool_function(**kwargs): return await execute_tool_server( url=tool_server_data['url'], @@ -455,7 +456,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr return tool_function - tool_function = await make_tool_function(function_name, tool_server_data, headers) + tool_function = await make_tool_function(function_name, tool_server_data, headers, cookies) callable = await get_async_tool_function_and_apply_extra_params( tool_function, @@ -541,9 +542,9 @@ async def get_builtin_tools( # Helper to check if a builtin tool category is enabled via meta.builtinTools # Defaults to True if not specified (backward compatible) - def is_builtin_tool_enabled(category: str) -> bool: + def is_builtin_tool_enabled(category: str, default: bool = True) -> bool: builtin_tools = model.get('info', {}).get('meta', {}).get('builtinTools', {}) - return builtin_tools.get(category, True) + return builtin_tools.get(category, default) # Helper to check user-level feature permission (admins always pass) user = extra_params.get('__user__', {}) @@ -583,6 +584,9 @@ async def get_builtin_tools( if is_builtin_tool_enabled('time'): builtin_functions.extend([get_current_timestamp, calculate_timestamp]) + if is_builtin_tool_enabled('user_input', True): + builtin_functions.append(ask_user) + metadata = extra_params.get('__metadata__') or {} chat_files = metadata.get('files') or extra_params.get('__files__') or [] has_chat_files = any( @@ -1166,15 +1170,17 @@ async def set_tool_servers(request: Request): async def get_tool_servers(request: Request): try: - tool_servers = [] + tool_servers = None if request.app.state.redis is not None: try: - tool_servers = JSONCodec.loads(await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:tool_servers')) - request.app.state.TOOL_SERVERS = tool_servers + data = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:tool_servers') + if data is not None: + tool_servers = JSONCodec.loads(data) + request.app.state.TOOL_SERVERS = tool_servers except Exception as e: log.error(f'Error fetching tool_servers from Redis: {e}') - if not tool_servers: + if tool_servers is None: tool_servers = await set_tool_servers(request) return tool_servers @@ -1309,15 +1315,17 @@ async def set_terminal_servers(request: Request): async def get_terminal_servers(request: Request): """Return cached terminal server specs, loading if needed.""" - terminal_servers = [] + terminal_servers = None if request.app.state.redis is not None: try: - terminal_servers = JSONCodec.loads(await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:terminal_servers')) - request.app.state.TERMINAL_SERVERS = terminal_servers + data = await request.app.state.redis.get(f'{REDIS_KEY_PREFIX}:terminal_servers') + if data is not None: + terminal_servers = JSONCodec.loads(data) + request.app.state.TERMINAL_SERVERS = terminal_servers except Exception as e: log.error(f'Error fetching terminal_servers from Redis: {e}') - if not terminal_servers: + if terminal_servers is None: terminal_servers = await set_terminal_servers(request) return terminal_servers diff --git a/src/app.css b/src/app.css index e1dde4bb77..a5be2c2d78 100644 --- a/src/app.css +++ b/src/app.css @@ -409,11 +409,15 @@ input[type='number'] { outline: none; } -/* unlayered, so it beats the outline-none utilities used across the app */ -:focus-visible, -.ProseMirror:focus-visible { +:focus-visible:not(.ProseMirror):not([contenteditable='true']):not(input):not(textarea), +.focus-ring:focus-visible { outline: 2px solid theme(--color-blue-500); - outline-offset: 2px; + outline-offset: -2px; +} + +html.high-contrast :is(input, textarea, [contenteditable='true'], .ProseMirror):focus-visible { + outline: 2px solid theme(--color-blue-500); + outline-offset: -2px; } .ProseMirror p.is-editor-empty:first-child::before { diff --git a/src/lib/apis/automations/index.ts b/src/lib/apis/automations/index.ts index 593a56cab1..1a427e81d3 100644 --- a/src/lib/apis/automations/index.ts +++ b/src/lib/apis/automations/index.ts @@ -5,11 +5,17 @@ export type AutomationTerminalConfig = { cwd?: string; }; +export type AutomationTarget = { + type: 'chat' | 'channel'; + channel_id?: string | null; +}; + export type AutomationData = { prompt: string; model_id: string; rrule: string; terminal?: AutomationTerminalConfig; + target?: AutomationTarget | null; }; export type AutomationForm = { diff --git a/src/lib/apis/chats/index.ts b/src/lib/apis/chats/index.ts index a659fd9b08..5dca7283ba 100644 --- a/src/lib/apis/chats/index.ts +++ b/src/lib/apis/chats/index.ts @@ -1369,6 +1369,46 @@ export const deleteChatMessageById = async (token: string, id: string, messageId return res; }; +export const resolveChatMessageToolCall = async ( + token: string, + id: string, + messageId: string, + callId: string, + action: 'approve' | 'reject' | 'answer', + options: { answers?: unknown; timed_out?: boolean } = {} +) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/chats/${id}/messages/${messageId}/resolve`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + ...(token && { authorization: `Bearer ${token}` }) + }, + body: JSON.stringify({ + call_id: callId, + action, + ...options + }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = getErrorDetail(err); + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const deleteChatById = async (token: string, id: string) => { let error = null; diff --git a/src/lib/apis/index.ts b/src/lib/apis/index.ts index 4f55ab14be..3409b88ed5 100644 --- a/src/lib/apis/index.ts +++ b/src/lib/apis/index.ts @@ -1758,6 +1758,7 @@ export interface ModelConfig { export interface ModelMeta { toolIds: never[]; description?: string; + hidden?: boolean; capabilities?: object; profile_image_url?: string; } diff --git a/src/lib/apis/knowledge/index.ts b/src/lib/apis/knowledge/index.ts index f0a6e4d140..3714e46752 100644 --- a/src/lib/apis/knowledge/index.ts +++ b/src/lib/apis/knowledge/index.ts @@ -955,6 +955,34 @@ export const reindexKnowledgeFiles = async (token: string) => { return res; }; +export const reindexKnowledgeMetadata = async (token: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/knowledge/metadata/reindex`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const exportKnowledgeById = async (token: string, id: string) => { let error = null; diff --git a/src/lib/apis/memories/index.ts b/src/lib/apis/memories/index.ts index 5ebee8fb47..61547a84b6 100644 --- a/src/lib/apis/memories/index.ts +++ b/src/lib/apis/memories/index.ts @@ -128,6 +128,34 @@ export const queryMemory = async (token: string, content: string) => { return res; }; +export const reindexMemoryVectors = async (token: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/memories/reindex`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = err.detail; + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const deleteMemoryById = async (token: string, id: string) => { let error = null; diff --git a/src/lib/apis/openai/index.ts b/src/lib/apis/openai/index.ts index d18565fec3..b89571428a 100644 --- a/src/lib/apis/openai/index.ts +++ b/src/lib/apis/openai/index.ts @@ -1,5 +1,18 @@ import { OPENAI_API_BASE_URL, WEBUI_API_BASE_URL, WEBUI_BASE_URL } from '$lib/constants'; +export const getErrorMessage = (err: any, fallback = 'Server connection failed') => { + const detail = err?.detail; + if (typeof detail === 'string') return detail; + + return ( + detail?.error?.message ?? + detail?.message ?? + err?.error?.message ?? + err?.message ?? + (typeof err === 'string' ? err : fallback) + ); +}; + export const getOpenAIConfig = async (token: string = '') => { let error = null; @@ -17,11 +30,7 @@ export const getOpenAIConfig = async (token: string = '') => { }) .catch((err) => { console.error(err); - if ('detail' in err) { - error = err.detail; - } else { - error = 'Server connection failed'; - } + error = getErrorMessage(err); return null; }); @@ -59,11 +68,7 @@ export const updateOpenAIConfig = async (token: string = '', config: OpenAIConfi }) .catch((err) => { console.error(err); - if ('detail' in err) { - error = err.detail; - } else { - error = 'Server connection failed'; - } + error = getErrorMessage(err); return null; }); @@ -131,9 +136,197 @@ export const getOpenAIModels = async (token: string, urlIdx?: number) => { return res; }; +export const getProviderModelCatalog = async (token: string, urlIdx: number) => { + let error = null; + + const res = await fetch(`${OPENAI_API_BASE_URL}/models/${urlIdx}/catalog`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + ...(token && { authorization: `Bearer ${token}` }) + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = getErrorMessage(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const downloadProviderModel = async ( + token: string, + urlIdx: number, + model: string, + signal?: AbortSignal +) => { + let error = null; + + const res = await fetch(`${OPENAI_API_BASE_URL}/models/${urlIdx}/download`, { + signal, + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + }, + body: JSON.stringify({ model }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = getErrorMessage(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const getProviderModelDownloadStatus = async ( + token: string, + urlIdx: number, + jobId: string, + signal?: AbortSignal +) => { + let error = null; + + const res = await fetch( + `${OPENAI_API_BASE_URL}/models/${urlIdx}/download/status/${encodeURIComponent(jobId)}`, + { + signal, + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + } + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = getErrorMessage(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const loadProviderModel = async (token: string, urlIdx: number, model: string) => { + let error = null; + + const res = await fetch(`${OPENAI_API_BASE_URL}/models/${urlIdx}/load`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + }, + body: JSON.stringify({ model }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = getErrorMessage(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const unloadProviderModel = async ( + token: string, + urlIdx: number, + model: string, + instanceId?: string +) => { + let error = null; + + const res = await fetch(`${OPENAI_API_BASE_URL}/models/${urlIdx}/unload`, { + method: 'POST', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + }, + body: JSON.stringify({ model, ...(instanceId ? { instance_id: instanceId } : {}) }) + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = getErrorMessage(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + +export const deleteProviderModel = async (token: string, urlIdx: number, model: string) => { + let error = null; + + const res = await fetch( + `${OPENAI_API_BASE_URL}/models/${urlIdx}?${new URLSearchParams({ model })}`, + { + method: 'DELETE', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + } + } + ) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .catch((err) => { + error = getErrorMessage(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const verifyOpenAIConnection = async ( token: string = '', - connection: dict = {}, + connection: Record = {}, direct: boolean = false ) => { const { url, key, config } = connection; @@ -246,7 +439,7 @@ export const generateOpenAIChatCompletion = async ( return res.json(); }) .catch((err) => { - error = err?.detail ?? err; + error = getErrorMessage(err); return null; }); diff --git a/src/lib/apis/retrieval/index.ts b/src/lib/apis/retrieval/index.ts index 99801dfedb..f913e3e868 100644 --- a/src/lib/apis/retrieval/index.ts +++ b/src/lib/apis/retrieval/index.ts @@ -188,6 +188,8 @@ type OpenAIConfigForm = { url: string; }; +type OllamaConfigForm = OpenAIConfigForm; + type AzureOpenAIConfigForm = { key: string; url: string; @@ -196,10 +198,13 @@ type AzureOpenAIConfigForm = { type EmbeddingModelUpdateForm = { openai_config?: OpenAIConfigForm; + ollama_config?: OllamaConfigForm; azure_openai_config?: AzureOpenAIConfigForm; - embedding_engine: string; - embedding_model: string; - embedding_batch_size?: number; + RAG_EMBEDDING_ENGINE: string; + RAG_EMBEDDING_MODEL: string; + RAG_EMBEDDING_BATCH_SIZE?: number; + ENABLE_ASYNC_EMBEDDING?: boolean; + RAG_EMBEDDING_CONCURRENT_REQUESTS?: number; }; export const updateEmbeddingConfig = async (token: string, payload: EmbeddingModelUpdateForm) => { diff --git a/src/lib/apis/skills/index.ts b/src/lib/apis/skills/index.ts index fc1dc24ce1..5185d5efaf 100644 --- a/src/lib/apis/skills/index.ts +++ b/src/lib/apis/skills/index.ts @@ -31,10 +31,13 @@ export const createNewSkill = async (token: string, skill: object) => { return res; }; -export const getSkills = async (token: string = '') => { +export const getSkills = async (token: string = '', query: string | null = null) => { let error = null; - const res = await fetch(`${WEBUI_API_BASE_URL}/skills/`, { + const searchParams = new URLSearchParams(); + if (query) searchParams.append('query', query); + + const res = await fetch(`${WEBUI_API_BASE_URL}/skills/?${searchParams.toString()}`, { method: 'GET', headers: { Accept: 'application/json', diff --git a/src/lib/apis/tools/index.ts b/src/lib/apis/tools/index.ts index 5d26e50fee..8378299923 100644 --- a/src/lib/apis/tools/index.ts +++ b/src/lib/apis/tools/index.ts @@ -65,10 +65,13 @@ export const loadToolByUrl = async (token: string = '', url: string) => { return res; }; -export const getTools = async (token: string = '') => { +export const getTools = async (token: string = '', query: string | null = null) => { let error = null; - const res = await fetch(`${WEBUI_API_BASE_URL}/tools/`, { + const searchParams = new URLSearchParams(); + if (query) searchParams.append('query', query); + + const res = await fetch(`${WEBUI_API_BASE_URL}/tools/?${searchParams.toString()}`, { method: 'GET', headers: { Accept: 'application/json', diff --git a/src/lib/apis/users/index.ts b/src/lib/apis/users/index.ts index b48c3b27cd..a6da772d18 100644 --- a/src/lib/apis/users/index.ts +++ b/src/lib/apis/users/index.ts @@ -286,7 +286,7 @@ export const getUserSettings = async (token: string, raw = false) => { }) .catch((err) => { console.error(err); - error = err.detail; + error = err?.detail ?? err; return null; }); diff --git a/src/lib/components/AddConnectionModal.svelte b/src/lib/components/AddConnectionModal.svelte index 68eb96b70f..2c7ed22a87 100644 --- a/src/lib/components/AddConnectionModal.svelte +++ b/src/lib/components/AddConnectionModal.svelte @@ -606,6 +606,7 @@ + diff --git a/src/lib/components/AutomationModal.svelte b/src/lib/components/AutomationModal.svelte index 23028b3bc8..6a709386df 100644 --- a/src/lib/components/AutomationModal.svelte +++ b/src/lib/components/AutomationModal.svelte @@ -8,9 +8,10 @@ import ScheduleDropdown from '$lib/components/automations/ScheduleDropdown.svelte'; import ModelDropdown from '$lib/components/automations/ModelDropdown.svelte'; - import FolderDropdown from '$lib/components/automations/FolderDropdown.svelte'; + import DestinationDropdown from '$lib/components/automations/DestinationDropdown.svelte'; import { getFolders } from '$lib/apis/folders'; - import { folders } from '$lib/stores'; + import { getChannels } from '$lib/apis/channels'; + import { channels, folders } from '$lib/stores'; import { createAutomation, @@ -30,10 +31,13 @@ let prompt = ''; let model_id = ''; let folder_id = ''; + let target_type: 'chat' | 'channel' = 'chat'; + let channel_id = ''; let is_active = true; let loading = false; let foldersLoaded = false; + let channelsLoaded = false; // Schedule dropdown ref let scheduleDropdown: ScheduleDropdown; @@ -43,6 +47,10 @@ toast.error($i18n.t('Name, prompt, and model are required')); return; } + if (target_type === 'channel' && !channel_id) { + toast.error($i18n.t('Channel is required')); + return; + } if (scheduleDropdown?.frequency === 'ONCE') { const scheduled = new Date(`${scheduleDropdown.onceDate}T${scheduleDropdown.onceTime}`); if (scheduled <= new Date()) { @@ -54,11 +62,12 @@ try { const form: AutomationForm = { name: name.trim(), - folder_id: folder_id || null, + folder_id: target_type === 'channel' ? null : folder_id || null, data: { prompt: prompt.trim(), model_id: model_id.trim(), - rrule: scheduleDropdown.buildRrule() + rrule: scheduleDropdown.buildRrule(), + target: target_type === 'channel' ? { type: 'channel', channel_id } : { type: 'chat' } }, is_active }; @@ -88,12 +97,19 @@ if (res) folders.set(res); foldersLoaded = true; } + if (!channelsLoaded && ($channels ?? []).length === 0) { + const res = await getChannels(localStorage.token).catch(() => null); + if (res) channels.set(res); + channelsLoaded = true; + } if (automation) { name = automation.name; prompt = automation.data.prompt; model_id = automation.data.model_id; folder_id = automation.folder_id ?? ''; + target_type = automation.data.target?.type === 'channel' ? 'channel' : 'chat'; + channel_id = automation.data.target?.channel_id ?? ''; is_active = automation.is_active; if (scheduleDropdown) { scheduleDropdown.parseRrule(automation.data.rrule); @@ -105,6 +121,12 @@ folder_id = ($folders ?? []).some((folder) => folder.id === cloneFrom.folder_id) ? (cloneFrom.folder_id ?? '') : ''; + target_type = cloneFrom.data.target?.type === 'channel' ? 'channel' : 'chat'; + channel_id = ($channels ?? []).some( + (channel) => channel.id === cloneFrom.data.target?.channel_id + ) + ? (cloneFrom.data.target?.channel_id ?? '') + : ''; is_active = true; if (scheduleDropdown) { scheduleDropdown.parseRrule(cloneFrom.data.rrule); @@ -114,6 +136,8 @@ prompt = ''; model_id = ''; folder_id = ''; + target_type = 'chat'; + channel_id = ''; is_active = true; } }; @@ -162,7 +186,15 @@ - +
diff --git a/src/lib/components/ChangelogModal.svelte b/src/lib/components/ChangelogModal.svelte index 5e429ab72d..aff4b15fce 100644 --- a/src/lib/components/ChangelogModal.svelte +++ b/src/lib/components/ChangelogModal.svelte @@ -154,7 +154,7 @@ class="mt-[0.6em] h-1 w-1 shrink-0 rounded-full bg-gray-300 dark:bg-gray-700" >
{@html DOMPurify.sanitize(entry?.raw)} diff --git a/src/lib/components/admin/Evaluations/Feedbacks.svelte b/src/lib/components/admin/Evaluations/Feedbacks.svelte index b61d65e1a8..d9bade0b7e 100644 --- a/src/lib/components/admin/Evaluations/Feedbacks.svelte +++ b/src/lib/components/admin/Evaluations/Feedbacks.svelte @@ -304,11 +304,10 @@
{#if (items ?? []).length === 0} -
-
-
😕
-
{$i18n.t('No feedback found')}
-
+
+
+
{$i18n.t('No feedback found')}
+
{$i18n.t('Try adjusting your search or filter to find what you are looking for.')}
diff --git a/src/lib/components/admin/Settings/Connections.svelte b/src/lib/components/admin/Settings/Connections.svelte index f7b43bfb4b..fe5278d703 100644 --- a/src/lib/components/admin/Settings/Connections.svelte +++ b/src/lib/components/admin/Settings/Connections.svelte @@ -14,6 +14,7 @@ import Switch from '$lib/components/common/Switch.svelte'; import Spinner from '$lib/components/common/Spinner.svelte'; import Tooltip from '$lib/components/common/Tooltip.svelte'; + import ArrowPath from '$lib/components/icons/ArrowPath.svelte'; import Plus from '$lib/components/icons/Plus.svelte'; import OpenAIConnection from './Connections/OpenAIConnection.svelte'; @@ -50,6 +51,7 @@ let pipelineUrls: Record = {}; let showAddOpenAIConnectionModal = false; let showAddOllamaConnectionModal = false; + let modelListRefreshing = false; const updateOpenAIHandler = async () => { if (ENABLE_OPENAI_API !== null) { @@ -120,6 +122,19 @@ } }; + const refreshModelListHandler = async () => { + modelListRefreshing = true; + + try { + await models.set(await getModels()); + toast.success($i18n.t('Model list refreshed')); + } catch (error) { + toast.error(`${error}`); + } finally { + modelListRefreshing = false; + } + }; + const addOpenAIConnectionHandler = async (connection: any) => { OPENAI_API_BASE_URLS = [...OPENAI_API_BASE_URLS, connection.url]; OPENAI_API_KEYS = [...OPENAI_API_KEYS, connection.key]; @@ -373,13 +388,33 @@ )} let:labelId > - { - updateConnectionsHandler(); - }} - ariaLabelledbyId={labelId} - /> +
+ {#if connectionsConfig.ENABLE_BASE_MODELS_CACHE} + + + + {/if} + + { + updateConnectionsHandler(); + }} + ariaLabelledbyId={labelId} + /> +
{:else} diff --git a/src/lib/components/admin/Settings/Documents.svelte b/src/lib/components/admin/Settings/Documents.svelte index 3f2911ff36..8425f338ea 100644 --- a/src/lib/components/admin/Settings/Documents.svelte +++ b/src/lib/components/admin/Settings/Documents.svelte @@ -17,12 +17,13 @@ updateRAGConfig } from '$lib/apis/retrieval'; - import { reindexKnowledgeFiles } from '$lib/apis/knowledge'; + import { reindexKnowledgeFiles, reindexKnowledgeMetadata } from '$lib/apis/knowledge'; + import { reindexMemoryVectors } from '$lib/apis/memories'; import { deleteAllFiles } from '$lib/apis/files'; import ResetUploadDirConfirmDialog from '$lib/components/common/ConfirmDialog.svelte'; import ResetVectorDBConfirmDialog from '$lib/components/common/ConfirmDialog.svelte'; - import ReindexKnowledgeFilesConfirmDialog from '$lib/components/common/ConfirmDialog.svelte'; + import ReindexEmbeddingDataConfirmDialog from '$lib/components/common/ConfirmDialog.svelte'; import SensitiveInput from '$lib/components/common/SensitiveInput.svelte'; import Tooltip from '$lib/components/common/Tooltip.svelte'; import Switch from '$lib/components/common/Switch.svelte'; @@ -121,26 +122,33 @@ }); updateEmbeddingModelLoading = true; - const res = await updateEmbeddingConfig(localStorage.token, { + const payload: Parameters[1] = { RAG_EMBEDDING_ENGINE: RAG_EMBEDDING_ENGINE, RAG_EMBEDDING_MODEL: RAG_EMBEDDING_MODEL, RAG_EMBEDDING_BATCH_SIZE: RAG_EMBEDDING_BATCH_SIZE, ENABLE_ASYNC_EMBEDDING: ENABLE_ASYNC_EMBEDDING, - RAG_EMBEDDING_CONCURRENT_REQUESTS: RAG_EMBEDDING_CONCURRENT_REQUESTS, - ollama_config: { + RAG_EMBEDDING_CONCURRENT_REQUESTS: RAG_EMBEDDING_CONCURRENT_REQUESTS + }; + + if (RAG_EMBEDDING_ENGINE === 'ollama') { + payload.ollama_config = { key: OllamaKey, url: OllamaUrl - }, - openai_config: { + }; + } else if (RAG_EMBEDDING_ENGINE === 'openai') { + payload.openai_config = { key: OpenAIKey, url: OpenAIUrl - }, - azure_openai_config: { + }; + } else if (RAG_EMBEDDING_ENGINE === 'azure_openai') { + payload.azure_openai_config = { key: AzureOpenAIKey, url: AzureOpenAIUrl, version: AzureOpenAIVersion - } - }).catch(async (error) => { + }; + } + + const res = await updateEmbeddingConfig(localStorage.token, payload).catch(async (error) => { toast.error(`${error}`); await setEmbeddingConfig(); return null; @@ -299,15 +307,15 @@ ENABLE_ASYNC_EMBEDDING = embeddingConfig.ENABLE_ASYNC_EMBEDDING ?? true; RAG_EMBEDDING_CONCURRENT_REQUESTS = embeddingConfig.RAG_EMBEDDING_CONCURRENT_REQUESTS ?? 0; - OpenAIKey = embeddingConfig.openai_config.key; - OpenAIUrl = embeddingConfig.openai_config.url; + OpenAIKey = embeddingConfig.openai_config.key ?? ''; + OpenAIUrl = embeddingConfig.openai_config.url ?? ''; - OllamaKey = embeddingConfig.ollama_config.key; - OllamaUrl = embeddingConfig.ollama_config.url; + OllamaKey = embeddingConfig.ollama_config.key ?? ''; + OllamaUrl = embeddingConfig.ollama_config.url ?? ''; - AzureOpenAIKey = embeddingConfig.azure_openai_config.key; - AzureOpenAIUrl = embeddingConfig.azure_openai_config.url; - AzureOpenAIVersion = embeddingConfig.azure_openai_config.version; + AzureOpenAIKey = embeddingConfig.azure_openai_config.key ?? ''; + AzureOpenAIUrl = embeddingConfig.azure_openai_config.url ?? ''; + AzureOpenAIVersion = embeddingConfig.azure_openai_config.version ?? ''; } }; onMount(async () => { @@ -371,15 +379,37 @@ }} /> - { - const res = await reindexKnowledgeFiles(localStorage.token).catch((error) => { + const knowledgeRes = await reindexKnowledgeFiles(localStorage.token).catch((error) => { + toast.error(`${error}`); + return null; + }); + if (!knowledgeRes) { + return; + } + + const knowledgeMetadataRes = await reindexKnowledgeMetadata(localStorage.token).catch( + (error) => { + toast.error(`${error}`); + return null; + } + ); + if (!knowledgeMetadataRes) { + return; + } + + const memoryRes = await reindexMemoryVectors(localStorage.token).catch((error) => { toast.error(`${error}`); return null; }); - if (res) { + if (memoryRes) { toast.success($i18n.t('Success')); } }} @@ -1116,7 +1146,7 @@
{$i18n.t( - 'After changing the embedding model, reindex the knowledge base for changes to take effect.' + 'After changing the embedding model, reindex knowledge, knowledge search, and memory vectors for changes to take effect.' )}
@@ -1526,8 +1556,10 @@ + +
+ +
+ + + + +
+ + {#if loading} +
+ +
+ {:else if providerModels.length === 0} +
+ {$i18n.t('No models found')} +
+ {:else} +
+ {#each providerModels as model} + {@const modelId = getModelId(model)} + {@const displayName = getDisplayName(model)} + {@const status = getStatus(model)} +
+
+
+ {displayName} +
+ {#if displayName !== modelId} +
{modelId}
+ {/if} +
+ + {status} + +
+
+ +
+ + + + + + + + + {#if supportsDelete} + + + + {/if} +
+
+ {/each} +
+ {/if} +
diff --git a/src/lib/components/admin/Settings/Models/ManageModelsModal.svelte b/src/lib/components/admin/Settings/Models/ManageModelsModal.svelte index ac68484c42..1a330c2b1c 100644 --- a/src/lib/components/admin/Settings/Models/ManageModelsModal.svelte +++ b/src/lib/components/admin/Settings/Models/ManageModelsModal.svelte @@ -1,38 +1,72 @@ - @@ -64,8 +98,22 @@ {:else if selected !== null}
+ {#if hasOllamaManagement && hasProviderManagement} +
+ + + + +
+ {/if} {#if selected === 'ollama'} + {:else if selected === 'provider'} + {/if}
diff --git a/src/lib/components/admin/Settings/Models/ModelMenu.svelte b/src/lib/components/admin/Settings/Models/ModelMenu.svelte index 361a759bbe..81c7a9ca15 100644 --- a/src/lib/components/admin/Settings/Models/ModelMenu.svelte +++ b/src/lib/components/admin/Settings/Models/ModelMenu.svelte @@ -18,7 +18,7 @@ import GlobeAlt from '$lib/components/icons/GlobeAlt.svelte'; import LockClosed from '$lib/components/icons/LockClosed.svelte'; - import { config, settings } from '$lib/stores'; + import { config, pinnedModels, settings } from '$lib/stores'; import Link from '$lib/components/icons/Link.svelte'; const i18n = getContext('i18n'); @@ -173,14 +173,14 @@ class="select-none flex w-full gap-2 items-center h-[1.6875rem] px-2 text-[0.8125rem] font-normal cursor-pointer hover:bg-gray-50/40 dark:hover:bg-gray-800/40 rounded-xl" on:click={() => runAndClose(() => pinModelHandler(model?.id))} > - {#if ($settings?.pinnedModels ?? []).includes(model?.id)} + {#if $pinnedModels.includes(model?.id)} {:else} {/if}
- {#if ($settings?.pinnedModels ?? []).includes(model?.id)} + {#if $pinnedModels.includes(model?.id)} {$i18n.t('Hide from Sidebar')} {:else} {$i18n.t('Keep in Sidebar')} diff --git a/src/lib/components/admin/Users/Groups.svelte b/src/lib/components/admin/Users/Groups.svelte index 843da668f2..d4301a49c7 100644 --- a/src/lib/components/admin/Users/Groups.svelte +++ b/src/lib/components/admin/Users/Groups.svelte @@ -184,11 +184,10 @@ {/each}
{:else} -
-
-
👥
-
{$i18n.t('No groups found')}
-
+
+
+
{$i18n.t('No groups found')}
+
{$i18n.t('Use groups to organize your users and assign permissions.')}
diff --git a/src/lib/components/automations/AutomationEditor.svelte b/src/lib/components/automations/AutomationEditor.svelte index 1bda35ec74..4e42626e7b 100644 --- a/src/lib/components/automations/AutomationEditor.svelte +++ b/src/lib/components/automations/AutomationEditor.svelte @@ -7,8 +7,9 @@ import localizedFormat from 'dayjs/plugin/localizedFormat'; import type i18nType from '$lib/i18n'; - import { WEBUI_NAME, folders } from '$lib/stores'; + import { WEBUI_NAME, channels, folders } from '$lib/stores'; import { getFolders } from '$lib/apis/folders'; + import { getChannels } from '$lib/apis/channels'; import { getAutomationById, @@ -44,6 +45,7 @@ let hasMoreRuns = true; let runsPage = 0; let foldersLoaded = false; + let channelsLoaded = false; const ensureFolders = async () => { if (foldersLoaded || ($folders ?? []).length > 0) return; @@ -52,11 +54,29 @@ foldersLoaded = true; }; + const ensureChannels = async () => { + if (channelsLoaded || ($channels ?? []).length > 0) return; + const res = await getChannels(localStorage.token).catch(() => null); + if (res) channels.set(res); + channelsLoaded = true; + }; + const getFolderName = (folderId: string | null): string => folderId ? (($folders ?? []).find((folder) => folder.id === folderId)?.name ?? $i18n.t('None')) : $i18n.t('None'); + const getDestinationName = (): string => { + const target = automation.data.target; + if (target?.type === 'channel') { + const channel = ($channels ?? []).find((channel) => channel.id === target.channel_id); + return channel?.name ? `#${channel.name}` : $i18n.t('Channel'); + } + return automation.folder_id + ? `${$i18n.t('Folder')}: ${getFolderName(automation.folder_id)}` + : $i18n.t('New chat'); + }; + const formatTime = (ts: number | null): string => { if (!ts) return '-'; return new Date(ts / 1_000_000).toLocaleString(undefined, { @@ -216,6 +236,7 @@ is_active = automation.is_active; await ensureFolders(); + await ensureChannels(); await loadRuns(); }); @@ -283,10 +304,10 @@
- {$i18n.t('Folder')} + {$i18n.t('Destination')} - {getFolderName(automation.folder_id)} + {getDestinationName()}
@@ -356,11 +377,19 @@ {/if} diff --git a/src/lib/components/automations/DestinationDropdown.svelte b/src/lib/components/automations/DestinationDropdown.svelte new file mode 100644 index 0000000000..1c2c1ab779 --- /dev/null +++ b/src/lib/components/automations/DestinationDropdown.svelte @@ -0,0 +1,308 @@ + + + { + if (!state) { + tab = ''; + folderSearch = ''; + channelSearch = ''; + } + }} +> + + +
+ + {#if tab === ''} +
+ + + + + +
+ {:else if tab === 'folders'} +
+ + +
+ + +
+ +
+ {#each filteredFolderOptions as folder (folder.id)} + {@const path = folderPath(folder)} + + {:else} +
+ {folderOptions.length > 0 ? $i18n.t('No results found') : $i18n.t('No folders')} +
+ {/each} +
+
+ {:else if tab === 'channels'} +
+ + +
+ + +
+ +
+ {#each filteredChannelOptions as channel (channel.id)} + + {:else} +
+ {channelOptions.length > 0 ? $i18n.t('No results found') : $i18n.t('No channels')} +
+ {/each} +
+
+ {/if} +
+
+
diff --git a/src/lib/components/calendar/CalendarEventModal.svelte b/src/lib/components/calendar/CalendarEventModal.svelte index 5eaa3c8d08..8c0b724683 100644 --- a/src/lib/components/calendar/CalendarEventModal.svelte +++ b/src/lib/components/calendar/CalendarEventModal.svelte @@ -121,6 +121,11 @@ return; } + if (!startDate) { + toast.error($i18n.t('Date is required')); + return; + } + loading = true; try { const startNs = dateTimeToNs(startDate, allDay ? '00:00' : startTime); diff --git a/src/lib/components/calendar/CreateCalendarModal.svelte b/src/lib/components/calendar/CreateCalendarModal.svelte index d1dafd558b..077568abe0 100644 --- a/src/lib/components/calendar/CreateCalendarModal.svelte +++ b/src/lib/components/calendar/CreateCalendarModal.svelte @@ -76,7 +76,7 @@
-
+
{$i18n.t('Name')}
diff --git a/src/lib/components/channel/MessageInput.svelte b/src/lib/components/channel/MessageInput.svelte index 4e48c50e62..e500414c08 100644 --- a/src/lib/components/channel/MessageInput.svelte +++ b/src/lib/components/channel/MessageInput.svelte @@ -575,7 +575,8 @@ i18n, triggerChar: '@', modelSuggestions: true, - userSuggestions + userSuggestions, + channelId: channel?.id }) }, ...(channelSuggestions diff --git a/src/lib/components/channel/MessageInput/MentionList.svelte b/src/lib/components/channel/MessageInput/MentionList.svelte index f3d14b7d23..8dab0e932f 100644 --- a/src/lib/components/channel/MessageInput/MentionList.svelte +++ b/src/lib/components/channel/MessageInput/MentionList.svelte @@ -2,12 +2,13 @@ import { getContext, onDestroy, onMount } from 'svelte'; const i18n = getContext('i18n'); - import { channels, models, user } from '$lib/stores'; + import { channels, models } from '$lib/stores'; import Tooltip from '$lib/components/common/Tooltip.svelte'; import Hashtag from '$lib/components/icons/Hashtag.svelte'; import Lock from '$lib/components/icons/Lock.svelte'; import { WEBUI_API_BASE_URL, WEBUI_BASE_URL } from '$lib/constants'; import { searchUsers } from '$lib/apis/users'; + import { getChannelMembersById } from '$lib/apis/channels'; export let query = ''; @@ -20,28 +21,47 @@ export let modelSuggestions = false; export let userSuggestions = false; export let channelSuggestions = false; + export let channelId: string | null = null; let _models = []; let _users = []; let _channels = []; + type UserSuggestion = { id: string; name: string }; + $: filteredItems = [..._users, ..._models, ..._channels].filter( (u) => u.label.toLowerCase().includes(query.toLowerCase()) || u.id.toLowerCase().includes(query.toLowerCase()) ); - const getUserList = async () => { - const res = await searchUsers(localStorage.token, query).catch((error) => { - console.error('Error searching users:', error); - return null; - }); + const toUserItems = (users: UserSuggestion[]) => + [...users] + .map((u) => ({ type: 'user', id: u.id, label: u.name })) + .sort((a, b) => a.label.localeCompare(b.label)); - if (res) { - _users = [...res.users.map((u) => ({ type: 'user', id: u.id, label: u.name }))].sort((a, b) => - a.label.localeCompare(b.label) - ); - } + const getUserList = async () => { + const [channelMembers, searchResults] = await Promise.all([ + channelId + ? getChannelMembersById(localStorage.token, channelId, query, 'name', 'asc').catch( + (error) => { + console.error('Error loading channel members:', error); + return null; + } + ) + : Promise.resolve(null), + searchUsers(localStorage.token, query).catch((error) => { + console.error('Error searching users:', error); + return null; + }) + ]); + + const memberUsers = (channelMembers?.users ?? []) as UserSuggestion[]; + const searchedUsers = (searchResults?.users ?? []) as UserSuggestion[]; + const memberIds = new Set(memberUsers.map((u) => u.id)); + const globalUsers = searchedUsers.filter((u) => !memberIds.has(u.id)); + + _users = [...toUserItems(memberUsers), ...toUserItems(globalUsers)]; }; $: if (query !== null && userSuggestions) { diff --git a/src/lib/components/channel/Messages/Message.svelte b/src/lib/components/channel/Messages/Message.svelte index b0ad2084e2..b14d79f9ec 100644 --- a/src/lib/components/channel/Messages/Message.svelte +++ b/src/lib/components/channel/Messages/Message.svelte @@ -10,7 +10,7 @@ dayjs.extend(isYesterday); dayjs.extend(localizedFormat); - import { getContext, onMount } from 'svelte'; + import { getContext } from 'svelte'; const i18n = getContext>('i18n'); import { formatDate } from '$lib/utils'; @@ -146,11 +146,9 @@ } }; - onMount(async () => { - if (message && message?.data === true) { - await loadMessageData(); - } - }); + $: if (message?.data === true) { + loadMessageData(); + } $: messageOutput = Array.isArray(message?.data?.output) ? message.data.output : []; $: hasStructuredOutput = buildOutputDisplayItems(messageOutput).length > 0; @@ -209,7 +207,7 @@ : 'transition: transform 0.3s cubic-bezier(0.2, 0.9, 0.3, 1);'}" > {#if !edit && !disabled} -
+
@@ -454,7 +452,6 @@ {/if} {#if message?.data === true} -
diff --git a/src/lib/components/channel/PinnedMessagesModal.svelte b/src/lib/components/channel/PinnedMessagesModal.svelte index ab171f9f96..45316c3609 100644 --- a/src/lib/components/channel/PinnedMessagesModal.svelte +++ b/src/lib/components/channel/PinnedMessagesModal.svelte @@ -94,7 +94,7 @@
{:else}
{#if pinnedMessages.length === 0}
diff --git a/src/lib/components/channel/Thread.svelte b/src/lib/components/channel/Thread.svelte index 2323a95992..39203eb7dd 100644 --- a/src/lib/components/channel/Thread.svelte +++ b/src/lib/components/channel/Thread.svelte @@ -188,7 +188,7 @@
-
+
{#if messages !== null}
{/if} +
-
- -
+
+
{/if} diff --git a/src/lib/components/chat/AskUserCard.svelte b/src/lib/components/chat/AskUserCard.svelte new file mode 100644 index 0000000000..bdd937ea48 --- /dev/null +++ b/src/lib/components/chat/AskUserCard.svelte @@ -0,0 +1,303 @@ + + +{#if show} +
+
+ {#if question} + {#key question.id} +
+
+
+
+ {question.header} +
+ {#if questions.length > 1} +
+ {questionIndex + 1}/{questions.length} +
+ {/if} +
+
+ {question.question} +
+
+ +
+
+ {#each question.options || [] as option, optionIndex} + + {/each} +
+ + {#if questionAllowsOther(question)} +
+
+ {$i18n.t('Other')} +
+ selectOther(question)} + on:input={(event) => + updateOther(question, (event.currentTarget as HTMLInputElement).value)} + on:keydown={(event) => { + if (event.key === 'Enter') { + event.preventDefault(); + advance(); + } + }} + /> +
+ {/if} +
+
+ {/key} + {/if} + +
+
+ + +
+ {#if questionIndex < questions.length - 1} + + {:else} + + {/if} +
+
+
+{/if} diff --git a/src/lib/components/chat/Chat.svelte b/src/lib/components/chat/Chat.svelte index 686613cca9..ed8d94d511 100644 --- a/src/lib/components/chat/Chat.svelte +++ b/src/lib/components/chat/Chat.svelte @@ -66,7 +66,7 @@ } from '$lib/utils'; import { AudioQueue } from '$lib/utils/audio'; import { createTemporaryChatId, isTemporaryChatId } from '$lib/utils/chatId'; - import { getOutputText } from './Messages/structuredOutput'; + import { applyResponseStreamEvent, getOutputText } from './Messages/structuredOutput'; import { archiveChatById, @@ -77,12 +77,18 @@ getAllTags, getChatById, getTagsById, + resolveChatMessageToolCall, updateChatById, updateChatFolderIdById } from '$lib/apis/chats'; import { generateOpenAIChatCompletion } from '$lib/apis/openai'; import { processUrl, processWebSearch } from '$lib/apis/retrieval'; - import { getAndUpdateUserLocation, getUserInfoById, getUserSettings } from '$lib/apis/users'; + import { + getAndUpdateUserLocation, + getUserInfoById, + getUserSettings, + updateUserSettings + } from '$lib/apis/users'; import { generateQueries, chatAction, @@ -162,7 +168,11 @@ let eventConfirmationInputValue = ''; let eventConfirmationInputType = ''; let eventConfirmationInputOptions: ({ label?: string; value: string } | string)[] = []; - let eventCallback = null; + let eventCallback: (value: any) => void = () => {}; + let showAskUserDialog = false; + let askUserQuestions: any[] = []; + let askUserAllowOther = true; + let askUserTimeoutMs: number | null = null; let selectedModels = ['']; let atSelectedModel: Model | undefined; @@ -410,6 +420,212 @@ let loadedChatIdProp = ''; let currentDraftKey = ''; + $: toolApprovalMode = + (params?.tool_approval_mode ?? $settings?.params?.tool_approval_mode) === 'ask' + ? 'ask' + : 'full'; + + const handleToolApprovalModeChange = async (mode: string) => { + const tool_approval_mode = mode === 'ask' ? 'ask' : 'full'; + params = { + ...params, + tool_approval_mode + }; + + settings.set({ + ...$settings, + params: { + ...($settings?.params ?? {}), + tool_approval_mode + } + }); + await updateUserSettings(localStorage.token, { ui: $settings }).catch((err) => { + console.error('[tool permissions settings]', err); + }); + + if ($chatId && !$temporaryChatEnabled && !isTemporaryChatId($chatId)) { + const res = await updateChatById(localStorage.token, $chatId, { params }).catch((err) => { + console.error('[tool permissions chat]', err); + return null; + }); + if (res) chat = res; + } + + if (tool_approval_mode === 'full') { + const messages = [...Object.values(history?.messages ?? {})].reverse() as any[]; + for (const message of messages) { + const output = (Array.isArray(message?.output) ? message.output : []) as any[]; + const resultCallIds = new Set( + output + .filter((item: any) => item?.type === 'function_call_output' && item?.call_id) + .map((item: any) => item.call_id) + ); + const pendingCall = output.find((item: any) => { + const callId = item?.call_id ?? item?.id; + return ( + item?.type === 'function_call' && + item?.name !== 'ask_user' && + (item?.status === 'pending' || item?.status === 'requires_approval') && + callId && + !resultCallIds.has(callId) + ); + }); + const callId = pendingCall?.call_id ?? pendingCall?.id; + if (!message?.id || !callId) { + continue; + } + + const res = await resolveChatMessageToolCall( + localStorage.token, + $chatId, + message.id, + callId, + 'approve' + ).catch(async (error) => { + toast.error(`${error}`); + await loadChat(); + return null; + }); + if (res) onToolCallResolved(res); + break; + } + } + }; + + const parseToolArguments = (args) => { + if (!args) { + return {}; + } + let value = args; + while (typeof value === 'string') { + try { + value = JSON.parse(value); + } catch { + break; + } + } + return typeof value === 'object' && value !== null && !Array.isArray(value) ? value : {}; + }; + + const getPendingAskUserFromMessage = (message) => { + if (message?.role !== 'assistant' || !Array.isArray(message.output)) { + return null; + } + + const call = message.output.find( + (item) => + item?.type === 'function_call' && + item?.name === 'ask_user' && + item?.status === 'pending' && + (item?.call_id || item?.id) + ); + + if (call) { + return { message, call, args: parseToolArguments(call.arguments) }; + } + + return null; + }; + + const findPendingAskUser = (chatHistory) => { + if (!chatHistory?.messages) { + return null; + } + const messages = chatHistory.currentId + ? createMessagesList(chatHistory, chatHistory.currentId) + : Object.values(chatHistory.messages); + for (const message of [...messages].reverse()) { + const pending = getPendingAskUserFromMessage(message); + if (pending) return pending; + } + return null; + }; + + const messageHasPendingAskUser = (message) => { + return !!getPendingAskUserFromMessage(message); + }; + + const answerPendingAskUser = async (messageId, callId, answers, timedOut = false) => { + if (!$chatId || !messageId || !callId) { + return; + } + + const res = await resolveChatMessageToolCall( + localStorage.token, + $chatId, + messageId, + callId, + 'answer', + { + answers, + timed_out: timedOut + } + ).catch(async (error) => { + toast.error(`${error}`); + await loadChat(); + }); + onToolCallResolved(res); + }; + + const rejectPendingAskUser = async (messageId, callId) => { + if (!$chatId || !messageId || !callId) { + return; + } + + const res = await resolveChatMessageToolCall( + localStorage.token, + $chatId, + messageId, + callId, + 'reject' + ).catch(async (error) => { + toast.error(`${error}`); + await loadChat(); + }); + onToolCallResolved(res); + }; + + $: pendingAskUser = findPendingAskUser(history); + $: savedAskUserPrompt = pendingAskUser + ? { + show: true, + questions: Array.isArray(pendingAskUser.args?.questions) + ? pendingAskUser.args.questions + : [], + allowOther: pendingAskUser.args?.allow_other !== false, + timeoutMs: null, + onConfirm: (value) => { + void answerPendingAskUser( + pendingAskUser.message.id, + pendingAskUser.call.call_id || pendingAskUser.call.id, + value?.answers ?? {}, + false + ); + }, + onCancel: () => { + void rejectPendingAskUser( + pendingAskUser.message.id, + pendingAskUser.call.call_id || pendingAskUser.call.id + ); + } + } + : null; + + $: socketAskUserPrompt = { + show: showAskUserDialog, + questions: askUserQuestions, + allowOther: askUserAllowOther, + timeoutMs: askUserTimeoutMs, + onConfirm: (value) => { + showAskUserDialog = false; + eventCallback(value); + }, + onCancel: () => { + showAskUserDialog = false; + eventCallback({ status: 'cancelled', answers: {} }); + } + }; + const mergeChatVariableSchemas = (modelIds = [], availableModels = []) => { const byKey: Record = {}; const conflicts: any[] = []; @@ -499,6 +715,13 @@ } }; + const onToolCallResolved = (res) => { + const newTaskIds = res?.task_ids ?? (res?.task_id ? [res.task_id] : []); + if (newTaskIds.length > 0) { + taskIds = [...(taskIds ?? []), ...newTaskIds]; + } + }; + let oldSelectedModelIds = ['']; $: if (!equal(selectedModelIds, oldSelectedModelIds)) { onSelectedModelIdsChange(); @@ -615,6 +838,9 @@ webSearchEnabled = input.webSearchEnabled; imageGenerationEnabled = input.imageGenerationEnabled; codeInterpreterEnabled = input.codeInterpreterEnabled; + if (input.toolApprovalMode) { + handleToolApprovalModeChange(input.toolApprovalMode); + } } } catch (e) {} } else { @@ -1000,6 +1226,8 @@ updateLastReadAt($chatId); } } + } else if (type === 'response:completion') { + responseCompletionEventHandler(data, message); } else if (type === 'chat:completion') { chatCompletionEventHandler(data, message, event.chat_id); } else if (type === 'chat:tasks:cancel') { @@ -1144,6 +1372,13 @@ eventConfirmationInputValue = data?.value ?? ''; eventConfirmationInputType = data?.input?.type ?? data?.type ?? ''; eventConfirmationInputOptions = data?.input?.options ?? data?.options ?? []; + } else if (type === 'request:user_input') { + eventCallback = cb; + askUserQuestions = data?.questions ?? []; + askUserAllowOther = data?.allow_other ?? true; + askUserTimeoutMs = + typeof data?.timeout_ms === 'number' && data.timeout_ms > 0 ? data.timeout_ms : null; + showAskUserDialog = true; } else if (type.startsWith('terminal:')) { terminalEventHandler(type, data); } else { @@ -1385,6 +1620,9 @@ webSearchEnabled = input.webSearchEnabled; imageGenerationEnabled = input.imageGenerationEnabled; codeInterpreterEnabled = input.codeInterpreterEnabled; + if (input.toolApprovalMode) { + handleToolApprovalModeChange(input.toolApprovalMode); + } } } catch (e) {} } @@ -1600,7 +1838,9 @@ fileItem.content_type = uploadedFile.meta?.content_type; fileItem.size = uploadedFile.meta?.size; fileItem.collection_name = - res.collection_name ?? uploadedFile.meta?.collection_name ?? uploadedFile.collection_name; + res.collection_name ?? + uploadedFile.meta?.collection_name ?? + uploadedFile.collection_name; } else { fileItem.type = 'text'; fileItem.file = { @@ -1793,6 +2033,22 @@ const defaultModels = $config?.default_models ? $config?.default_models.split(',') : []; + const openModelSelectorWithSearch = async (modelId: string) => { + const modelSelectorButton = document.getElementById('model-selector-model-button'); + modelSelectorButton?.click(); + + await tick(); + + const modelSelectorInput = document.getElementById( + 'model-search-input' + ) as HTMLInputElement | null; + if (modelSelectorInput) { + modelSelectorInput.focus(); + modelSelectorInput.value = modelId; + modelSelectorInput.dispatchEvent(new Event('input', { bubbles: true })); + } + }; + if ($page.url.searchParams.get('models') || $page.url.searchParams.get('model')) { const urlModels = ( $page.url.searchParams.get('models') || @@ -1803,18 +2059,7 @@ if (urlModels.length === 1) { if (!$models.find((m) => m.id === urlModels[0])) { // Model not found; open model selector and prefill - const modelSelectorButton = document.getElementById('model-selector-0-button'); - if (modelSelectorButton) { - modelSelectorButton.click(); - await tick(); - - const modelSelectorInput = document.getElementById('model-search-input'); - if (modelSelectorInput) { - modelSelectorInput.focus(); - modelSelectorInput.value = urlModels[0]; - modelSelectorInput.dispatchEvent(new Event('input')); - } - } + await openModelSelectorWithSearch(urlModels[0]); } else { // Model found; set it as selected selectedModels = urlModels; @@ -2142,7 +2387,11 @@ } else { taskIds = null; // No active tasks and message incomplete → generation was interrupted - if (currentMessage?.role === 'assistant' && !currentMessage.done) { + if ( + currentMessage?.role === 'assistant' && + !currentMessage.done && + !messageHasPendingAskUser(currentMessage) + ) { currentMessage.done = true; } } @@ -2236,16 +2485,24 @@ } processingQueueChats.add(targetChatId); + const queuedMessages = [...queue]; + const queuedMessageIds = new Set(queuedMessages.map((m) => m.id)); try { - const combinedPrompt = queue.map((m) => m.prompt).join('\n\n'); - const combinedFiles = queue.flatMap((m) => m.files); + const combinedPrompt = queuedMessages.map((m) => m.prompt).join('\n\n'); + const combinedFiles = queuedMessages.flatMap((m) => m.files); - chatRequestQueues.update((q) => { - const { [targetChatId]: _, ...rest } = q; - return rest; - }); + chatRequestQueues.update((q) => ({ + ...q, + [targetChatId]: (q[targetChatId] ?? []).filter((m) => !queuedMessageIds.has(m.id)) + })); await submitPrompt(combinedPrompt, combinedFiles); + } catch (error) { + console.error(error); + chatRequestQueues.update((q) => ({ + ...q, + [targetChatId]: [...queuedMessages, ...(q[targetChatId] ?? [])] + })); } finally { processingQueueChats.delete(targetChatId); } @@ -2486,6 +2743,27 @@ } }; + const responseCompletionEventHandler = (data, message) => { + message.output = applyResponseStreamEvent(message.output ?? [], data); + + if (data?.type === 'response.output_text.delta') { + const value = data.delta ?? ''; + if (!(message.content == '' && value == '\n')) { + message.content += value; + + if (navigator.vibrate && ($settings?.hapticFeedback ?? false)) { + navigator.vibrate(5); + } + dispatchCallOverlayAudio(message); + } + } else if (data?.type === 'response.completed' || data?.type?.endsWith('.done')) { + message.content = getOutputText(message.output) || message.content; + } + + history.messages[message.id] = message; + history = history; + }; + const chatCompletionEventHandler = async (data, message, chatId) => { const { id, done, choices, content, output, sources, selected_model_id, error, usage } = data; @@ -2545,6 +2823,7 @@ } history.messages[message.id] = message; + history = history; if (done) { message.done = true; @@ -3376,10 +3655,7 @@ await handleOpenAIError(res.error, responseMessage); } else { // Backend returns task_ids (multi-model) or task_id (single model) - const newTaskIds = res.task_ids ?? (res.task_id ? [res.task_id] : []); - if (newTaskIds.length > 0) { - taskIds = [...(taskIds ?? []), ...newTaskIds]; - } + onToolCallResolved(res); // Backend returns chat_id for new chats — set store + URL. // Only update if the user hasn't navigated to a different chat @@ -3463,7 +3739,13 @@ }; const stopResponse = async (processQueue = true) => { - if (taskIds) { + const responseMessage = history.currentId ? history.messages[history.currentId] : null; + const hasTaskIds = (taskIds?.length ?? 0) > 0; + const hasPendingAssistantResponse = + !!$chatId && + (hasTaskIds || (responseMessage?.role === 'assistant' && responseMessage?.done !== true)); + + if (hasTaskIds || hasPendingAssistantResponse) { if ($chatId) { await stopTasksByChatId(localStorage.token, $chatId).catch((error) => { toast.error(`${error}`); @@ -3480,15 +3762,16 @@ taskIds = null; - const responseMessage = history.messages[history.currentId]; // Set all response messages to done - if (responseMessage.parentId && history.messages[responseMessage.parentId]) { + if (responseMessage?.parentId && history.messages[responseMessage.parentId]) { for (const messageId of history.messages[responseMessage.parentId].childrenIds) { history.messages[messageId].done = true; } } - history.messages[history.currentId] = responseMessage; + if (responseMessage) { + history.messages[history.currentId] = responseMessage; + } if (shouldAutoScrollResponse()) { scrollToBottom(); @@ -3940,7 +4223,7 @@ : 'h-screen max-h-[100dvh]'} transition-width duration-200 ease-in-out {$showSidebar && !embedded ? ' md:max-w-[calc(100%-var(--sidebar-width))]' - : ' '} w-full max-w-full flex flex-col" + : ' '} w-full max-w-full min-w-0 flex flex-col" id={chatContainerId} > {#if !loading} @@ -4095,6 +4378,7 @@ {mergeResponses} {chatActionHandler} {addMessages} + {onToolCallResolved} allowDelete={!(generating || taskIds?.length)} forkHandler={handleForkChat} topPadding={!embedded} @@ -4142,7 +4426,8 @@ compactHandler={handleManualCompact} statusHandler={handleStatusCommand} forkHandler={handleForkChat} - toolServers={$toolServers} + {toolApprovalMode} + onToolApprovalModeChange={handleToolApprovalModeChange} {generating} {stopResponse} {createMessagePair} @@ -4150,6 +4435,7 @@ {onUpdate} messageQueue={$chatRequestQueues[$chatId] ?? []} {chatTasks} + askUser={savedAskUserPrompt ?? socketAskUserPrompt} onQueueSendNow={sendQueuedMessageNow} onQueueEdit={editQueuedMessage} onQueueDelete={deleteQueuedMessage} @@ -4204,10 +4490,7 @@
{/if} -
+
chatContext(terminal) !== false; const chatContextNeedsSavedChat = (terminal: any) => chatContext(terminal)?.context_id === 'chat_id'; + function showFilesOnTerminalSelect(largeScreen: boolean) { + activeTab = 'files'; + if (largeScreen) { + showControls.set($settings?.showFilesOnTerminalSelect ?? true); + } + } $: selectedSystemTerminal = ($terminalServers ?? []).find( (t) => t.id && t.id === $selectedTerminalId ); @@ -118,10 +124,7 @@ // Auto-open Files tab when a terminal is selected (suppress panel open when full-screen) $: if ($selectedTerminalId && terminalFilesAvailable) { - activeTab = 'files'; - if (largeScreen) { - showControls.set($settings?.showFilesOnTerminalSelect ?? true); - } + showFilesOnTerminalSelect(largeScreen); } // Clear selected direct terminal if user lost permission diff --git a/src/lib/components/chat/FileNav.svelte b/src/lib/components/chat/FileNav.svelte index 9a9a568d69..9574f014be 100644 --- a/src/lib/components/chat/FileNav.svelte +++ b/src/lib/components/chat/FileNav.svelte @@ -1,6 +1,9 @@ -
- +
+ {$i18n.t('{{count}} selected', { count })}
diff --git a/src/lib/components/chat/FileNav/FileEntryRow.svelte b/src/lib/components/chat/FileNav/FileEntryRow.svelte index 2abde52d1e..2aa511a301 100644 --- a/src/lib/components/chat/FileNav/FileEntryRow.svelte +++ b/src/lib/components/chat/FileNav/FileEntryRow.svelte @@ -6,36 +6,49 @@ import Dropdown from '$lib/components/common/Dropdown.svelte'; import DropdownMenu from '$lib/components/common/DropdownMenu.svelte'; - import Folder from '../../icons/Folder.svelte'; - import EllipsisHorizontal from '../../icons/EllipsisHorizontal.svelte'; - import GarbageBin from '../../icons/GarbageBin.svelte'; - import Pencil from '../../icons/Pencil.svelte'; - import Clipboard from '../../icons/Clipboard.svelte'; + import FileTypeIcon from './FileTypeIcon.svelte'; + import Icon from './Icon.svelte'; - const i18n = getContext('i18n'); + const i18n: any = getContext('i18n'); export let entry: FileEntry; export let currentPath: string; + export let fullPath: string | null = null; + export let depth = 0; + export let rowIndex = 0; + export let expanded = false; + export let loadingChildren = false; export let terminalUrl: string = ''; export let terminalKey: string = ''; export let onOpen: (entry: FileEntry) => void = () => {}; export let onDownload: (path: string) => void = () => {}; export let onDelete: (path: string, name: string) => void = () => {}; - export let onMove: (source: string, destFolder: string) => void = () => {}; + export let onMove: (sources: string[], destFolder: string) => void | Promise = () => {}; export let onRename: (oldPath: string, newName: string) => void = () => {}; // ── Selection ───────────────────────────────────────────────────────── export let selected: boolean = false; export let selectionMode: boolean = false; export let selectedPaths: Set = new Set(); - export let onSelect: (entry: FileEntry, event: MouseEvent) => void = () => {}; + export let onSelect: ( + entry: FileEntry, + event: MouseEvent, + path: string, + index: number + ) => void = () => {}; export let onLongPress: () => void = () => {}; + export let onToggleExpand: (path: string) => void = () => {}; export let showDate: boolean = false; export let parentWritable = true; + $: entryPath = + fullPath ?? + (entry.type === 'directory' ? `${currentPath}${entry.name}/` : `${currentPath}${entry.name}`); + $: directoryPath = entryPath.endsWith('/') ? entryPath : `${entryPath}/`; $: writable = entry.writable !== false; $: canMutate = parentWritable && writable; + $: rowIndent = `${8 + depth * 16}px`; const formatRelativeTime = (epoch: number): string => { const diff = Math.floor(Date.now() / 1000) - epoch; @@ -48,6 +61,14 @@ }; let dragOverFolder = false; + let expandTimer: ReturnType | null = null; + let menuOpen = false; + + const clearExpandTimer = () => { + if (!expandTimer) return; + clearTimeout(expandTimer); + expandTimer = null; + }; // ── Rename state ───────────────────────────────────────────────────── let renaming = false; @@ -71,7 +92,7 @@ const newName = renameValue.trim(); renaming = false; if (!newName || newName === entry.name) return; - onRename(`${currentPath}${entry.name}`, newName); + onRename(entryPath.replace(/\/$/, ''), newName); }; const cancelRename = () => { @@ -89,7 +110,7 @@ longPressTimer = setTimeout(() => { didLongPress = true; onLongPress(); - onSelect(entry, e as any); + onSelect(entry, e as any, entryPath, rowIndex); }, 500); }; @@ -109,6 +130,7 @@ onDestroy(() => { if (longPressTimer) clearTimeout(longPressTimer); + clearExpandTimer(); }); // ── Click handler ──────────────────────────────────────────────────── @@ -122,13 +144,13 @@ // Modifier click → toggle/range select if (e.metaKey || e.ctrlKey || e.shiftKey) { e.preventDefault(); - onSelect(entry, e); + onSelect(entry, e, entryPath, rowIndex); return; } // In selection mode (touch) → toggle select if (selectionMode) { - onSelect(entry, e); + onSelect(entry, e, entryPath, rowIndex); return; } @@ -137,14 +159,14 @@ }; -
  • +
  • { if (entry.type !== 'directory') return; if (!writable) return; @@ -152,13 +174,20 @@ e.preventDefault(); e.stopPropagation(); dragOverFolder = true; + if (!expanded && !expandTimer) { + expandTimer = setTimeout(() => { + onToggleExpand(directoryPath); + expandTimer = null; + }, 600); + } }} on:dragleave={(e) => { if (entry.type !== 'directory') return; e.stopPropagation(); dragOverFolder = false; + clearExpandTimer(); }} - on:drop={(e) => { + on:drop={async (e) => { if (entry.type !== 'directory') return; if (!writable) return; const raw = e.dataTransfer?.getData('application/x-terminal-file-move'); @@ -166,26 +195,47 @@ e.preventDefault(); e.stopPropagation(); dragOverFolder = false; + clearExpandTimer(); try { const data = JSON.parse(raw); - const paths = data.paths || (data.path ? [data.path] : []); - const destFolder = `${currentPath}${entry.name}/`; - for (const p of paths) { - if (p + '/' === destFolder || p === destFolder) continue; - onMove(p, destFolder); - } + const paths = (data.paths || (data.path ? [data.path] : [])) as string[]; + const destFolder = directoryPath; + await onMove( + paths.filter((p) => p + '/' !== destFolder && p !== destFolder), + destFolder + ); } catch {} }} > + {#if entry.type === 'directory'} + + {:else} + + {/if} +
    {/if} - {#if entry.type === 'directory'} - - {:else} - - - - {/if} + {#if renaming} { if (e.key === 'Enter') { e.preventDefault(); @@ -303,93 +328,119 @@ {/if} {formatFileSize(entry.size)} {:else if entry.type === 'directory' && showDate && entry.modified && !renaming} - {formatRelativeTime(entry.modified)} + {formatRelativeTime(entry.modified)} {/if} - +
    + {#if entry.type === 'directory'} + + {:else} + + {/if} + diff --git a/src/lib/components/chat/FileNav/FileNavToolbar.svelte b/src/lib/components/chat/FileNav/FileNavToolbar.svelte index 8ff8af08c8..5a1356ae3a 100644 --- a/src/lib/components/chat/FileNav/FileNavToolbar.svelte +++ b/src/lib/components/chat/FileNav/FileNavToolbar.svelte @@ -1,15 +1,12 @@ -
    - - - - + + + - - - - + + + +
    {#each breadcrumbs as crumb, i} - {#if i > 1} - / + {#if showSeparator(i)} + / {/if} {/each} {#if selectedFile} - / - + / + {selectedFile.split('/').pop()} {/if}
    {#if !writable} - - Read-only - + Read-only {/if} {#if !selectedFile} - + @@ -199,117 +163,140 @@ +
    - - - - - - - - - - - - + + + + +
    + + + + + + + +
    +
    { panzoomRef?.reset(); pptxPreviewRef?.resetView(); @@ -288,6 +289,7 @@ {:else if fileImageUrl !== null} @@ -471,6 +473,57 @@ {$i18n.t('Could not read file.')}
  • {/if} + + {#if !fileLoading && fileImageUrl !== null} +
    + + + +
    + {/if}
    diff --git a/src/lib/components/common/Select.svelte b/src/lib/components/common/Select.svelte index e4f01e5138..efa8456d05 100644 --- a/src/lib/components/common/Select.svelte +++ b/src/lib/components/common/Select.svelte @@ -122,7 +122,7 @@ + + + {:else} + +
    + {#if open} + + {:else} + + {/if} +
    + {/if}
    @@ -189,8 +240,8 @@
    - {#if args} +
    { + const res = await markChatUnreadById(localStorage.token, id).catch((error) => { + toast.error(`${error}`); + return null; + }); + + if (res) { + await refreshSidebar(); + } + }; + const archiveChatHandler = async (id) => { try { await archiveChatById(localStorage.token, id); @@ -818,6 +830,9 @@ renameHandler={() => { renameHandler(chat.id); }} + markUnreadHandler={() => { + markUnreadHandler(chat.id); + }} deleteHandler={() => { menuChatId = chat.id; menuChatTitle = chat.title; diff --git a/src/lib/components/layout/Sidebar.svelte b/src/lib/components/layout/Sidebar.svelte index d61359a430..16aa3fab7d 100644 --- a/src/lib/components/layout/Sidebar.svelte +++ b/src/lib/components/layout/Sidebar.svelte @@ -22,7 +22,7 @@ socket, config, isApp, - models, + visiblePinnedModels, selectedFolder, WEBUI_NAME, sidebarWidth @@ -120,9 +120,7 @@ let showCreateFolderModal = false; - let pinnedModels = []; - - let showPinnedModels = false; + let showPinnedModels = true; let showPinnedNotes = false; let showChannels = false; let showFolders = false; @@ -130,6 +128,7 @@ let showChatsMenu = false; let folders = {}; + type SelectedSidebarFolder = { id: string } | null; let folderRegistry: Record< string, { @@ -145,6 +144,14 @@ let sharedFolders: any[] = []; + const initSelectedFolderChats = (folder: SelectedSidebarFolder) => { + if (!folder?.id) { + return; + } + + folderRegistry[folder.id]?.setFolderItems?.(); + }; + $: pinnedItems = $settings?.pinnedMenuItems ?? DEFAULT_PINNED_ITEMS; const isMenuItemVisible = (id) => { @@ -230,29 +237,33 @@ } }; - $: if ($selectedFolder) { - initFolders(); - } + $: initSelectedFolderChats($selectedFolder as SelectedSidebarFolder); const initFolders = async () => { if ($config?.features?.enable_folders === false) { return; } - const folderList = await getFolders(localStorage.token).catch((error) => { - return []; - }); + const [folderList, sharedFolderList] = await Promise.all([ + getFolders(localStorage.token).catch((error) => { + return []; + }), + getSharedFolders(localStorage.token).catch((error) => { + return []; + }) + ]); _folders.set(folderList.sort((a, b) => b.updated_at - a.updated_at)); - folders = {}; + sharedFolders = sharedFolderList; + const folderMap: Record = {}; // First pass: Initialize all folder entries for (const folder of folderList) { // Ensure folder is added to folders with its data - folders[folder.id] = { ...(folders[folder.id] || {}), ...folder }; + folderMap[folder.id] = { ...(folderMap[folder.id] || {}), ...folder }; if (newFolderId && folder.id === newFolderId) { - folders[folder.id].new = true; + folderMap[folder.id].new = true; newFolderId = null; } } @@ -261,42 +272,38 @@ for (const folder of folderList) { if (folder.parent_id) { // Ensure the parent folder is initialized if it doesn't exist - if (!folders[folder.parent_id]) { - folders[folder.parent_id] = {}; // Create a placeholder if not already present + if (!folderMap[folder.parent_id]) { + folderMap[folder.parent_id] = {}; // Create a placeholder if not already present } // Initialize childrenIds array if it doesn't exist and add the current folder id - folders[folder.parent_id].childrenIds = folders[folder.parent_id].childrenIds - ? [...folders[folder.parent_id].childrenIds, folder.id] + folderMap[folder.parent_id].childrenIds = folderMap[folder.parent_id].childrenIds + ? [...folderMap[folder.parent_id].childrenIds, folder.id] : [folder.id]; // Sort the children by updated_at field - folders[folder.parent_id].childrenIds.sort((a, b) => { - return folders[b].updated_at - folders[a].updated_at; + folderMap[folder.parent_id].childrenIds.sort((a, b) => { + return folderMap[b].updated_at - folderMap[a].updated_at; }); } } // Merge shared folders into the same structure - try { - sharedFolders = await getSharedFolders(localStorage.token); - } catch (e) { - sharedFolders = []; - } - for (const sf of sharedFolders) { - if (folders[sf.id]) continue; // Already owned by user - folders[sf.id] = { ...sf, shared: true }; + if (folderMap[sf.id]) continue; // Already owned by user + folderMap[sf.id] = { ...sf, shared: true }; } // Build parent-child relationships for shared folders for (const sf of sharedFolders) { - if (folders[sf.id]?.shared && sf.parent_id && folders[sf.parent_id]) { - folders[sf.parent_id].childrenIds = folders[sf.parent_id].childrenIds - ? [...new Set([...folders[sf.parent_id].childrenIds, sf.id])] + if (folderMap[sf.id]?.shared && sf.parent_id && folderMap[sf.parent_id]) { + folderMap[sf.parent_id].childrenIds = folderMap[sf.parent_id].childrenIds + ? [...new Set([...folderMap[sf.parent_id].childrenIds, sf.id])] : [sf.id]; } } + + folders = folderMap; }; const initSharedFolders = async () => { @@ -674,12 +681,6 @@ navElement.style['-webkit-app-region'] = 'drag'; } } - }), - settings.subscribe((value) => { - if (pinnedModels != value?.pinnedModels ?? []) { - pinnedModels = value?.pinnedModels ?? []; - showPinnedModels = pinnedModels.length > 0; - } }) ]; @@ -700,7 +701,12 @@ socketInstance?.on('events', chatActiveEventHandler); socketInstance?.on('connect', refreshChatRows); - const unregisterFolderRefreshHandler = registerFolderRefreshHandler((folderId, chat) => { + const unregisterFolderRefreshHandler = registerFolderRefreshHandler(async (folderId, chat) => { + // null refreshes the folder tree; undefined refreshes all folder chat lists. + if (folderId === null) { + return initFolders(); + } + if (folderId) { if (chat) { return folderRegistry[folderId]?.upsertChat?.(chat); @@ -1177,7 +1183,7 @@
    { if (e.target.scrollTop === 0) { scrollTop = 0; @@ -1276,11 +1282,10 @@
    - {#if ($models ?? []).length > 0 && (($settings?.pinnedModels ?? []).length > 0 || $config?.default_pinned_models)} + {#if $visiblePinnedModels.length > 0} @@ -1292,7 +1297,6 @@ { @@ -1311,7 +1315,6 @@ { showCreateFolderModal = true; @@ -1397,7 +1399,6 @@ { selectedFolder.set(null); diff --git a/src/lib/components/layout/Sidebar/ChannelItem.svelte b/src/lib/components/layout/Sidebar/ChannelItem.svelte index 834996ef4d..7cc786d2cc 100644 --- a/src/lib/components/layout/Sidebar/ChannelItem.svelte +++ b/src/lib/components/layout/Sidebar/ChannelItem.svelte @@ -44,6 +44,12 @@ } return hasPublicReadGrant(channel?.access_grants); }; + + const formatUnreadCount = (count: number) => + new Intl.NumberFormat(undefined, { + notation: 'compact', + compactDisplay: 'short' + }).format(count); { console.log(channel); @@ -104,7 +110,7 @@ }} draggable="false" > -
    +
    {#if channel?.type === 'dm'} {#if channel?.users} @@ -152,15 +158,13 @@ {/if}
    -
    +
    {#if channel?.name} - + {channel.name} {:else} - + {channel?.users ?.filter((u) => u.id !== $user?.id) .map((u) => u.name) @@ -171,39 +175,35 @@ {@const dmUser = channel.users.find((u) => u.id !== $user?.id)} {#if dmUser?.status_emoji || dmUser?.status_message} - + {#if dmUser?.status_emoji}
    {/if} -
    +
    {dmUser?.status_message}
    {/if} {/if} {/if} -
    -
    -
    - {#if channel?.unread_count > 0} -
    - {new Intl.NumberFormat($i18n.locale, { - notation: 'compact', - compactDisplay: 'short' - }).format(channel.unread_count)} -
    - {/if} + {#if channel?.unread_count > 0} +
    + {formatUnreadCount(channel.unread_count)} +
    + {/if} +
    {#if ['dm'].includes(channel?.type)} -
    +
    {:else if $user?.role === 'admin' || channel.user_id === $user?.id} -
    +
    - {#if (folders[folderId]?.childrenIds ?? []).length > 0 || (chats ?? []).length > 0 || hasMoreChats} + {#if (folders[folderId]?.childrenIds ?? []).length > 0 || chats !== null || hasMoreChats || chatsLoading}
    @@ -925,6 +928,24 @@ {/each} {/if} + {#if chats === null && chatsLoading} +
    + + + +
    + {/if} + + {#if chats !== null && chats.length === 0 && !chatsLoading} +
    + {$i18n.t('No chats')} +
    + {/if} + {#each chats ?? [] as chat (chat.id)} {/if} - - {#if chats === null && chatsLoading} -
    - - - -
    - {/if}
    diff --git a/src/lib/components/workspace/Knowledge/KnowledgeBase/NewDirectoryModal.svelte b/src/lib/components/workspace/Knowledge/KnowledgeBase/NewDirectoryModal.svelte index 20f7c7d145..a85f446370 100644 --- a/src/lib/components/workspace/Knowledge/KnowledgeBase/NewDirectoryModal.svelte +++ b/src/lib/components/workspace/Knowledge/KnowledgeBase/NewDirectoryModal.svelte @@ -43,7 +43,7 @@
    -
    +
    {$i18n.t('Name')}
    { - let pinnedModels = $settings?.pinnedModels ?? []; - - if (pinnedModels.includes(modelId)) { - pinnedModels = pinnedModels.filter((id) => id !== modelId); - } else { - pinnedModels = [...new Set([...pinnedModels, modelId])]; - } - - settings.set({ ...$settings, pinnedModels: pinnedModels }); + settings.set({ + ...$settings, + pinnedModels: $pinnedModels.includes(modelId) + ? $pinnedModels.filter((id) => id !== modelId) + : [...$pinnedModels, modelId] + }); await updateUserSettings(localStorage.token, { ui: $settings }); }; diff --git a/src/lib/components/workspace/Models/BuiltinTools.svelte b/src/lib/components/workspace/Models/BuiltinTools.svelte index f437b6f5db..01754c56b6 100644 --- a/src/lib/components/workspace/Models/BuiltinTools.svelte +++ b/src/lib/components/workspace/Models/BuiltinTools.svelte @@ -1,16 +1,22 @@
    {$i18n.t('Builtin Tools')}
    {#each allTools as tool} -
    -
    - - {$i18n.t(toolLabels[tool].label)} - -
    +
    { - if (e.detail === 'checked') { - delete builtinTools[tool]; - } else { - builtinTools[tool] = false; - } - builtinTools = builtinTools; + setBuiltinTool(tool, e.detail === 'checked'); }} /> +
    {/each}
    diff --git a/src/lib/components/workspace/Models/Capabilities.svelte b/src/lib/components/workspace/Models/Capabilities.svelte index 825d65b39b..04be7c3939 100644 --- a/src/lib/components/workspace/Models/Capabilities.svelte +++ b/src/lib/components/workspace/Models/Capabilities.svelte @@ -1,10 +1,12 @@
    {$i18n.t('Default Features')}
    {#each availableFeatures as feature} -
    -
    - - {$i18n.t(featureLabels[feature].label)} - -
    +
    { - if (e.detail === 'checked') { - featureIds = [...featureIds, feature]; - } else { - featureIds = featureIds.filter((id) => id !== feature); - } + setFeature(feature, e.detail === 'checked'); }} /> +
    {/each}
    diff --git a/src/lib/components/workspace/Models/ModelMenu.svelte b/src/lib/components/workspace/Models/ModelMenu.svelte index 0d99a31510..e44307776d 100644 --- a/src/lib/components/workspace/Models/ModelMenu.svelte +++ b/src/lib/components/workspace/Models/ModelMenu.svelte @@ -15,7 +15,7 @@ import Pin from '$lib/components/icons/Pin.svelte'; import PinSlash from '$lib/components/icons/PinSlash.svelte'; - import { config, user as currentUser, settings } from '$lib/stores'; + import { config, user as currentUser, pinnedModels, settings } from '$lib/stores'; import Link from '$lib/components/icons/Link.svelte'; const i18n = getContext('i18n'); @@ -131,14 +131,14 @@ class="select-none flex h-[1.6875rem] w-full cursor-pointer items-center gap-2 rounded-xl bg-transparent px-2 text-[0.8125rem] hover:text-gray-900 dark:hover:text-gray-100" on:click={() => runAndClose(() => pinModelHandler(model?.id))} > - {#if ($settings?.pinnedModels ?? []).includes(model?.id)} + {#if $pinnedModels.includes(model?.id)} {:else} {/if}
    - {#if ($settings?.pinnedModels ?? []).includes(model?.id)} + {#if $pinnedModels.includes(model?.id)} {$i18n.t('Hide from Sidebar')} {:else} {$i18n.t('Keep in Sidebar')} diff --git a/src/lib/components/workspace/Skills.svelte b/src/lib/components/workspace/Skills.svelte index fdc6f2a465..1395b56350 100644 --- a/src/lib/components/workspace/Skills.svelte +++ b/src/lib/components/workspace/Skills.svelte @@ -21,7 +21,7 @@ deleteSkillById, toggleSkillById } from '$lib/apis/skills'; - import { capitalizeFirstLetter, parseFrontmatter, formatSkillName } from '$lib/utils'; + import { capitalizeFirstLetter, parseFrontmatter, formatSkillName, slugify } from '$lib/utils'; import TagInput from '$lib/components/common/Tags/TagInput.svelte'; import Tooltip from '../common/Tooltip.svelte'; @@ -306,7 +306,7 @@ const displayName = formatSkillName(rawName); sessionStorage.skill = JSON.stringify({ name: displayName, - id: fm.name || '', + id: slugify(rawName), description: fm.description || '', content: mdContent, is_active: true, diff --git a/src/lib/components/workspace/Skills/SkillEditor.svelte b/src/lib/components/workspace/Skills/SkillEditor.svelte index 894107bc19..06495ac2f5 100644 --- a/src/lib/components/workspace/Skills/SkillEditor.svelte +++ b/src/lib/components/workspace/Skills/SkillEditor.svelte @@ -39,7 +39,6 @@ const fm = parseFrontmatter(content); if (fm.name && !name) { name = formatSkillName(fm.name); - id = fm.name; } if (fm.description && !description) { description = fm.description; @@ -52,6 +51,7 @@ return; } loading = true; + if (!edit) id = slugify(id); await onSubmit({ id, diff --git a/src/lib/constants.ts b/src/lib/constants.ts index f18e3620a8..3cff98d10e 100644 --- a/src/lib/constants.ts +++ b/src/lib/constants.ts @@ -1,4 +1,3 @@ -import { browser, dev } from '$app/environment'; // import { version } from '../../package.json'; // LICENSE covers this Open WebUI branding surface, including name, logo, @@ -7,8 +6,8 @@ import { browser, dev } from '$app/environment'; // https://docs.openwebui.com/license. export const APP_NAME = 'Open WebUI'; -export const WEBUI_HOSTNAME = browser ? (dev ? `${location.hostname}:8080` : ``) : ''; -export const WEBUI_BASE_URL = browser ? (dev ? `http://${WEBUI_HOSTNAME}` : ``) : ``; +export const WEBUI_HOSTNAME = ''; +export const WEBUI_BASE_URL = ''; export const WEBUI_API_BASE_URL = `${WEBUI_BASE_URL}/api/v1`; export const OLLAMA_API_BASE_URL = `${WEBUI_BASE_URL}/ollama`; diff --git a/src/lib/i18n/locales/en-GB/translation.json b/src/lib/i18n/locales/en-GB/translation.json index 7f67b5f81f..f543d59b18 100644 --- a/src/lib/i18n/locales/en-GB/translation.json +++ b/src/lib/i18n/locales/en-GB/translation.json @@ -1716,6 +1716,7 @@ "Model Parameters": "", "Model Params": "", "Model Permissions": "", + "Model list refreshed": "", "Model removed from pinned models": "", "Model removed from selected models": "", "Model Response Mode": "", diff --git a/src/lib/i18n/locales/en-US/translation.json b/src/lib/i18n/locales/en-US/translation.json index a01be4d9f3..95fb38127e 100644 --- a/src/lib/i18n/locales/en-US/translation.json +++ b/src/lib/i18n/locales/en-US/translation.json @@ -709,6 +709,7 @@ "Database": "", "Datalab Marker API": "", "Datalab Marker service endpoint used for document parsing.": "", + "Date is required": "", "Date Modified": "", "Day": "", "DD/MM/YYYY": "", @@ -1719,6 +1720,7 @@ "Model Parameters": "", "Model Params": "", "Model Permissions": "", + "Model list refreshed": "", "Model removed from pinned models": "", "Model removed from selected models": "", "Model Response Mode": "", diff --git a/src/lib/shortcuts.ts b/src/lib/shortcuts.ts index 1e60ae7928..0214891d3d 100644 --- a/src/lib/shortcuts.ts +++ b/src/lib/shortcuts.ts @@ -46,6 +46,8 @@ export enum Shortcut { //Message GENERATE_MESSAGE_PAIR = 'generateMessagePair', REGENERATE_RESPONSE = 'regenerateResponse', + ALLOW_TOOL_CALL = 'allowToolCall', + DENY_TOOL_CALL = 'denyToolCall', COPY_LAST_CODE_BLOCK = 'copyLastCodeBlock', COPY_LAST_RESPONSE = 'copyLastResponse', STOP_GENERATING = 'stopGenerating', @@ -71,6 +73,8 @@ export const CONFIGURABLE_SHORTCUTS = [ Shortcut.FOCUS_INPUT, Shortcut.GENERATE_MESSAGE_PAIR, Shortcut.REGENERATE_RESPONSE, + Shortcut.ALLOW_TOOL_CALL, + Shortcut.DENY_TOOL_CALL, Shortcut.COPY_LAST_CODE_BLOCK, Shortcut.COPY_LAST_RESPONSE ] as const; @@ -95,6 +99,8 @@ export const DEFAULT_KEYBINDINGS: KeybindingsMap = { [Shortcut.FOCUS_INPUT]: 'Shift+Escape', [Shortcut.GENERATE_MESSAGE_PAIR]: 'Cmd+Shift+Enter', [Shortcut.REGENERATE_RESPONSE]: 'Cmd+R', + [Shortcut.ALLOW_TOOL_CALL]: 'Cmd+Alt+Enter', + [Shortcut.DENY_TOOL_CALL]: 'Cmd+Alt+Backspace', [Shortcut.COPY_LAST_CODE_BLOCK]: 'Cmd+Shift+;', [Shortcut.COPY_LAST_RESPONSE]: 'Cmd+Shift+C' }; @@ -333,6 +339,20 @@ export const shortcuts: ShortcutRegistry = { category: 'Message', configurable: true }, + [Shortcut.ALLOW_TOOL_CALL]: { + name: 'Allow Tool Call', + keys: ['mod', 'alt', 'Enter'], + category: 'Message', + configurable: true, + tooltip: 'Only active when a tool call is waiting for approval.' + }, + [Shortcut.DENY_TOOL_CALL]: { + name: 'Deny Tool Call', + keys: ['mod', 'alt', 'Backspace'], + category: 'Message', + configurable: true, + tooltip: 'Only active when a tool call is waiting for approval.' + }, [Shortcut.STOP_GENERATING]: { name: 'Stop Generating', keys: ['Escape'], diff --git a/src/lib/stores/chatList.ts b/src/lib/stores/chatList.ts index 0fbfea960c..1be7826c1e 100644 --- a/src/lib/stores/chatList.ts +++ b/src/lib/stores/chatList.ts @@ -61,10 +61,7 @@ export const refreshChatList = async ( return { accepted: true, allLoaded }; }; -// The sidebar's folders keep their chat lists in local component state, refreshed -// through a registry of per-folder callbacks that only the sidebar can reach. -// Handlers registered here let other components (e.g. the open chat's menu) -// request that refresh without a reference to the sidebar. +// The sidebar owns folder state. This bridge lets other components refresh it. type FolderRefreshHandler = (folderId?: string | null, chat?: ChatListItem | null) => unknown; const folderRefreshHandlers = new Set(); diff --git a/src/lib/stores/index.ts b/src/lib/stores/index.ts index abbd2c556d..9277a30d0b 100644 --- a/src/lib/stores/index.ts +++ b/src/lib/stores/index.ts @@ -1,5 +1,5 @@ import { APP_NAME } from '$lib/constants'; -import { type Writable, writable } from 'svelte/store'; +import { type Writable, derived, writable } from 'svelte/store'; import type { ModelConfig } from '$lib/apis'; import type { Banner } from '$lib/types'; import type { Socket } from 'socket.io-client'; @@ -103,6 +103,20 @@ export const banners: Writable = writable([]); export const settings: Writable = writable({}); +// Users who never pinned a model follow the admin default, so changes to it keep reaching them +export const pinnedModels = derived([settings, config], ([$settings, $config]) => + $settings?.pinnedModels === undefined + ? ($config?.default_pinned_models ?? '').split(',').filter((id) => id) + : $settings.pinnedModels +); + +// Pins for models the user cannot see are kept in their settings but left out of the sidebar +export const visiblePinnedModels = derived([pinnedModels, models], ([$pinnedModels, $models]) => + $pinnedModels.filter((id) => + $models.some((model) => model.id === id && !model.info?.meta?.hidden) + ) +); + export const audioQueue = writable(null); export const chatRequestQueues: Writable< Record @@ -200,7 +214,7 @@ type OllamaModelDetails = { }; type Settings = { - pinnedModels?: never[]; + pinnedModels?: string[]; toolServers?: never[]; detectArtifacts?: boolean; showUpdateToast?: boolean; @@ -310,6 +324,7 @@ type Config = { version: string; default_locale: string; default_models: string; + default_pinned_models?: string | null; default_prompt_suggestions: PromptSuggestion[]; features: { auth: boolean; @@ -327,6 +342,7 @@ type Config = { enable_admin_chat_access: boolean; enable_admin_analytics: boolean; enable_context_compaction?: boolean; + enable_tool_permissions?: boolean; enable_community_sharing: boolean; enable_memories: boolean; enable_plugins?: boolean; diff --git a/src/lib/utils/pptxToHtml.ts b/src/lib/utils/pptxToHtml.ts index d3eec432fe..cac78b5f95 100644 --- a/src/lib/utils/pptxToHtml.ts +++ b/src/lib/utils/pptxToHtml.ts @@ -5,7 +5,7 @@ * directly to canvas, returning PNG data URLs. * * Uses jszip (dynamically imported) and the browser Canvas 2D API. - * No theme resolution, charts, SmartArt, or animations — preview only. + * No full theme resolution, SmartArt, or animations — preview only. */ const EMU_PER_PX = 9525; @@ -38,6 +38,7 @@ type Placeholder = { type TextToken = { text: string; + fontFace: string; fontPt: number; bold: boolean; italic: boolean; @@ -45,6 +46,26 @@ type TextToken = { width: number; }; +type TextStyle = Omit; + +type TextInsets = { + left: number; + right: number; + top: number; + bottom: number; +}; + +type ChartSeries = { + name: string; + categories: string[]; + values: number[]; + color: string; +}; + +type ThemeColors = Record; + +let activeThemeColors: ThemeColors = {}; + const getTag = (el: Element) => el.tagName.split(':').pop(); const normalizePath = (path: string) => { @@ -96,6 +117,34 @@ const readRels = async ( return rels; }; +const readThemeColors = async (zip: ZipLike): Promise => { + const themePath = + Object.keys(zip.files) + .filter((path) => /^ppt\/theme\/theme\d+\.xml$/.test(path)) + .sort()[0] ?? ''; + const themeDoc = themePath ? await readXml(zip, themePath) : null; + const clrScheme = themeDoc?.getElementsByTagName('a:clrScheme')[0]; + const colors: ThemeColors = {}; + if (!clrScheme) return colors; + + for (const child of Array.from(clrScheme.children)) { + const key = getTag(child); + if (!key) continue; + + const srgb = child.getElementsByTagName('a:srgbClr')[0]?.getAttribute('val'); + if (srgb) { + colors[key] = `#${srgb}`; + continue; + } + + const sys = child.getElementsByTagName('a:sysClr')[0]; + const lastClr = sys?.getAttribute('lastClr'); + if (lastClr) colors[key] = `#${lastClr}`; + } + + return colors; +}; + const getRelationshipTarget = ( rels: Record, typeSuffix: string @@ -214,6 +263,18 @@ const paragraphAlign = ( }; const schemeColor = (val: string | null) => { + const themeKey = + val === 'bg1' + ? 'lt1' + : val === 'tx1' + ? 'dk1' + : val === 'bg2' + ? 'lt2' + : val === 'tx2' + ? 'dk2' + : val; + if (themeKey && activeThemeColors[themeKey]) return activeThemeColors[themeKey]; + switch (val) { case 'bg1': case 'lt1': @@ -238,17 +299,119 @@ const schemeColor = (val: string | null) => { } }; +const clampByte = (value: number) => Math.max(0, Math.min(255, Math.round(value))); + +const hexToRgb = (color: string) => { + const hex = color.replace('#', ''); + if (!/^[\da-f]{6}$/i.test(hex)) return null; + return { + r: parseInt(hex.slice(0, 2), 16), + g: parseInt(hex.slice(2, 4), 16), + b: parseInt(hex.slice(4, 6), 16) + }; +}; + +const rgbToHex = ({ r, g, b }: { r: number; g: number; b: number }) => + `#${[r, g, b].map((value) => clampByte(value).toString(16).padStart(2, '0')).join('')}`; + +const colorWithTransforms = (color: string | null, colorEl: Element | undefined) => { + if (!color || !colorEl) return color; + const rgb = hexToRgb(color); + if (!rgb) return color; + + const lumMod = colorEl.getElementsByTagName('a:lumMod')[0]?.getAttribute('val'); + const lumOff = colorEl.getElementsByTagName('a:lumOff')[0]?.getAttribute('val'); + const alpha = colorEl.getElementsByTagName('a:alpha')[0]?.getAttribute('val'); + const mod = lumMod ? parseInt(lumMod, 10) / 100000 : 1; + const off = lumOff ? (parseInt(lumOff, 10) / 100000) * 255 : 0; + const transformed = { + r: rgb.r * mod + off, + g: rgb.g * mod + off, + b: rgb.b * mod + off + }; + + if (alpha) { + const opacity = Math.max(0, Math.min(1, parseInt(alpha, 10) / 100000)); + return `rgba(${clampByte(transformed.r)}, ${clampByte(transformed.g)}, ${clampByte( + transformed.b + )}, ${opacity})`; + } + + return rgbToHex(transformed); +}; + +const prstColor = (val: string | null) => { + switch (val) { + case 'black': + return '#000000'; + case 'white': + return '#ffffff'; + case 'gray': + case 'grey': + return '#808080'; + default: + return null; + } +}; + +const colorFromElement = (el: Element | undefined): string | null => { + const srgb = el?.getElementsByTagName('a:srgbClr')[0]; + const srgbVal = srgb?.getAttribute('val'); + if (srgbVal) return colorWithTransforms(`#${srgbVal}`, srgb); + + const prst = el?.getElementsByTagName('a:prstClr')[0]; + const prstVal = prst?.getAttribute('val'); + if (prstVal) return colorWithTransforms(prstColor(prstVal), prst); + + const scheme = el?.getElementsByTagName('a:schemeClr')[0]; + return colorWithTransforms(schemeColor(scheme?.getAttribute('val') ?? null), scheme); +}; + const solidFillColor = (el: Element | Document | null): string | null => { if (!el) return null; const fill = el.getElementsByTagName('a:solidFill')[0]; if (!fill) return null; - const srgb = fill.getElementsByTagName('a:srgbClr')[0]; - const srgbVal = srgb?.getAttribute('val'); - if (srgbVal) return `#${srgbVal}`; + return colorFromElement(fill); +}; - const scheme = fill.getElementsByTagName('a:schemeClr')[0]; - return schemeColor(scheme?.getAttribute('val') ?? null); +const renderFill = (ctx: CanvasRenderingContext2D, el: Element | Document | null, rect: Rect) => { + if (!el) return false; + + const gradFill = el.getElementsByTagName('a:gradFill')[0]; + if (gradFill) { + const stops = Array.from(gradFill.getElementsByTagName('a:gs')) + .map((stop) => ({ + pos: parseInt(stop.getAttribute('pos') ?? '0', 10) / 100000, + color: colorFromElement(stop) + })) + .filter((stop): stop is { pos: number; color: string } => Boolean(stop.color)); + + if (stops.length > 0) { + const lin = gradFill.getElementsByTagName('a:lin')[0]; + const angle = + (((parseInt(lin?.getAttribute('ang') ?? '0', 10) / 60000) % 360) * Math.PI) / 180; + const dx = Math.cos(angle) * rect.w; + const dy = Math.sin(angle) * rect.h; + const gradient = ctx.createLinearGradient( + rect.x + rect.w / 2 - dx / 2, + rect.y + rect.h / 2 - dy / 2, + rect.x + rect.w / 2 + dx / 2, + rect.y + rect.h / 2 + dy / 2 + ); + for (const stop of stops) + gradient.addColorStop(Math.max(0, Math.min(1, stop.pos)), stop.color); + ctx.fillStyle = gradient; + ctx.fillRect(rect.x, rect.y, rect.w, rect.h); + return true; + } + } + + const fill = solidFillColor(el); + if (!fill) return false; + ctx.fillStyle = fill; + ctx.fillRect(rect.x, rect.y, rect.w, rect.h); + return true; }; const renderBackground = ( @@ -257,31 +420,84 @@ const renderBackground = ( slideW: number, slideH: number ) => { - const bgColor = solidFillColor(doc.getElementsByTagName('p:bg')[0]); - if (!bgColor) return; - ctx.fillStyle = bgColor; - ctx.fillRect(0, 0, slideW, slideH); + renderFill(ctx, doc.getElementsByTagName('p:bg')[0], { x: 0, y: 0, w: slideW, h: slideH }); }; -const fontString = ({ italic, bold, fontPt }: Pick) => - `${italic ? 'italic ' : ''}${bold ? 'bold ' : ''}${fontPt}pt Calibri, Arial, sans-serif`; +const directChild = (el: Element, tag: string) => + Array.from(el.children).find((child) => getTag(child) === tag); -const readRunStyle = (run: Element, defaultFontSize: number) => { - const rPr = run.getElementsByTagName('a:rPr')[0]; - let fontPt = defaultFontSize; - let bold = false; - let italic = false; - let color = '#000000'; +const defaultTextStyle = (fontPt: number): TextStyle => ({ + fontFace: 'Calibri', + fontPt, + bold: false, + italic: false, + color: '#000000' +}); - if (rPr) { - if (rPr.getAttribute('b') === '1') bold = true; - if (rPr.getAttribute('i') === '1') italic = true; - const sz = rPr.getAttribute('sz'); - if (sz) fontPt = parseInt(sz, 10) / 100; - color = solidFillColor(rPr) ?? color; - } +const fontString = ({ italic, bold, fontPt, fontFace }: TextStyle) => { + const family = fontFace.includes(' ') ? `"${fontFace}"` : fontFace; + return `${italic ? 'italic ' : ''}${bold ? 'bold ' : ''}${fontPt}pt ${family}, Calibri, Arial, sans-serif`; +}; - return { fontPt, bold, italic, color }; +const readTextStyle = (rPr: Element | undefined, base: TextStyle): TextStyle => { + if (!rPr) return base; + + const latin = rPr.getElementsByTagName('a:latin')[0]; + const typeface = latin?.getAttribute('typeface'); + const sz = rPr.getAttribute('sz'); + + return { + fontFace: typeface && !typeface.startsWith('+') ? typeface : base.fontFace, + fontPt: sz ? parseInt(sz, 10) / 100 : base.fontPt, + bold: rPr.getAttribute('b') === '1' ? true : rPr.getAttribute('b') === '0' ? false : base.bold, + italic: + rPr.getAttribute('i') === '1' ? true : rPr.getAttribute('i') === '0' ? false : base.italic, + color: solidFillColor(rPr) ?? base.color + }; +}; + +const readParagraphStyle = (para: Element, defaultFontSize: number) => { + const pPr = directChild(para, 'pPr'); + const defRPr = pPr?.getElementsByTagName('a:defRPr')[0]; + const endParaRPr = para.getElementsByTagName('a:endParaRPr')[0]; + return readTextStyle(defRPr ?? endParaRPr, defaultTextStyle(defaultFontSize)); +}; + +const readRunStyle = (run: Element, base: TextStyle) => { + const rPr = directChild(run, 'rPr') as Element | undefined; + return readTextStyle(rPr, base); +}; + +const textBodyInsets = ( + txBody: Element, + fallback: TextInsets = { left: 10, right: 10, top: 5, bottom: 5 } +) => { + const bodyPr = txBody.getElementsByTagName('a:bodyPr')[0]; + return { + left: bodyPr?.hasAttribute('lIns') + ? emuToPx(parseEmu(bodyPr.getAttribute('lIns'))) + : fallback.left, + right: bodyPr?.hasAttribute('rIns') + ? emuToPx(parseEmu(bodyPr.getAttribute('rIns'))) + : fallback.right, + top: bodyPr?.hasAttribute('tIns') + ? emuToPx(parseEmu(bodyPr.getAttribute('tIns'))) + : fallback.top, + bottom: bodyPr?.hasAttribute('bIns') + ? emuToPx(parseEmu(bodyPr.getAttribute('bIns'))) + : fallback.bottom + }; +}; + +const paragraphBullet = (para: Element) => { + const pPr = directChild(para, 'pPr'); + if (!pPr || pPr.getElementsByTagName('a:buNone')[0]) return ''; + + const buChar = pPr.getElementsByTagName('a:buChar')[0]?.getAttribute('char'); + if (buChar) return `${buChar} `; + + if (pPr.getElementsByTagName('a:buAutoNum')[0]) return '1. '; + return ''; }; const drawTextLine = ( @@ -293,9 +509,9 @@ const drawTextLine = ( w: number ) => { const lineWidth = line.reduce((sum, token) => sum + token.width, 0); - let cursorX = x + 4; + let cursorX = x; if (align === 'center') cursorX = x + Math.max(0, (w - lineWidth) / 2); - if (align === 'right') cursorX = x + w - lineWidth - 4; + if (align === 'right') cursorX = x + w - lineWidth; for (const token of line) { ctx.font = fontString(token); @@ -305,6 +521,395 @@ const drawTextLine = ( } }; +const paragraphTextParts = (para: Element, defaultFontSize: number) => { + const paragraphStyle = readParagraphStyle(para, defaultFontSize); + const textParts: Array<{ text: string; style: TextStyle } | { newline: true }> = []; + const bullet = paragraphBullet(para); + + for (const child of Array.from(para.children)) { + const tag = getTag(child); + if (tag === 'br') { + textParts.push({ newline: true }); + continue; + } + if (tag !== 'r' && tag !== 'fld') continue; + + const style = tag === 'r' ? readRunStyle(child, paragraphStyle) : paragraphStyle; + const text = child.getElementsByTagName('a:t')[0]?.textContent ?? ''; + if (text) textParts.push({ text, style }); + } + + if (bullet && textParts.some((part) => 'text' in part && part.text.trim())) { + textParts.unshift({ text: bullet, style: paragraphStyle }); + } + + return { paragraphStyle, textParts }; +}; + +const estimateTextHeight = ( + ctx: CanvasRenderingContext2D, + paragraphs: HTMLCollectionOf, + defaultFontSize: number, + textW: number +) => { + let height = 0; + for (let pi = 0; pi < paragraphs.length; pi++) { + const { paragraphStyle, textParts } = paragraphTextParts(paragraphs[pi], defaultFontSize); + const maxFontPt = textParts.reduce( + (max, part) => ('text' in part ? Math.max(max, part.style.fontPt) : max), + paragraphStyle.fontPt + ); + const lineHeight = maxFontPt * 1.4; + let lines = 1; + let lineWidth = 0; + + for (const part of textParts) { + if ('newline' in part) { + lines++; + lineWidth = 0; + continue; + } + ctx.font = fontString(part.style); + for (const word of part.text.split(/(\n|[^\S\n]+)/)) { + if (word === '') continue; + if (word === '\n') { + lines++; + lineWidth = 0; + continue; + } + const width = ctx.measureText(word).width; + if (lineWidth + width > textW && lineWidth > 0) { + lines++; + lineWidth = word.trim() ? width : 0; + } else if (lineWidth > 0 || word.trim()) { + lineWidth += width; + } + } + } + + height += lines * lineHeight; + if (pi < paragraphs.length - 1) height += lineHeight * 0.4; + } + return height; +}; + +const renderTextBody = ( + ctx: CanvasRenderingContext2D, + txBody: Element, + rect: Rect, + defaultFontSize: number, + ph: { type: string; idx: string } | Placeholder | null = null, + fallbackInsets?: TextInsets +) => { + const { x, y, w, h } = rect; + ctx.save(); + ctx.rect(x, y, w, h); + ctx.clip(); + + const paragraphs = txBody.getElementsByTagName('a:p'); + const bodyPr = txBody.getElementsByTagName('a:bodyPr')[0]; + const insets = textBodyInsets(txBody, fallbackInsets); + const textX = x + insets.left; + const textW = Math.max(1, w - insets.left - insets.right); + const textBottom = y + h - insets.bottom; + const textH = Math.max(1, textBottom - (y + insets.top)); + const estimatedHeight = estimateTextHeight(ctx, paragraphs, defaultFontSize, textW); + const anchor = bodyPr?.getAttribute('anchor'); + const anchorOffset = + anchor === 'b' + ? Math.max(0, textH - estimatedHeight) + : anchor === 'ctr' + ? Math.max(0, (textH - estimatedHeight) / 2) + : 0; + const textY = y + insets.top + anchorOffset; + let cursorY = textY; + + for (let pi = 0; pi < paragraphs.length; pi++) { + const para = paragraphs[pi]; + const align = paragraphAlign(para, ph); + const { paragraphStyle, textParts } = paragraphTextParts(para, defaultFontSize); + + const maxFontPt = textParts.reduce( + (max, part) => ('text' in part ? Math.max(max, part.style.fontPt) : max), + paragraphStyle.fontPt + ); + const lineHeight = maxFontPt * 1.4; + cursorY += maxFontPt; + + let line: TextToken[] = []; + let lineWidth = 0; + const flushLine = () => { + if (line.length === 0) return; + if (cursorY <= textBottom) { + drawTextLine(ctx, line, align, textX, cursorY, textW); + } + line = []; + lineWidth = 0; + cursorY += lineHeight; + }; + + if (textParts.length === 0) { + cursorY += lineHeight; + continue; + } + + for (const part of textParts) { + if ('newline' in part) { + if (line.length > 0) flushLine(); + else cursorY += lineHeight; + continue; + } + const { text, style } = part; + ctx.font = fontString(style); + ctx.textBaseline = 'alphabetic'; + + const words = text.split(/(\n|[^\S\n]+)/); + for (const word of words) { + if (word === '') continue; + if (word === '\n') { + flushLine(); + continue; + } + if (cursorY > textBottom) break; + + const width = ctx.measureText(word).width; + if (lineWidth + width > textW && line.length > 0) { + flushLine(); + if (word.trim() === '') continue; + } + if (line.length === 0 && word.trim() === '') continue; + + line.push({ ...style, text: word, width }); + lineWidth += width; + } + } + + flushLine(); + cursorY -= lineHeight * 0.6; + } + + ctx.restore(); +}; + +const renderTable = (ctx: CanvasRenderingContext2D, frame: Element, rect: Rect) => { + const table = frame.getElementsByTagName('a:tbl')[0]; + if (!table) return false; + + const grid = table.getElementsByTagName('a:tblGrid')[0]; + const colEmus = Array.from(grid?.getElementsByTagName('a:gridCol') ?? []).map((col) => + Math.max(1, parseEmu(col.getAttribute('w'))) + ); + const rows = Array.from(table.getElementsByTagName('a:tr')); + if (colEmus.length === 0 || rows.length === 0) return true; + + const rowEmus = rows.map((row) => Math.max(1, parseEmu(row.getAttribute('h')))); + const colTotal = colEmus.reduce((sum, val) => sum + val, 0); + const rowTotal = rowEmus.reduce((sum, val) => sum + val, 0); + const colWidths = colEmus.map((val) => (val / colTotal) * rect.w); + const rowHeights = rowEmus.map((val) => (val / rowTotal) * rect.h); + + let cy = rect.y; + for (let ri = 0; ri < rows.length; ri++) { + const row = rows[ri]; + const cells = Array.from(row.children).filter((child) => getTag(child) === 'tc'); + let cx = rect.x; + for (let ci = 0; ci < colWidths.length; ci++) { + const cw = colWidths[ci]; + const ch = rowHeights[ri]; + const cell = cells[ci]; + if (cell) { + const tcPr = cell.getElementsByTagName('a:tcPr')[0]; + const fill = solidFillColor(tcPr); + if (fill) { + ctx.fillStyle = fill; + ctx.fillRect(cx, cy, cw, ch); + } + + ctx.strokeStyle = 'rgba(17, 24, 39, 0.12)'; + ctx.lineWidth = 1; + ctx.beginPath(); + ctx.moveTo(cx, cy + ch); + ctx.lineTo(cx + cw, cy + ch); + ctx.stroke(); + + const txBody = cell.getElementsByTagName('a:txBody')[0]; + if (txBody) { + const marL = tcPr?.hasAttribute('marL') + ? emuToPx(parseEmu(tcPr.getAttribute('marL'))) + : 4; + const marR = tcPr?.hasAttribute('marR') + ? emuToPx(parseEmu(tcPr.getAttribute('marR'))) + : 4; + const marT = tcPr?.hasAttribute('marT') + ? emuToPx(parseEmu(tcPr.getAttribute('marT'))) + : 2; + const marB = tcPr?.hasAttribute('marB') + ? emuToPx(parseEmu(tcPr.getAttribute('marB'))) + : 2; + renderTextBody(ctx, txBody, { x: cx, y: cy, w: cw, h: ch }, 10, null, { + left: marL, + right: marR, + top: marT, + bottom: marB + }); + } + } + cx += cw; + } + cy += rowHeights[ri]; + } + + return true; +}; + +const chartPointValues = (parent: Element, tag: string) => { + const container = parent.getElementsByTagName(tag)[0]; + const pts = Array.from(container?.getElementsByTagName('c:pt') ?? []); + return pts + .sort((a, b) => parseInt(a.getAttribute('idx') ?? '0') - parseInt(b.getAttribute('idx') ?? '0')) + .map((pt) => pt.getElementsByTagName('c:v')[0]?.textContent ?? ''); +}; + +const chartSeries = (chartDoc: Document, chartType: 'bar' | 'line'): ChartSeries[] => { + const chartEl = chartDoc.getElementsByTagName( + chartType === 'bar' ? 'c:barChart' : 'c:lineChart' + )[0]; + if (!chartEl) return []; + const colors = ['#4472c4', '#70ad47', '#5b9bd5', '#ed7d31', '#a5a5a5', '#ffc000']; + + return Array.from(chartEl.getElementsByTagName('c:ser')).map((ser, index) => { + const name = chartPointValues(ser, 'c:tx')[0] || `Series ${index + 1}`; + const categories = chartPointValues(ser, 'c:cat'); + const values = chartPointValues(ser, 'c:val').map((value) => Number(value) || 0); + return { + name, + categories, + values, + color: solidFillColor(ser.getElementsByTagName('c:spPr')[0]) ?? colors[index % colors.length] + }; + }); +}; + +const renderChart = (ctx: CanvasRenderingContext2D, chartDoc: Document, rect: Rect) => { + const type = chartDoc.getElementsByTagName('c:barChart')[0] + ? 'bar' + : chartDoc.getElementsByTagName('c:lineChart')[0] + ? 'line' + : null; + if (!type) return false; + + const series = chartSeries(chartDoc, type); + const categories = series[0]?.categories ?? []; + if (series.length === 0 || categories.length === 0) return true; + + const values = series.flatMap((item) => item.values); + const maxValue = Math.max(1, ...values) * 1.15; + const left = rect.x + 42; + const top = rect.y + 28; + const right = rect.x + rect.w - 14; + const bottom = rect.y + rect.h - 34; + const plotW = Math.max(1, right - left); + const plotH = Math.max(1, bottom - top); + + ctx.save(); + ctx.fillStyle = '#ffffff'; + ctx.fillRect(rect.x, rect.y, rect.w, rect.h); + ctx.strokeStyle = 'rgba(17, 24, 39, 0.16)'; + ctx.strokeRect(rect.x, rect.y, rect.w, rect.h); + + ctx.font = '9px Arial, sans-serif'; + ctx.fillStyle = '#6b7280'; + ctx.textAlign = 'right'; + ctx.textBaseline = 'middle'; + for (let i = 0; i <= 4; i++) { + const value = (maxValue / 4) * i; + const y = bottom - (value / maxValue) * plotH; + ctx.strokeStyle = i === 0 ? '#9ca3af' : 'rgba(17, 24, 39, 0.08)'; + ctx.beginPath(); + ctx.moveTo(left, y); + ctx.lineTo(right, y); + ctx.stroke(); + ctx.fillText(Math.round(value).toString(), left - 6, y); + } + + ctx.textAlign = 'center'; + ctx.textBaseline = 'top'; + categories.forEach((category, index) => { + const x = left + ((index + 0.5) / categories.length) * plotW; + ctx.fillStyle = '#6b7280'; + ctx.fillText(category, x, bottom + 7); + }); + + if (type === 'bar') { + const groupW = plotW / categories.length; + const barW = (groupW * 0.7) / Math.max(1, series.length); + series.forEach((item, si) => { + ctx.fillStyle = item.color; + item.values.forEach((value, vi) => { + const x = left + vi * groupW + groupW * 0.15 + si * barW; + const h = (value / maxValue) * plotH; + const y = bottom - h; + ctx.fillRect(x, y, Math.max(1, barW - 2), h); + ctx.fillStyle = '#111827'; + ctx.fillText(String(value), x + barW / 2, y - 12); + ctx.fillStyle = item.color; + }); + }); + } else { + series.forEach((item) => { + ctx.strokeStyle = item.color; + ctx.fillStyle = item.color; + ctx.lineWidth = 2; + ctx.beginPath(); + item.values.forEach((value, vi) => { + const x = left + ((vi + 0.5) / categories.length) * plotW; + const y = bottom - (value / maxValue) * plotH; + if (vi === 0) ctx.moveTo(x, y); + else ctx.lineTo(x, y); + }); + ctx.stroke(); + item.values.forEach((value, vi) => { + const x = left + ((vi + 0.5) / categories.length) * plotW; + const y = bottom - (value / maxValue) * plotH; + ctx.beginPath(); + ctx.arc(x, y, 2.5, 0, Math.PI * 2); + ctx.fill(); + ctx.fillStyle = '#111827'; + ctx.fillText(String(value), x, y - 14); + ctx.fillStyle = item.color; + }); + }); + } + + ctx.textAlign = 'left'; + ctx.textBaseline = 'middle'; + let legendX = rect.x + rect.w * 0.48; + const legendY = rect.y + 13; + series.forEach((item) => { + ctx.fillStyle = item.color; + ctx.fillRect(legendX, legendY - 3, 6, 6); + ctx.fillStyle = '#6b7280'; + ctx.fillText(item.name, legendX + 9, legendY); + legendX += ctx.measureText(item.name).width + 24; + }); + ctx.restore(); + + return true; +}; + +const renderConnector = (ctx: CanvasRenderingContext2D, shape: Element, rect: Rect) => { + const spPr = shape.getElementsByTagName('p:spPr')[0]; + const line = spPr?.getElementsByTagName('a:ln')[0]; + ctx.save(); + ctx.strokeStyle = solidFillColor(line) ?? '#9ca3af'; + ctx.lineWidth = Math.max(1, emuToPx(parseEmu(line?.getAttribute('w') ?? '9525'))); + ctx.beginPath(); + ctx.moveTo(rect.x, rect.y); + ctx.lineTo(rect.x + rect.w, rect.y + rect.h); + ctx.stroke(); + ctx.restore(); +}; + /** Load a data URI into an Image element and wait for it. */ const loadImage = (src: string): Promise => new Promise((resolve, reject) => { @@ -322,6 +927,7 @@ export async function pptxToImages( ): Promise<{ images: string[]; width: number; height: number }> { const JSZip = (await import('jszip')).default; const zip = (await JSZip.loadAsync(buffer)) as ZipLike; + activeThemeColors = await readThemeColors(zip); // ── Read slide dimensions from presentation.xml ────────────────── let slideW = 960; @@ -402,7 +1008,7 @@ export async function pptxToImages( } const shapes = Array.from(spTree.children).filter((el) => - ['sp', 'pic'].includes(getTag(el) ?? '') + ['sp', 'pic', 'graphicFrame', 'cxnSp'].includes(getTag(el) ?? '') ); for (const shape of shapes) { @@ -414,13 +1020,24 @@ export async function pptxToImages( const { x, y, w, h } = rect; if (w === 0 && h === 0) continue; + const tag = getTag(shape); - if (getTag(shape) === 'sp') { - const shapeFill = solidFillColor(shape.getElementsByTagName('p:spPr')[0]); - if (shapeFill) { - ctx.fillStyle = shapeFill; - ctx.fillRect(x, y, w, h); - } + if (tag === 'cxnSp') { + renderConnector(ctx, shape, rect); + continue; + } + + if (tag === 'graphicFrame') { + if (renderTable(ctx, shape, rect)) continue; + + const chartRelId = shape.getElementsByTagName('c:chart')[0]?.getAttribute('r:id') ?? ''; + const chartPath = rels[chartRelId]?.target; + const chartDoc = chartPath ? await readXml(zip, chartPath) : null; + if (chartDoc && renderChart(ctx, chartDoc, rect)) continue; + } + + if (tag === 'sp') { + renderFill(ctx, shape.getElementsByTagName('p:spPr')[0], rect); } // ── Picture ────────────────────────────────────────────── @@ -447,89 +1064,8 @@ export async function pptxToImages( const txBody = shape.getElementsByTagName('p:txBody')[0]; if (!txBody) continue; - ctx.save(); - ctx.rect(x, y, w, h); - ctx.clip(); - - const paragraphs = txBody.getElementsByTagName('a:p'); - let cursorY = y; const defaultFontSize = placeholderFontSize(ph ?? placeholder ?? null); - - for (let pi = 0; pi < paragraphs.length; pi++) { - const para = paragraphs[pi]; - const runs = para.getElementsByTagName('a:r'); - const align = paragraphAlign(para, ph ?? placeholder ?? null); - - if (runs.length === 0) { - cursorY += defaultFontSize * 1.5; - continue; - } - - // Calculate max font size in this paragraph for line height - let maxFontPt = defaultFontSize; - for (let ri = 0; ri < runs.length; ri++) { - const rPr = runs[ri].getElementsByTagName('a:rPr')[0]; - if (rPr) { - const sz = rPr.getAttribute('sz'); - if (sz) { - const pt = parseInt(sz, 10) / 100; - if (pt > maxFontPt) maxFontPt = pt; - } - } - } - - const lineHeight = maxFontPt * 1.4; - cursorY += maxFontPt; // baseline offset - - let line: TextToken[] = []; - let lineWidth = 0; - const flushLine = () => { - if (line.length === 0) return; - if (cursorY <= y + h) { - drawTextLine(ctx, line, align, x, cursorY, w); - } - line = []; - lineWidth = 0; - cursorY += lineHeight; - }; - - for (let ri = 0; ri < runs.length; ri++) { - const run = runs[ri]; - const text = run.getElementsByTagName('a:t')[0]?.textContent ?? ''; - if (!text) continue; - - const style = readRunStyle(run, defaultFontSize); - ctx.font = fontString(style); - - ctx.textBaseline = 'alphabetic'; - - // Simple word-wrap within the shape bounds - const words = text.split(/(\n|[^\S\n]+)/); - for (const word of words) { - if (word === '') continue; - if (word === '\n') { - flushLine(); - continue; - } - if (cursorY > y + h) break; - - const width = ctx.measureText(word).width; - if (lineWidth + width > w - 8 && line.length > 0) { - flushLine(); - if (word.trim() === '') continue; - } - if (line.length === 0 && word.trim() === '') continue; - - line.push({ ...style, text: word, width }); - lineWidth += width; - } - } - - flushLine(); - cursorY -= lineHeight * 0.6; // paragraph spacing - } - - ctx.restore(); + renderTextBody(ctx, txBody, rect, defaultFontSize, ph ?? placeholder ?? null); } images.push(canvas.toDataURL('image/png')); diff --git a/src/routes/(app)/+layout.svelte b/src/routes/(app)/+layout.svelte index e9931a5839..eec8049922 100644 --- a/src/routes/(app)/+layout.svelte +++ b/src/routes/(app)/+layout.svelte @@ -90,19 +90,7 @@ }; const setUserSettings = async (cb?: () => Promise) => { - let userSettings = await getUserSettings(localStorage.token).catch((error) => { - console.error(error); - return null; - }); - - if (!userSettings) { - try { - userSettings = JSON.parse(localStorage.getItem('settings') ?? '{}'); - } catch (e: unknown) { - console.error('Failed to parse settings from localStorage', e); - userSettings = {}; - } - } + const userSettings = await getUserSettings(localStorage.token); if (userSettings?.ui) { settings.set(userSettings.ui); @@ -254,14 +242,20 @@ } clearChatInputStorage(); - await Promise.all([ - checkLocalDBChats(), - setBanners().catch((e) => console.error('Failed to load banners:', e)), - setTools().catch((e) => console.error('Failed to load tools:', e)), - setUserSettings(async () => { - await setModels().catch((e) => console.error('Failed to load models:', e)); - }).catch((e) => console.error('Failed to load user settings:', e)) - ]); + try { + await Promise.all([ + checkLocalDBChats(), + setBanners().catch((e) => console.error('Failed to load banners:', e)), + setTools().catch((e) => console.error('Failed to load tools:', e)), + setUserSettings(async () => { + await setModels().catch((e) => console.error('Failed to load models:', e)); + }) + ]); + } catch (e) { + console.error('Failed to load user settings:', e); + toast.error($i18n.t('Failed to load Interface settings')); + return; + } selectedTerminalId.set(localStorage.selectedTerminalId ?? null); @@ -338,7 +332,7 @@ } else if (shortcut === Shortcut.OPEN_MODEL_SELECTOR) { console.log('Shortcut triggered: OPEN_MODEL_SELECTOR'); event.preventDefault(); - document.getElementById('model-selector-0-button')?.click(); + document.getElementById('model-selector-model-button')?.click(); } else if (shortcut === Shortcut.NEW_TEMPORARY_CHAT) { console.log('Shortcut triggered: NEW_TEMPORARY_CHAT'); event.preventDefault(); @@ -355,6 +349,24 @@ console.log('Shortcut triggered: GENERATE_MESSAGE_PAIR'); event.preventDefault(); document.getElementById('generate-message-pair-button')?.click(); + } else if (shortcut === Shortcut.ALLOW_TOOL_CALL) { + const button = [...document.getElementsByClassName('tool-call-allow-button')] + .reverse() + .find((el) => !(el as HTMLButtonElement).disabled) as HTMLButtonElement | undefined; + if (button) { + console.log('Shortcut triggered: ALLOW_TOOL_CALL'); + event.preventDefault(); + button.click(); + } + } else if (shortcut === Shortcut.DENY_TOOL_CALL) { + const button = [...document.getElementsByClassName('tool-call-deny-button')] + .reverse() + .find((el) => !(el as HTMLButtonElement).disabled) as HTMLButtonElement | undefined; + if (button) { + console.log('Shortcut triggered: DENY_TOOL_CALL'); + event.preventDefault(); + button.click(); + } } else if ( shortcut === Shortcut.REGENERATE_RESPONSE && document.activeElement?.id === 'chat-input' diff --git a/src/routes/(app)/automations/+page.svelte b/src/routes/(app)/automations/+page.svelte index 74e31daac8..5c3aa8e11f 100644 --- a/src/routes/(app)/automations/+page.svelte +++ b/src/routes/(app)/automations/+page.svelte @@ -5,8 +5,9 @@ import relativeTime from 'dayjs/plugin/relativeTime'; import { toast } from 'svelte-sonner'; import { goto } from '$app/navigation'; - import { WEBUI_NAME, user, config, folders } from '$lib/stores'; + import { WEBUI_NAME, user, config, channels, folders } from '$lib/stores'; import { getFolders } from '$lib/apis/folders'; + import { getChannels } from '$lib/apis/channels'; import { createAutomation, @@ -67,6 +68,7 @@ let importFiles: FileList | null = null; let automationsImportInputElement: HTMLInputElement; let foldersLoaded = false; + let channelsLoaded = false; const syncHeader = () => { automationsLayout?.setHeader({ @@ -151,6 +153,13 @@ foldersLoaded = true; }; + const ensureChannels = async () => { + if (channelsLoaded || ($channels ?? []).length > 0) return; + const res = await getChannels(localStorage.token).catch(() => null); + if (res) channels.set(res); + channelsLoaded = true; + }; + const toggleHandler = async (automation: AutomationResponse) => { const res = await toggleAutomationById(localStorage.token, automation.id).catch((err) => { toast.error(`${err}`); @@ -216,6 +225,16 @@ : $i18n.t('Never'); }; + const formatDestination = (automation: AutomationResponse): string => { + if (automation.data.target?.type === 'channel') { + const channel = ($channels ?? []).find( + (channel) => channel.id === automation.data.target?.channel_id + ); + return channel?.name ? `#${channel.name}` : $i18n.t('Channel'); + } + return automation.folder_id ? $i18n.t('Folder') : $i18n.t('New chat'); + }; + const getAllAutomations = async () => { let currentPage = 1; let allAutomations: AutomationResponse[] = []; @@ -341,6 +360,7 @@ loaded = true; syncHeader(); + ensureChannels(); return () => { clearTimeout(searchDebounceTimer); @@ -580,11 +600,14 @@
    diff --git a/src/routes/+layout.svelte b/src/routes/+layout.svelte index 59dbcd808f..3f39c62430 100644 --- a/src/routes/+layout.svelte +++ b/src/routes/+layout.svelte @@ -62,7 +62,7 @@ removeTerminalConnection } from '$lib/utils/connections'; - import { WEBUI_API_BASE_URL, WEBUI_BASE_URL, WEBUI_HOSTNAME } from '$lib/constants'; + import { WEBUI_API_BASE_URL, WEBUI_BASE_URL } from '$lib/constants'; import { bestMatchingLanguage, cleanText, diff --git a/vite.config.ts b/vite.config.ts index df802c081c..5052e5d6fa 100644 --- a/vite.config.ts +++ b/vite.config.ts @@ -3,6 +3,8 @@ import { defineConfig } from 'vite'; import { viteStaticCopy } from 'vite-plugin-static-copy'; +const backendTarget = process.env.WEBUI_BACKEND_URL || 'http://localhost:8080'; + export default defineConfig({ plugins: [ sveltekit(), @@ -23,6 +25,32 @@ export default defineConfig({ build: { sourcemap: true }, + server: { + proxy: { + '/api': { + target: backendTarget, + changeOrigin: true, + ws: true + }, + '/ollama': { + target: backendTarget, + changeOrigin: true + }, + '/openai': { + target: backendTarget, + changeOrigin: true + }, + '/oauth': { + target: backendTarget, + changeOrigin: true + }, + '/ws': { + target: backendTarget, + changeOrigin: true, + ws: true + } + } + }, worker: { format: 'es' },