diff --git a/.env.example b/.env.example index 313024511a..81cd826ada 100644 --- a/.env.example +++ b/.env.example @@ -25,8 +25,15 @@ ENABLE_KNOWLEDGE_FILE_RETENTION=false # Comma-separated chunk metadata keys to expose to the model alongside retrieved content. RAG_SOURCE_METADATA_KEYS='' -# Set to false to disable workspace Tools and Functions. +# Master switch for internal Tools/Functions and external OpenAPI/MCP/Open Terminal plugins. +# All plugin switches require a restart; the master overrides all feature switches. ENABLE_PLUGINS=true +# Set false to disable workspace Tools and their dependency installation. +ENABLE_TOOLS=true +# Set false to disable Functions (filters, pipes, actions, event functions) and their dependencies. +ENABLE_FUNCTIONS=true +# Set false to disable external tools and terminals, including personal direct connections. +ENABLE_TOOL_SERVERS=true # For production you should set this to match the proxy configuration (127.0.0.1) FORWARDED_ALLOW_IPS='*' diff --git a/backend/open_webui/constants.py b/backend/open_webui/constants.py index c6c78de35a..2e5b0ff60c 100644 --- a/backend/open_webui/constants.py +++ b/backend/open_webui/constants.py @@ -58,6 +58,7 @@ class ERROR_MESSAGES(str, Enum): INVALID_TOKEN = 'Your session has expired or the token is invalid. Please sign in again.' INVALID_CRED = 'The email or password provided is incorrect. Please check for typos and try logging in again.' + OAUTH_LOGIN_FAILED = 'Sign-in with your identity provider failed. Please contact your administrator for assistance.' INVALID_EMAIL_FORMAT = "The email format you entered is invalid. Please double-check and make sure you're using a valid email address (e.g., yourname@example.com)." INCORRECT_PASSWORD = 'The password provided is incorrect. Please check for typos and try again.' INVALID_TRUSTED_HEADER = ( diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index dc160e1592..1f3a719bcf 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -159,7 +159,7 @@ ENABLE_DB_MIGRATIONS = os.getenv('ENABLE_DB_MIGRATIONS', 'True').lower() == 'tru # Swap the JSON encoder/decoder used across the app (HTTP request bodies, JSONResponse # bodies, upstream provider responses, socket.io payloads) from the stdlib `json` module # to orjson. Faster, but stricter: see open_webui/utils/json_codec.py for the differences. -ENABLE_ORJSON = os.getenv('ENABLE_ORJSON', 'False').lower() == 'true' +ENABLE_ORJSON = os.getenv('ENABLE_ORJSON', 'True').lower() == 'true' # Function to parse each section @@ -1192,6 +1192,10 @@ VIEW_FILE_DEFAULT_MAX_CHARS = _int_env('VIEW_FILE_DEFAULT_MAX_CHARS', 10_000) #################################### ENABLE_PLUGINS = os.getenv('ENABLE_PLUGINS', 'True').lower() == 'true' +# Deployment controls: the master switch always overrides all feature switches. +ENABLE_TOOLS = ENABLE_PLUGINS and os.getenv('ENABLE_TOOLS', 'True').lower() == 'true' +ENABLE_FUNCTIONS = ENABLE_PLUGINS and os.getenv('ENABLE_FUNCTIONS', 'True').lower() == 'true' +ENABLE_TOOL_SERVERS = ENABLE_PLUGINS and os.getenv('ENABLE_TOOL_SERVERS', 'True').lower() == 'true' ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS = ( os.getenv('ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS', 'True').lower() == 'true' diff --git a/backend/open_webui/events.py b/backend/open_webui/events.py index 81b8cc3968..aa4a49dccd 100644 --- a/backend/open_webui/events.py +++ b/backend/open_webui/events.py @@ -8,9 +8,10 @@ import uuid from types import SimpleNamespace from typing import Any -from open_webui.env import ENABLE_PLUGINS, VERSION -from open_webui.models.config import Config from pydantic import BaseModel, ConfigDict, Field, model_validator + +from open_webui.env import ENABLE_FUNCTIONS, VERSION +from open_webui.models.config import Config from open_webui.retrieval.web.utils import validate_url from open_webui.utils.webhook import post_webhook @@ -1103,7 +1104,7 @@ class SocketSessionEventSink: async def dispatch_event_functions( app: Any, event: Event, request: Any | None = None, extra_function_ids: list[str] | None = None ) -> None: - if not ENABLE_PLUGINS: + if not ENABLE_FUNCTIONS: return from open_webui.models.functions import Functions diff --git a/backend/open_webui/functions.py b/backend/open_webui/functions.py index 41e1a42d9e..a0ed278af2 100644 --- a/backend/open_webui/functions.py +++ b/backend/open_webui/functions.py @@ -19,7 +19,7 @@ from starlette.responses import Response, StreamingResponse from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL from open_webui.constants import ERROR_MESSAGES -from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, ENABLE_PLUGINS, GLOBAL_LOG_LEVEL +from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, ENABLE_FUNCTIONS, GLOBAL_LOG_LEVEL from open_webui.models.functions import Functions from open_webui.models.models import Models from open_webui.models.users import UserModel @@ -69,7 +69,7 @@ async def get_function_module_by_id(request: Request, pipe_id: str): async def get_function_models(request): - if not ENABLE_PLUGINS: + if not ENABLE_FUNCTIONS: return [] pipes = await Functions.get_functions_by_type('pipe', active_only=True) diff --git a/backend/open_webui/internal/db.py b/backend/open_webui/internal/db.py index e2f523cfc2..ed691663fb 100644 --- a/backend/open_webui/internal/db.py +++ b/backend/open_webui/internal/db.py @@ -134,7 +134,7 @@ class JSONField(types.TypeDecorator): # TEXT-backed JSON storage cache_ok = True def process_bind_param(self, value: _T | None, dialect: Dialect) -> Any: - return JSONCodec.dumps(value) if value is not None else None + return JSONCodec.dumps(value, ensure_ascii=False) if value is not None else None def process_result_value(self, value: _T | None, dialect: Dialect) -> Any: return JSONCodec.loads(value) if value is not None else None @@ -265,7 +265,7 @@ def _json_codec_kwargs(kwargs: dict) -> dict: Unlike ``JSONField``, those serialize through the engine, which otherwise uses stdlib ``json``. With ``ENABLE_ORJSON`` off JSONCodec is stdlib ``json`` anyway. """ - kwargs.setdefault('json_serializer', JSONCodec.dumps) + kwargs.setdefault('json_serializer', lambda value: JSONCodec.dumps(value, ensure_ascii=False)) kwargs.setdefault('json_deserializer', JSONCodec.loads) return kwargs diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index ecba6407f3..b7ec53ab7d 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -74,9 +74,7 @@ from open_webui.config import ( seed_registered_defaults, ) from open_webui.constants import ERROR_MESSAGES, TASKS -from open_webui.utils.recurrence import RecurrenceEvaluationTimeout from open_webui.env import ( - USE_SLIM, AIOHTTP_CLIENT_SESSION_SSL, AUDIT_EXCLUDED_PATHS, AUDIT_INCLUDED_PATHS, @@ -88,6 +86,7 @@ from open_webui.env import ( ENABLE_COMPRESSION_MIDDLEWARE, ENABLE_CUSTOM_MODEL_FALLBACK, ENABLE_EASTER_EGGS, + ENABLE_FUNCTIONS, # OAuth Back-Channel Logout ENABLE_OAUTH_BACKCHANNEL_LOGOUT, ENABLE_OTEL, @@ -98,6 +97,8 @@ from open_webui.env import ( ENABLE_SCIM, ENABLE_SIGNUP_PASSWORD_CONFIRMATION, ENABLE_STAR_SESSIONS_MIDDLEWARE, + ENABLE_TOOL_SERVERS, + ENABLE_TOOLS, ENABLE_VERSION_UPDATE_CHECK, ENABLE_WEBSOCKET_SUPPORT, EXTERNAL_PWA_MANIFEST_URL, @@ -113,6 +114,7 @@ from open_webui.env import ( RESET_CONFIG_ON_START, SAFE_MODE, SCIM_TOKEN, + USE_SLIM, VERSION, WEBSOCKET_HEARTBEAT_INTERVAL, WEBSOCKET_MANAGER, @@ -193,6 +195,7 @@ from open_webui.socket.main import ( periodic_session_pool_cleanup, periodic_usage_pool_cleanup, redis_event_listener, + sio, ) from open_webui.socket.main import ( app as socket_app, @@ -221,6 +224,7 @@ from open_webui.utils.auth import ( get_http_authorization_cred, get_license_data, get_verified_user, + is_valid_token, ) from open_webui.utils.chat import ( chat_completed as chat_completed_handler, @@ -269,6 +273,7 @@ from open_webui.utils.oauth import ( resolve_oauth_client_info, ) from open_webui.utils.plugin import install_tool_and_function_dependencies +from open_webui.utils.recurrence import RecurrenceEvaluationTimeout from open_webui.utils.redis import get_redis_client from open_webui.utils.session_pool import cleanup_response, get_client_timeout, get_session, stream_wrapper from open_webui.utils.tool_approval import ( @@ -397,6 +402,9 @@ async def lifespan(app: FastAPI): if WEBSOCKET_MANAGER == 'redis': app.state.redis_event_listener = asyncio.create_task(redis_event_listener()) + # socket.io only starts listening on its first connect; event call answers need it earlier + sio.manager_initialized = True + sio.manager.initialize() 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()) @@ -430,7 +438,9 @@ async def lifespan(app: FastAPI): log.warning(f'Failed to pre-fetch models at startup: {e}') # Pre-fetch tool server specs so the first request doesn't pay the latency cost - if len(await Config.get('tool_server.connections', []) or []) > 0: + if ENABLE_TOOL_SERVERS and ( + await Config.get('tool_server.connections', []) or await Config.get('terminal_server.connections', []) + ): mock_request = Request( { 'type': 'http', @@ -497,7 +507,7 @@ async def lifespan(app: FastAPI): await publish_event(app, EVENTS.SYSTEM_SHUTDOWN_COMPLETED, source='system') -# Opt-in (ENABLE_ORJSON): orjson for request-body parsing and JSONResponse bodies; +# ENABLE_ORJSON: orjson for request-body parsing and JSONResponse bodies; # response_model routes keep FastAPI's Pydantic fast path either way. apply_orjson_http_json() @@ -629,7 +639,7 @@ async def initialize_runtime_config(app: FastAPI): migrate_access_control(connection.get('config', {})) await Config.upsert({'tool_server.connections': connections}) - for tool_server_connection in connections: + for tool_server_connection in connections if ENABLE_TOOL_SERVERS else []: if tool_server_connection.get('type', 'openapi') == 'mcp': server_id = (tool_server_connection.get('info') or {}).get('id') auth_type = tool_server_connection.get('auth_type', 'none') @@ -1001,6 +1011,7 @@ async def unload_model(request: Request, form_data: ModelUnloadForm, user=Depend data=payload, headers=headers, cookies=cookies, + ssl=AIOHTTP_CLIENT_SESSION_SSL, ) as r: if not r.ok: errors.append({'url_idx': idx, 'error': await r.text()}) @@ -1040,6 +1051,7 @@ async def unload_model(request: Request, form_data: ModelUnloadForm, user=Depend json={'model': actual_model}, headers=headers, cookies=cookies, + ssl=AIOHTTP_CLIENT_SESSION_SSL, ) as r: if not r.ok: detail = await r.text() @@ -2242,7 +2254,7 @@ async def get_app_config(request: Request): status_code=status.HTTP_401_UNAUTHORIZED, detail='Invalid token', ) - if data is not None and 'id' in data: + if data is not None and 'id' in data and await is_valid_token(data, request.app.state.redis): user = await Users.get_user_by_id(data['id']) onboarding = False @@ -2347,8 +2359,12 @@ async def get_app_config(request: Request): 'enable_public_active_users_count': ENABLE_PUBLIC_ACTIVE_USERS_COUNT, 'enable_easter_eggs': ENABLE_EASTER_EGGS, 'enable_direct_connections': config.get('direct.enable'), - 'enable_direct_integrations': config.get('direct.integrations.enable', False), + 'enable_direct_integrations': ENABLE_TOOL_SERVERS + and config.get('direct.integrations.enable', False), 'enable_plugins': ENABLE_PLUGINS, + 'enable_tools': ENABLE_TOOLS, + 'enable_functions': ENABLE_FUNCTIONS, + 'enable_tool_servers': ENABLE_TOOL_SERVERS, 'enable_folders': config.get('folders.enable'), 'folder_max_file_count': config.get('folders.max_file_count'), 'enable_channels': config.get('channels.enable'), @@ -2673,6 +2689,9 @@ except Exception as e: async def register_client(request, client_id: str) -> bool: + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + server_type, server_id = client_id.split(':', 1) connection = None @@ -2775,6 +2794,9 @@ async def oauth_client_authorize( user=Depends(get_verified_user), ): # ensure_valid_client_registration + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + client = await oauth_client_manager.get_client(client_id) client_info = await oauth_client_manager.get_client_info(client_id) if client is None or client_info is None: @@ -2816,6 +2838,9 @@ async def oauth_client_callback( request: Request, response: Response, ): + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + return await oauth_client_manager.handle_callback( request, client_id=client_id, diff --git a/backend/open_webui/models/automations.py b/backend/open_webui/models/automations.py index 954026cfc8..0aab02e5f6 100644 --- a/backend/open_webui/models/automations.py +++ b/backend/open_webui/models/automations.py @@ -244,6 +244,15 @@ class AutomationTable: await db.commit() return AutomationModel.model_validate(row) + async def update_last_run_at(self, id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationModel]: + async with get_async_db_context(db) as db: + row = await db.get(Automation, id) + if not row: + return None + row.last_run_at = int(time.time_ns()) + await db.commit() + return AutomationModel.model_validate(row) + async def clear_folder_ids( self, user_id: str, diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index c5620f0cfe..a02e4327a4 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -660,11 +660,13 @@ class ChatTable: db: AsyncSession | None = None, ) -> list[ChatModel]: async with get_async_db_context(db) as session: + from open_webui.utils.access_control.folders import has_folder_write_access + # Validate folder_id references — clear any that don't exist folder_ids = {f.folder_id for f in chat_import_forms if f.folder_id} existing = set() for fid in folder_ids: - if await Folders.get_folder_by_id_and_user_id(fid, user_id, db=session): + if await has_folder_write_access(user_id, fid, db=session): existing.add(fid) cleared = 0 @@ -2583,11 +2585,11 @@ class ChatTable: if not file_ids: return None - chat_message_file_ids = { - item.id for item in await self.get_chat_files_by_chat_id_and_message_id(chat_id, message_id, db=db) - } + async with get_async_db_context(db) as session: + result = await session.execute(select(ChatFile.file_id).filter_by(chat_id=chat_id)) + chat_file_ids = set(result.scalars().all()) # Remove duplicates and existing file_ids - file_ids = list({file_id for file_id in file_ids if file_id and file_id not in chat_message_file_ids}) + file_ids = list({file_id for file_id in file_ids if file_id and file_id not in chat_file_ids}) if not file_ids: return None diff --git a/backend/open_webui/retrieval/models/colbert.py b/backend/open_webui/retrieval/models/colbert.py index 94e6f8af5e..0a12f13b8b 100644 --- a/backend/open_webui/retrieval/models/colbert.py +++ b/backend/open_webui/retrieval/models/colbert.py @@ -4,12 +4,29 @@ import os import numpy as np import torch from colbert.infra import ColBERTConfig +from colbert.modeling import base_colbert, hf_colbert from colbert.modeling.checkpoint import Checkpoint from open_webui.retrieval.models.base_reranker import BaseReranker +from transformers import PretrainedConfig log = logging.getLogger(__name__) +def class_factory(name_or_path: str) -> type: + hf_colbert_class = hf_colbert.class_factory(name_or_path) + + # colbert-ai never calls post_init, which transformers 5 needs to finish setting up the model + class HF_ColBERT(hf_colbert_class): + def __init__(self, config: PretrainedConfig, colbert_config: ColBERTConfig) -> None: + super().__init__(config, colbert_config) + self.post_init() + + return HF_ColBERT + + +base_colbert.class_factory = class_factory + + class ColBERT(BaseReranker): def __init__(self, name, **kwargs) -> None: log.info('ColBERT: Loading model %s', name) diff --git a/backend/open_webui/retrieval/vector/dbs/milvus.py b/backend/open_webui/retrieval/vector/dbs/milvus.py index d09f01c16d..0bf4a5955d 100644 --- a/backend/open_webui/retrieval/vector/dbs/milvus.py +++ b/backend/open_webui/retrieval/vector/dbs/milvus.py @@ -4,6 +4,8 @@ NOTE: This vector database integration is community-supported and maintained on import logging import re +from collections.abc import Callable, Iterable +from concurrent.futures import ThreadPoolExecutor from typing import Any, Optional from open_webui.config import ( @@ -18,24 +20,32 @@ from open_webui.config import ( MILVUS_TOKEN, MILVUS_URI, ) +from open_webui.env import ENABLE_DB_MIGRATIONS from open_webui.retrieval.vector.main import ( GetResult, SearchResult, VectorDBBase, VectorItem, ) -from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata +from open_webui.retrieval.vector.utils import iter_filter_conditions, merge_hybrid_search_results, process_metadata from open_webui.utils.json_codec import JSONCodec -from pymilvus import DataType +from pymilvus import CollectionSchema, DataType, Function, FunctionType from pymilvus import MilvusClient as Client +from pymilvus.client.types import LoadState from pymilvus.exceptions import MilvusException log = logging.getLogger(__name__) -# Milvus caps stored text length (here the chunk lives under the JSON `data` +# Milvus caps stored text length (here the chunk lives in `text` or the JSON `data` # field). Clamp long chunks before insert so one oversized chunk can't fail the # whole batch and leave the file with zero embeddings. MILVUS_TEXT_MAX_LENGTH = 65535 +# Milvus cannot add BM25 to an existing collection, so migration copies each one into a new collection. +BM25_STAGING_SUFFIX = '_bm25_staging' +# Even rows at Milvus's field size limits keep a batch this size below its 64 MB message limit. +BM25_BACKFILL_BATCH_SIZE = 128 +BM25_BACKFILL_WORKERS = 8 +BM25_BACKFILL_INSERTS_IN_FLIGHT = 4 _SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$') @@ -68,6 +78,163 @@ def _metadata_exprs(filter: Optional[dict]) -> list[str]: return exprs +def _chunk_text(entity: dict) -> Optional[str]: + return entity['text'] if 'text' in entity else entity.get('data', {}).get('text') + + +def _truncate_text(text: str) -> str: + return text.encode()[:MILVUS_TEXT_MAX_LENGTH].decode(errors='ignore') + + +def _bm25_rows(rows: list[dict]) -> list[dict]: + return [ + { + 'id': row['id'], + 'vector': row['vector'], + 'text': row['data']['text'], + 'metadata': row['metadata'], + } + for row in rows + ] + + +def _add_bm25_fields(schema: CollectionSchema): + schema.add_field(field_name='sparse', datatype=DataType.SPARSE_FLOAT_VECTOR) + schema.add_function( + Function( + name='text_bm25', + function_type=FunctionType.BM25, + input_field_names=['text'], + output_field_names=['sparse'], + ) + ) + + +def _supports_bm25(client: Client) -> bool: + return tuple(int(part) for part in re.findall(r'\d+', client.get_server_version())[:2]) >= (2, 5) + + +def _has_bm25_field(client: Client, collection: str) -> bool: + return any(field['name'] == 'sparse' for field in client.describe_collection(collection)['fields']) + + +def _pending_bm25_collections(client: Client, collections: Iterable[str], output_fields: list[str]) -> dict[str, int]: + """Finishes or clears an interrupted run, then maps each collection left to migrate to its dimension.""" + existing_collections = set(client.list_collections()) + pending_collections = {} + for collection in collections: + staging_collection = f'{collection}{BM25_STAGING_SUFFIX}' + if staging_collection in existing_collections: + if collection not in existing_collections: + # An earlier start dropped the original after a complete copy but stopped before this rename. + client.rename_collection(staging_collection, collection) + client.load_collection(collection) + continue + client.drop_collection(staging_collection) + if collection not in existing_collections: + continue + fields = client.describe_collection(collection)['fields'] + field_names = {field['name'] for field in fields} + if 'sparse' not in field_names and field_names.issuperset(output_fields): + pending_collections[collection] = next( + field['params']['dim'] for field in fields if field['name'] == 'vector' + ) + return pending_collections + + +def _backfill_bm25_collections( + client: Client, + collections: Iterable[str], + create_bm25_collection: Callable[[str, int], None], + output_fields: list[str], + to_bm25_rows: Callable[[list[dict]], list[dict]], +): + if not _supports_bm25(client): + log.info('Milvus has no BM25 (needs 2.5+), native hybrid search stays off.') + return + pending_collections = _pending_bm25_collections(client, collections, output_fields) + if not pending_collections: + return + + log.info('Migrating %s Milvus collections to native hybrid search.', len(pending_collections)) + # Milvus Lite (a local .db file) breaks under concurrent collection changes. + max_workers = 1 if MILVUS_URI.endswith('.db') else BM25_BACKFILL_WORKERS + with ThreadPoolExecutor(max_workers=max_workers) as executor: + copies = [ + executor.submit( + _copy_to_bm25_collection, + client, + collection, + dimension, + create_bm25_collection, + output_fields, + to_bm25_rows, + ) + for collection, dimension in pending_collections.items() + ] + copied = all(copy.result() for copy in copies) + if copied: + swaps = [ + executor.submit(_swap_in_bm25_collection, client, collection) for collection in pending_collections + ] + for swap in swaps: + swap.result() + else: + executor.shutdown(cancel_futures=True) + log.error('Milvus migration to native hybrid search failed, all collections are kept unchanged.') + for collection in pending_collections: + client.drop_collection(f'{collection}{BM25_STAGING_SUFFIX}') + return + log.info('Migrated %s Milvus collections to native hybrid search.', len(pending_collections)) + + +def _swap_in_bm25_collection(client: Client, collection: str): + client.drop_collection(collection) + client.rename_collection(f'{collection}{BM25_STAGING_SUFFIX}', collection) + client.load_collection(collection) + + +def _copy_to_bm25_collection( + client: Client, + collection: str, + dimension: int, + create_bm25_collection: Callable[[str, int], None], + output_fields: list[str], + to_bm25_rows: Callable[[list[dict]], list[dict]], +) -> bool: + staging_collection = f'{collection}{BM25_STAGING_SUFFIX}' + try: + create_bm25_collection(staging_collection, dimension) + was_released = client.get_load_state(collection)['state'] == LoadState.NotLoad + try: + client.load_collection(collection) + iterator = client.query_iterator( + collection_name=collection, + output_fields=output_fields, + batch_size=BM25_BACKFILL_BATCH_SIZE, + ) + with ThreadPoolExecutor(max_workers=BM25_BACKFILL_INSERTS_IN_FLIGHT) as insert_executor: + inserts = [] + while batch := iterator.next(): + if len(inserts) == BM25_BACKFILL_INSERTS_IN_FLIGHT: + inserts.pop(0).result() + inserts.append( + insert_executor.submit( + client.insert, collection_name=staging_collection, data=to_bm25_rows(batch) + ) + ) + for insert in inserts: + insert.result() + iterator.close() + finally: + if was_released: + client.release_collection(collection) + return True + except Exception as e: + log.error('Error copying Milvus collection %s: %s', collection, e) + return False + + class MilvusClient(VectorDBBase): def __init__(self): self.collection_prefix = 'open_webui' @@ -75,6 +242,19 @@ class MilvusClient(VectorDBBase): self.client = Client(uri=MILVUS_URI, db_name=MILVUS_DB) else: self.client = Client(uri=MILVUS_URI, db_name=MILVUS_DB, token=MILVUS_TOKEN) + if ENABLE_DB_MIGRATIONS: + collections = { + collection_name_full.removesuffix(BM25_STAGING_SUFFIX) + for collection_name_full in self.client.list_collections() + if collection_name_full.startswith(f'{self.collection_prefix}_') + } + _backfill_bm25_collections( + self.client, + collections, + self._create_unloaded_collection, + ['id', 'vector', 'data', 'metadata'], + _bm25_rows, + ) def _result_to_get_result(self, result) -> GetResult: ids = [] @@ -86,7 +266,7 @@ class MilvusClient(VectorDBBase): _metadatas = [] for item in match: _ids.append(item.get('id')) - _documents.append(item.get('data', {}).get('text')) + _documents.append(_chunk_text(item)) _metadatas.append(item.get('metadata')) ids.append(_ids) documents.append(_documents) @@ -115,7 +295,7 @@ class MilvusClient(VectorDBBase): # https://milvus.io/docs/de/metric.md _dist = (item.get('distance') + 1.0) / 2.0 _distances.append(_dist) - _documents.append(item.get('entity', {}).get('data', {}).get('text')) + _documents.append(_chunk_text(item.get('entity', {}))) _metadatas.append(item.get('entity', {}).get('metadata')) ids.append(_ids) distances.append(_distances) @@ -131,6 +311,12 @@ class MilvusClient(VectorDBBase): ) def _create_collection(self, collection_name: str, dimension: int): + collection_name_full = f'{self.collection_prefix}_{collection_name}' + self._create_unloaded_collection(collection_name_full, dimension) + self.client.load_collection(collection_name_full) + + def _create_unloaded_collection(self, collection_name_full: str, dimension: int): + supports_bm25 = _supports_bm25(self.client) schema = self.client.create_schema( auto_id=False, enable_dynamic_field=True, @@ -147,7 +333,17 @@ class MilvusClient(VectorDBBase): dim=dimension, description='vector', ) - schema.add_field(field_name='data', datatype=DataType.JSON, description='data') + if supports_bm25: + schema.add_field( + field_name='text', + datatype=DataType.VARCHAR, + max_length=MILVUS_TEXT_MAX_LENGTH, + enable_analyzer=True, + description='text', + ) + _add_bm25_fields(schema) + else: + schema.add_field(field_name='data', datatype=DataType.JSON, description='data') schema.add_field(field_name='metadata', datatype=DataType.JSON, description='metadata') index_params = self.client.prepare_index_params() @@ -192,15 +388,14 @@ class MilvusClient(VectorDBBase): params=index_creation_params, ) - self.client.create_collection( - collection_name=f'{self.collection_prefix}_{collection_name}', - schema=schema, - index_params=index_params, - ) + if supports_bm25: + index_params.add_index(field_name='sparse', index_type='SPARSE_INVERTED_INDEX', metric_type='BM25') + + self.client.create_collection(collection_name=collection_name_full, schema=schema) + self.client.create_index(collection_name=collection_name_full, index_params=index_params) log.info( - "Successfully created collection '%s_%s' with index type '%s' and metric '%s'.", - self.collection_prefix, - collection_name, + "Successfully created collection '%s' with index type '%s' and metric '%s'.", + collection_name_full, index_type, metric_type, ) @@ -233,13 +428,59 @@ class MilvusClient(VectorDBBase): result = self.client.search( collection_name=f'{self.collection_prefix}_{collection_name}', data=vectors, + anns_field='vector', limit=limit, - output_fields=['data', 'metadata'], + output_fields=['data', 'text', 'metadata'], **kwargs, # search_params=search_params # Potentially add later if needed ) return self._result_to_search_result(result) + def hybrid_search( + self, + collection_name: str, + query: str, + vectors: list[list[float | int]], + filter: Optional[dict] = None, + limit: int = 10, + hybrid_bm25_weight: float = 0.5, + ) -> Optional[SearchResult]: + collection_name = collection_name.replace('-', '_') + collection_name_full = f'{self.collection_prefix}_{collection_name}' + if not self.client.has_collection(collection_name_full) or not _has_bm25_field( + self.client, collection_name_full + ): + return None + self.client.load_collection(f'{self.collection_prefix}_{collection_name}') + + vector_result = None + if hybrid_bm25_weight < 1 and vectors: + vector_result = self.search(collection_name=collection_name, vectors=vectors, filter=filter, limit=limit) + + fts_results = [] + if hybrid_bm25_weight > 0 and query.strip(): + metadata_exprs = _metadata_exprs(filter) + result = self.client.search( + collection_name=collection_name_full, + data=[query], + anns_field='sparse', + limit=limit, + filter=' and '.join(metadata_exprs), + output_fields=['text', 'metadata'], + ) + fts_results = [ + {'id': hit['id'], 'text': hit['entity']['text'], 'vmetadata': hit['entity']['metadata']} + for hit in result[0] + ] + + return merge_hybrid_search_results( + vector_result=vector_result, + fts_results=fts_results, + num_queries=len(vectors) or 1, + limit=limit, + hybrid_bm25_weight=hybrid_bm25_weight, + ) + def query(self, collection_name: str, filter: dict, limit: int = -1): collection_name = collection_name.replace('-', '_') if not self.has_collection(collection_name): @@ -272,6 +513,7 @@ class MilvusClient(VectorDBBase): output_fields=[ 'id', 'data', + 'text', 'metadata', ], limit=limit if limit > 0 else -1, @@ -317,20 +559,26 @@ class MilvusClient(VectorDBBase): self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector'])) log.info('Inserting %s items into collection %s_%s.', len(items), self.collection_prefix, collection_name) + has_bm25 = _has_bm25_field(self.client, f'{self.collection_prefix}_{collection_name}') data = [] for item in items: text = item['text'] or '' - if len(text) > MILVUS_TEXT_MAX_LENGTH: - log.warning(f'Milvus: truncating text id={item["id"]} {len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars') - text = text[:MILVUS_TEXT_MAX_LENGTH] - data.append( - { - 'id': item['id'], - 'vector': item['vector'], - 'data': {'text': text}, - 'metadata': process_metadata(item['metadata']), - } - ) + text_bytes = len(text.encode()) + if text_bytes > MILVUS_TEXT_MAX_LENGTH: + log.warning( + 'Milvus: truncating text id=%s %s->%s bytes', item['id'], text_bytes, MILVUS_TEXT_MAX_LENGTH + ) + text = _truncate_text(text) + row = { + 'id': item['id'], + 'vector': item['vector'], + 'metadata': process_metadata(item['metadata']), + } + if has_bm25: + row['text'] = text + else: + row['data'] = {'text': text} + data.append(row) try: return self.client.insert( collection_name=f'{self.collection_prefix}_{collection_name}', @@ -357,20 +605,26 @@ class MilvusClient(VectorDBBase): self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector'])) log.info('Upserting %s items into collection %s_%s.', len(items), self.collection_prefix, collection_name) + has_bm25 = _has_bm25_field(self.client, f'{self.collection_prefix}_{collection_name}') data = [] for item in items: text = item['text'] or '' - if len(text) > MILVUS_TEXT_MAX_LENGTH: - log.warning(f'Milvus: truncating text id={item["id"]} {len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars') - text = text[:MILVUS_TEXT_MAX_LENGTH] - data.append( - { - 'id': item['id'], - 'vector': item['vector'], - 'data': {'text': text}, - 'metadata': process_metadata(item['metadata']), - } - ) + text_bytes = len(text.encode()) + if text_bytes > MILVUS_TEXT_MAX_LENGTH: + log.warning( + 'Milvus: truncating text id=%s %s->%s bytes', item['id'], text_bytes, MILVUS_TEXT_MAX_LENGTH + ) + text = _truncate_text(text) + row = { + 'id': item['id'], + 'vector': item['vector'], + 'metadata': process_metadata(item['metadata']), + } + if has_bm25: + row['text'] = text + else: + row['data'] = {'text': text} + data.append(row) try: return self.client.upsert( collection_name=f'{self.collection_prefix}_{collection_name}', diff --git a/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py b/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py index de0a227ea3..d0924fa9ae 100644 --- a/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py +++ b/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py @@ -17,14 +17,23 @@ from open_webui.config import ( MILVUS_TOKEN, MILVUS_URI, ) -from open_webui.retrieval.vector.dbs.milvus import _metadata_exprs +from open_webui.env import ENABLE_DB_MIGRATIONS +from open_webui.retrieval.vector.dbs.milvus import ( + BM25_STAGING_SUFFIX, + _add_bm25_fields, + _backfill_bm25_collections, + _has_bm25_field, + _metadata_exprs, + _supports_bm25, + _truncate_text, +) from open_webui.retrieval.vector.main import ( GetResult, SearchResult, VectorDBBase, VectorItem, ) -from open_webui.retrieval.vector.utils import process_metadata +from open_webui.retrieval.vector.utils import merge_hybrid_search_results, process_metadata from pymilvus import DataType from pymilvus import MilvusClient as Client from pymilvus.exceptions import MilvusException @@ -81,6 +90,14 @@ class MilvusClient(VectorDBBase): self.WEB_SEARCH_COLLECTION, self.HASH_BASED_COLLECTION, ] + if ENABLE_DB_MIGRATIONS: + _backfill_bm25_collections( + self.client, + self.shared_collections, + self._create_shared_collection, + ['id', 'vector', 'text', 'metadata', RESOURCE_ID_FIELD], + lambda rows: rows, + ) def _get_collection_and_resource_id(self, collection_name: str) -> Tuple[str, str]: """ @@ -107,10 +124,18 @@ class MilvusClient(VectorDBBase): return self.KNOWLEDGE_COLLECTION, resource_id def _create_shared_collection(self, mt_collection_name: str, dimension: int): + supports_bm25 = _supports_bm25(self.client) schema = self.client.create_schema(auto_id=False, description='Shared collection for multi-tenancy') schema.add_field(field_name='id', datatype=DataType.VARCHAR, is_primary=True, max_length=36) schema.add_field(field_name='vector', datatype=DataType.FLOAT_VECTOR, dim=dimension) - schema.add_field(field_name='text', datatype=DataType.VARCHAR, max_length=MILVUS_TEXT_MAX_LENGTH) + schema.add_field( + field_name='text', + datatype=DataType.VARCHAR, + max_length=MILVUS_TEXT_MAX_LENGTH, + enable_analyzer=supports_bm25, + ) + if supports_bm25: + _add_bm25_fields(schema) schema.add_field(field_name='metadata', datatype=DataType.JSON) schema.add_field(field_name=RESOURCE_ID_FIELD, datatype=DataType.VARCHAR, max_length=255) @@ -132,6 +157,17 @@ class MilvusClient(VectorDBBase): self.client.create_collection(collection_name=mt_collection_name, schema=schema) self.client.create_index(collection_name=mt_collection_name, index_params=vector_index) + if supports_bm25: + self.client.create_index( + collection_name=mt_collection_name, + index_params=self.client.prepare_index_params( + field_name='sparse', index_type='SPARSE_INVERTED_INDEX', metric_type='BM25' + ), + ) + self._create_resource_id_index(mt_collection_name) + log.info('Created shared collection: %s', mt_collection_name) + + def _create_resource_id_index(self, mt_collection_name: str): try: # A Milvus server auto-selects the scalar index type from a parameterless call. self.client.create_index( @@ -148,7 +184,6 @@ class MilvusClient(VectorDBBase): # The index only accelerates resource_id filters; never fail # collection creation over it. log.warning(f'Could not create {RESOURCE_ID_FIELD} index on {mt_collection_name}: {e}') - log.info('Created shared collection: %s', mt_collection_name) def _ensure_collection(self, mt_collection_name: str, dimension: int): if not self.client.has_collection(mt_collection_name): @@ -180,13 +215,17 @@ class MilvusClient(VectorDBBase): entities = [] for item in items: text = item['text'] or '' - if len(text) > MILVUS_TEXT_MAX_LENGTH: + text_bytes = len(text.encode()) + if text_bytes > MILVUS_TEXT_MAX_LENGTH: log.warning( - f'Milvus: truncating text id={item["id"]} ' - f'{len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars ' - f'(collection={mt_collection}, resource_id={resource_id})' + 'Milvus: truncating text id=%s %s->%s bytes (collection=%s, resource_id=%s)', + item['id'], + text_bytes, + MILVUS_TEXT_MAX_LENGTH, + mt_collection, + resource_id, ) - text = text[:MILVUS_TEXT_MAX_LENGTH] + text = _truncate_text(text) entities.append( { 'id': item['id'], @@ -250,6 +289,49 @@ class MilvusClient(VectorDBBase): return SearchResult(ids=ids, documents=documents, metadatas=metadatas, distances=distances) + def hybrid_search( + self, + collection_name: str, + query: str, + vectors: List[List[float]], + filter: Optional[Dict] = None, + limit: int = 10, + hybrid_bm25_weight: float = 0.5, + ) -> Optional[SearchResult]: + mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) + _validate_resource_id(resource_id) + if not self.client.has_collection(mt_collection) or not _has_bm25_field(self.client, mt_collection): + return None + + vector_result = None + if hybrid_bm25_weight < 1 and vectors: + vector_result = self.search(collection_name=collection_name, vectors=vectors, filter=filter, limit=limit) + + fts_results = [] + if hybrid_bm25_weight > 0 and query.strip(): + self.client.load_collection(mt_collection) + expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'", *_metadata_exprs(filter)] + results = self.client.search( + collection_name=mt_collection, + data=[query], + anns_field='sparse', + limit=limit, + filter=' and '.join(expr), + output_fields=['text', 'metadata'], + ) + fts_results = [ + {'id': hit['id'], 'text': hit['entity']['text'], 'vmetadata': hit['entity']['metadata']} + for hit in results[0] + ] + + return merge_hybrid_search_results( + vector_result=vector_result, + fts_results=fts_results, + num_queries=len(vectors) or 1, + limit=limit, + hybrid_bm25_weight=hybrid_bm25_weight, + ) + def delete( self, collection_name: str, @@ -278,6 +360,7 @@ class MilvusClient(VectorDBBase): for collection_name in self.shared_collections: if self.client.has_collection(collection_name): self.client.drop_collection(collection_name) + self.client.drop_collection(f'{collection_name}{BM25_STAGING_SUFFIX}') def delete_collection(self, collection_name: str): mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) diff --git a/backend/open_webui/retrieval/vector/dbs/weaviate.py b/backend/open_webui/retrieval/vector/dbs/weaviate.py index 6d398d4ced..8ccb159721 100644 --- a/backend/open_webui/retrieval/vector/dbs/weaviate.py +++ b/backend/open_webui/retrieval/vector/dbs/weaviate.py @@ -212,7 +212,7 @@ class WeaviateClient(VectorDBBase): # Weaviate has cosine distance, 2 (worst) -> 0 (best). Re-ordering to 0 -> 1 raw_distances = [ - (obj.metadata.distance if obj.metadata and obj.metadata.distance else 2.0) + (obj.metadata.distance if obj.metadata and obj.metadata.distance is not None else 2.0) for obj in response.objects ] distances = [(2 - dist) / 2 for dist in raw_distances] diff --git a/backend/open_webui/retrieval/web/utils.py b/backend/open_webui/retrieval/web/utils.py index c2b8be5159..5cd9a8973a 100644 --- a/backend/open_webui/retrieval/web/utils.py +++ b/backend/open_webui/retrieval/web/utils.py @@ -307,6 +307,12 @@ _DROPPED_RESPONSE_HEADERS = {'connection', 'content-encoding', 'content-length', # The Playwright loader only reads the page HTML, which none of these feed. _DROPPED_RESOURCE_TYPES = {'font', 'image', 'media'} +# unstructured keeps only the first
, so text in any others would be dropped. +_UNWRAP_EXTRA_MAINS = ( + '() => { const mains = document.querySelectorAll("main"); ' + 'if (mains.length > 1) mains.forEach(main => main.replaceWith(...main.childNodes)); }' +) + 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} @@ -878,6 +884,7 @@ class SafePlaywrightURLLoader(BaseLoader, RateLimitMixin, URLProcessingMixin): for element in page.locator(selector).all(): if element.is_visible(): element.evaluate('element => element.remove()') + page.evaluate(_UNWRAP_EXTRA_MAINS) text = self._extract_html(page.content()) page.unroute_all(behavior='ignoreErrors') metadata = {'source': url} @@ -918,6 +925,7 @@ class SafePlaywrightURLLoader(BaseLoader, RateLimitMixin, URLProcessingMixin): for element in await page.locator(selector).all(): if await element.is_visible(): await element.evaluate('element => element.remove()') + await page.evaluate(_UNWRAP_EXTRA_MAINS) text = await asyncio.to_thread(self._extract_html, await page.content()) await page.unroute_all(behavior='ignoreErrors') metadata = {'source': url} diff --git a/backend/open_webui/routers/automations.py b/backend/open_webui/routers/automations.py index 4b44a8b8ed..f5ab9b9536 100644 --- a/backend/open_webui/routers/automations.py +++ b/backend/open_webui/routers/automations.py @@ -341,6 +341,7 @@ async def run_automation_by_id( await check_automations_permission(request, user) automation = await Automations.get_by_id(id, db=db) check_automation_access(automation, user) + automation = await Automations.update_last_run_at(automation.id, db=db) asyncio.create_task(execute_automation(request.app, automation)) await publish_event( request, diff --git a/backend/open_webui/routers/calendar.py b/backend/open_webui/routers/calendar.py index 86d9824803..acd434451d 100644 --- a/backend/open_webui/routers/calendar.py +++ b/backend/open_webui/routers/calendar.py @@ -26,6 +26,7 @@ from open_webui.models.users import UserModel from open_webui.utils.access_control import filter_allowed_access_grants, has_permission from open_webui.utils.auth import get_verified_user from open_webui.utils.calendar import expand_recurring_event +from open_webui.utils.recurrence import schedule_start_ns log = logging.getLogger(__name__) @@ -206,13 +207,22 @@ async def get_events( if not rrule_str: continue + start_at = auto.next_run_at or 0 + upper_rrule = rrule_str.upper() + if 'COUNT=' in upper_rrule and 'DTSTART' in upper_rrule: + # COUNT runs from DTSTART, anchoring on the next run would restart it + try: + start_at = schedule_start_ns(rrule_str, user.timezone) + except ValueError: + pass + virtual = { 'id': f'auto_{auto.id}', 'calendar_id': SCHEDULED_TASKS_CALENDAR_ID, 'user_id': user.id, 'title': auto.name, 'description': auto.data.get('prompt', '') if auto.data else '', - 'start_at': auto.next_run_at or 0, + 'start_at': start_at, 'end_at': None, 'all_day': False, 'rrule': rrule_str, diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index c9d38bf2ce..2a16e23be2 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -642,6 +642,9 @@ async def add_members_by_id( if channel.user_id != user.id and user.role != 'admin': raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) + if channel.type == 'dm': + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) + try: memberships = await Channels.add_members_to_channel( channel.id, user.id, form_data.user_ids, form_data.group_ids, db=db @@ -686,9 +689,12 @@ async def remove_members_by_id( if channel.user_id != user.id and user.role != 'admin': raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) + if channel.type == 'dm': + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()) + try: deleted = await Channels.remove_members_from_channel(channel.id, form_data.user_ids, db=db) - if channel.type in ['group', 'dm']: + if channel.type == 'group': await leave_room_for_users(f'channel:{channel.id}', form_data.user_ids) await publish_event( diff --git a/backend/open_webui/routers/configs.py b/backend/open_webui/routers/configs.py index 35f1b23cf4..f48e5cf68f 100644 --- a/backend/open_webui/routers/configs.py +++ b/backend/open_webui/routers/configs.py @@ -7,7 +7,7 @@ import aiohttp from fastapi import APIRouter, Depends, HTTPException, Request from mcp.shared.auth import OAuthMetadata from open_webui.config import BannerModel -from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT +from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT, ENABLE_TOOL_SERVERS from open_webui.events import EVENTS, publish_event from open_webui.models.config import Config from open_webui.models.oauth_sessions import OAuthSessions @@ -176,6 +176,9 @@ async def register_oauth_client( type: str | None = None, user=Depends(get_admin_user), ): + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + try: oauth_client_id = form_data.client_id if type: @@ -264,7 +267,7 @@ async def set_tool_servers_config( await set_tool_servers(request) - for connection in connections: + for connection in connections if ENABLE_TOOL_SERVERS else []: server_type = connection.get('type', 'openapi') if server_type == 'mcp': server_id = (connection.get('info') or {}).get('id') @@ -362,6 +365,9 @@ async def verify_terminal_server_connection( Tries GET {url}/api/v1/policies (orchestrator) then GET {url}/api/config (plain terminal). Returns ``{status: true, type: "orchestrator"|"terminal"}``. """ + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + base_url = (form_data.url or '').rstrip('/') if not base_url: raise HTTPException(status_code=400, detail='Terminal server URL is required') @@ -432,6 +438,9 @@ async def put_terminal_server_policy( request: Request, form_data: TerminalServerPolicyForm, user=Depends(get_admin_user) ): """Proxy a policy read or update to an orchestrator terminal server.""" + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + base_url = (form_data.url or '').rstrip('/') if not base_url: raise HTTPException(status_code=400, detail='Terminal server URL is required') @@ -469,6 +478,9 @@ async def put_terminal_server_lifecycle( request: Request, form_data: TerminalServerLifecycleForm, user=Depends(get_admin_user) ): """Proxy a lifecycle read or update to an orchestrator terminal server.""" + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + base_url = (form_data.url or '').rstrip('/') if not base_url: raise HTTPException(status_code=400, detail='Terminal server URL is required') @@ -508,6 +520,9 @@ async def refresh_terminal_server_terminals( """ Proxy a terminal refresh request to an orchestrator terminal server. """ + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + base_url = (form_data.url or '').rstrip('/') if not base_url: raise HTTPException(status_code=400, detail='Terminal server URL is required') @@ -553,6 +568,9 @@ async def verify_tool_servers_config(request: Request, form_data: ToolServerConn """ Verify the connection to the tool server. """ + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + try: if form_data.type == 'mcp': if form_data.auth_type in ('oauth_2.1', 'oauth_2.1_static'): diff --git a/backend/open_webui/routers/functions.py b/backend/open_webui/routers/functions.py index 212a419500..18fa4ed433 100644 --- a/backend/open_webui/routers/functions.py +++ b/backend/open_webui/routers/functions.py @@ -10,7 +10,7 @@ import aiohttp from fastapi import APIRouter, Depends, HTTPException, Request, status from open_webui.config import CACHE_DIR from open_webui.constants import ERROR_MESSAGES -from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT, ENABLE_PLUGINS +from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT, ENABLE_FUNCTIONS from open_webui.events import EVENTS, build_event, dispatch_event_functions, publish_event, schedule_webhook_dispatch from open_webui.internal.db import get_async_session from open_webui.models.functions import ( @@ -24,8 +24,8 @@ from open_webui.models.functions import ( from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.plugin import ( get_function_contents_cache, - get_functions_cache, get_function_module_from_cache, + get_functions_cache, load_function_module_by_id, replace_imports, resolve_valves_schema_options, @@ -47,7 +47,7 @@ router = APIRouter() @router.get('/', response_model=list[FunctionResponse]) async def get_functions(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): - if not ENABLE_PLUGINS: + if not ENABLE_FUNCTIONS: return [] return await Functions.get_functions(db=db) @@ -55,7 +55,7 @@ async def get_functions(user=Depends(get_verified_user), db: AsyncSession = Depe @router.get('/list', response_model=list[FunctionUserResponse]) async def get_function_list(user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session)): - if not ENABLE_PLUGINS: + if not ENABLE_FUNCTIONS: return [] return await Functions.get_function_list(db=db) @@ -72,7 +72,7 @@ async def get_functions( user=Depends(get_admin_user), db: AsyncSession = Depends(get_async_session), ): - if not ENABLE_PLUGINS: + if not ENABLE_FUNCTIONS: return [] return await Functions.get_functions(include_valves=include_valves, db=db) diff --git a/backend/open_webui/routers/images.py b/backend/open_webui/routers/images.py index c894f2c651..28a037a993 100644 --- a/backend/open_webui/routers/images.py +++ b/backend/open_webui/routers/images.py @@ -7,6 +7,7 @@ import logging import mimetypes import re import uuid +from contextlib import nullcontext from pathlib import Path from types import SimpleNamespace from typing import Optional @@ -488,14 +489,18 @@ async def get_image_data(data: str, headers=None, trusted_base_url: str | None = # that would follow arbitrary redirects. if trusted_base_url and _is_same_origin(data, trusted_base_url): log.debug('Skipping URL validation for trusted backend: %s', data) + session_context = nullcontext(await get_session()) else: await asyncio.to_thread(validate_url, data) - session = await get_session() - async with session.get( - data, - headers=headers, - ssl=AIOHTTP_CLIENT_SESSION_SSL, - ) as r: + session_context = get_ssrf_safe_session() + async with ( + session_context as session, + session.get( + data, + headers=headers, + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) as r, + ): r.raise_for_status() content_type = r.headers.get('content-type', '') if content_type.split('/')[0] == 'image': diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index b8a75c2bfb..66351100cd 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -1109,7 +1109,8 @@ def get_azure_allowed_params(api_version: str) -> set[str]: def is_openai_new_model(model: str) -> bool: - model_lower = model.lower() + # Amazon Bedrock ids carry a provider prefix, e.g. us.openai.gpt-6-sol + model_lower = re.sub(r'^(?:[a-z-]+\.)?openai\.', '', model.lower()) # o-series models (o1, o3, o4, o5, ...) if re.match(r'^o\d+', model_lower): return True diff --git a/backend/open_webui/routers/scim.py b/backend/open_webui/routers/scim.py index a2d093dba8..7301e9066b 100644 --- a/backend/open_webui/routers/scim.py +++ b/backend/open_webui/routers/scim.py @@ -116,6 +116,8 @@ class SCIMPhoto(BaseModel): class SCIMGroupMember(BaseModel): """SCIM Group Member""" + model_config = ConfigDict(populate_by_name=True) + value: str # User ID ref: Optional[str] = Field(None, alias='$ref') type: Optional[str] = 'User' diff --git a/backend/open_webui/routers/terminals.py b/backend/open_webui/routers/terminals.py index 6b5cd6e1d3..447bede9d3 100644 --- a/backend/open_webui/routers/terminals.py +++ b/backend/open_webui/routers/terminals.py @@ -14,7 +14,7 @@ import aiohttp from fastapi import APIRouter, Depends, Request, Response, WebSocket from fastapi.responses import JSONResponse, StreamingResponse from open_webui.config import TERMINAL_PROXY_HEADERS -from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL +from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, ENABLE_TOOL_SERVERS from open_webui.events import EVENTS, publish_event from open_webui.models.config import Config from open_webui.models.groups import Groups @@ -87,6 +87,9 @@ def _sanitize_proxy_path(path: str) -> str | None: @router.get('/') async def list_terminal_servers(request: Request, user=Depends(get_verified_user)): """Return terminal servers the authenticated user has access to.""" + if not ENABLE_TOOL_SERVERS: + return [] + connections = await Config.get('terminal_server.connections', []) or [] user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} @@ -114,6 +117,9 @@ async def proxy_terminal( user=Depends(get_verified_user), ): """Proxy a request to the admin terminal server identified by *server_id*.""" + if not ENABLE_TOOL_SERVERS: + return JSONResponse({'error': 'Tool servers are disabled'}, status_code=403) + connections = await Config.get('terminal_server.connections', []) or [] connection = next((c for c in connections if c.get('id') == server_id), None) @@ -288,6 +294,10 @@ async def _resolve_authenticated_connection(ws: WebSocket, server_id: str): async def _resolve_terminal_access(ws: WebSocket, server_id: str, token: str): """Resolve current access for both the handshake and an open terminal session.""" + if not ENABLE_TOOL_SERVERS: + await ws.close(code=4003, reason='Tool servers are disabled') + return None + try: user = await get_verified_user_by_token(token, getattr(ws.app.state, 'redis', None)) if user is None: diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index 12c3771420..2191f82db0 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -10,7 +10,12 @@ import aiohttp from fastapi import APIRouter, Depends, HTTPException, Request, status from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, CACHE_DIR from open_webui.constants import ERROR_MESSAGES -from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL, AIOHTTP_CLIENT_TIMEOUT, ENABLE_PLUGINS +from open_webui.env import ( + AIOHTTP_CLIENT_SESSION_SSL, + AIOHTTP_CLIENT_TIMEOUT, + ENABLE_TOOL_SERVERS, + ENABLE_TOOLS, +) from open_webui.events import EVENTS, publish_event from open_webui.internal.db import get_async_session from open_webui.models.access_grants import AccessGrants @@ -33,8 +38,8 @@ from open_webui.utils.access_control import ( from open_webui.utils.auth import get_admin_user, get_verified_user from open_webui.utils.plugin import ( get_tool_contents_cache, - get_tools_cache, get_tool_module_from_cache, + get_tools_cache, load_tool_module_by_id, replace_imports, resolve_valves_schema_options, @@ -78,7 +83,7 @@ async def get_tools( ) # Local Tools - if ENABLE_PLUGINS: + if ENABLE_TOOLS: tools_cache = get_tools_cache(request) for tool in await Tools.get_tools( defer_content=True, @@ -132,7 +137,7 @@ async def get_tools( ) # MCP Tool Servers - for server in await Config.get('tool_server.connections', []): + for server in (await Config.get('tool_server.connections', [])) if ENABLE_TOOL_SERVERS else []: if server.get('type', 'openapi') == 'mcp' and (server.get('config') or {}).get('enable'): info = server.get('info') or {} server_id = info.get('id') @@ -198,7 +203,7 @@ async def get_tools( @router.get('/list', response_model=list[ToolAccessResponse]) async def get_tool_list(user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): - if not ENABLE_PLUGINS: + if not ENABLE_TOOLS: return [] bypass_access_control = user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index f157ed2f81..b901c658c4 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -36,6 +36,7 @@ from open_webui.env import ( from open_webui.models.access_grants import AccessGrants from open_webui.models.channels import Channels from open_webui.models.chats import Chats +from open_webui.models.config import Config from open_webui.models.folders import Folders from open_webui.models.notes import Notes, NoteUpdateForm from open_webui.models.users import UserNameResponse, Users @@ -392,8 +393,12 @@ async def enter_room_for_users(room: str, user_ids: list[str]): user_ids (list[str]): The target user's IDs. """ try: - for user_id in user_ids: - session_ids = get_session_ids_from_room(f'user:{user_id}') + default_permissions = await Config.get('user.permissions') + for user in await Users.get_users_by_user_ids(user_ids): + if user.role != 'admin' and not await has_permission(user.id, 'features.channels', default_permissions): + continue + + session_ids = get_session_ids_from_room(f'user:{user.id}') for sid in session_ids: await sio.enter_room(sid, room) except Exception as e: @@ -501,7 +506,7 @@ async def user_join(sid, data): await sio.enter_room(sid, f'user:{user.id}') # Join all the channels only if user has channels permission - if user.role == 'admin' or await has_permission(user.id, 'features.channels'): + if user.role == 'admin' or await has_permission(user.id, 'features.channels', await Config.get('user.permissions')): channels = await Channels.get_channels_by_user_id(user.id) log.debug('channels=%r', channels) for channel in channels: @@ -533,7 +538,7 @@ async def join_channel(sid, data): return # Join all the channels only if user has channels permission - if user.role == 'admin' or await has_permission(user.id, 'features.channels'): + if user.role == 'admin' or await has_permission(user.id, 'features.channels', await Config.get('user.permissions')): channels = await Channels.get_channels_by_user_id(user.id) log.debug('channels=%r', channels) for channel in channels: @@ -554,6 +559,11 @@ async def join_note(sid, data): if not user: return + if user.role != 'admin' and not await has_permission( + user.id, 'features.notes', await Config.get('user.permissions') + ): + return + note = await Notes.get_note_by_id(data['note_id']) if not note: log.error(f'Note {data["note_id"]} not found for user {user.id}') @@ -687,6 +697,11 @@ async def ydoc_document_join(sid, data): document_id = normalize_document_id(data['document_id']) if document_id.startswith('note:'): + if user.get('role') != 'admin' and not await has_permission( + user.get('id'), 'features.notes', await Config.get('user.permissions') + ): + return + note_id = document_id.split(':')[1] note = await Notes.get_note_by_id(note_id) if not note: diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index c1d1b164d6..c974c68b46 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -4413,10 +4413,10 @@ async def update_calendar_event( return JSONCodec.dumps({'error': 'Event not found'}) # Check write access to the event's calendar - if event.user_id != user_id and __user__.get('role') != 'admin': - cal = await Calendars.get_calendar_by_id(event.calendar_id) - if not cal: - return JSONCodec.dumps({'error': 'Access denied'}) + cal = await Calendars.get_calendar_by_id(event.calendar_id) + if not cal: + return JSONCodec.dumps({'error': 'Access denied'}) + if cal.user_id != user_id and __user__.get('role') != 'admin': user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id)] if not await AccessGrants.has_access( user_id=user_id, @@ -4517,10 +4517,10 @@ async def delete_calendar_event( return JSONCodec.dumps({'error': 'Event not found'}) # Check write access - if event.user_id != user_id and __user__.get('role') != 'admin': - cal = await Calendars.get_calendar_by_id(event.calendar_id) - if not cal: - return JSONCodec.dumps({'error': 'Access denied'}) + cal = await Calendars.get_calendar_by_id(event.calendar_id) + if not cal: + return JSONCodec.dumps({'error': 'Access denied'}) + if cal.user_id != user_id and __user__.get('role') != 'admin': user_group_ids = [g.id for g in await Groups.get_groups_by_member_id(user_id)] if not await AccessGrants.has_access( user_id=user_id, diff --git a/backend/open_webui/utils/actions.py b/backend/open_webui/utils/actions.py index 1249dd60cb..a39d44bda5 100644 --- a/backend/open_webui/utils/actions.py +++ b/backend/open_webui/utils/actions.py @@ -4,7 +4,7 @@ import sys from typing import Any from fastapi import Request -from open_webui.env import ENABLE_PLUGINS, GLOBAL_LOG_LEVEL +from open_webui.env import ENABLE_FUNCTIONS, GLOBAL_LOG_LEVEL from open_webui.models.functions import Functions from open_webui.models.users import UserModel from open_webui.socket.main import get_event_call, get_event_emitter @@ -17,8 +17,8 @@ log = logging.getLogger(__name__) async def chat_action(request: Request, action_id: str, form_data: dict, user: Any): - if not ENABLE_PLUGINS: - raise Exception('Plugins are disabled by ENABLE_PLUGINS=false') + if not ENABLE_FUNCTIONS: + raise Exception('Functions are disabled by ENABLE_PLUGINS or ENABLE_FUNCTIONS') if '.' in action_id: action_id, sub_action_id = action_id.split('.') diff --git a/backend/open_webui/utils/audit.py b/backend/open_webui/utils/audit.py index 7354b2c641..177036acb6 100644 --- a/backend/open_webui/utils/audit.py +++ b/backend/open_webui/utils/audit.py @@ -282,11 +282,12 @@ class AuditLoggingMiddleware: response_body = context.response_body.decode('utf-8', errors='replace') # Redact sensitive information - if 'password' in request_body: + if 'password' in request_body.lower(): request_body = re.sub( - r'"password":\s*"(.*?)"', - '"password": "********"', + r'"(\w*password)":\s*".*?"', + r'"\1": "********"', request_body, + flags=re.IGNORECASE, ) entry = AuditLogEntry( diff --git a/backend/open_webui/utils/automations.py b/backend/open_webui/utils/automations.py index be5d6868f1..d6c0f4dfe4 100644 --- a/backend/open_webui/utils/automations.py +++ b/backend/open_webui/utils/automations.py @@ -261,6 +261,13 @@ async def _execute_channel_automation( if not channel_id or not await Config.get('channels.enable'): raise ValueError('Channel not found') + from open_webui.utils.access_control import has_permission + + if user.role != 'admin' and not await has_permission( + user.id, 'features.channels', await Config.get('user.permissions') + ): + raise ValueError('Owner no longer permitted to use channels') + model = getattr(app.state, 'MODELS', {}).get(model_id, {}) request = _build_request(app, token=token) diff --git a/backend/open_webui/utils/calendar.py b/backend/open_webui/utils/calendar.py index 0406db38f6..b7815da628 100644 --- a/backend/open_webui/utils/calendar.py +++ b/backend/open_webui/utils/calendar.py @@ -40,7 +40,6 @@ def expand_recurring_event( range_start = to_local_datetime(range_start_ns) range_end = to_local_datetime(range_end_ns) - scan_start = range_start - dt.timedelta(days=1) original_start_ns = event_dict['start_at'] original_start = to_local_datetime(original_start_ns) @@ -55,24 +54,30 @@ def expand_recurring_event( original_end_ns = event_dict.get('end_at') duration_ns = (original_end_ns - original_start_ns) if original_end_ns else None + # Look back by the event length so occurrences still running at range start are found + event_length = dt.timedelta(microseconds=max(duration_ns or 0, 0) // 1000) + scan_start = range_start - dt.timedelta(days=1) - event_length instances = [] previous_start = None - for occurrence_start in rule.xafter(scan_start, count=max_instances, inc=True): + for occurrence_start in rule.xafter(scan_start, inc=True): if occurrence_start >= range_end or occurrence_start == previous_start: break previous_start = occurrence_start instance_start_ns = int(occurrence_start.replace(tzinfo=user_timezone).timestamp() * 1_000_000_000) + instance_end_ns = (instance_start_ns + duration_ns) if duration_ns else None - if instance_start_ns >= range_start_ns: + if instance_start_ns >= range_start_ns or (instance_end_ns and instance_end_ns > range_start_ns): instance = { **event_dict, 'start_at': instance_start_ns, - 'end_at': (instance_start_ns + duration_ns) if duration_ns else None, + 'end_at': instance_end_ns, 'instance_id': f'{event_dict["id"]}_{instance_start_ns}', } instances.append(instance) + if len(instances) >= max_instances: + break return instances diff --git a/backend/open_webui/utils/chat_variables.py b/backend/open_webui/utils/chat_variables.py index c1776f333f..aea222ddcd 100644 --- a/backend/open_webui/utils/chat_variables.py +++ b/backend/open_webui/utils/chat_variables.py @@ -185,7 +185,7 @@ def validate_user_variables(variables: Any) -> dict[str, str]: raise ChatVariablesError('User variables must be an object.') try: - if len(JSONCodec.dumps(variables)) > MAX_VARIABLES_JSON_LENGTH: + if len(JSONCodec.dumps(variables, ensure_ascii=False)) > MAX_VARIABLES_JSON_LENGTH: raise ChatVariablesError('User variables are too large.') except TypeError: raise ChatVariablesError('User variables must be JSON serializable.') @@ -214,7 +214,7 @@ def validate_chat_variables( variables = normalize_chat_variables(variables) try: - if len(JSONCodec.dumps(variables)) > MAX_VARIABLES_JSON_LENGTH: + if len(JSONCodec.dumps(variables, ensure_ascii=False)) > MAX_VARIABLES_JSON_LENGTH: raise ChatVariablesError('Chat variables are too large.') except TypeError: raise ChatVariablesError('Chat variables must be JSON serializable.') diff --git a/backend/open_webui/utils/filter.py b/backend/open_webui/utils/filter.py index c4d9754e39..3bc5346b1d 100644 --- a/backend/open_webui/utils/filter.py +++ b/backend/open_webui/utils/filter.py @@ -1,7 +1,7 @@ import inspect import logging -from open_webui.env import ENABLE_PLUGINS +from open_webui.env import ENABLE_FUNCTIONS from open_webui.models.functions import Functions from open_webui.utils.plugin import get_function_module_from_cache @@ -66,7 +66,7 @@ def get_model_filter_ids(model, active_filters): async def resolve_filter_pipeline(request, model: dict, enabled_filter_ids: list = None): - if not ENABLE_PLUGINS: + if not ENABLE_FUNCTIONS: return [], [] active_filters = await get_filter_context(request).get_active_filters() @@ -217,7 +217,7 @@ async def process_filter_functions( form_data, extra_params, ): - if not ENABLE_PLUGINS: + if not ENABLE_FUNCTIONS: return form_data, {} skip_files = None diff --git a/backend/open_webui/utils/mcp/client.py b/backend/open_webui/utils/mcp/client.py index 29e67f04fb..35256b89b2 100644 --- a/backend/open_webui/utils/mcp/client.py +++ b/backend/open_webui/utils/mcp/client.py @@ -1,6 +1,7 @@ import asyncio import logging from contextlib import AsyncExitStack +from datetime import timedelta from typing import Optional log = logging.getLogger(__name__) @@ -116,7 +117,11 @@ class MCPClient: if not self.session: raise RuntimeError('MCP client is not connected.') - result = await self.session.call_tool(function_name, function_args) + tool_call_timeout = None + if AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER is not None and AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER > 0: + tool_call_timeout = timedelta(seconds=AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER) + + result = await self.session.call_tool(function_name, function_args, read_timeout_seconds=tool_call_timeout) if not result: raise Exception('No result returned from MCP tool call.') diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 55f036ab5a..fbe45b4b2e 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -37,10 +37,11 @@ from open_webui.env import ( ENABLE_API_OUTLET_FILTERS, ENABLE_CHAT_RESPONSE_BASE64_IMAGE_URL_CONVERSION, ENABLE_CHAT_RESPONSE_STREAM_INPLACE_APPEND, - ENABLE_PLUGINS, + ENABLE_FUNCTIONS, ENABLE_QUERIES_CACHE, ENABLE_REALTIME_CHAT_SAVE, ENABLE_RESPONSES_API_STATEFUL, + ENABLE_TOOL_SERVERS, GLOBAL_LOG_LEVEL, RAG_SYSTEM_CONTEXT, ) @@ -116,8 +117,8 @@ from open_webui.utils.misc import ( get_message_list, get_output_text, get_paired_tool_call_ids, - get_response_error_detail, get_reasoning_details, + get_response_error_detail, get_system_message, is_raster_image_content_type, is_string_allowed, @@ -2331,6 +2332,10 @@ async def connect_mcp_server( Returns None if the server is not found or access is denied. """ + if not ENABLE_TOOL_SERVERS: + log.debug('MCP resolution skipped: external plugins are disabled') + return None + mcp_server_connection = None for server_connection in await Config.get('tool_server.connections', []): if server_connection.get('type', '') == 'mcp' and (server_connection.get('info') or {}).get('id') == server_id: @@ -2649,8 +2654,8 @@ async def process_chat_payload(request, form_data, user, metadata, model): raise e filter_functions = [] - filter_context = get_filter_context(request) if ENABLE_PLUGINS else None - if ENABLE_PLUGINS: + filter_context = get_filter_context(request) if ENABLE_FUNCTIONS else None + if ENABLE_FUNCTIONS: try: filter_functions = await get_filter_functions(request, model, metadata.get('filter_ids', [])) @@ -2748,6 +2753,11 @@ async def process_chat_payload(request, form_data, user, metadata, model): tool_ids = form_data.pop('tool_ids', None) terminal_id = form_data.pop('terminal_id', None) + if not ENABLE_TOOL_SERVERS: + if terminal_id or metadata.get('tool_servers'): + log.debug('Excluded external plugins disabled by plugin configuration') + terminal_id = None + metadata['tool_servers'] = None files = form_data.pop('files', None) form_data.pop('folder_id', None) metadata['terminal_id'] = terminal_id @@ -2927,7 +2937,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): mcp_tools_dict = {} if tool_ids: - db_tool_ids = [] + resolved_tool_ids = [] for tool_id in tool_ids: if tool_id.startswith('server:mcp:'): try: @@ -2979,13 +2989,13 @@ async def process_chat_payload(request, form_data, user, metadata, model): } ) continue - elif ENABLE_PLUGINS: - db_tool_ids.append(tool_id) + else: + resolved_tool_ids.append(tool_id) - if db_tool_ids: + if resolved_tool_ids: tools_dict = await get_tools( request, - db_tool_ids, + resolved_tool_ids, user, { **extra_params, @@ -3199,7 +3209,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): } ) - if ENABLE_PLUGINS: + if ENABLE_FUNCTIONS: try: form_data, _ = await process_filter_functions( request=request, @@ -3527,7 +3537,7 @@ async def drain_approved_tool_calls(request, form_data, user, model, metadata) - ) form_data['messages'] = sanitize_tool_pairs(form_data['messages']) - if not paused and ENABLE_PLUGINS: + if not paused and ENABLE_FUNCTIONS: filter_functions = await get_filter_functions(request, model, metadata.get('filter_ids', [])) if filter_functions: filtered_form_data, _ = await process_filter_functions( @@ -4026,7 +4036,7 @@ async def outlet_filter_handler(ctx): is_unsaved_chat = not is_saved_chat_id(chat_id) try: filter_functions = ( - await get_filter_functions(request, model, metadata.get('filter_ids', [])) if ENABLE_PLUGINS else [] + await get_filter_functions(request, model, metadata.get('filter_ids', [])) if ENABLE_FUNCTIONS else [] ) model_id = model.get('id') if isinstance(model, dict) else model models = request.app.state.MODELS @@ -4408,7 +4418,7 @@ async def streaming_chat_response_handler(response, ctx): } filter_functions = ( - await get_filter_functions(request, model, metadata.get('filter_ids', [])) if ENABLE_PLUGINS else [] + await get_filter_functions(request, model, metadata.get('filter_ids', [])) if ENABLE_FUNCTIONS else [] ) # Standard streaming response handler @@ -4760,6 +4770,9 @@ async def streaming_chat_response_handler(response, ctx): 'content': [{'type': 'output_text', 'text': initial_content}], } ] + if not continuing: + # the client needs this item before the first delta arrives + await event_emitter({'type': 'chat:completion', 'data': {'output': output}}) else: output = [] diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index bba1729a44..778441de48 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -8,22 +8,22 @@ from open_webui.config import ( BYPASS_ADMIN_ACCESS_CONTROL, DEFAULT_ARENA_MODEL, ) -from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, ENABLE_PLUGINS, GLOBAL_LOG_LEVEL, REDIS_KEY_PREFIX +from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, ENABLE_FUNCTIONS, GLOBAL_LOG_LEVEL, REDIS_KEY_PREFIX from open_webui.functions import get_function_models from open_webui.models.access_grants import AccessGrants from open_webui.models.config import Config from open_webui.models.functions import Functions from open_webui.models.groups import Groups from open_webui.models.models import Models -from open_webui.utils.chat_variables import get_chat_variables_schema from open_webui.models.users import UserModel from open_webui.routers import ollama, openai from open_webui.socket.utils import RedisDict from open_webui.utils.access_control import has_access, has_base_model_access +from open_webui.utils.chat_variables import get_chat_variables_schema from open_webui.utils.json_codec import JSONCodec from open_webui.utils.plugin import ( - get_functions_cache, get_function_module_from_cache, + get_functions_cache, ) logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) @@ -150,7 +150,7 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) # One query per type: the global sets are subsets of the active sets, so # deriving them from the same rows halves the function-table queries. - if ENABLE_PLUGINS: + if ENABLE_FUNCTIONS: active_actions = await Functions.get_active_function_ids_by_type('action') global_action_ids = {function_id for function_id, is_global in active_actions if is_global} enabled_action_ids = {function_id for function_id, _ in active_actions} @@ -205,7 +205,7 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) if 'info' in model: if 'meta' in model['info']: - if ENABLE_PLUGINS: + if ENABLE_FUNCTIONS: action_ids.extend(model['info']['meta'].get('actionIds', [])) filter_ids.extend(model['info']['meta'].get('filterIds', [])) @@ -265,10 +265,10 @@ async def get_all_models(request, refresh: bool = False, user: UserModel = None) if custom_model.meta: meta = custom_model.meta.model_dump() - if ENABLE_PLUGINS and 'actionIds' in meta: + if ENABLE_FUNCTIONS and 'actionIds' in meta: action_ids.extend(meta['actionIds']) - if ENABLE_PLUGINS and 'filterIds' in meta: + if ENABLE_FUNCTIONS and 'filterIds' in meta: filter_ids.extend(meta['filterIds']) model['action_ids'] = action_ids diff --git a/backend/open_webui/utils/notifications.py b/backend/open_webui/utils/notifications.py index 8b6f090b4c..a35ad9b88f 100644 --- a/backend/open_webui/utils/notifications.py +++ b/backend/open_webui/utils/notifications.py @@ -268,9 +268,6 @@ def _notification_webhook_content(event: Any) -> tuple[str, str, dict[str, Any], title = str(data.get('title') or event.message or 'Chat finished') content = str(data.get('message') or '') url = str(data.get('url') or '') - chat_id = str(data.get('chat_id') or '') - if chat_id and url.endswith(f'/c/{chat_id}'): - url = f'{url[: -len(f"/c/{chat_id}")].rstrip("/")}/{chat_id}' body = '\n'.join(part for part in (content, url) if part) return ( f'**{title}**', @@ -288,9 +285,6 @@ def _notification_webhook_content(event: Any) -> tuple[str, str, dict[str, Any], title = str(event.message or 'Chat failed') content = str(data.get('message') or '') url = str(data.get('url') or '') - chat_id = str(data.get('chat_id') or '') - if chat_id and url.endswith(f'/c/{chat_id}'): - url = f'{url[: -len(f"/c/{chat_id}")].rstrip("/")}/{chat_id}' body = '\n'.join(part for part in (content, url) if part) return ( f'**{title}**', diff --git a/backend/open_webui/utils/oauth.py b/backend/open_webui/utils/oauth.py index 061737a35b..9155787ec7 100644 --- a/backend/open_webui/utils/oauth.py +++ b/backend/open_webui/utils/oauth.py @@ -70,6 +70,7 @@ from open_webui.env import ( AIOHTTP_CLIENT_SESSION_SSL, ENABLE_OAUTH_EMAIL_FALLBACK, ENABLE_OAUTH_ID_TOKEN_COOKIE, + ENABLE_TOOL_SERVERS, OAUTH_CLIENT_INFO_ENCRYPTION_KEY, OAUTH_MAX_SESSIONS_PER_USER, REDIS_KEY_PREFIX, @@ -85,8 +86,8 @@ from open_webui.models.users import Users from open_webui.retrieval.web.utils import get_ssrf_safe_session, validate_url from open_webui.utils.auth import ( create_token, - get_password_hash, get_optional_verified_user_from_request, + get_password_hash, get_verified_user_by_id, revoke_user_tokens, ) @@ -789,6 +790,11 @@ def should_send_oauth_resource(client_info: OAuthClientInformationFull | None) - return not scope_has_resource_indicator(client_info.scope) +def uses_google_authorization_server(client_info: OAuthClientInformationFull) -> bool: + server_metadata = client_info.server_metadata + return server_metadata is not None and server_metadata.authorization_endpoint.host == 'accounts.google.com' + + def build_oauth_request_params(client_info: OAuthClientInformationFull | None) -> dict: if not client_info: return {} @@ -798,6 +804,10 @@ def build_oauth_request_params(client_info: OAuthClientInformationFull | None) - params['scope'] = client_info.scope if should_send_oauth_resource(client_info): params['resource'] = client_info.resource + # Google only issues a refresh token for offline access, and only re-issues it on a fresh consent. + if uses_google_authorization_server(client_info): + params['access_type'] = 'offline' + params['prompt'] = 'consent' return params @@ -887,6 +897,9 @@ class OAuthClientManager: Lazy-load an OAuth client from the current TOOL_SERVER_CONNECTIONS config if it hasn't been registered on this node yet. """ + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + if client_id in self.clients: return self.clients[client_id]['client'] @@ -1019,6 +1032,9 @@ class OAuthClientManager: return True async def get_client(self, client_id): + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + if client_id not in self.clients: await self.ensure_client_from_config(client_id) @@ -1026,6 +1042,9 @@ class OAuthClientManager: return client['client'] if client else None async def get_client_info(self, client_id): + if not ENABLE_TOOL_SERVERS: + raise HTTPException(status_code=403, detail='Tool servers are disabled') + if client_id not in self.clients: await self.ensure_client_from_config(client_id) @@ -1917,7 +1936,7 @@ class OAuthManager: detailed_error, exc_info=True, ) - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) except Exception as e: detailed_error = _build_oauth_callback_error_message(e) log.warning( @@ -1926,7 +1945,7 @@ class OAuthManager: detailed_error, exc_info=True, ) - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) # Try to get userinfo from the token first, some providers include it there user_data: UserInfo = token.get('userinfo') @@ -1949,7 +1968,7 @@ class OAuthManager: user_data = user_data['data'] if not user_data: log.warning('OAuth callback failed for provider %s, user data is missing', provider) - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) # Extract the "sub" claim, using custom claim if configured if auth_config.OAUTH_SUB_CLAIM: @@ -1959,7 +1978,7 @@ class OAuthManager: sub = user_data.get(OAUTH_PROVIDERS[provider].get('sub_claim', 'sub')) if not sub: log.warning(f'OAuth callback failed, sub is missing: {user_data}') - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) sub = str(sub) oauth_data = {} @@ -1994,18 +2013,18 @@ class OAuthManager: email = primary_email else: log.warning('No primary email found in GitHub response') - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) else: log.warning('Failed to fetch GitHub email') - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) except Exception as e: log.warning(f'Error fetching GitHub email: {e}') - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) elif ENABLE_OAUTH_EMAIL_FALLBACK: email = f'{provider}@{sub}.local' else: log.warning(f'OAuth callback failed, email is missing: {user_data}') - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) email = email.lower() # If allowed domains are configured, check if the email domain is in the list @@ -2014,7 +2033,7 @@ class OAuthManager: and email.split('@')[-1] not in auth_config.OAUTH_ALLOWED_DOMAINS ): log.warning(f'OAuth callback failed, e-mail domain is not in the list of allowed domains: {user_data}') - raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED) + raise HTTPException(400, detail=ERROR_MESSAGES.OAUTH_LOGIN_FAILED) # Check if the user exists user = await Users.get_user_by_oauth_sub(provider, sub, db=db) diff --git a/backend/open_webui/utils/plugin.py b/backend/open_webui/utils/plugin.py index 1c6be3858f..cc9dc27b47 100644 --- a/backend/open_webui/utils/plugin.py +++ b/backend/open_webui/utils/plugin.py @@ -12,8 +12,9 @@ from importlib import util from typing import Any from open_webui.env import ( + ENABLE_FUNCTIONS, ENABLE_PIP_INSTALL_FRONTMATTER_REQUIREMENTS, - ENABLE_PLUGINS, + ENABLE_TOOLS, OFFLINE_MODE, PIP_OPTIONS, PIP_PACKAGE_INDEX_OPTIONS, @@ -204,8 +205,8 @@ def replace_imports(content): # May the intent of the one who wrote it survive every # import and transformation, as a deed survives the generations. async def load_tool_module_by_id(tool_id, content=None): - if not ENABLE_PLUGINS: - raise RuntimeError('Plugins are disabled by ENABLE_PLUGINS=false') + if not ENABLE_TOOLS: + raise RuntimeError('Tools are disabled by ENABLE_PLUGINS or ENABLE_TOOLS') frontmatter = None if content is None: @@ -257,8 +258,8 @@ async def load_tool_module_by_id(tool_id, content=None): async def load_function_module_by_id(function_id: str, content: str | None = None): - if not ENABLE_PLUGINS: - raise RuntimeError('Plugins are disabled by ENABLE_PLUGINS=false') + if not ENABLE_FUNCTIONS: + raise RuntimeError('Functions are disabled by ENABLE_PLUGINS or ENABLE_FUNCTIONS') frontmatter = None if content is None: @@ -338,6 +339,9 @@ def get_function_contents_cache(request) -> dict: async def get_tool_module_from_cache(request, tool_id, load_from_db=True): + if not ENABLE_TOOLS: + raise RuntimeError('Tools are disabled by ENABLE_PLUGINS or ENABLE_TOOLS') + tools_cache = get_tools_cache(request) tool_contents_cache = get_tool_contents_cache(request) content = None @@ -375,6 +379,9 @@ async def get_tool_module_from_cache(request, tool_id, load_from_db=True): async def get_function_module_from_cache( request, function_id, function: FunctionModel | None = None, load_from_db=True ): + if not ENABLE_FUNCTIONS: + raise RuntimeError('Functions are disabled by ENABLE_PLUGINS or ENABLE_FUNCTIONS') + functions_cache = get_functions_cache(request) function_contents_cache = get_function_contents_cache(request) content = None @@ -458,12 +465,12 @@ async def install_tool_and_function_dependencies(): and then installing them using pip. Duplicates or similar version specifications are handled by pip as much as possible. """ - if not ENABLE_PLUGINS: - log.info('ENABLE_PLUGINS is disabled, skipping tool and function dependencies.') + if not ENABLE_TOOLS and not ENABLE_FUNCTIONS: + log.info('Tools and Functions are disabled, skipping their dependencies.') return - function_list = await Functions.get_functions(active_only=True) - tool_list = await Tools.get_tools() + function_list = await Functions.get_functions(active_only=True) if ENABLE_FUNCTIONS else [] + tool_list = await Tools.get_tools() if ENABLE_TOOLS else [] all_dependencies = '' try: diff --git a/backend/open_webui/utils/recurrence.py b/backend/open_webui/utils/recurrence.py index c4817eca07..717953cc6b 100644 --- a/backend/open_webui/utils/recurrence.py +++ b/backend/open_webui/utils/recurrence.py @@ -148,6 +148,18 @@ async def next_run_ns(s: str, tz: str = None) -> Optional[int]: return int(dt.timestamp() * 1_000_000_000) +def schedule_start_ns(s: str, tz: str = None) -> int: + """DTSTART the scheduler anchors the rule to, as epoch nanoseconds.""" + zi = _resolve_tz(tz) + now = datetime.now(zi).replace(tzinfo=None) if zi else datetime.now() + parsed = _parse_rule(s, now) + rule = parsed._rrule[0] if isinstance(parsed, rruleset) else parsed + dt = rule._dtstart + if zi: + dt = dt.replace(tzinfo=zi) + return int(dt.timestamp() * 1_000_000_000) + + async def next_n_runs_ns(s: str, n: int = 5, tz: str = None) -> list[int]: """Compute next N occurrences for UI preview. diff --git a/backend/open_webui/utils/subagents.py b/backend/open_webui/utils/subagents.py index 13b9699216..b2f9693af8 100644 --- a/backend/open_webui/utils/subagents.py +++ b/backend/open_webui/utils/subagents.py @@ -219,6 +219,7 @@ async def process_pending_internal_messages( history['messages'] = messages history['currentId'] = assistant_message_id chat.chat = {**(chat.chat or {}), 'history': history} + chat.current_message_id = assistant_message_id chat.updated_at = int(time.time()) await db.commit() @@ -614,10 +615,16 @@ async def delegate( updated_chat = copy.deepcopy(parent.chat or {}) updated_history = updated_chat.setdefault('history', {}) updated_messages = updated_history.setdefault('messages', {}) + parent_message = updated_messages.get(parent_message_id) done_assistants = [ message - for message in updated_messages.values() - if message.get('role') == 'assistant' and message.get('done') is not False + for message_id, message in updated_messages.items() + if message.get('role') == 'assistant' + and message.get('done') is not False + and ( + parent_message is None + or any(entry is parent_message for entry in get_message_list(updated_messages, message_id)) + ) ] result_parent_id = ( max(done_assistants, key=lambda message: message.get('timestamp', 0)).get('id') diff --git a/backend/open_webui/utils/telemetry/instrumentors.py b/backend/open_webui/utils/telemetry/instrumentors.py index 7daf0a8354..3f72339a9d 100644 --- a/backend/open_webui/utils/telemetry/instrumentors.py +++ b/backend/open_webui/utils/telemetry/instrumentors.py @@ -178,7 +178,7 @@ class Instrumentor(BaseInstrumentor): SQLAlchemyInstrumentor().instrument(engine=self.db_engine) RedisInstrumentor().instrument(request_hook=redis_request_hook) RequestsInstrumentor().instrument(request_hook=requests_hook, response_hook=response_hook) - LoggingInstrumentor().instrument() + LoggingInstrumentor().instrument(enable_log_auto_instrumentation=False) HTTPXClientInstrumentor().instrument( request_hook=httpx_request_hook, response_hook=httpx_response_hook, diff --git a/backend/open_webui/utils/terminals.py b/backend/open_webui/utils/terminals.py index 4dbd7e6b3f..f2c945ba63 100644 --- a/backend/open_webui/utils/terminals.py +++ b/backend/open_webui/utils/terminals.py @@ -6,6 +6,7 @@ import ntpath import posixpath from urllib.parse import quote +from open_webui.env import ENABLE_TOOL_SERVERS from open_webui.utils.chat_id import is_saved_chat_id TERMINAL_CONTEXT_HEADER = 'X-Terminal-Context-Id' @@ -122,8 +123,10 @@ def terminal_chat_uploads(connection: dict) -> str: async def get_terminal_json(request, user, metadata: dict, path: str, extra_params: dict | None = None): """Read from an admin terminal on the backend or a personal terminal in its browser.""" - import aiohttp + if not ENABLE_TOOL_SERVERS: + return None + import aiohttp from open_webui.env import AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL, AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER_DATA from open_webui.models.config import Config from open_webui.models.groups import Groups diff --git a/backend/open_webui/utils/timers.py b/backend/open_webui/utils/timers.py index 15a6b898db..c2a91139ef 100644 --- a/backend/open_webui/utils/timers.py +++ b/backend/open_webui/utils/timers.py @@ -351,6 +351,7 @@ async def execute_due_timer(app, timer_id: str, claim_id: str | None = None) -> parent.chat = parent_chat history['currentId'] = assistant_message_id + parent.current_message_id = assistant_message_id parent.updated_at = int(time.time()) timer_row = await db.get(Chat, timer_id) if timer_row: diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 7bd6a8dd69..7edd7f401a 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -34,7 +34,8 @@ from open_webui.env import ( AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER, AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER_DATA, ENABLE_FORWARD_USER_INFO_HEADERS, - ENABLE_PLUGINS, + ENABLE_TOOL_SERVERS, + ENABLE_TOOLS, FORWARD_SESSION_INFO_HEADER_CHAT_ID, FORWARD_SESSION_INFO_HEADER_MESSAGE_ID, REDIS_KEY_PREFIX, @@ -101,7 +102,7 @@ from open_webui.tools.builtin import ( write_note, ) from open_webui.utils.access_control import has_access, has_connection_access, has_permission -from open_webui.utils.chat_id import is_saved_chat_id +from open_webui.utils.chat_id import is_saved_chat_id, is_temporary_chat_id from open_webui.utils.headers import ( bearer_auth_header, get_custom_headers, @@ -266,9 +267,17 @@ async def get_updated_tool_function(function: Callable, extra_params: dict): async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extra_params: dict) -> dict[str, dict]: """Load tools for the given tool_ids, checking access control.""" - if not ENABLE_PLUGINS: + if not tool_ids: return {} + enabled_ids = [ + tool_id + for tool_id in tool_ids + if (ENABLE_TOOL_SERVERS if tool_id.startswith('server:') else ENABLE_TOOLS) + ] + if len(enabled_ids) != len(tool_ids): + log.debug('Excluded tools disabled by plugin configuration') + tool_ids = enabled_ids if not tool_ids: return {} @@ -278,7 +287,8 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)} # Batch-fetch all DB tools in one query instead of one per tool_id - tool_models = await Tools.get_tools_by_ids(tool_ids) + local_tool_ids = [tool_id for tool_id in tool_ids if not tool_id.startswith('server:')] + tool_models = await Tools.get_tools_by_ids(local_tool_ids) if local_tool_ids else {} for tool_id in tool_ids: tool = tool_models.get(tool_id) @@ -650,6 +660,7 @@ async def get_builtin_tools( and config.get('subagents.enable') and getattr(request.state, 'internal', False) is not True and getattr(request.state, 'direct', False) is not True + and not is_temporary_chat_id(metadata.get('chat_id')) ): builtin_functions.extend([delegate_task, timer]) @@ -1168,6 +1179,9 @@ def convert_openapi_to_tool_payload(openapi_spec): async def set_tool_servers(request: Request): + if not ENABLE_TOOL_SERVERS: + return [] + try: request.app.state.TOOL_SERVERS = await get_tool_servers_data(await Config.get('tool_server.connections', [])) except Exception as e: @@ -1186,6 +1200,9 @@ async def set_tool_servers(request: Request): async def get_tool_servers(request: Request): + if not ENABLE_TOOL_SERVERS: + return [] + try: tool_servers = None if request.app.state.redis is not None: @@ -1270,6 +1287,9 @@ async def get_terminal_system_prompt( async def set_terminal_servers(request: Request): """Load and cache OpenAPI specs from all TERMINAL_SERVER_CONNECTIONS.""" + if not ENABLE_TOOL_SERVERS: + return [] + connections = await Config.get('terminal_server.connections', []) or [] # Build server configs compatible with get_tool_servers_data @@ -1332,6 +1352,9 @@ async def set_terminal_servers(request: Request): async def get_terminal_servers(request: Request): """Return cached terminal server specs, loading if needed.""" + if not ENABLE_TOOL_SERVERS: + return [] + terminal_servers = None if request.app.state.redis is not None: try: @@ -1367,6 +1390,9 @@ async def get_terminal_tools( - Loads specs from cache - Builds callables that route through the terminal proxy """ + if not ENABLE_TOOL_SERVERS: + return {} + connections = await Config.get('terminal_server.connections', []) or [] connection = next( (terminal_connection for terminal_connection in connections if terminal_connection.get('id') == terminal_id), @@ -1471,6 +1497,9 @@ async def get_terminal_tools( async def get_tool_server_data(url: str, headers: dict | None) -> dict[str, Any]: + if not ENABLE_TOOL_SERVERS: + raise RuntimeError('Tool servers are disabled') + _headers = { 'Accept': 'application/json', 'Content-Type': 'application/json', @@ -1519,6 +1548,9 @@ async def get_tool_server_data(url: str, headers: dict | None) -> dict[str, Any] async def get_tool_servers_data(servers: list[dict[str, Any]]) -> list[dict[str, Any]]: # Prepare list of enabled servers along with their original index + if not ENABLE_TOOL_SERVERS: + return [] + tasks = [] server_entries = [] for idx, server in enumerate(servers): @@ -1625,6 +1657,9 @@ async def execute_tool_server( params: dict[str, Any], server_data: dict[str, Any], ) -> tuple[dict[str, Any], dict[str, Any | None]]: + if not ENABLE_TOOL_SERVERS: + raise RuntimeError('Tool servers are disabled') + error = None try: openapi = server_data.get('openapi', {}) diff --git a/backend/requirements-slim.txt b/backend/requirements-slim.txt index c688c89d23..e5f83c4faf 100644 --- a/backend/requirements-slim.txt +++ b/backend/requirements-slim.txt @@ -36,7 +36,7 @@ aiosqlite==0.22.1 psycopg[binary]==3.3.4 alembic==1.18.4 -pycrdt==0.13.1 +pycrdt==0.14.8 redis==8.0.1 hiredis==3.4.2 diff --git a/backend/requirements.txt b/backend/requirements.txt index adbc0e560e..d3050f6464 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -32,7 +32,7 @@ aiosqlite==0.22.1 psycopg[binary]==3.3.4 alembic==1.18.4 -pycrdt==0.13.1 +pycrdt==0.14.8 redis==8.0.1 hiredis==3.4.2 diff --git a/docs/SECURITY.md b/docs/SECURITY.md index 16536236da..4eec9614f1 100644 --- a/docs/SECURITY.md +++ b/docs/SECURITY.md @@ -130,7 +130,7 @@ Your remediation guidance can include, for example: > Similar to rule "Default Configuration Testing": If you believe you have found a vulnerability that affects admins and is NOT caused by admin negligence or intentionally malicious actions, > **then we absolutely want to hear about it.** This policy is intended to filter social engineering attacks on admins, malicious plugins being deployed by admins and similar malicious actions, not to discourage legitimate security research. -10. **Tools & Functions Code Execution Is Intended Behavior:** Open WebUI's Tools and Functions feature is **designed** to execute user-provided Python code on the server. This is core, intentional functionality — not a vulnerability (see also 'Threat Model Understanding'). Function creation is **restricted to administrators only**. Tool creation is controlled by the `workspace.tools` permission, which is **disabled by default** for non-admin users and should only be granted to fully trusted users who are equivalent to system administrators in terms of trust. **Granting a user the ability to create Tools is equivalent to giving them shell access to the server**. If an administrator grants this permission to untrusted users, this constitutes intentional misconfiguration and is additionally covered by 'Admin Actions Are Out of Scope'. Deployments that do not need `workspace.tools` or Functions plugin execution can set `ENABLE_PLUGINS=false`. More generally, **reports describing ANY attack chain that involves Tools or Functions — including but not limited to code execution, file access, network requests, or environment variable access — will be closed as not a vulnerability / intended behavior.** This applies to both direct code execution and frontmatter-based package installation (`pip install`). +10. **Tools & Functions Code Execution Is Intended Behavior:** Open WebUI's Tools and Functions feature is **designed** to execute user-provided Python code on the server. This is core, intentional functionality — not a vulnerability (see also 'Threat Model Understanding'). Function creation is **restricted to administrators only**. Tool creation is controlled by the `workspace.tools` permission, which is **disabled by default** for non-admin users and should only be granted to fully trusted users who are equivalent to system administrators in terms of trust. **Granting a user the ability to create Tools is equivalent to giving them shell access to the server**. If an administrator grants this permission to untrusted users, this constitutes intentional misconfiguration and is additionally covered by 'Admin Actions Are Out of Scope'. Deployments that do not need internal Tools or Functions execution can set `ENABLE_TOOLS=false` and `ENABLE_FUNCTIONS=false` while keeping external plugins available. `ENABLE_PLUGINS=false` is the master switch and disables both internal and external plugins, including OpenAPI/MCP tool servers and Open Terminal. `ENABLE_TOOL_SERVERS=false` disables only external plugins. All four settings default to `true`, are environment-only, and require a restart; the master always overrides all feature switches. These controls do not disable built-in tools, the code interpreter, external knowledge, or model-provider connections, which have their own controls. More generally, **reports describing ANY attack chain that involves Tools or Functions — including but not limited to code execution, file access, network requests, or environment variable access — will be closed as not a vulnerability / intended behavior.** This applies to both direct code execution and frontmatter-based package installation (`pip install`). > [!IMPORTANT] > **For administrators:** Treat the `workspace.tools` permission as **root-equivalent access**. Only grant it to users you would trust with direct access to your server. If you enable this permission for untrusted users, you are accepting the risk of arbitrary code execution on your host. For more details, see our [Plugin Security documentation](https://docs.openwebui.com/features/extensibility/plugin/). diff --git a/pyproject.toml b/pyproject.toml index e410b2b788..776dc418ce 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,7 +38,7 @@ dependencies = [ "psycopg[binary]==3.3.4", "alembic==1.18.4", - "pycrdt==0.13.1", + "pycrdt==0.14.8", "redis==8.0.1", "hiredis==3.4.2", # "valkey-glide-sync==2.3.1", # optional: install manually if VECTOR_DB=valkey diff --git a/src/lib/apis/index.ts b/src/lib/apis/index.ts index 9f9ed0b837..1aa65e9a2a 100644 --- a/src/lib/apis/index.ts +++ b/src/lib/apis/index.ts @@ -1,3 +1,5 @@ +import { get } from 'svelte/store'; +import { config } from '$lib/stores'; import { WEBUI_BASE_URL } from '$lib/constants'; import { convertOpenApiToToolPayload, resolveSchema } from '$lib/utils'; import { normalizeTags } from '$lib/utils/tags'; @@ -392,6 +394,7 @@ export const getTaskIdsByChatId = async (token: string, chat_id: string) => { }; export const getToolServerData = async (token: string, url: string) => { + if (!get(config)?.features?.enable_tool_servers) throw new Error('Tool servers are disabled'); let error = null; const res = await fetch(`${url}`, { @@ -435,6 +438,7 @@ export const getToolServerData = async (token: string, url: string) => { }; export const getToolServersData = async (servers: object[]) => { + if (!get(config)?.features?.enable_tool_servers) return []; return ( await Promise.all( servers @@ -546,6 +550,7 @@ export const executeToolServer = async ( serverData: { openapi: any; info: any; specs: any }, sessionId?: string ) => { + if (!get(config)?.features?.enable_tool_servers) throw new Error('Tool servers are disabled'); let error = null; try { diff --git a/src/lib/apis/terminal/index.ts b/src/lib/apis/terminal/index.ts index 2bdbfeab21..dae872b6cc 100644 --- a/src/lib/apis/terminal/index.ts +++ b/src/lib/apis/terminal/index.ts @@ -1,3 +1,5 @@ +import { get } from 'svelte/store'; +import { config } from '$lib/stores'; export type FileEntry = { name: string; type: 'file' | 'directory'; @@ -98,7 +100,7 @@ export const resolveTerminalConnection = ( directServers: any[], token: string ): TerminalConnection | null => { - if (!selector) return null; + if (!get(config)?.features?.enable_tool_servers || !selector) return null; if (servers.some((server) => server.id === selector)) { return { selector, @@ -119,6 +121,7 @@ export const terminalRequest = async ( path: string, options: RequestInit = {} ): Promise => { + if (!get(config)?.features?.enable_tool_servers) throw new Error('Tool servers are disabled'); const response = await fetch(`${connection.baseUrl}${path}`, { ...options, headers: { @@ -131,9 +134,10 @@ export const terminalRequest = async ( return response.json(); }; -const bearerHeaders = (apiKey: string): Record => ({ - Authorization: `Bearer ${apiKey.trim()}` -}); +const bearerHeaders = (apiKey: string): Record => { + if (!get(config)?.features?.enable_tool_servers) throw new Error('Tool servers are disabled'); + return { Authorization: `Bearer ${apiKey.trim()}` }; +}; export const joinTerminalPath = (base: string, child: string) => { if (!child) return base; @@ -162,6 +166,7 @@ export type TerminalServer = { }; export const getTerminalServers = async (token: string): Promise => { + if (!get(config)?.features?.enable_tool_servers) return []; const res = await fetch(`${WEBUI_API_BASE_URL}/terminals/`, { headers: { Authorization: `Bearer ${token}` @@ -622,6 +627,7 @@ export const getListeningPorts = async ( }; export const getPortProxyUrl = (baseUrl: string, port: number, path: string = ''): string => { + if (!get(config)?.features?.enable_tool_servers) return ''; return `${baseUrl.replace(/\/$/, '')}/proxy/${port}/${path}`; }; diff --git a/src/lib/components/AddConnectionModal.svelte b/src/lib/components/AddConnectionModal.svelte index 0ada5f55c3..4a9a112fa2 100644 --- a/src/lib/components/AddConnectionModal.svelte +++ b/src/lib/components/AddConnectionModal.svelte @@ -219,6 +219,7 @@ } headers = JSON.stringify(_headers, null, 2); } catch (error) { + loading = false; toast.error($i18n.t('Headers must be a valid JSON object')); return; } diff --git a/src/lib/components/admin/Analytics/Dashboard.svelte b/src/lib/components/admin/Analytics/Dashboard.svelte index 423269a985..22fea3afc7 100644 --- a/src/lib/components/admin/Analytics/Dashboard.svelte +++ b/src/lib/components/admin/Analytics/Dashboard.svelte @@ -18,6 +18,7 @@ import { WEBUI_API_BASE_URL } from '$lib/constants'; import { formatNumber } from '$lib/utils'; import { goto } from '$app/navigation'; + import type { Instance } from 'tippy.js'; const i18n = getContext('i18n'); @@ -83,6 +84,7 @@ { input_tokens: number; output_tokens: number; total_tokens: number } > = {}; let totalTokens = { input: 0, output: 0, total: 0 }; + let tokenTooltip: Instance | null = null; let loading = true; @@ -232,6 +234,12 @@ } + { + if (event.key === 'Escape') tokenTooltip?.hide(); + }} +/> +

{$i18n.t('Analytics')} @@ -291,8 +299,25 @@ > {$i18n.t('messages')} - - document.body, + aria: { content: 'describedby', expanded: true }, + onCreate: (instance) => (tokenTooltip = instance), + onDestroy: () => (tokenTooltip = null) + }} + content={`
+
${$i18n.t('Input')}${totalTokens.input.toLocaleString()}
+
${$i18n.t('Output')}${totalTokens.output.toLocaleString()}
+
${$i18n.t('Token counts are estimates and may not reflect actual API usage')}
+
`} + > + {formatNumber(totalTokens.total)} diff --git a/src/lib/components/admin/Settings/Images.svelte b/src/lib/components/admin/Settings/Images.svelte index 831667244e..7e3f03f04c 100644 --- a/src/lib/components/admin/Settings/Images.svelte +++ b/src/lib/components/admin/Settings/Images.svelte @@ -176,6 +176,26 @@ const saveHandler = async () => { loading = true; + if ( + typeof config?.IMAGES_OPENAI_API_PARAMS === 'string' && + config.IMAGES_OPENAI_API_PARAMS.trim() !== '' && + !validateJSON(config.IMAGES_OPENAI_API_PARAMS) + ) { + toast.error($i18n.t('Invalid JSON format for Parameters')); + loading = false; + return; + } + + if ( + typeof config?.AUTOMATIC1111_PARAMS === 'string' && + config.AUTOMATIC1111_PARAMS.trim() !== '' && + !validateJSON(config.AUTOMATIC1111_PARAMS) + ) { + toast.error($i18n.t('Invalid JSON format for Parameters')); + loading = false; + return; + } + if (config?.COMFYUI_WORKFLOW) { if (!validateJSON(config?.COMFYUI_WORKFLOW)) { toast.error($i18n.t('Invalid JSON format for ComfyUI Workflow.')); diff --git a/src/lib/components/admin/Settings/Integrations.svelte b/src/lib/components/admin/Settings/Integrations.svelte index a6e63a997b..9a6e943a22 100644 --- a/src/lib/components/admin/Settings/Integrations.svelte +++ b/src/lib/components/admin/Settings/Integrations.svelte @@ -186,155 +186,160 @@
{#if servers !== null && connectionsConfig !== null} - -
-
-
- {$i18n.t('settings.admin.integrations.externalToolServers.label')} + +
+
+
+ {$i18n.t('settings.admin.integrations.externalToolServers.label')} +
+ + + +
- - - -
- -
- {#each servers ?? [] as server, idx} - { - updateHandler(); - }} - onDelete={() => { - servers = (servers ?? []).filter((_, i) => i !== idx); - updateHandler(); - }} - /> - {/each} -
- - {#if (servers ?? []).length === 0} -
- {$i18n.t('No tool server connections configured.')} -
- {/if} - -
- {$i18n.t('Connect to your own OpenAPI compatible external tool servers.')} -
-
- - - -
-
-
- {$i18n.t('settings.admin.integrations.openTerminal.label')} +
+ {#each servers ?? [] as server, idx} + { + updateHandler(); + }} + onDelete={() => { + servers = (servers ?? []).filter((_, i) => i !== idx); + updateHandler(); + }} + /> + {/each}
- - - + {#if (servers ?? []).length === 0} +
+ {$i18n.t('No tool server connections configured.')} +
+ {/if} + +
+ {$i18n.t('Connect to your own OpenAPI compatible external tool servers.')} +
+ -
- {#each terminalConnections as connection, idx} -
- -
-
- - - + +
+
+
+ {$i18n.t('settings.admin.integrations.openTerminal.label')} +
+ + + +
+ +
+ {#each terminalConnections as connection, idx} +
+ +
- {connection.name || connection.url || $i18n.t('New Terminal')} + + + + +
+ {connection.name || connection.url || $i18n.t('New Terminal')} +
-
- + -
- - + + + - - - - - - { - terminalConnections = terminalConnections.map((c, i) => - i === idx ? { ...c, enabled: !(c?.enabled !== false) } : c - ); - saveTerminalServers(); - }} - /> - + { + terminalConnections = terminalConnections.map((c, i) => + i === idx ? { ...c, enabled: !(c?.enabled !== false) } : c + ); + saveTerminalServers(); + }} + /> + +
-
- {/each} -
- - {#if terminalConnections.length === 0} -
- {$i18n.t('No terminal connections configured.')} + {/each}
- {/if} -
- {$i18n.t( - 'Connect to Open Terminal instances. Admins and users granted access can use file browsing and terminal tools through these servers.' - )} + {#if terminalConnections.length === 0} +
+ {$i18n.t('No terminal connections configured.')} +
+ {/if} + +
+ {$i18n.t( + 'Connect to Open Terminal instances. Admins and users granted access can use file browsing and terminal tools through these servers.' + )} +
+ {$i18n.t('Learn more about Open Terminal')} ↗
- {$i18n.t('Learn more about Open Terminal')} ↗ -
- + + @@ -349,6 +354,7 @@ let:labelId > - {#if $config?.features?.enable_plugins} + {#if $config?.features?.enable_tools}
loadRuns(false), 2000); } diff --git a/src/lib/components/channel/Channel.svelte b/src/lib/components/channel/Channel.svelte index 0117b27232..a2d2d335ec 100644 --- a/src/lib/components/channel/Channel.svelte +++ b/src/lib/components/channel/Channel.svelte @@ -116,6 +116,7 @@ messages = null; channel = null; threadId = null; + replyToMessage = null; typingUsers = []; typingUsersTimeout = {}; diff --git a/src/lib/components/channel/MessageInput.svelte b/src/lib/components/channel/MessageInput.svelte index ff432ae89a..b9cfaef618 100644 --- a/src/lib/components/channel/MessageInput.svelte +++ b/src/lib/components/channel/MessageInput.svelte @@ -9,6 +9,7 @@ import { config, mobile, settings, socket, user } from '$lib/stores'; import { convertHeicToJpeg, + isHeicImage, compressImage, extractInputVariables, getAge, @@ -24,7 +25,7 @@ import { getSessionUser } from '$lib/apis/auths'; import { uploadFile } from '$lib/apis/files'; - import { WEBUI_API_BASE_URL } from '$lib/constants'; + import { WEBUI_API_BASE_URL, PASTED_TEXT_CHARACTER_LIMIT } from '$lib/constants'; import { getSuggestionRenderer } from '../common/RichTextInput/suggestions'; import CommandSuggestionList from '../chat/MessageInput/CommandSuggestionList.svelte'; @@ -377,7 +378,7 @@ return; } - if (file['type'].startsWith('image/')) { + if (file['type'].startsWith('image/') || isHeicImage(file)) { const compressImageHandler = async (imageUrl, settings = {}, config = {}) => { // Quick shortcut so we don’t do unnecessary work. const settingsCompression = @@ -415,6 +416,7 @@ return imageUrl; }; + const imageFile = isHeicImage(file) ? await convertHeicToJpeg(file) : file; let reader = new FileReader(); reader.onload = async (event) => { @@ -424,12 +426,12 @@ imageUrl = await compressImageHandler(imageUrl, $settings, $config); const blob = await (await fetch(imageUrl)).blob(); - const compressedFile = new File([blob], file.name, { type: file.type }); + const compressedFile = new File([blob], imageFile.name, { type: imageFile.type }); uploadFileHandler(compressedFile, false); }; - reader.readAsDataURL(file['type'] === 'image/heic' ? await convertHeicToJpeg(file) : file); + reader.readAsDataURL(imageFile); } else { uploadFileHandler(file); } @@ -971,10 +973,26 @@ if (clipboardData && clipboardData.items) { for (const item of clipboardData.items) { - const file = item.getAsFile(); - if (file) { - await inputFilesHandler([file]); - e.preventDefault(); + if (item.type === 'text/plain') { + if ($settings?.largeTextAsFile ?? false) { + const text = clipboardData.getData('text/plain'); + + if (text.length > PASTED_TEXT_CHARACTER_LIMIT) { + e.preventDefault(); + const blob = new Blob([text], { type: 'text/plain' }); + const file = new File([blob], `Pasted_Text_${Date.now()}.txt`, { + type: 'text/plain' + }); + + await uploadFileHandler(file); + } + } + } else { + const file = item.getAsFile(); + if (file) { + await inputFilesHandler([file]); + e.preventDefault(); + } } } } diff --git a/src/lib/components/chat/Artifacts.svelte b/src/lib/components/chat/Artifacts.svelte index d0a6e195ee..98b6604b03 100644 --- a/src/lib/components/chat/Artifacts.svelte +++ b/src/lib/components/chat/Artifacts.svelte @@ -190,7 +190,7 @@